From e5ae8b8ae173071b60e957d1ed300deba146607c Mon Sep 17 00:00:00 2001 From: Alex Schapiro Date: Tue, 28 Jul 2026 15:40:59 +0000 Subject: [PATCH 1/4] fix(cost): capture OpenRouter streamed usage.cost (fixes $0 kimi-k3 cost) --- strix/config/models.py | 45 +++++++++++++++++ strix/report/state.py | 64 ++++++++++++++++++++++++ tests/test_cost_tracking.py | 97 ++++++++++++++++++++++++++++++++++++- 3 files changed, 204 insertions(+), 2 deletions(-) diff --git a/strix/config/models.py b/strix/config/models.py index b6e6b0f0d..dceb6654a 100644 --- a/strix/config/models.py +++ b/strix/config/models.py @@ -277,6 +277,51 @@ def _configure_litellm_compatibility() -> None: litellm.suppress_debug_info = True _register_litellm_cost_callback() + _install_openrouter_stream_cost_capture() + + +def _install_openrouter_stream_cost_capture() -> None: + """Preserve OpenRouter's per-stream cost, which LiteLLM drops when streaming. + + OpenRouter reports the real charge in ``usage.cost`` of the final stream + chunk, but LiteLLM rebuilds streamed responses from token-only fields and + discards it (its non-streamed path stashes the cost in hidden params; the + streaming path does not). Every scan streams, so without this the cost is + lost and Strix falls back to a cost-map estimate that is missing entirely + for new models (e.g. kimi-k3), reporting $0. Subclass the OpenRouter + streaming handler to record the cost keyed by response id so the cost + callback can recover the exact charge for the matching rebuilt response. + """ + import litellm + from litellm.llms.openrouter.chat.transformation import ( + OpenRouterChatCompletionStreamingHandler, + OpenrouterConfig, + ) + + from strix.report.state import remember_streamed_openrouter_cost + + class _StrixOpenRouterStreamingHandler(OpenRouterChatCompletionStreamingHandler): + def chunk_parser(self, chunk: dict[str, Any]) -> Any: + stream = super().chunk_parser(chunk) + remember_streamed_openrouter_cost( + chunk.get("id") or getattr(stream, "id", None), chunk.get("usage") + ) + return stream + + class _StrixOpenrouterConfig(OpenrouterConfig): + def get_model_response_iterator( + self, streaming_response: Any, sync_stream: bool, json_mode: bool | None = False + ) -> Any: + return _StrixOpenRouterStreamingHandler( + streaming_response=streaming_response, + sync_stream=sync_stream, + json_mode=json_mode, + ) + + # LiteLLM's provider-config factory reads litellm.OpenrouterConfig at call + # time, so overriding the attribute is enough for the subclass to take + # effect. (type: ignore — mypy rejects reassigning a class attribute.) + litellm.OpenrouterConfig = _StrixOpenrouterConfig # type: ignore[misc] _OPENROUTER_ATTRIBUTION_HEADERS = { diff --git a/strix/report/state.py b/strix/report/state.py index 1475ccf33..501673ce8 100644 --- a/strix/report/state.py +++ b/strix/report/state.py @@ -1,6 +1,8 @@ import json import logging import subprocess +import threading +from collections import OrderedDict from collections.abc import Callable from datetime import UTC, datetime from importlib.metadata import PackageNotFoundError, version @@ -507,6 +509,63 @@ def _hydrate_llm_usage(self, raw_usage: Any) -> None: self._sync_llm_usage_record() +# LiteLLM rebuilds streamed responses from token-only chunks and drops the +# provider-reported ``usage.cost`` that OpenRouter sends in its final stream +# chunk (unlike the non-streamed path, which stashes it in hidden params). Since +# every scan streams, that cost never reaches the callback below. The OpenRouter +# streaming handler (see strix.config.models) stashes the cost here keyed by the +# response id so the callback can recover the exact charge for the matching +# rebuilt response. +_STREAMED_OPENROUTER_COST_LIMIT = 4096 +_streamed_openrouter_costs: OrderedDict[str, float] = OrderedDict() +_streamed_openrouter_costs_lock = threading.Lock() + + +def openrouter_stream_cost(usage: Any) -> float | None: + """Total OpenRouter-reported cost from a raw stream ``usage`` block, or None. + + Non-BYOK responses bill everything to ``usage.cost``. BYOK responses put the + OpenRouter fee in ``usage.cost`` (often 0) and the provider charge in + ``usage.cost_details.upstream_inference_cost``, so BYOK totals sum the two. + """ + if not isinstance(usage, dict): + return None + total = 0.0 + cost = usage.get("cost") + if isinstance(cost, int | float) and cost > 0: + total += float(cost) + if bool(usage.get("is_byok")): + details = usage.get("cost_details") + upstream = details.get("upstream_inference_cost") if isinstance(details, dict) else None + if isinstance(upstream, int | float) and upstream > 0: + total += float(upstream) + return total if total > 0 else None + + +def remember_streamed_openrouter_cost(response_id: Any, usage: Any) -> None: + """Record an OpenRouter stream's reported cost so the cost callback can read it.""" + if not isinstance(response_id, str) or not response_id: + return + cost = openrouter_stream_cost(usage) + if cost is None: + return + with _streamed_openrouter_costs_lock: + _streamed_openrouter_costs[response_id] = cost + _streamed_openrouter_costs.move_to_end(response_id) + while len(_streamed_openrouter_costs) > _STREAMED_OPENROUTER_COST_LIMIT: + _streamed_openrouter_costs.popitem(last=False) + + +def _take_streamed_openrouter_cost(completion_response: Any) -> float | None: + response_id = getattr(completion_response, "id", None) + if response_id is None and isinstance(completion_response, dict): + response_id = cast("dict[str, Any]", completion_response).get("id") + if not isinstance(response_id, str) or not response_id: + return None + with _streamed_openrouter_costs_lock: + return _streamed_openrouter_costs.pop(response_id, None) + + def litellm_cost_callback( kwargs: Any, completion_response: Any, @@ -541,6 +600,11 @@ def litellm_cost_callback( if cost is None: cost = _usage_reported_cost(completion_response) + # Recover the exact OpenRouter cost the streaming handler stashed for this + # response — LiteLLM drops it from streamed usage, so nothing above sees it. + if cost is None: + cost = _take_streamed_openrouter_cost(completion_response) + if cost is None: cost = _estimate_response_cost(kwargs, completion_response) diff --git a/tests/test_cost_tracking.py b/tests/test_cost_tracking.py index 543d3fd5c..a4039c52e 100644 --- a/tests/test_cost_tracking.py +++ b/tests/test_cost_tracking.py @@ -8,8 +8,21 @@ import litellm import pytest -from strix.config.models import _configure_litellm_compatibility -from strix.report.state import litellm_cost_callback +import strix.report.state as state_module +from strix.config.models import ( + _configure_litellm_compatibility, + _install_openrouter_stream_cost_capture, +) +from strix.report.state import ( + litellm_cost_callback, + openrouter_stream_cost, + remember_streamed_openrouter_cost, +) + + +@pytest.fixture(autouse=True) +def _clear_streamed_costs() -> None: + state_module._streamed_openrouter_costs.clear() def test_streaming_logging_stays_enabled_for_cost_callback() -> None: @@ -151,3 +164,83 @@ def test_cost_callback_records_nothing_when_no_cost_available() -> None: litellm_cost_callback({"response_cost": None, "model": "x/y"}, response) report_state.record_observed_llm_cost.assert_not_called() + + +def test_openrouter_stream_cost_extracts_plain_and_byok_totals() -> None: + assert openrouter_stream_cost({"cost": 0.003168}) == pytest.approx(0.003168) + assert openrouter_stream_cost( + {"cost": 0.01, "is_byok": True, "cost_details": {"upstream_inference_cost": 0.2}} + ) == pytest.approx(0.21) + # Upstream cost is only added for BYOK responses. + assert openrouter_stream_cost( + {"cost": 0.05, "is_byok": False, "cost_details": {"upstream_inference_cost": 0.04}} + ) == pytest.approx(0.05) + assert openrouter_stream_cost({"prompt_tokens": 10}) is None + assert openrouter_stream_cost(None) is None + + +def test_cost_callback_recovers_streamed_openrouter_cost_by_response_id() -> None: + report_state = MagicMock() + remember_streamed_openrouter_cost("gen-abc", {"cost": 0.42}) + # LiteLLM strips cost from the rebuilt streamed usage; only the id survives. + response = SimpleNamespace(id="gen-abc", usage=SimpleNamespace(cost=None), _hidden_params={}) + + with ( + patch("strix.report.state.get_global_report_state", return_value=report_state), + patch("litellm.completion_cost", side_effect=ValueError("unknown model")), + ): + litellm_cost_callback({"response_cost": None, "model": "moonshotai/kimi-k3"}, response) + + report_state.record_observed_llm_cost.assert_called_once_with(0.42) + # The entry is consumed so a later response cannot double-count it. + assert "gen-abc" not in state_module._streamed_openrouter_costs + + +def test_streamed_openrouter_cost_prefers_provider_report_over_estimate() -> None: + report_state = MagicMock() + remember_streamed_openrouter_cost("gen-xyz", {"cost": 0.9}) + response = SimpleNamespace( + id="gen-xyz", + usage=SimpleNamespace(prompt_tokens=10, completion_tokens=5, total_tokens=15), + _hidden_params={}, + ) + + with ( + patch("strix.report.state.get_global_report_state", return_value=report_state), + patch("litellm.completion_cost", return_value=0.1) as estimate, + ): + litellm_cost_callback({"response_cost": None, "model": "moonshotai/kimi-k3"}, response) + + report_state.record_observed_llm_cost.assert_called_once_with(0.9) + estimate.assert_not_called() + + +def test_remember_streamed_openrouter_cost_evicts_oldest_over_limit() -> None: + limit = state_module._STREAMED_OPENROUTER_COST_LIMIT + for i in range(limit + 5): + remember_streamed_openrouter_cost(f"gen-{i}", {"cost": 0.001}) + + assert len(state_module._streamed_openrouter_costs) == limit + assert "gen-0" not in state_module._streamed_openrouter_costs + assert f"gen-{limit + 4}" in state_module._streamed_openrouter_costs + + +def test_openrouter_stream_handler_records_cost() -> None: + _install_openrouter_stream_cost_capture() + handler_cls = ( + litellm.OpenrouterConfig() + .get_model_response_iterator(streaming_response=iter([]), sync_stream=True) + .__class__ + ) + + chunk = { + "id": "gen-stream", + "created": 1, + "model": "moonshotai/kimi-k3", + "choices": [{"index": 0, "delta": {"content": None}}], + "usage": {"prompt_tokens": 89, "completion_tokens": 138, "cost": 0.0035055}, + } + handler = handler_cls(streaming_response=iter([]), sync_stream=True) + handler.chunk_parser(chunk) + + assert state_module._streamed_openrouter_costs["gen-stream"] == pytest.approx(0.0035055) From 4b76c73c3ff3a80480173be22e17cf6d2533d7e2 Mon Sep 17 00:00:00 2001 From: Alex Schapiro Date: Tue, 28 Jul 2026 15:59:22 +0000 Subject: [PATCH 2/4] refactor(cost): encapsulate streamed OpenRouter cost cache, clear per run --- strix/config/models.py | 4 +- strix/report/state.py | 76 +++++++++++++++++++++---------------- tests/test_cost_tracking.py | 33 +++++++++------- 3 files changed, 64 insertions(+), 49 deletions(-) diff --git a/strix/config/models.py b/strix/config/models.py index dceb6654a..1401dc5a8 100644 --- a/strix/config/models.py +++ b/strix/config/models.py @@ -298,12 +298,12 @@ def _install_openrouter_stream_cost_capture() -> None: OpenrouterConfig, ) - from strix.report.state import remember_streamed_openrouter_cost + from strix.report.state import streamed_openrouter_costs class _StrixOpenRouterStreamingHandler(OpenRouterChatCompletionStreamingHandler): def chunk_parser(self, chunk: dict[str, Any]) -> Any: stream = super().chunk_parser(chunk) - remember_streamed_openrouter_cost( + streamed_openrouter_costs.remember( chunk.get("id") or getattr(stream, "id", None), chunk.get("usage") ) return stream diff --git a/strix/report/state.py b/strix/report/state.py index 501673ce8..490afa961 100644 --- a/strix/report/state.py +++ b/strix/report/state.py @@ -2,7 +2,6 @@ import logging import subprocess import threading -from collections import OrderedDict from collections.abc import Callable from datetime import UTC, datetime from importlib.metadata import PackageNotFoundError, version @@ -97,6 +96,8 @@ def get_global_report_state() -> Optional["ReportState"]: def set_global_report_state(report_state: "ReportState") -> None: global _global_report_state # noqa: PLW0603 _global_report_state = report_state + # New run: drop any streamed-cost entries a prior run left unconsumed. + streamed_openrouter_costs.clear() class ReportState: @@ -509,18 +510,6 @@ def _hydrate_llm_usage(self, raw_usage: Any) -> None: self._sync_llm_usage_record() -# LiteLLM rebuilds streamed responses from token-only chunks and drops the -# provider-reported ``usage.cost`` that OpenRouter sends in its final stream -# chunk (unlike the non-streamed path, which stashes it in hidden params). Since -# every scan streams, that cost never reaches the callback below. The OpenRouter -# streaming handler (see strix.config.models) stashes the cost here keyed by the -# response id so the callback can recover the exact charge for the matching -# rebuilt response. -_STREAMED_OPENROUTER_COST_LIMIT = 4096 -_streamed_openrouter_costs: OrderedDict[str, float] = OrderedDict() -_streamed_openrouter_costs_lock = threading.Lock() - - def openrouter_stream_cost(usage: Any) -> float | None: """Total OpenRouter-reported cost from a raw stream ``usage`` block, or None. @@ -542,28 +531,49 @@ def openrouter_stream_cost(usage: Any) -> float | None: return total if total > 0 else None -def remember_streamed_openrouter_cost(response_id: Any, usage: Any) -> None: - """Record an OpenRouter stream's reported cost so the cost callback can read it.""" - if not isinstance(response_id, str) or not response_id: - return - cost = openrouter_stream_cost(usage) - if cost is None: - return - with _streamed_openrouter_costs_lock: - _streamed_openrouter_costs[response_id] = cost - _streamed_openrouter_costs.move_to_end(response_id) - while len(_streamed_openrouter_costs) > _STREAMED_OPENROUTER_COST_LIMIT: - _streamed_openrouter_costs.popitem(last=False) - - -def _take_streamed_openrouter_cost(completion_response: Any) -> float | None: +def _response_id(completion_response: Any) -> str | None: response_id = getattr(completion_response, "id", None) if response_id is None and isinstance(completion_response, dict): response_id = cast("dict[str, Any]", completion_response).get("id") - if not isinstance(response_id, str) or not response_id: - return None - with _streamed_openrouter_costs_lock: - return _streamed_openrouter_costs.pop(response_id, None) + return response_id if isinstance(response_id, str) and response_id else None + + +class StreamedOpenRouterCosts: + """Correlates OpenRouter's per-stream cost from the parser to the cost callback. + + LiteLLM rebuilds streamed responses from token-only chunks and drops the + ``usage.cost`` OpenRouter reports in its final stream chunk (its non-streamed + path preserves it; streaming snapshots hidden params at stream start). Every + scan streams, so the OpenRouter streaming handler (see strix.config.models) + records the cost here keyed by response id, and the callback takes it back out + for the matching rebuilt response. Entries are removed on read; ``clear()`` + runs per scan so nothing accumulates across runs. + """ + + def __init__(self) -> None: + self._costs: dict[str, float] = {} + self._lock = threading.Lock() + + def remember(self, response_id: Any, usage: Any) -> None: + cost = openrouter_stream_cost(usage) + if cost is None or not (isinstance(response_id, str) and response_id): + return + with self._lock: + self._costs[response_id] = cost + + def take(self, completion_response: Any) -> float | None: + response_id = _response_id(completion_response) + if response_id is None: + return None + with self._lock: + return self._costs.pop(response_id, None) + + def clear(self) -> None: + with self._lock: + self._costs.clear() + + +streamed_openrouter_costs = StreamedOpenRouterCosts() def litellm_cost_callback( @@ -603,7 +613,7 @@ def litellm_cost_callback( # Recover the exact OpenRouter cost the streaming handler stashed for this # response — LiteLLM drops it from streamed usage, so nothing above sees it. if cost is None: - cost = _take_streamed_openrouter_cost(completion_response) + cost = streamed_openrouter_costs.take(completion_response) if cost is None: cost = _estimate_response_cost(kwargs, completion_response) diff --git a/tests/test_cost_tracking.py b/tests/test_cost_tracking.py index a4039c52e..5d8b70790 100644 --- a/tests/test_cost_tracking.py +++ b/tests/test_cost_tracking.py @@ -8,21 +8,22 @@ import litellm import pytest -import strix.report.state as state_module from strix.config.models import ( _configure_litellm_compatibility, _install_openrouter_stream_cost_capture, ) from strix.report.state import ( + ReportState, litellm_cost_callback, openrouter_stream_cost, - remember_streamed_openrouter_cost, + set_global_report_state, + streamed_openrouter_costs, ) @pytest.fixture(autouse=True) def _clear_streamed_costs() -> None: - state_module._streamed_openrouter_costs.clear() + streamed_openrouter_costs.clear() def test_streaming_logging_stays_enabled_for_cost_callback() -> None: @@ -181,7 +182,7 @@ def test_openrouter_stream_cost_extracts_plain_and_byok_totals() -> None: def test_cost_callback_recovers_streamed_openrouter_cost_by_response_id() -> None: report_state = MagicMock() - remember_streamed_openrouter_cost("gen-abc", {"cost": 0.42}) + streamed_openrouter_costs.remember("gen-abc", {"cost": 0.42}) # LiteLLM strips cost from the rebuilt streamed usage; only the id survives. response = SimpleNamespace(id="gen-abc", usage=SimpleNamespace(cost=None), _hidden_params={}) @@ -193,12 +194,12 @@ def test_cost_callback_recovers_streamed_openrouter_cost_by_response_id() -> Non report_state.record_observed_llm_cost.assert_called_once_with(0.42) # The entry is consumed so a later response cannot double-count it. - assert "gen-abc" not in state_module._streamed_openrouter_costs + assert streamed_openrouter_costs.take(response) is None def test_streamed_openrouter_cost_prefers_provider_report_over_estimate() -> None: report_state = MagicMock() - remember_streamed_openrouter_cost("gen-xyz", {"cost": 0.9}) + streamed_openrouter_costs.remember("gen-xyz", {"cost": 0.9}) response = SimpleNamespace( id="gen-xyz", usage=SimpleNamespace(prompt_tokens=10, completion_tokens=5, total_tokens=15), @@ -215,14 +216,16 @@ def test_streamed_openrouter_cost_prefers_provider_report_over_estimate() -> Non estimate.assert_not_called() -def test_remember_streamed_openrouter_cost_evicts_oldest_over_limit() -> None: - limit = state_module._STREAMED_OPENROUTER_COST_LIMIT - for i in range(limit + 5): - remember_streamed_openrouter_cost(f"gen-{i}", {"cost": 0.001}) +def test_streamed_openrouter_costs_ignores_entries_without_cost() -> None: + streamed_openrouter_costs.remember("gen-none", {"prompt_tokens": 10}) + streamed_openrouter_costs.remember("", {"cost": 0.5}) + assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-none")) is None - assert len(state_module._streamed_openrouter_costs) == limit - assert "gen-0" not in state_module._streamed_openrouter_costs - assert f"gen-{limit + 4}" in state_module._streamed_openrouter_costs + +def test_streamed_openrouter_costs_cleared_on_new_run() -> None: + streamed_openrouter_costs.remember("gen-stale", {"cost": 0.7}) + set_global_report_state(ReportState.__new__(ReportState)) + assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stale")) is None def test_openrouter_stream_handler_records_cost() -> None: @@ -243,4 +246,6 @@ def test_openrouter_stream_handler_records_cost() -> None: handler = handler_cls(streaming_response=iter([]), sync_stream=True) handler.chunk_parser(chunk) - assert state_module._streamed_openrouter_costs["gen-stream"] == pytest.approx(0.0035055) + assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stream")) == pytest.approx( + 0.0035055 + ) From d2fd1d610330a6caf89242400be0231daea2ce98 Mon Sep 17 00:00:00 2001 From: Alex Schapiro Date: Tue, 28 Jul 2026 16:36:12 +0000 Subject: [PATCH 3/4] test(cost): resolve OpenRouter handler via LiteLLM provider pipeline --- tests/test_cost_tracking.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/tests/test_cost_tracking.py b/tests/test_cost_tracking.py index 5d8b70790..065b7cb8b 100644 --- a/tests/test_cost_tracking.py +++ b/tests/test_cost_tracking.py @@ -7,6 +7,8 @@ import litellm import pytest +from litellm.types.utils import LlmProviders +from litellm.utils import ProviderConfigManager from strix.config.models import ( _configure_litellm_compatibility, @@ -230,11 +232,15 @@ def test_streamed_openrouter_costs_cleared_on_new_run() -> None: def test_openrouter_stream_handler_records_cost() -> None: _install_openrouter_stream_cost_capture() - handler_cls = ( - litellm.OpenrouterConfig() - .get_model_response_iterator(streaming_response=iter([]), sync_stream=True) - .__class__ + # Resolve the config the way LiteLLM does in production so we prove the + # override is actually reachable through provider resolution, not just as a + # directly-constructed class. + config = ProviderConfigManager.get_provider_chat_config( + model="moonshotai/kimi-k3", provider=LlmProviders.OPENROUTER ) + assert config is not None + assert type(config).__name__ == "_StrixOpenrouterConfig" + handler = config.get_model_response_iterator(streaming_response=iter([]), sync_stream=True) chunk = { "id": "gen-stream", @@ -243,7 +249,6 @@ def test_openrouter_stream_handler_records_cost() -> None: "choices": [{"index": 0, "delta": {"content": None}}], "usage": {"prompt_tokens": 89, "completion_tokens": 138, "cost": 0.0035055}, } - handler = handler_cls(streaming_response=iter([]), sync_stream=True) handler.chunk_parser(chunk) assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stream")) == pytest.approx( From b861413d10a9850b4fc5aed9ac5a2751370d0237 Mon Sep 17 00:00:00 2001 From: Alex Schapiro Date: Wed, 29 Jul 2026 03:27:40 +0000 Subject: [PATCH 4/4] fix(toolchoice): default required tool choice on OpenAI-compatible custom endpoints Reasoning models on OpenAI-compatible custom endpoints (e.g. GLM / Kimi) often reason and reply with prose instead of emitting native tool calls under the default auto tool choice, stalling the tool-driven scan loop while burning tokens. Default force_required_tool_choice to auto (None): enable tool_choice=required when an api_base is set (still gated to models that accept it), off otherwise. Explicit 0/1 overrides. --- strix/config/settings.py | 7 ++++-- strix/core/inputs.py | 12 ++++++++-- strix/core/runner.py | 1 + tests/test_inputs.py | 38 ++++++++++++++++++++++++++++++++ tests/test_runner_rate_limit.py | 1 + tests/test_runner_root_prompt.py | 1 + 6 files changed, 56 insertions(+), 4 deletions(-) diff --git a/strix/config/settings.py b/strix/config/settings.py index 78d52273b..13b96c197 100644 --- a/strix/config/settings.py +++ b/strix/config/settings.py @@ -36,8 +36,11 @@ class LlmSettings(BaseSettings): ), ) reasoning_effort: ReasoningEffort = Field(default="high", alias="STRIX_REASONING_EFFORT") - force_required_tool_choice: bool = Field( - default=False, + # None = auto: force required tool choice on OpenAI-compatible custom + # endpoints (where reasoning models otherwise burn tokens thinking without + # ever calling a tool), off elsewhere. True/False overrides explicitly. + force_required_tool_choice: bool | None = Field( + default=None, alias="STRIX_FORCE_REQUIRED_TOOL_CHOICE", ) prompt_cache: bool = Field( diff --git a/strix/core/inputs.py b/strix/core/inputs.py index aef1fe137..a39b6523e 100644 --- a/strix/core/inputs.py +++ b/strix/core/inputs.py @@ -129,7 +129,8 @@ def make_model_settings( reasoning_effort: ReasoningEffort | None, *, model_name: str, - force_required_tool_choice: bool = False, + force_required_tool_choice: bool | None = None, + custom_api_base: bool = False, request_timeout: float | None = None, prompt_cache: bool = True, ) -> ModelSettings: @@ -147,7 +148,14 @@ def make_model_settings( model_settings = model_settings.resolve( ModelSettings(reasoning=Reasoning(effort=reasoning_effort)), ) - if force_required_tool_choice and _accepts_required_tool_choice(model_name): + # Strix is fully tool-driven, so a turn that returns prose instead of a tool + # call stalls the scan. Reasoning models on OpenAI-compatible custom + # endpoints (e.g. GLM / Kimi via cortecs) do exactly that under the default + # ``auto`` tool choice, so default to ``required`` there; ``None`` means auto. + use_required = ( + custom_api_base if force_required_tool_choice is None else force_required_tool_choice + ) + if use_required and _accepts_required_tool_choice(model_name): model_settings = model_settings.resolve(ModelSettings(tool_choice="required")) cache_extra_args = _prompt_cache_extra_args(model_name) if prompt_cache else None diff --git a/strix/core/runner.py b/strix/core/runner.py index c5f51b154..3b29aed6f 100644 --- a/strix/core/runner.py +++ b/strix/core/runner.py @@ -248,6 +248,7 @@ async def _spill_to_workspace(output_id: str, text: str) -> str | None: settings.llm.reasoning_effort, model_name=resolved_model, force_required_tool_choice=settings.llm.force_required_tool_choice, + custom_api_base=bool(settings.llm.api_base), request_timeout=settings.llm.timeout, prompt_cache=settings.llm.prompt_cache, ) diff --git a/tests/test_inputs.py b/tests/test_inputs.py index da914879a..c0c0f275e 100644 --- a/tests/test_inputs.py +++ b/tests/test_inputs.py @@ -255,6 +255,44 @@ def test_make_model_settings_forces_required_for_anyllm_routed_openai_model() -> assert settings.tool_choice == "required" +def test_make_model_settings_auto_forces_required_on_custom_openai_endpoint() -> None: + # GLM / Kimi via cortecs: openai/-routed reasoning model on a custom base. + settings = make_model_settings( + None, + model_name="openai/glm-5.2", + custom_api_base=True, + ) + + assert settings.tool_choice == "required" + + +def test_make_model_settings_auto_skips_required_without_custom_endpoint() -> None: + settings = make_model_settings(None, model_name="openai/glm-5.2") + + assert settings.tool_choice is None + + +def test_make_model_settings_explicit_false_overrides_custom_endpoint() -> None: + settings = make_model_settings( + None, + model_name="openai/glm-5.2", + force_required_tool_choice=False, + custom_api_base=True, + ) + + assert settings.tool_choice is None + + +def test_make_model_settings_auto_skips_required_for_non_openai_custom_endpoint() -> None: + settings = make_model_settings( + None, + model_name="anthropic/claude-3-7-sonnet-latest", + custom_api_base=True, + ) + + assert settings.tool_choice is None + + def test_make_model_settings_sets_request_timeout() -> None: settings = make_model_settings( "none", diff --git a/tests/test_runner_rate_limit.py b/tests/test_runner_rate_limit.py index 482c4730b..c5e22ac99 100644 --- a/tests/test_runner_rate_limit.py +++ b/tests/test_runner_rate_limit.py @@ -38,6 +38,7 @@ async def test_persistent_rate_limit_stops_gracefully( model="openai/gpt-4o", reasoning_effort="high", force_required_tool_choice=False, + api_base=None, timeout=300, prompt_cache=True, ), diff --git a/tests/test_runner_root_prompt.py b/tests/test_runner_root_prompt.py index 56d7caa6e..ef8be1a02 100644 --- a/tests/test_runner_root_prompt.py +++ b/tests/test_runner_root_prompt.py @@ -46,6 +46,7 @@ def _patch_engine_scaffold( model="openai/gpt-4o", reasoning_effort="high", force_required_tool_choice=False, + api_base=None, timeout=300, prompt_cache=True, ),