Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
45 changes: 45 additions & 0 deletions strix/config/models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 streamed_openrouter_costs

class _StrixOpenRouterStreamingHandler(OpenRouterChatCompletionStreamingHandler):
def chunk_parser(self, chunk: dict[str, Any]) -> Any:
stream = super().chunk_parser(chunk)
streamed_openrouter_costs.remember(
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 = {
Expand Down
7 changes: 5 additions & 2 deletions strix/config/settings.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
12 changes: 10 additions & 2 deletions strix/core/inputs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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
Expand Down
1 change: 1 addition & 0 deletions strix/core/runner.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
)
Expand Down
74 changes: 74 additions & 0 deletions strix/report/state.py
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
import json
import logging
import subprocess
import threading
from collections.abc import Callable
from datetime import UTC, datetime
from importlib.metadata import PackageNotFoundError, version
Expand Down Expand Up @@ -95,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:
Expand Down Expand Up @@ -507,6 +510,72 @@ def _hydrate_llm_usage(self, raw_usage: Any) -> None:
self._sync_llm_usage_record()


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 _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")
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(
kwargs: Any,
completion_response: Any,
Expand Down Expand Up @@ -541,6 +610,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 = streamed_openrouter_costs.take(completion_response)

if cost is None:
cost = _estimate_response_cost(kwargs, completion_response)

Expand Down
107 changes: 105 additions & 2 deletions tests/test_cost_tracking.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,9 +7,25 @@

import litellm
import pytest
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager

from strix.config.models import _configure_litellm_compatibility
from strix.report.state import litellm_cost_callback
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,
set_global_report_state,
streamed_openrouter_costs,
)


@pytest.fixture(autouse=True)
def _clear_streamed_costs() -> None:
streamed_openrouter_costs.clear()


def test_streaming_logging_stays_enabled_for_cost_callback() -> None:
Expand Down Expand Up @@ -151,3 +167,90 @@ 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()
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={})

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 streamed_openrouter_costs.take(response) is None


def test_streamed_openrouter_cost_prefers_provider_report_over_estimate() -> None:
report_state = MagicMock()
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),
_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_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


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:
_install_openrouter_stream_cost_capture()
# 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",
"created": 1,
"model": "moonshotai/kimi-k3",
"choices": [{"index": 0, "delta": {"content": None}}],
"usage": {"prompt_tokens": 89, "completion_tokens": 138, "cost": 0.0035055},
}
handler.chunk_parser(chunk)

assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stream")) == pytest.approx(
0.0035055
)
38 changes: 38 additions & 0 deletions tests/test_inputs.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
1 change: 1 addition & 0 deletions tests/test_runner_rate_limit.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
),
Expand Down
1 change: 1 addition & 0 deletions tests/test_runner_root_prompt.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
),
Expand Down