From 4b76c73c3ff3a80480173be22e17cf6d2533d7e2 Mon Sep 17 00:00:00 2001 From: Alex Schapiro Date: Tue, 28 Jul 2026 15:59:22 +0000 Subject: [PATCH] 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 dceb6654..1401dc5a 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 501673ce..490afa96 100644 --- a/strix/report/state.py +++ b/strix/report/state.py @@ -2,7 +2,6 @@ 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 @@ -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 @@ class ReportState: 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 a4039c52..5d8b7079 100644 --- a/tests/test_cost_tracking.py +++ b/tests/test_cost_tracking.py @@ -8,21 +8,22 @@ from unittest.mock import MagicMock, patch 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 + )