mirror of
https://github.com/usestrix/strix.git
synced 2026-08-16 09:26:39 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
b861413d10 | ||
|
|
d2fd1d6103 | ||
|
|
4b76c73c3f | ||
|
|
e5ae8b8ae1 |
@@ -277,6 +277,51 @@ def _configure_litellm_compatibility() -> None:
|
|||||||
litellm.suppress_debug_info = True
|
litellm.suppress_debug_info = True
|
||||||
|
|
||||||
_register_litellm_cost_callback()
|
_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 = {
|
_OPENROUTER_ATTRIBUTION_HEADERS = {
|
||||||
|
|||||||
@@ -36,8 +36,11 @@ class LlmSettings(BaseSettings):
|
|||||||
),
|
),
|
||||||
)
|
)
|
||||||
reasoning_effort: ReasoningEffort = Field(default="high", alias="STRIX_REASONING_EFFORT")
|
reasoning_effort: ReasoningEffort = Field(default="high", alias="STRIX_REASONING_EFFORT")
|
||||||
force_required_tool_choice: bool = Field(
|
# None = auto: force required tool choice on OpenAI-compatible custom
|
||||||
default=False,
|
# 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",
|
alias="STRIX_FORCE_REQUIRED_TOOL_CHOICE",
|
||||||
)
|
)
|
||||||
prompt_cache: bool = Field(
|
prompt_cache: bool = Field(
|
||||||
|
|||||||
+10
-2
@@ -129,7 +129,8 @@ def make_model_settings(
|
|||||||
reasoning_effort: ReasoningEffort | None,
|
reasoning_effort: ReasoningEffort | None,
|
||||||
*,
|
*,
|
||||||
model_name: str,
|
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,
|
request_timeout: float | None = None,
|
||||||
prompt_cache: bool = True,
|
prompt_cache: bool = True,
|
||||||
) -> ModelSettings:
|
) -> ModelSettings:
|
||||||
@@ -147,7 +148,14 @@ def make_model_settings(
|
|||||||
model_settings = model_settings.resolve(
|
model_settings = model_settings.resolve(
|
||||||
ModelSettings(reasoning=Reasoning(effort=reasoning_effort)),
|
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"))
|
model_settings = model_settings.resolve(ModelSettings(tool_choice="required"))
|
||||||
|
|
||||||
cache_extra_args = _prompt_cache_extra_args(model_name) if prompt_cache else None
|
cache_extra_args = _prompt_cache_extra_args(model_name) if prompt_cache else None
|
||||||
|
|||||||
@@ -248,6 +248,7 @@ async def run_strix_scan(
|
|||||||
settings.llm.reasoning_effort,
|
settings.llm.reasoning_effort,
|
||||||
model_name=resolved_model,
|
model_name=resolved_model,
|
||||||
force_required_tool_choice=settings.llm.force_required_tool_choice,
|
force_required_tool_choice=settings.llm.force_required_tool_choice,
|
||||||
|
custom_api_base=bool(settings.llm.api_base),
|
||||||
request_timeout=settings.llm.timeout,
|
request_timeout=settings.llm.timeout,
|
||||||
prompt_cache=settings.llm.prompt_cache,
|
prompt_cache=settings.llm.prompt_cache,
|
||||||
)
|
)
|
||||||
|
|||||||
@@ -1,6 +1,7 @@
|
|||||||
import json
|
import json
|
||||||
import logging
|
import logging
|
||||||
import subprocess
|
import subprocess
|
||||||
|
import threading
|
||||||
from collections.abc import Callable
|
from collections.abc import Callable
|
||||||
from datetime import UTC, datetime
|
from datetime import UTC, datetime
|
||||||
from importlib.metadata import PackageNotFoundError, version
|
from importlib.metadata import PackageNotFoundError, version
|
||||||
@@ -95,6 +96,8 @@ def get_global_report_state() -> Optional["ReportState"]:
|
|||||||
def set_global_report_state(report_state: "ReportState") -> None:
|
def set_global_report_state(report_state: "ReportState") -> None:
|
||||||
global _global_report_state # noqa: PLW0603
|
global _global_report_state # noqa: PLW0603
|
||||||
_global_report_state = report_state
|
_global_report_state = report_state
|
||||||
|
# New run: drop any streamed-cost entries a prior run left unconsumed.
|
||||||
|
streamed_openrouter_costs.clear()
|
||||||
|
|
||||||
|
|
||||||
class ReportState:
|
class ReportState:
|
||||||
@@ -507,6 +510,72 @@ class ReportState:
|
|||||||
self._sync_llm_usage_record()
|
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(
|
def litellm_cost_callback(
|
||||||
kwargs: Any,
|
kwargs: Any,
|
||||||
completion_response: Any,
|
completion_response: Any,
|
||||||
@@ -541,6 +610,11 @@ def litellm_cost_callback(
|
|||||||
if cost is None:
|
if cost is None:
|
||||||
cost = _usage_reported_cost(completion_response)
|
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:
|
if cost is None:
|
||||||
cost = _estimate_response_cost(kwargs, completion_response)
|
cost = _estimate_response_cost(kwargs, completion_response)
|
||||||
|
|
||||||
|
|||||||
+105
-2
@@ -7,9 +7,25 @@ from unittest.mock import MagicMock, patch
|
|||||||
|
|
||||||
import litellm
|
import litellm
|
||||||
import pytest
|
import pytest
|
||||||
|
from litellm.types.utils import LlmProviders
|
||||||
|
from litellm.utils import ProviderConfigManager
|
||||||
|
|
||||||
from strix.config.models import _configure_litellm_compatibility
|
from strix.config.models import (
|
||||||
from strix.report.state import litellm_cost_callback
|
_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:
|
def test_streaming_logging_stays_enabled_for_cost_callback() -> None:
|
||||||
@@ -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)
|
litellm_cost_callback({"response_cost": None, "model": "x/y"}, response)
|
||||||
|
|
||||||
report_state.record_observed_llm_cost.assert_not_called()
|
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
|
||||||
|
)
|
||||||
|
|||||||
@@ -255,6 +255,44 @@ def test_make_model_settings_forces_required_for_anyllm_routed_openai_model() ->
|
|||||||
assert settings.tool_choice == "required"
|
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:
|
def test_make_model_settings_sets_request_timeout() -> None:
|
||||||
settings = make_model_settings(
|
settings = make_model_settings(
|
||||||
"none",
|
"none",
|
||||||
|
|||||||
@@ -38,6 +38,7 @@ async def test_persistent_rate_limit_stops_gracefully(
|
|||||||
model="openai/gpt-4o",
|
model="openai/gpt-4o",
|
||||||
reasoning_effort="high",
|
reasoning_effort="high",
|
||||||
force_required_tool_choice=False,
|
force_required_tool_choice=False,
|
||||||
|
api_base=None,
|
||||||
timeout=300,
|
timeout=300,
|
||||||
prompt_cache=True,
|
prompt_cache=True,
|
||||||
),
|
),
|
||||||
|
|||||||
@@ -46,6 +46,7 @@ def _patch_engine_scaffold(
|
|||||||
model="openai/gpt-4o",
|
model="openai/gpt-4o",
|
||||||
reasoning_effort="high",
|
reasoning_effort="high",
|
||||||
force_required_tool_choice=False,
|
force_required_tool_choice=False,
|
||||||
|
api_base=None,
|
||||||
timeout=300,
|
timeout=300,
|
||||||
prompt_cache=True,
|
prompt_cache=True,
|
||||||
),
|
),
|
||||||
|
|||||||
Reference in New Issue
Block a user