diff --git a/solstone/think/providers/local.py b/solstone/think/providers/local.py index 22518cf03..c7ce6d7f3 100644 --- a/solstone/think/providers/local.py +++ b/solstone/think/providers/local.py @@ -13,16 +13,18 @@ import asyncio import contextlib import copy import logging +import re import time import traceback import uuid from collections.abc import Callable from dataclasses import dataclass -from typing import Any +from typing import Any, Literal from solstone.think.models import LOCAL_MODEL from solstone.think.providers._image import encode_image_part, is_image_part from solstone.think.providers.local_endpoint import ( + ENDPOINT_ERROR_BODY_CAP_CHARS, LOCAL_ENDPOINT_CONTRACT_COPY, LOCAL_ENDPOINT_UNREACHABLE_COPY, classify_byo_cogitate_error, @@ -30,6 +32,7 @@ from solstone.think.providers.local_endpoint import ( is_byo_network_error, local_endpoint_reason_copy, redact_local_endpoint_credential, + resolve_endpoint_served_window, resolve_local_endpoint, ) from solstone.think.providers.shared import ( @@ -73,6 +76,18 @@ _LOCAL_UNSUPPORTED_FINISH_REASONS = frozenset({"tool_calls", "function_call"}) _LOCAL_CAPACITY_EXHAUSTED_MESSAGE = ( "The local model was busy and could not finish this request. Try again in a moment." ) +_ENDPOINT_CONTEXT_WINDOW_MESSAGE = ( + "The configured endpoint rejected the request: prompt and completion exceed " + "the served context window." +) +_ENDPOINT_MIN_COMPLETION_TOKENS = 256 +_ENDPOINT_RECLAMP_SLACK_TOKENS = 16 +_ENDPOINT_COMPLETION_ANCHOR = "tokens for the completion" +_ENDPOINT_LIMIT_RE = re.compile(r"maximum context length of\s+(?P\d+)\s+tokens") +_ENDPOINT_INPUT_RE = re.compile( + r"(?P\d+)\s+tokens?\s+from\s+the\s+input\s+messages?\s+and\s+" + r"\d+\s+tokens?\s+for\s+the\s+completion" +) @dataclass(frozen=True) @@ -127,6 +142,12 @@ class LocalCapacityExhausted(LocalProviderError): super().__init__("local_capacity_exhausted", _LOCAL_CAPACITY_EXHAUSTED_MESSAGE) +@dataclass(frozen=True) +class _EndpointOverflowDecision: + kind: Literal["retry", "context", "budget", "contract"] + max_tokens: int | None = None + + def normalize_model_id(model: str | None) -> str: model_id = str(model or LOCAL_MODEL) if model_id.startswith("openai/"): @@ -248,6 +269,7 @@ def _build_request_body( json_output: bool, json_schema: dict | None, is_bundled: bool, + is_confidential: bool = False, ) -> dict[str, Any]: body: dict[str, Any] = { "model": model_id, @@ -256,7 +278,7 @@ def _build_request_body( "max_tokens": max_output_tokens, "stream": False, } - if is_bundled: + if is_bundled or is_confidential: body.update( { "chat_template_kwargs": {"enable_thinking": False}, @@ -280,6 +302,31 @@ def _build_request_body( return body +def _count_image_parts(value: Any) -> int: + if is_image_part(value): + return 1 + if isinstance(value, dict): + return sum(_count_image_parts(item) for item in value.values()) + if isinstance(value, list | tuple): + return sum(_count_image_parts(item) for item in value) + return 0 + + +def _serialized_message_text(messages: list[dict[str, Any]]) -> str: + text_parts: list[str] = [] + for message in messages: + content = message.get("content") + if isinstance(content, str): + text_parts.append(content) + elif isinstance(content, list): + for part in content: + if isinstance(part, dict) and part.get("type") == "text": + text = part.get("text") + if isinstance(text, str): + text_parts.append(text) + return "\n".join(text_parts) + + def _extract_usage(data: dict[str, Any]) -> dict[str, int] | None: usage = data.get("usage") if not isinstance(usage, dict): @@ -435,7 +482,10 @@ def _telemetry_record( return record -def _classify_byo_generate_error(exc: BaseException) -> LocalProviderError: +def _classify_byo_generate_error( + exc: BaseException, + endpoint: Any, +) -> LocalProviderError: if is_byo_capacity_error(exc): return LocalCapacityExhausted() if is_byo_network_error(exc): @@ -443,6 +493,15 @@ def _classify_byo_generate_error(exc: BaseException) -> LocalProviderError: "local_endpoint_unreachable", LOCAL_ENDPOINT_UNREACHABLE_COPY, ) + response = getattr(exc, "response", None) + body_text = getattr(response, "text", None) + if isinstance(body_text, str) and body_text: + excerpt = redact_local_endpoint_credential( + body_text[:ENDPOINT_ERROR_BODY_CAP_CHARS], + endpoint, + ) + if _contains_any(excerpt.lower(), _CONTEXT_WINDOW_PATTERNS): + return _context_window_exceeded_error() return LocalProviderError( "local_endpoint_contract_failed", LOCAL_ENDPOINT_CONTRACT_COPY, @@ -460,6 +519,41 @@ def _remaining_timeout(started: float, timeout_s: float) -> float: return remaining +def _context_window_exceeded_error() -> LocalProviderError: + return LocalProviderError( + "context_window_exceeded", + _ENDPOINT_CONTEXT_WINDOW_MESSAGE, + ) + + +def _endpoint_overflow_decision( + body_text: str, + served_window: int | None, + attempt: int, +) -> _EndpointOverflowDecision: + body_lower = body_text.lower() + if _ENDPOINT_COMPLETION_ANCHOR in body_lower: + limit_match = _ENDPOINT_LIMIT_RE.search(body_lower) + input_match = _ENDPOINT_INPUT_RE.search(body_lower) + limit = ( + int(limit_match.group("limit")) + if limit_match is not None + else served_window + ) + if limit is not None and input_match is not None: + reported_input = int(input_match.group("input")) + new_max = limit - reported_input - _ENDPOINT_RECLAMP_SLACK_TOKENS + if attempt == 0 and new_max >= _ENDPOINT_MIN_COMPLETION_TOKENS: + return _EndpointOverflowDecision("retry", new_max) + if attempt == 0: + return _EndpointOverflowDecision("budget") + return _EndpointOverflowDecision("context") + + if _contains_any(body_lower, _CONTEXT_WINDOW_PATTERNS): + return _EndpointOverflowDecision("context") + return _EndpointOverflowDecision("contract") + + def _prepare_bundled_request( *, server: Any, @@ -496,6 +590,101 @@ def _prepare_bundled_request( ) +def _prepare_endpoint_request( + *, + endpoint: Any, + served_window: int | None, + contents: str | list[Any], + system_instruction: str | None, + temperature: float, + max_output_tokens: int, + json_output: bool, + json_schema: dict | None, +) -> tuple[dict[str, Any], dict[str, Any] | None, dict[str, int | None]]: + from solstone.think.providers import local_budget + + if served_window is None: + messages = _build_messages(contents, system_instruction) + return ( + _build_request_body( + endpoint.served_model_id, + messages, + temperature, + max_output_tokens, + json_output, + json_schema, + endpoint.is_bundled, + endpoint.is_confidential, + ), + None, + { + "served_window": None, + "estimated_prompt_tokens": None, + "clamped_max_tokens": max_output_tokens, + "requested_max_output_tokens": max_output_tokens, + }, + ) + + fitted_contents, input_budget = local_budget.fit_contents( + contents, + system_instruction, + max_output_tokens, + count=local_budget.estimate_tokens, + window=served_window, + ) + messages = _build_messages(fitted_contents, system_instruction) + estimated_prompt_tokens = local_budget.estimate_tokens( + _serialized_message_text(messages) + ) + local_budget._ESTIMATED_IMAGE_TOKENS * _count_image_parts(fitted_contents) + room = served_window - estimated_prompt_tokens - local_budget._SAFETY_MARGIN_TOKENS + if room < _ENDPOINT_MIN_COMPLETION_TOKENS: + raise ContextBudgetExceeded( + "Local endpoint request prompt content exceeds the served context window." + ) + clamped_max_tokens = min(max_output_tokens, room) + return ( + _build_request_body( + endpoint.served_model_id, + messages, + temperature, + clamped_max_tokens, + json_output, + json_schema, + endpoint.is_bundled, + endpoint.is_confidential, + ), + input_budget, + { + "served_window": served_window, + "estimated_prompt_tokens": estimated_prompt_tokens, + "clamped_max_tokens": clamped_max_tokens, + "requested_max_output_tokens": max_output_tokens, + }, + ) + + +def _prepare_endpoint_request_with_resolution( + *, + endpoint: Any, + contents: str | list[Any], + system_instruction: str | None, + temperature: float, + max_output_tokens: int, + json_output: bool, + json_schema: dict | None, +) -> tuple[dict[str, Any], dict[str, Any] | None, dict[str, int | None]]: + return _prepare_endpoint_request( + endpoint=endpoint, + served_window=resolve_endpoint_served_window(endpoint), + contents=contents, + system_instruction=system_instruction, + temperature=temperature, + max_output_tokens=max_output_tokens, + json_output=json_output, + json_schema=json_schema, + ) + + def _raise_bundled_status(response: Any) -> None: import httpx @@ -644,15 +833,14 @@ def run_generate( ) raise - messages = _build_messages(contents, system_instruction) - body = _build_request_body( - endpoint.served_model_id, - messages, - temperature, - max_output_tokens, - json_output, - json_schema, - endpoint.is_bundled, + body, input_budget, endpoint_budget = _prepare_endpoint_request_with_resolution( + endpoint=endpoint, + contents=contents, + system_instruction=system_instruction, + temperature=temperature, + max_output_tokens=max_output_tokens, + json_output=json_output, + json_schema=json_schema, ) import httpx @@ -686,19 +874,59 @@ def run_generate( try: with admission: base_url = confidential_egress_base_url(endpoint.base_url) - response = httpx.post( - f"{base_url}/v1/chat/completions", - timeout=post_timeout, - **post_kwargs, - ) - response.raise_for_status() - return _parse_response(response.json()) + attempt = 0 + while True: + response = httpx.post( + f"{base_url}/v1/chat/completions", + timeout=( + post_timeout + if attempt == 0 + else _remaining_timeout(started, timeout) + ), + **post_kwargs, + ) + try: + response.raise_for_status() + break + except httpx.HTTPStatusError as exc: + if response.status_code != 400: + raise + decision = _endpoint_overflow_decision( + response.text, + endpoint_budget.get("served_window"), + attempt, + ) + if decision.kind == "retry" and decision.max_tokens is not None: + post_kwargs["json"] = { + **post_kwargs["json"], + "max_tokens": decision.max_tokens, + } + endpoint_budget = { + **endpoint_budget, + "clamped_max_tokens": decision.max_tokens, + } + attempt += 1 + continue + if decision.kind == "budget": + raise ContextBudgetExceeded( + "Local endpoint request exceeded the served context " + "window after completion re-clamp." + ) from exc + if decision.kind == "context": + raise _context_window_exceeded_error() from exc + raise + result = _parse_response(response.json()) + if input_budget is not None: + result["input_budget"] = input_budget + if endpoint_budget.get("served_window") is not None: + result["endpoint_budget"] = endpoint_budget + return result except LocalAdmissionTimeout: raise except LocalProviderError: raise except Exception as exc: - raise _classify_byo_generate_error(exc) from exc + raise _classify_byo_generate_error(exc, endpoint) from exc async def run_agenerate( @@ -725,7 +953,6 @@ async def run_agenerate( import httpx if not endpoint.is_bundled: - messages = _build_messages(contents, system_instruction) from solstone.think.providers.local_admission import ( LocalAdmissionTimeout, acquire_local_slot_async, @@ -734,14 +961,15 @@ async def run_agenerate( confidential_egress_base_url, ) - body = _build_request_body( - endpoint.served_model_id, - messages, - temperature, - max_output_tokens, - json_output, - json_schema, - False, + body, input_budget, endpoint_budget = await asyncio.to_thread( + _prepare_endpoint_request_with_resolution, + endpoint=endpoint, + contents=contents, + system_instruction=system_instruction, + temperature=temperature, + max_output_tokens=max_output_tokens, + json_output=json_output, + json_schema=json_schema, ) post_kwargs: dict[str, Any] = { "json": body, @@ -766,14 +994,57 @@ async def run_agenerate( try: async with admission: base_url = confidential_egress_base_url(endpoint.base_url) + attempt = 0 async with httpx.AsyncClient() as client: - response = await client.post( - f"{base_url}/v1/chat/completions", - timeout=post_timeout, - **post_kwargs, - ) - response.raise_for_status() - return _parse_response(response.json()) + while True: + response = await client.post( + f"{base_url}/v1/chat/completions", + timeout=( + post_timeout + if attempt == 0 + else _remaining_timeout(started, timeout) + ), + **post_kwargs, + ) + try: + response.raise_for_status() + break + except httpx.HTTPStatusError as exc: + if response.status_code != 400: + raise + decision = _endpoint_overflow_decision( + response.text, + endpoint_budget.get("served_window"), + attempt, + ) + if ( + decision.kind == "retry" + and decision.max_tokens is not None + ): + post_kwargs["json"] = { + **post_kwargs["json"], + "max_tokens": decision.max_tokens, + } + endpoint_budget = { + **endpoint_budget, + "clamped_max_tokens": decision.max_tokens, + } + attempt += 1 + continue + if decision.kind == "budget": + raise ContextBudgetExceeded( + "Local endpoint request exceeded the served " + "context window after completion re-clamp." + ) from exc + if decision.kind == "context": + raise _context_window_exceeded_error() from exc + raise + result = _parse_response(response.json()) + if input_budget is not None: + result["input_budget"] = input_budget + if endpoint_budget.get("served_window") is not None: + result["endpoint_budget"] = endpoint_budget + return result except asyncio.CancelledError: raise except LocalAdmissionTimeout: @@ -781,7 +1052,7 @@ async def run_agenerate( except LocalProviderError: raise except Exception as exc: - raise _classify_byo_generate_error(exc) from exc + raise _classify_byo_generate_error(exc, endpoint) from exc from solstone.think.providers import local_server from solstone.think.providers.local_admission import ( diff --git a/solstone/think/providers/openhands.py b/solstone/think/providers/openhands.py index b8904a783..2ca458f79 100644 --- a/solstone/think/providers/openhands.py +++ b/solstone/think/providers/openhands.py @@ -89,8 +89,6 @@ _SHELL_STDOUT_CAP = 6000 _SHELL_STDERR_CAP = 6000 _SHELL_TIMEOUT_SECONDS = 30 _COST_WARNING_TEXT = "Cost calculation failed" -_LOCAL_OUTPUT_RESERVE_TOKENS = LOCAL_MIN_CONTEXT_TOKENS // 4 -_LOCAL_CONDENSER_MAX_TOKENS = LOCAL_MIN_CONTEXT_TOKENS * 11 // 16 _LOCAL_CONDENSER_KEEP_FIRST = 4 _GENERATE_NUM_RETRIES = 2 _GEMINI_MAX_OUTPUT_TOKENS = 65_535 @@ -151,6 +149,14 @@ def _resolve_provider_key(provider: str, api_key: str | None = None) -> str: return effective +def _local_output_reserve_tokens(window: int) -> int: + return window // 4 + + +def _local_condenser_max_tokens(window: int) -> int: + return window * 11 // 16 + + def _resolve_allowed_roots(config: dict[str, Any]) -> list[Path]: journal = Path(get_journal()).resolve() project_root = Path(get_project_root()).resolve() @@ -206,30 +212,47 @@ def _cogitate_budgets(config: dict[str, Any]) -> tuple[int, float, float]: return max_turns, cost_cap, timeout_seconds -def _build_llm(provider: str, model: str, *, num_retries: int | None = None) -> Any: +def _build_llm( + provider: str, + model: str, + *, + num_retries: int | None = None, + endpoint: Any | None = None, + served_window: int | None = None, +) -> Any: from openhands.sdk import LLM retry_count = LLM_NUM_RETRIES if num_retries is None else num_retries if provider == "local": - from solstone.think.providers.local_endpoint import resolve_local_endpoint from solstone.think.services.spp_transport import confidential_egress_base_url - endpoint = resolve_local_endpoint() + if endpoint is None: + raise ValueError("Resolved local endpoint is required to build local LLM.") if not endpoint.is_bundled: base_url = confidential_egress_base_url(endpoint.base_url) - return LLM( - model=f"openai/{endpoint.served_model_id}", - base_url=f"{base_url}/v1", - api_key=endpoint.credential or "EMPTY", - native_tool_calling=False, - timeout=LLM_TIMEOUT_S, - num_retries=retry_count, - retry_min_wait=1, - retry_max_wait=2, - retry_multiplier=1.0, - input_cost_per_token=0, - output_cost_per_token=0, - ) + llm_kwargs: dict[str, Any] = { + "model": f"openai/{endpoint.served_model_id}", + "base_url": f"{base_url}/v1", + "api_key": endpoint.credential or "EMPTY", + "native_tool_calling": False, + "timeout": LLM_TIMEOUT_S, + "num_retries": retry_count, + "retry_min_wait": 1, + "retry_max_wait": 2, + "retry_multiplier": 1.0, + "input_cost_per_token": 0, + "output_cost_per_token": 0, + } + if served_window is not None and served_window >= LOCAL_MIN_CONTEXT_TOKENS: + llm_kwargs["max_input_tokens"] = served_window + llm_kwargs["max_output_tokens"] = _local_output_reserve_tokens( + served_window + ) + if endpoint.is_confidential: + llm_kwargs["litellm_extra_body"] = { + "chat_template_kwargs": {"enable_thinking": False} + } + return LLM(**llm_kwargs) from solstone.think.providers import local_server @@ -242,7 +265,9 @@ def _build_llm(provider: str, model: str, *, num_retries: int | None = None) -> timeout=LLM_TIMEOUT_S, num_retries=retry_count, max_input_tokens=local_server.LOCAL_MIN_CONTEXT_TOKENS, - max_output_tokens=_LOCAL_OUTPUT_RESERVE_TOKENS, + max_output_tokens=_local_output_reserve_tokens( + local_server.LOCAL_MIN_CONTEXT_TOKENS + ), input_cost_per_token=0, output_cost_per_token=0, litellm_extra_body={"chat_template_kwargs": {"enable_thinking": False}}, @@ -677,8 +702,8 @@ async def _run_agenerate( return _generate_result(response, model) -def _build_local_condenser(llm: Any) -> Any: - """LLM-summarizing condenser for the bundled-local floor window. +def _build_local_condenser(llm: Any, *, max_tokens: int) -> Any: + """LLM-summarizing condenser for an explicitly resolved local window. Reuses the agent's own LLM (shared usage_id is accepted in openhands-sdk 1.27.1) so there is no separate summarization endpoint. @@ -687,7 +712,7 @@ def _build_local_condenser(llm: Any) -> Any: return LLMSummarizingCondenser( llm=llm, - max_tokens=_LOCAL_CONDENSER_MAX_TOKENS, + max_tokens=max_tokens, keep_first=_LOCAL_CONDENSER_KEEP_FIRST, ) @@ -695,14 +720,18 @@ def _build_local_condenser(llm: Any) -> Any: def _build_cogitate_agent( *, llm: Any, - is_bundled_local: bool, + condenser_max_tokens: int | None, tool_specs: list[Any], include_default_tools: list[Any], system_prompt: str, ) -> Any: from openhands.sdk import Agent - condenser = _build_local_condenser(llm) if is_bundled_local else None + condenser = ( + _build_local_condenser(llm, max_tokens=condenser_max_tokens) + if condenser_max_tokens is not None + else None + ) return Agent( llm=llm, tools=tool_specs, @@ -1597,13 +1626,17 @@ async def run_cogitate( model = str(config["model"]) effective_on_event = on_event byo_endpoint = None + byo_served_window = None if provider == "local": from solstone.think.providers.local_endpoint import ( + resolve_endpoint_served_window, resolve_local_endpoint, wrap_on_event_redacting, ) byo_endpoint = resolve_local_endpoint() + if not byo_endpoint.is_bundled: + byo_served_window = resolve_endpoint_served_window(byo_endpoint) if not byo_endpoint.is_bundled and byo_endpoint.credential: effective_on_event = wrap_on_event_redacting( on_event, @@ -1642,7 +1675,13 @@ async def run_cogitate( config.get("read_call_budget", DEFAULT_READ_CALL_BUDGET) or 0 ) journal = Path(get_journal()) - llm = _build_llm(provider, model, num_retries=_llm_num_retries(config)) + llm = _build_llm( + provider, + model, + num_retries=_llm_num_retries(config), + endpoint=byo_endpoint, + served_window=byo_served_window, + ) usage_start = _usage_snapshot(llm) tool_specs = [] sol_executor = None @@ -1679,10 +1718,20 @@ async def run_cogitate( tool_specs.append(Tool(name="emit_final")) default_tools = [] - is_bundled_local = byo_endpoint is not None and byo_endpoint.is_bundled + condenser_max_tokens = None + if byo_endpoint is not None: + if byo_endpoint.is_bundled: + condenser_max_tokens = _local_condenser_max_tokens( + LOCAL_MIN_CONTEXT_TOKENS + ) + elif ( + byo_served_window is not None + and byo_served_window >= LOCAL_MIN_CONTEXT_TOKENS + ): + condenser_max_tokens = _local_condenser_max_tokens(byo_served_window) agent = _build_cogitate_agent( llm=llm, - is_bundled_local=is_bundled_local, + condenser_max_tokens=condenser_max_tokens, tool_specs=tool_specs, include_default_tools=default_tools, system_prompt=system_instruction, diff --git a/tests/test_cogitate_local_condenser.py b/tests/test_cogitate_local_condenser.py index ca8c5d0e6..8af99689a 100644 --- a/tests/test_cogitate_local_condenser.py +++ b/tests/test_cogitate_local_condenser.py @@ -20,21 +20,26 @@ def _local_llm() -> LLM: api_key="EMPTY", native_tool_calling=False, max_input_tokens=LOCAL_MIN_CONTEXT_TOKENS, - max_output_tokens=openhands._LOCAL_OUTPUT_RESERVE_TOKENS, + max_output_tokens=openhands._local_output_reserve_tokens( + LOCAL_MIN_CONTEXT_TOKENS + ), ) def test_build_cogitate_agent_adds_bundled_local_condenser(): + condenser_max_tokens = openhands._local_condenser_max_tokens( + LOCAL_MIN_CONTEXT_TOKENS + ) agent = openhands._build_cogitate_agent( llm=_local_llm(), - is_bundled_local=True, + condenser_max_tokens=condenser_max_tokens, tool_specs=[], include_default_tools=[], system_prompt="sys", ) assert isinstance(agent.condenser, LLMSummarizingCondenser) - assert agent.condenser.max_tokens == openhands._LOCAL_CONDENSER_MAX_TOKENS + assert agent.condenser.max_tokens == condenser_max_tokens assert agent.condenser.keep_first == openhands._LOCAL_CONDENSER_KEEP_FIRST assert agent.condenser.llm is agent.llm @@ -42,7 +47,26 @@ def test_build_cogitate_agent_adds_bundled_local_condenser(): def test_build_cogitate_agent_skips_condenser_for_non_bundled_local(): agent = openhands._build_cogitate_agent( llm=_local_llm(), - is_bundled_local=False, + condenser_max_tokens=None, + tool_specs=[], + include_default_tools=[], + system_prompt="sys", + ) + + assert agent.condenser is None + + +def test_build_cogitate_agent_skips_condenser_for_small_endpoint_window(): + window = 8192 + condenser_max_tokens = ( + openhands._local_condenser_max_tokens(window) + if window >= LOCAL_MIN_CONTEXT_TOKENS + else None + ) + + agent = openhands._build_cogitate_agent( + llm=_local_llm(), + condenser_max_tokens=condenser_max_tokens, tool_specs=[], include_default_tools=[], system_prompt="sys", @@ -52,17 +76,18 @@ def test_build_cogitate_agent_skips_condenser_for_non_bundled_local(): def test_local_condenser_window_invariants(): - assert ( - openhands._LOCAL_CONDENSER_MAX_TOKENS // 2 - + openhands._LOCAL_OUTPUT_RESERVE_TOKENS - < LOCAL_MIN_CONTEXT_TOKENS + condenser_max_tokens = openhands._local_condenser_max_tokens( + LOCAL_MIN_CONTEXT_TOKENS ) - assert ( - openhands._LOCAL_CONDENSER_MAX_TOKENS + openhands._LOCAL_OUTPUT_RESERVE_TOKENS - <= LOCAL_MIN_CONTEXT_TOKENS + output_reserve_tokens = openhands._local_output_reserve_tokens( + LOCAL_MIN_CONTEXT_TOKENS ) - assert openhands._LOCAL_CONDENSER_MAX_TOKENS < LOCAL_MIN_CONTEXT_TOKENS - assert 11000 <= openhands._LOCAL_CONDENSER_MAX_TOKENS <= 11500 + assert condenser_max_tokens // 2 + output_reserve_tokens < LOCAL_MIN_CONTEXT_TOKENS + assert condenser_max_tokens + output_reserve_tokens <= LOCAL_MIN_CONTEXT_TOKENS + assert condenser_max_tokens < LOCAL_MIN_CONTEXT_TOKENS + assert 11000 <= condenser_max_tokens <= 11500 + assert openhands._local_output_reserve_tokens(16384) == 4096 + assert openhands._local_condenser_max_tokens(16384) == 11264 assert openhands._LOCAL_CONDENSER_KEEP_FIRST < 240 // 2 - 1 diff --git a/tests/test_local.py b/tests/test_local.py index e6effae4f..40764b206 100644 --- a/tests/test_local.py +++ b/tests/test_local.py @@ -38,6 +38,26 @@ def _isolate_local_admission(monkeypatch, tmp_path): monkeypatch.setattr(local_admission, "record_local_inference", lambda _record: None) +@pytest.fixture(autouse=True) +def _default_endpoint_models_unknown(monkeypatch): + import httpx + + from solstone.think.providers import local_endpoint + + local_endpoint.reset_endpoint_served_window_cache() + + def fake_get(url, **_kwargs): + return httpx.Response( + 404, + request=httpx.Request("GET", url), + text="not found", + ) + + monkeypatch.setattr(httpx, "get", fake_get) + yield + local_endpoint.reset_endpoint_served_window_cache() + + def _provider(): providers_pkg = importlib.import_module("solstone.think.providers") if hasattr(providers_pkg, "local_budget"): @@ -93,6 +113,146 @@ class _ChatResponse: } +_FAKE_MODELS_BODY = ( + '{"object":"list","data":[{"id":"Qwen/Qwen3.5-4B","object":"model",' + '"created":1784825047,"owned_by":"sglang","root":"Qwen/Qwen3.5-4B",' + '"parent":null,"max_model_len":16384}]}' +) +_FAKE_F1_COMPLETION_OVERFLOW_BODY = ( + '{"object":"error","message":"Requested token count exceeds the model\'s ' + "maximum context length of 16384 tokens. You requested a total of 16397 " + "tokens: 13 tokens from the input messages and 16384 tokens for the " + "completion. Please reduce the number of tokens in the input messages or " + 'the completion to fit within the limit.","type":"BadRequestError",' + '"param":null,"code":400}' +) +_FAKE_F2_PROMPT_OVERFLOW_BODY = ( + '{"object":"error","message":"The input (18010 tokens) is longer than ' + 'the model\'s context length (16384 tokens).","type":"BadRequestError",' + '"param":null,"code":400}' +) +_FAKE_REJECT_WINDOW = 16384 + + +class _FakeRejectEndpoint: + def __init__( + self, + *, + prompt_tokens: int = 13, + models_body: str = _FAKE_MODELS_BODY, + completion_overflow_body: str = _FAKE_F1_COMPLETION_OVERFLOW_BODY, + prompt_overflow_body: str = _FAKE_F2_PROMPT_OVERFLOW_BODY, + force_completion_overflow: bool = False, + force_prompt_overflow: bool = False, + force_non_context_400: bool = False, + force_second_context_400: bool = False, + ) -> None: + self.prompt_tokens = prompt_tokens + self.models_body = models_body + self.completion_overflow_body = completion_overflow_body + self.prompt_overflow_body = prompt_overflow_body + self.force_completion_overflow = force_completion_overflow + self.force_prompt_overflow = force_prompt_overflow + self.force_non_context_400 = force_non_context_400 + self.force_second_context_400 = force_second_context_400 + self.gets: list[dict] = [] + self.posts: list[dict] = [] + + @property + def max_tokens(self) -> list[int]: + return [int(post["json"]["max_tokens"]) for post in self.posts] + + def install(self, monkeypatch) -> None: + import httpx + + monkeypatch.setattr(httpx, "get", self.get) + monkeypatch.setattr(httpx, "post", self.post) + + fake_endpoint = self + + class AsyncClient: + async def __aenter__(self): + return self + + async def __aexit__(self, *_args): + return None + + async def post(self, url, **kwargs): + return fake_endpoint.post(url, **kwargs) + + monkeypatch.setattr(httpx, "AsyncClient", AsyncClient) + + def get(self, url, **kwargs): + import httpx + + self.gets.append({"url": url, **kwargs}) + request = httpx.Request("GET", url) + if str(url).endswith("/v1/models"): + return httpx.Response(200, request=request, text=self.models_body) + return httpx.Response(404, request=request, text="not found") + + def post(self, url, **kwargs): + import httpx + + request = httpx.Request("POST", url) + if str(url).endswith("/tokenize"): + content = str((kwargs.get("json") or {}).get("content") or "") + return httpx.Response( + 200, + request=request, + json={"tokens": list(range(max(1, len(content) // 3)))}, + ) + if not str(url).endswith("/v1/chat/completions"): + raise AssertionError(f"unexpected local provider URL: {url}") + + body = kwargs["json"] + self.posts.append({"url": url, **kwargs}) + if self.force_non_context_400: + return httpx.Response(400, request=request, text="invalid temperature") + if self.force_prompt_overflow or self.prompt_tokens >= _FAKE_REJECT_WINDOW: + return httpx.Response( + 400, + request=request, + text=self.prompt_overflow_body, + ) + if ( + self.force_completion_overflow + or self.force_second_context_400 + and len(self.posts) > 1 + or self.prompt_tokens + int(body["max_tokens"]) > _FAKE_REJECT_WINDOW + ): + return httpx.Response( + 400, + request=request, + text=self.completion_overflow_body, + ) + return httpx.Response( + 200, + request=request, + json={ + "choices": [ + { + "message": {"content": "ok"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": self.prompt_tokens, + "completion_tokens": int(body["max_tokens"]), + "total_tokens": self.prompt_tokens + int(body["max_tokens"]), + }, + }, + ) + + +def _http_status_error(body: str): + import httpx + + request = httpx.Request("POST", "http://byo.example/openai/v1/chat/completions") + response = httpx.Response(400, request=request, text=body) + return httpx.HTTPStatusError("bad request", request=request, response=response) + + def test_local_model_prefix_maps_to_provider(): assert get_model_provider(LOCAL_MODEL) == "local" @@ -854,7 +1014,11 @@ def test_openhands_local_llm_kwargs(monkeypatch): lambda: SimpleNamespace(port=9876, served_model_id=served_model_id), ) - llm = openhands._build_llm("local", LOCAL_MODEL) + llm = openhands._build_llm( + "local", + LOCAL_MODEL, + endpoint=_bundled_endpoint(), + ) assert isinstance(llm, FakeLLM) assert captured == { @@ -865,7 +1029,9 @@ def test_openhands_local_llm_kwargs(monkeypatch): "timeout": openhands.LLM_TIMEOUT_S, "num_retries": openhands.LLM_NUM_RETRIES, "max_input_tokens": local_server.LOCAL_MIN_CONTEXT_TOKENS, - "max_output_tokens": openhands._LOCAL_OUTPUT_RESERVE_TOKENS, + "max_output_tokens": openhands._local_output_reserve_tokens( + local_server.LOCAL_MIN_CONTEXT_TOKENS + ), "input_cost_per_token": 0, "output_cost_per_token": 0, "litellm_extra_body": {"chat_template_kwargs": {"enable_thinking": False}}, @@ -896,6 +1062,20 @@ def _byo_endpoint( ) +def _qwen_byo_endpoint(): + from solstone.think.providers.local_endpoint import ( + LocalEndpoint, + normalize_local_endpoint_url, + ) + + return LocalEndpoint( + base_url=normalize_local_endpoint_url("http://byo.example/openai/v1/"), + served_model_id="Qwen/Qwen3.5-4B", + credential="test-token-PLACEHOLDER", + is_bundled=False, + ) + + def _bundled_endpoint(): from solstone.think.providers.local_endpoint import LocalEndpoint @@ -919,6 +1099,407 @@ def _patch_bundled_server(monkeypatch): ) +def test_run_generate_endpoint_clamps_large_default_budget_against_served_window( + monkeypatch, +): + provider = _provider() + fake_endpoint = _FakeRejectEndpoint() + fake_endpoint.install(monkeypatch) + monkeypatch.setattr(provider, "resolve_local_endpoint", _qwen_byo_endpoint) + + result = provider.run_generate( + "hello", + model=LOCAL_MODEL, + max_output_tokens=49152, + ) + + assert result["text"] == "ok" + assert fake_endpoint.gets + assert fake_endpoint.max_tokens + assert all(token <= _FAKE_REJECT_WINDOW for token in fake_endpoint.max_tokens) + + +def test_run_agenerate_endpoint_clamps_window_sized_budget_against_served_window( + monkeypatch, +): + provider = _provider() + fake_endpoint = _FakeRejectEndpoint() + fake_endpoint.install(monkeypatch) + monkeypatch.setattr(provider, "resolve_local_endpoint", _qwen_byo_endpoint) + + result = asyncio.run( + provider.run_agenerate( + "hello", + model=LOCAL_MODEL, + max_output_tokens=16384, + ) + ) + + assert result["text"] == "ok" + assert fake_endpoint.gets + assert fake_endpoint.max_tokens + assert all(token <= _FAKE_REJECT_WINDOW for token in fake_endpoint.max_tokens) + + +def test_run_generate_endpoint_preflight_floor_skips_chat_post(monkeypatch): + provider = _provider() + fake_endpoint = _FakeRejectEndpoint() + fake_endpoint.install(monkeypatch) + monkeypatch.setattr(provider, "resolve_local_endpoint", _qwen_byo_endpoint) + + with pytest.raises(provider.ContextBudgetExceeded) as exc: + provider.run_generate( + [{"role": "user", "content": "x" * 50000}], + model=LOCAL_MODEL, + max_output_tokens=1024, + ) + + assert exc.value.reason_code == "context_budget_exceeded" + assert fake_endpoint.gets + assert fake_endpoint.posts == [] + + +def test_validate_key_known_window_keeps_tiny_max_tokens(monkeypatch): + provider = _provider() + fake_endpoint = _FakeRejectEndpoint() + fake_endpoint.install(monkeypatch) + monkeypatch.setattr(provider, "resolve_local_endpoint", _qwen_byo_endpoint) + + assert provider.validate_key("local", "") == {"valid": True} + assert fake_endpoint.max_tokens == [8] + + +def test_run_generate_endpoint_known_window_truncates_fittable_input(monkeypatch): + provider = _provider() + fake_endpoint = _FakeRejectEndpoint() + fake_endpoint.install(monkeypatch) + monkeypatch.setattr(provider, "resolve_local_endpoint", _qwen_byo_endpoint) + chunks = [ + f"## 2026-06-23 09:{minute:02d}:00\n### Transcript\n" + + (str(minute) * 3000) + + "\n" + for minute in range(20) + ] + + result = provider.run_generate( + "".join(chunks), + model=LOCAL_MODEL, + max_output_tokens=1024, + ) + + from solstone.think.providers import local_budget + + posted_content = fake_endpoint.posts[0]["json"]["messages"][0]["content"] + assert local_budget.TRUNCATION_MARKER in posted_content + assert result["input_budget"]["clipped"] is True + assert result["input_budget"]["dropped_entries"] > 0 + assert result["endpoint_budget"]["served_window"] == _FAKE_REJECT_WINDOW + assert result["endpoint_budget"]["clamped_max_tokens"] == 1024 + + +def test_endpoint_generate_branches_route_shared_prep(monkeypatch): + provider = _provider() + calls = [] + original = provider._prepare_endpoint_request + monkeypatch.setattr(provider, "resolve_local_endpoint", _byo_endpoint) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + ) + + def spy_prepare(**kwargs): + calls.append(kwargs) + return original(**kwargs) + + def fake_post(url, **kwargs): + return _ChatResponse("ok") + + class AsyncClient: + async def __aenter__(self): + return self + + async def __aexit__(self, *_args): + return None + + async def post(self, url, **kwargs): + return _ChatResponse("ok") + + import httpx + + monkeypatch.setattr(provider, "_prepare_endpoint_request", spy_prepare) + monkeypatch.setattr(httpx, "post", fake_post) + monkeypatch.setattr(httpx, "AsyncClient", AsyncClient) + + provider.run_generate("hello", model=LOCAL_MODEL, max_output_tokens=7) + asyncio.run(provider.run_agenerate("hello", model=LOCAL_MODEL, max_output_tokens=7)) + + assert [call["max_output_tokens"] for call in calls] == [7, 7] + assert [call["served_window"] for call in calls] == [None, None] + + +def _run_endpoint_generate(provider, mode: str, **kwargs): + if mode == "sync": + return provider.run_generate(**kwargs) + return asyncio.run(provider.run_agenerate(**kwargs)) + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +def test_endpoint_completion_overflow_reclamps_once(monkeypatch, mode): + provider = _provider() + fake_endpoint = _FakeRejectEndpoint() + fake_endpoint.install(monkeypatch) + monkeypatch.setattr(provider, "resolve_local_endpoint", _byo_endpoint) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + ) + + result = _run_endpoint_generate( + provider, + mode, + contents="hello", + model=LOCAL_MODEL, + max_output_tokens=16384, + ) + + assert result["text"] == "ok" + assert fake_endpoint.max_tokens == [16384, 16355] + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +def test_endpoint_completion_overflow_below_floor_is_budget_terminal( + monkeypatch, + mode, +): + provider = _provider() + low_room_body = _FAKE_F1_COMPLETION_OVERFLOW_BODY.replace( + "13 tokens from the input messages and 16384 tokens for the completion", + "16200 tokens from the input messages and 500 tokens for the completion", + ) + fake_endpoint = _FakeRejectEndpoint( + completion_overflow_body=low_room_body, + force_completion_overflow=True, + ) + fake_endpoint.install(monkeypatch) + monkeypatch.setattr(provider, "resolve_local_endpoint", _byo_endpoint) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + ) + + with pytest.raises(provider.ContextBudgetExceeded) as exc: + _run_endpoint_generate( + provider, + mode, + contents="hello", + model=LOCAL_MODEL, + max_output_tokens=500, + ) + + assert exc.value.reason_code == "context_budget_exceeded" + assert fake_endpoint.max_tokens == [500] + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +def test_endpoint_prompt_overflow_is_terminal_context_window(monkeypatch, mode): + provider = _provider() + fake_endpoint = _FakeRejectEndpoint(force_prompt_overflow=True) + fake_endpoint.install(monkeypatch) + monkeypatch.setattr(provider, "resolve_local_endpoint", _byo_endpoint) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + ) + + with pytest.raises(provider.LocalProviderError) as exc: + _run_endpoint_generate( + provider, + mode, + contents="hello", + model=LOCAL_MODEL, + max_output_tokens=1024, + ) + + assert exc.value.reason_code == "context_window_exceeded" + assert fake_endpoint.max_tokens == [1024] + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +def test_endpoint_non_context_400_stays_contract_failed(monkeypatch, mode): + provider = _provider() + fake_endpoint = _FakeRejectEndpoint(force_non_context_400=True) + fake_endpoint.install(monkeypatch) + monkeypatch.setattr(provider, "resolve_local_endpoint", _byo_endpoint) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + ) + + with pytest.raises(provider.LocalProviderError) as exc: + _run_endpoint_generate( + provider, + mode, + contents="hello", + model=LOCAL_MODEL, + max_output_tokens=1024, + ) + + assert exc.value.reason_code == "local_endpoint_contract_failed" + assert fake_endpoint.max_tokens == [1024] + + +@pytest.mark.parametrize("mode", ["sync", "async"]) +def test_endpoint_second_context_400_after_reclamp_is_terminal(monkeypatch, mode): + provider = _provider() + fake_endpoint = _FakeRejectEndpoint(force_second_context_400=True) + fake_endpoint.install(monkeypatch) + monkeypatch.setattr(provider, "resolve_local_endpoint", _byo_endpoint) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + ) + + with pytest.raises(provider.LocalProviderError) as exc: + _run_endpoint_generate( + provider, + mode, + contents="hello", + model=LOCAL_MODEL, + max_output_tokens=16384, + ) + + assert exc.value.reason_code == "context_window_exceeded" + assert fake_endpoint.max_tokens == [16384, 16355] + + +def test_run_generate_byo_unknown_window_request_body_matches_golden( + monkeypatch, +): + provider = _provider() + monkeypatch.setattr(provider, "resolve_local_endpoint", _byo_endpoint) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + raising=False, + ) + captured = {} + + def fake_post(url, **kwargs): + captured.update({"url": url, **kwargs}) + return _ChatResponse("ok") + + import httpx + + monkeypatch.setattr(httpx, "post", fake_post) + + provider.run_generate( + "hello", + model=LOCAL_MODEL, + temperature=0.4, + max_output_tokens=7, + ) + + assert captured["json"] == { + "model": "served-model", + "messages": [{"role": "user", "content": "hello"}], + "temperature": 0.4, + "max_tokens": 7, + "stream": False, + } + + +def test_run_generate_confidential_body_carries_qwen_block(monkeypatch): + provider = _provider() + from solstone.think.providers.local_endpoint import LocalEndpoint + + endpoint = LocalEndpoint( + base_url="https://spp.example.test", + served_model_id="confidential-model", + credential="confidential-token", + is_bundled=False, + is_confidential=True, + ) + captured = {} + monkeypatch.setattr(provider, "resolve_local_endpoint", lambda: endpoint) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + ) + monkeypatch.setattr( + "solstone.think.services.spp_transport.confidential_egress_base_url", + lambda _base_url: "http://127.0.0.1:4567", + ) + + def fake_post(url, **kwargs): + captured.update({"url": url, **kwargs}) + return _ChatResponse("ok") + + import httpx + + monkeypatch.setattr(httpx, "post", fake_post) + + provider.run_generate( + "hello", + model=LOCAL_MODEL, + temperature=0.4, + max_output_tokens=7, + ) + + assert captured["url"] == "http://127.0.0.1:4567/v1/chat/completions" + assert captured["json"]["chat_template_kwargs"] == {"enable_thinking": False} + assert captured["json"]["top_p"] == 0.8 + assert captured["json"]["top_k"] == 20 + assert captured["json"]["min_p"] == 0.0 + assert captured["json"]["presence_penalty"] == 1.5 + + +def test_run_generate_bundled_request_body_matches_golden(monkeypatch): + provider = _provider() + monkeypatch.setattr(provider, "resolve_local_endpoint", _bundled_endpoint) + _patch_bundled_server(monkeypatch) + captured = {} + + def fake_post(url, **kwargs): + if str(url).endswith("/tokenize"): + return _FakeRejectEndpoint().post(url, **kwargs) + if str(url).endswith("/v1/chat/completions"): + captured.update({"url": url, **kwargs}) + return _ChatResponse("ok") + raise AssertionError(f"unexpected local provider URL: {url}") + + import httpx + + monkeypatch.setattr(httpx, "post", fake_post) + + provider.run_generate( + "hello", + model=LOCAL_MODEL, + temperature=0.4, + max_output_tokens=7, + ) + + assert captured["json"] == { + "model": LOCAL_MODEL, + "messages": [{"role": "user", "content": "hello"}], + "temperature": 0.4, + "max_tokens": 7, + "stream": False, + "chat_template_kwargs": {"enable_thinking": False}, + "top_p": 0.8, + "top_k": 20, + "min_p": 0.0, + "presence_penalty": 1.5, + } + + def test_run_generate_bundled_encodes_image_once(monkeypatch): provider = _provider() monkeypatch.setattr(provider, "resolve_local_endpoint", _bundled_endpoint) @@ -1040,6 +1621,11 @@ def test_run_generate_byo_acquires_and_releases_permit(monkeypatch): "resolve_local_endpoint", lambda: _byo_endpoint(parallel_slots=1), ) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + ) captured = {} def fake_post(url, **kwargs): @@ -1069,6 +1655,11 @@ def test_run_generate_byo_queue_timeout_preserves_exact_type_and_skips_post( "resolve_local_endpoint", lambda: _byo_endpoint(parallel_slots=1), ) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + ) def fake_post(*_args, **_kwargs): raise AssertionError("httpx.post must not run after queue timeout") @@ -1096,6 +1687,11 @@ def test_run_generate_byo_http_timeout_uses_remaining_deadline(monkeypatch): "resolve_local_endpoint", lambda: _byo_endpoint(parallel_slots=1), ) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + ) captured = {} times = iter([100.0, 100.0, 100.35]) @@ -1312,6 +1908,11 @@ def test_run_agenerate_byo_acquires_and_releases_permit(monkeypatch): "resolve_local_endpoint", lambda: _byo_endpoint(parallel_slots=1), ) + monkeypatch.setattr( + provider, + "resolve_endpoint_served_window", + lambda _endpoint: None, + ) captured = {} class AsyncClient: @@ -1716,7 +2317,7 @@ def test_classify_byo_generate_error_capacity_names_are_non_blocking(exc_name): exc = RuntimeError("outer") exc.__cause__ = inner - classified = provider._classify_byo_generate_error(exc) + classified = provider._classify_byo_generate_error(exc, _byo_endpoint()) assert classified.reason_code == "local_capacity_exhausted" assert is_blocking_reason(classified.reason_code) is False @@ -1740,7 +2341,7 @@ def test_classify_byo_generate_error_unreachable_names_are_blocking(exc_name): exc = RuntimeError("outer") exc.__cause__ = inner - classified = provider._classify_byo_generate_error(exc) + classified = provider._classify_byo_generate_error(exc, _byo_endpoint()) assert classified.reason_code == "local_endpoint_unreachable" assert str(classified) == provider.LOCAL_ENDPOINT_UNREACHABLE_COPY @@ -1756,7 +2357,7 @@ def test_classify_byo_generate_error_capacity_wins_mixed_chain(): outer.__cause__ = unreachable unreachable.__cause__ = capacity - classified = provider._classify_byo_generate_error(outer) + classified = provider._classify_byo_generate_error(outer, _byo_endpoint()) assert classified.reason_code == "local_capacity_exhausted" @@ -1768,13 +2369,61 @@ def test_classify_byo_generate_error_500_stays_contract_failed(): status_code = 500 classified = provider._classify_byo_generate_error( - InternalServerError("server failed") + InternalServerError("server failed"), + _byo_endpoint(), ) assert classified.reason_code == "local_endpoint_contract_failed" assert str(classified) == provider.LOCAL_ENDPOINT_CONTRACT_COPY +@pytest.mark.parametrize( + "body", + [ + _FAKE_F1_COMPLETION_OVERFLOW_BODY, + _FAKE_F2_PROMPT_OVERFLOW_BODY, + ], +) +def test_classify_byo_generate_error_context_body_maps_to_context_window(body): + provider = _provider() + endpoint = _byo_endpoint() + + classified = provider._classify_byo_generate_error( + _http_status_error(body), + endpoint, + ) + + assert classified.reason_code == "context_window_exceeded" + assert ( + str(classified) + == "The configured endpoint rejected the request: prompt and completion " + "exceed the served context window." + ) + + +@pytest.mark.parametrize( + "body", + [ + _FAKE_F1_COMPLETION_OVERFLOW_BODY, + _FAKE_F2_PROMPT_OVERFLOW_BODY, + ], +) +def test_classify_byo_cogitate_error_context_body_maps_to_context_window(body): + from solstone.think.providers import local_endpoint + + class BadRequestError(RuntimeError): + status_code = 400 + + def __init__(self, message: str, payload: str) -> None: + super().__init__(message) + self.message = message + self.body = payload + + exc = BadRequestError("litellm bad request wrapper", body) + + assert local_endpoint.classify_byo_cogitate_error(exc) == "context_window_exceeded" + + def test_run_generate_byo_http_status_maps_to_contract_failed(monkeypatch): provider = _provider() monkeypatch.setattr(provider, "resolve_local_endpoint", _byo_endpoint) @@ -2278,7 +2927,7 @@ def test_run_cogitate_talent_hook_error_bypasses_local_error_event(monkeypatch): ], ) def test_openhands_local_byo_llm_kwargs(monkeypatch, credential, expected_key): - from solstone.think.providers import local_endpoint, openhands + from solstone.think.providers import openhands captured = {} @@ -2289,17 +2938,17 @@ def test_openhands_local_byo_llm_kwargs(monkeypatch, credential, expected_key): sdk_module = types.ModuleType("openhands.sdk") sdk_module.LLM = FakeLLM monkeypatch.setitem(sys.modules, "openhands.sdk", sdk_module) - monkeypatch.setattr( - local_endpoint, - "resolve_local_endpoint", - lambda: _byo_endpoint(credential), - ) monkeypatch.setattr( "solstone.think.providers.local_server.connect", lambda: (_ for _ in ()).throw(AssertionError("connect not expected")), ) - llm = openhands._build_llm("local", LOCAL_MODEL) + llm = openhands._build_llm( + "local", + LOCAL_MODEL, + endpoint=_byo_endpoint(credential), + served_window=None, + ) assert isinstance(llm, FakeLLM) assert captured == { @@ -2345,15 +2994,12 @@ def test_openhands_local_confidential_llm_uses_forwarder(monkeypatch): sdk_module = types.ModuleType("openhands.sdk") sdk_module.LLM = FakeLLM monkeypatch.setitem(sys.modules, "openhands.sdk", sdk_module) - monkeypatch.setattr( - local_endpoint, - "resolve_local_endpoint", - lambda: local_endpoint.LocalEndpoint( - base_url=configured_endpoint, - served_model_id="confidential-model", - credential="confidential-token", - is_bundled=False, - ), + endpoint = local_endpoint.LocalEndpoint( + base_url=configured_endpoint, + served_model_id="confidential-model", + credential="confidential-token", + is_bundled=False, + is_confidential=True, ) monkeypatch.setattr( "solstone.think.services.spp_transport.confidential_egress_base_url", @@ -2364,13 +3010,106 @@ def test_openhands_local_confidential_llm_uses_forwarder(monkeypatch): lambda: (_ for _ in ()).throw(AssertionError("connect not expected")), ) - llm = openhands._build_llm("local", LOCAL_MODEL) + llm = openhands._build_llm( + "local", + LOCAL_MODEL, + endpoint=endpoint, + served_window=None, + ) assert isinstance(llm, FakeLLM) assert captured["base_url"] == f"{forwarder}/v1" assert configured_endpoint not in captured["base_url"] +def test_openhands_local_confidential_known_window_caps_and_extra_body(monkeypatch): + from solstone.think.providers import local_endpoint, openhands + + captured = {} + + class FakeLLM: + def __init__(self, **kwargs): + captured.update(kwargs) + + sdk_module = types.ModuleType("openhands.sdk") + sdk_module.LLM = FakeLLM + monkeypatch.setitem(sys.modules, "openhands.sdk", sdk_module) + endpoint = local_endpoint.LocalEndpoint( + base_url="https://spp.example.test", + served_model_id="confidential-model", + credential="confidential-token", + is_bundled=False, + is_confidential=True, + ) + monkeypatch.setattr( + "solstone.think.services.spp_transport.confidential_egress_base_url", + lambda _base_url: "http://127.0.0.1:4567", + ) + + openhands._build_llm( + "local", + LOCAL_MODEL, + endpoint=endpoint, + served_window=16384, + ) + + assert captured["max_input_tokens"] == 16384 + assert captured["max_output_tokens"] == 4096 + assert captured["litellm_extra_body"] == { + "chat_template_kwargs": {"enable_thinking": False} + } + + +def test_openhands_local_known_window_caps_without_extra_body(monkeypatch): + from solstone.think.providers import openhands + + captured = {} + + class FakeLLM: + def __init__(self, **kwargs): + captured.update(kwargs) + + sdk_module = types.ModuleType("openhands.sdk") + sdk_module.LLM = FakeLLM + monkeypatch.setitem(sys.modules, "openhands.sdk", sdk_module) + + openhands._build_llm( + "local", + LOCAL_MODEL, + endpoint=_qwen_byo_endpoint(), + served_window=32768, + ) + + assert captured["max_input_tokens"] == 32768 + assert captured["max_output_tokens"] == 8192 + assert "litellm_extra_body" not in captured + + +def test_openhands_local_small_window_omits_caps_and_extra_body(monkeypatch): + from solstone.think.providers import openhands + + captured = {} + + class FakeLLM: + def __init__(self, **kwargs): + captured.update(kwargs) + + sdk_module = types.ModuleType("openhands.sdk") + sdk_module.LLM = FakeLLM + monkeypatch.setitem(sys.modules, "openhands.sdk", sdk_module) + + openhands._build_llm( + "local", + LOCAL_MODEL, + endpoint=_qwen_byo_endpoint(), + served_window=8192, + ) + + assert "max_input_tokens" not in captured + assert "max_output_tokens" not in captured + assert "litellm_extra_body" not in captured + + def test_local_context_window_split_floor_vs_tier(): import inspect diff --git a/tests/test_log_policy.py b/tests/test_log_policy.py index d11aabdf0..cd81fe8f5 100644 --- a/tests/test_log_policy.py +++ b/tests/test_log_policy.py @@ -66,7 +66,11 @@ def test_run_cogitate_invokes_restore(monkeypatch): real_apply_http_logging_policy = apply_http_logging_policy def fail_build_llm( - provider: str, model: str, *, num_retries: int | None = None + provider: str, + model: str, + *, + num_retries: int | None = None, + **_kwargs: Any, ) -> Any: raise RuntimeError("sentinel") diff --git a/tests/test_openhands_errors.py b/tests/test_openhands_errors.py index 6489a0574..0669c2de9 100644 --- a/tests/test_openhands_errors.py +++ b/tests/test_openhands_errors.py @@ -165,7 +165,7 @@ def test_run_cogitate_error_before_usage_baseline_omits_usage( ): build_exc = RuntimeError("llm exploded") - def fail_build(_provider, _model, *, num_retries=None): + def fail_build(_provider, _model, *, num_retries=None, **_kwargs): raise build_exc monkeypatch.setattr(openhands, "_build_llm", fail_build) @@ -249,6 +249,114 @@ def test_run_cogitate_local_byo_error_event_uses_fixed_copy_and_redacts( assert token not in events[0]["trace"] +def test_run_cogitate_local_byo_context_body_event_redacts( + fake_openhands, + run_env, + monkeypatch, +): + from tests.test_local import _FAKE_F2_PROMPT_OVERFLOW_BODY + + token = "SENTINEL-BYO-CONTEXT-CRED-219a" + endpoint = LocalEndpoint( + base_url="http://byo.example/openai", + served_model_id="served-model", + credential=token, + is_bundled=False, + ) + + class BadRequestError(RuntimeError): + status_code = 400 + + def __init__(self) -> None: + super().__init__(f"bad request with {token}") + self.message = f"bad request with {token}" + self.body = _FAKE_F2_PROMPT_OVERFLOW_BODY + f" {token}" + ("x" * 6000) + + async def fail(_conversation): + raise BadRequestError() + + monkeypatch.setattr( + "solstone.think.providers.local_endpoint.resolve_local_endpoint", + lambda: endpoint, + ) + fake_openhands.Conversation.arun_impl = fail + events: list[dict] = [] + local_env = {**run_env, "provider": "local", "model": LOCAL_MODEL} + + with pytest.raises(BadRequestError) as raised: + asyncio.run(openhands.run_cogitate(local_env, events.append)) + + assert len(events) == 1 + assert events[0]["reason_code"] == "context_window_exceeded" + assert token not in json.dumps(events) + assert token not in str(raised.value) + + +def test_run_cogitate_local_window_threaded_once( + fake_openhands, + run_env, + monkeypatch, +): + endpoint = LocalEndpoint( + base_url="https://spp.example.test", + served_model_id="confidential-model", + credential="confidential-token", + is_bundled=False, + is_confidential=True, + ) + window_calls = [] + llm_calls = [] + agent_calls = [] + original_build_llm = openhands._build_llm + + def resolve_window(resolved_endpoint): + window_calls.append(resolved_endpoint) + return 16384 + + def spy_build_llm(*args, **kwargs): + llm_calls.append(kwargs) + return original_build_llm(*args, **kwargs) + + def spy_build_agent(**kwargs): + agent_calls.append(kwargs) + return SimpleNamespace( + llm=kwargs["llm"], + tools=kwargs["tool_specs"], + include_default_tools=kwargs["include_default_tools"], + system_prompt=kwargs["system_prompt"], + condenser=SimpleNamespace(max_tokens=kwargs["condenser_max_tokens"]), + ) + + monkeypatch.setattr( + "solstone.think.providers.local_endpoint.resolve_local_endpoint", + lambda: endpoint, + ) + monkeypatch.setattr( + "solstone.think.providers.local_endpoint.resolve_endpoint_served_window", + resolve_window, + ) + monkeypatch.setattr( + "solstone.think.services.spp_transport.confidential_egress_base_url", + lambda _base_url: "http://127.0.0.1:4567", + ) + monkeypatch.setattr(openhands, "_build_llm", spy_build_llm) + monkeypatch.setattr(openhands, "_build_cogitate_agent", spy_build_agent) + local_env = {**run_env, "provider": "local", "model": LOCAL_MODEL} + + asyncio.run(openhands.run_cogitate(local_env, lambda _event: None)) + + assert window_calls == [endpoint] + assert llm_calls[0]["endpoint"] is endpoint + assert llm_calls[0]["served_window"] == 16384 + assert agent_calls[0]["condenser_max_tokens"] == 11264 + llm = fake_openhands.LLM.instances[0] + assert llm.max_input_tokens == 16384 + assert llm.max_output_tokens == 4096 + assert llm.litellm_extra_body == { + "chat_template_kwargs": {"enable_thinking": False} + } + + @pytest.mark.parametrize("exc_type", [httpx.ConnectError, httpx.ConnectTimeout]) def test_run_cogitate_byo_connection_error_classifies_unreachable_no_wall_clock( fake_openhands,