diff --git a/strix/config/models.py b/strix/config/models.py index 6eedf3788..e33f30604 100644 --- a/strix/config/models.py +++ b/strix/config/models.py @@ -8,6 +8,7 @@ import inspect import logging import os import time +import uuid from collections.abc import AsyncGenerator from typing import TYPE_CHECKING, Any, cast @@ -746,6 +747,11 @@ def _configure_litellm_compatibility() -> None: _install_openrouter_stream_cost_capture() +# Agent ids are 8 hex characters and can repeat across runs; the session id +# OpenRouter pins a provider to must not, so each agent gets its own UUID. +_OPENROUTER_SESSION_IDS: dict[str, str] = {} + + def _install_openrouter_stream_cost_capture() -> None: """Preserve OpenRouter's per-stream cost, which LiteLLM drops when streaming. @@ -764,14 +770,16 @@ def _install_openrouter_stream_cost_capture() -> None: OpenrouterConfig, ) - from strix.report.state import streamed_openrouter_costs + from strix.report.state import record_openrouter_provider, 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") - ) + usage = chunk.get("usage") + response_id = chunk.get("id") or getattr(stream, "id", None) + streamed_openrouter_costs.remember(response_id, usage) + if usage: + record_openrouter_provider(chunk.get("provider"), usage) return stream class _StrixOpenrouterConfig(OpenrouterConfig): @@ -784,6 +792,16 @@ def _install_openrouter_stream_cost_capture() -> None: json_mode=json_mode, ) + def transform_request(self, *args: Any, **kwargs: Any) -> dict[str, Any]: + # Pin each agent's calls to one upstream provider so its prompt cache + # survives between turns. + body = super().transform_request(*args, **kwargs) + agent_id = request_log.current_call_context().agent_id + if agent_id: + session_id = _OPENROUTER_SESSION_IDS.setdefault(agent_id, str(uuid.uuid4())) + body.setdefault("session_id", session_id) + return body + # 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.) diff --git a/strix/config/settings.py b/strix/config/settings.py index 7bf1de7f4..e9ed84272 100644 --- a/strix/config/settings.py +++ b/strix/config/settings.py @@ -58,6 +58,10 @@ class LlmSettings(BaseSettings): default=True, alias="STRIX_PROMPT_CACHE", ) + # Providers cache prompts in fixed-size token blocks, so a fully cached prompt + # reads back rounded down to a multiple of this. 64 is what the GLM calls in + # local runs showed; it's a per-deployment setting (vLLM defaults to 16). + cache_block_tokens: int = Field(default=64, ge=1, alias="STRIX_CACHE_BLOCK_TOKENS") disable_streaming: bool = Field( default=False, alias="LLM_DISABLE_STREAMING", diff --git a/strix/report/state.py b/strix/report/state.py index 5d13483e2..fd0ebd944 100644 --- a/strix/report/state.py +++ b/strix/report/state.py @@ -657,6 +657,35 @@ class ReportState: def record_observed_llm_cost(self, cost: float) -> None: self._llm_usage.record_observed_cost(cost) + def record_llm_provider( + self, + provider: str, + *, + agent_id: str | None, + input_tokens: int, + cached_tokens: int, + cost: float, + ) -> None: + self._llm_usage.record_provider( + provider, + agent_id=agent_id, + input_tokens=input_tokens, + cached_tokens=cached_tokens, + cost=cost, + cache_block_tokens=load_settings().llm.cache_block_tokens, + ) + + def get_process_llm_providers(self) -> dict[str, dict[str, float]]: + """Per-provider usage since this process started, like get_process_llm_usage.""" + baseline = self._telemetry_llm_usage_baseline.get("providers") or {} + providers: dict[str, dict[str, float]] = {} + for name, tally in (self._llm_usage.to_record().get("providers") or {}).items(): + before = baseline.get(name) or {} + delta = {key: max(0, value - _number(before.get(key))) for key, value in tally.items()} + if delta["requests"]: + providers[name] = delta + return providers + def get_total_llm_usage(self) -> dict[str, Any]: return dict(self.run_record.get("llm_usage") or self._build_llm_usage_record()) @@ -990,6 +1019,30 @@ class StreamedOpenRouterCosts: streamed_openrouter_costs = StreamedOpenRouterCosts() +def record_openrouter_provider(provider: Any, usage: Any) -> None: + """Tally which upstream provider served a stream, from its final usage chunk. + + OpenRouter spreads one model across many providers whose prices, quantization + and prompt caching differ, so this is what shows where a scan's tokens went. + """ + # Deferred: request_log pulls in the agents SDK, which strix.report must not import. + from strix.llm.request_log import current_call_context + + report_state = get_global_report_state() + if report_state is None or not isinstance(usage, dict): + return + details = usage.get("prompt_tokens_details") + report_state.record_llm_provider( + provider if isinstance(provider, str) and provider else "unknown", + agent_id=current_call_context().agent_id, + input_tokens=int(_number(usage.get("prompt_tokens"))), + cached_tokens=int(_number(details.get("cached_tokens"))) + if isinstance(details, dict) + else 0, + cost=openrouter_stream_cost(usage) or 0.0, + ) + + def litellm_cost_callback( kwargs: Any, completion_response: Any, diff --git a/strix/report/usage.py b/strix/report/usage.py index 3d6be050f..f84d02e8e 100644 --- a/strix/report/usage.py +++ b/strix/report/usage.py @@ -6,6 +6,7 @@ import logging from typing import Any from agents.usage import Usage, deserialize_usage, serialize_usage +from pydantic import BaseModel, TypeAdapter, ValidationError from strix.report.pricing import resolve_litellm_model @@ -13,6 +14,22 @@ from strix.report.pricing import resolve_litellm_model logger = logging.getLogger(__name__) +class ProviderUsage(BaseModel): + """Running totals for one upstream provider OpenRouter routed calls to.""" + + requests: int = 0 + input_tokens: int = 0 + cached_tokens: int = 0 + cost: float = 0.0 + # Calls that didn't find the agent's whole previous prompt cached, and the + # previous-prompt tokens they had to pay for again. + cache_misses: int = 0 + missed_tokens: int = 0 + + +_PROVIDER_USAGE = TypeAdapter(dict[str, ProviderUsage]) + + class LLMUsageLedger: """Aggregate SDK ``Usage`` objects and attach best-effort cost estimates.""" @@ -23,6 +40,10 @@ class LLMUsageLedger: self._observed_cost = 0.0 self._estimated_cost = 0.0 self._has_observed_cost = False + # Keyed by upstream provider name, e.g. "Z.AI" or "DeepInfra". + self._providers: dict[str, ProviderUsage] = {} + # Each agent's last prompt size, which its next call should find cached. + self._last_input_tokens: dict[str, int] = {} # When True, tokens are still tracked but cost stays $0 — the run is on a # model subscription, so there is no metered per-token charge to report. self.zero_cost = False @@ -62,6 +83,34 @@ class LLMUsageLedger: self._observed_cost += float(cost) self._has_observed_cost = True + def record_provider( + self, + provider: str, + *, + agent_id: str | None, + input_tokens: int, + cached_tokens: int, + cost: float, + cache_block_tokens: int, + ) -> None: + tally = self._providers.setdefault(provider, ProviderUsage()) + tally.requests += 1 + tally.input_tokens += input_tokens + tally.cached_tokens += cached_tokens + if agent_id: + previous = self._last_input_tokens.get(agent_id, 0) + # The previous prompt is a prefix of this one, so every full block of it + # should read back cached. A shrinking prompt means compaction rewrote + # it, so a miss is expected. + expected = previous - (previous % cache_block_tokens) + missed = expected - cached_tokens + if input_tokens >= previous and missed > 0: + tally.cache_misses += 1 + tally.missed_tokens += missed + self._last_input_tokens[agent_id] = input_tokens + if not self.zero_cost: + tally.cost = _round_cost(tally.cost + cost) + @property def total_cost(self) -> float: if self.zero_cost: @@ -71,6 +120,7 @@ class LLMUsageLedger: def to_record(self) -> dict[str, Any]: record = serialize_usage(self._total_usage) record["cost"] = self.total_cost + record["providers"] = {name: tally.model_dump() for name, tally in self._providers.items()} record["agents"] = [] agent_tokens = {aid: _resolve_total_tokens(u) for aid, u in self._agent_usage.items()} @@ -102,10 +152,16 @@ class LLMUsageLedger: self._observed_cost = 0.0 self._estimated_cost = 0.0 self._has_observed_cost = False + self._providers = {} if not isinstance(raw_usage, dict): return + try: + self._providers = _PROVIDER_USAGE.validate_python(raw_usage.get("providers") or {}) + except ValidationError: + logger.exception("Failed to hydrate llm_usage providers from run.json") + try: self._total_usage = deserialize_usage(raw_usage) except Exception: diff --git a/strix/telemetry/posthog.py b/strix/telemetry/posthog.py index 72927fbcf..a7cc284d4 100644 --- a/strix/telemetry/posthog.py +++ b/strix/telemetry/posthog.py @@ -118,6 +118,7 @@ def end(report_state: "ReportState", exit_reason: str = "completed") -> None: } except (TypeError, ValueError, AttributeError): pass + providers = report_state.get_process_llm_providers() report_state.posthog_scan_ended_sent = _send( "scan_ended", @@ -129,6 +130,7 @@ def end(report_state: "ReportState", exit_reason: str = "completed") -> None: "vulnerabilities_total": len(report_state.vulnerability_reports), **{f"vulnerabilities_{k}": v for k, v in vulnerabilities_counts.items()}, **llm_props, + **({"llm_providers": providers} if providers else {}), "skills": get_loaded_skill_names(), }, ) diff --git a/tests/test_cost_tracking.py b/tests/test_cost_tracking.py index 6db311456..c9bc283b5 100644 --- a/tests/test_cost_tracking.py +++ b/tests/test_cost_tracking.py @@ -2,7 +2,9 @@ from __future__ import annotations +import uuid from types import SimpleNamespace +from typing import TYPE_CHECKING, Any from unittest.mock import MagicMock, patch import litellm @@ -14,6 +16,7 @@ from strix.config.models import ( _configure_litellm_compatibility, _install_openrouter_stream_cost_capture, ) +from strix.llm import request_log from strix.report.state import ( ReportState, litellm_cost_callback, @@ -21,6 +24,11 @@ from strix.report.state import ( set_global_report_state, streamed_openrouter_costs, ) +from strix.report.usage import LLMUsageLedger + + +if TYPE_CHECKING: + from litellm.types.llms.openai import AllMessageValues @pytest.fixture(autouse=True) @@ -257,3 +265,107 @@ def test_openrouter_stream_handler_records_cost() -> None: assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stream")) == pytest.approx( 0.0035055 ) + + +def test_openrouter_stream_handler_tallies_provider() -> None: + _install_openrouter_stream_cost_capture() + config = ProviderConfigManager.get_provider_chat_config( + model="z-ai/glm-5.3", provider=LlmProviders.OPENROUTER + ) + assert config is not None + handler = config.get_model_response_iterator(streaming_response=iter([]), sync_stream=True) + report_state = MagicMock() + usage = { + "prompt_tokens": 1000, + "completion_tokens": 10, + "cost": 0.002, + "prompt_tokens_details": {"cached_tokens": 900}, + } + with patch("strix.report.state.get_global_report_state", return_value=report_state): + handler.chunk_parser( + { + "id": "gen-a", + "created": 1, + "model": "z-ai/glm-5.3", + "provider": "Together", + "choices": [{"index": 0, "delta": {"content": None}}], + "usage": usage, + } + ) + + report_state.record_llm_provider.assert_called_once_with( + "Together", agent_id=None, input_tokens=1000, cached_tokens=900, cost=0.002 + ) + + +def test_provider_tally_survives_run_record_round_trip() -> None: + ledger = LLMUsageLedger() + for input_tokens, cached_tokens, cost in [(1000, 900, 0.002), (500, 0, 0.001)]: + ledger.record_provider( + "Together", + agent_id=None, + input_tokens=input_tokens, + cached_tokens=cached_tokens, + cost=cost, + cache_block_tokens=64, + ) + + restored = LLMUsageLedger() + restored.hydrate(ledger.to_record()) + + assert restored.to_record()["providers"] == { + "Together": { + "requests": 2, + "input_tokens": 1500, + "cached_tokens": 900, + "cost": 0.003, + "cache_misses": 0, + "missed_tokens": 0, + } + } + + +def test_provider_tally_counts_cache_misses_per_agent() -> None: + ledger = LLMUsageLedger() + calls = [ + ("Z.AI", "a1", 1000, 0), # first call: nothing to miss + ("Z.AI", "a1", 1200, 960), # previous 1000 cached, rounded down to 64s + ("DeepInfra", "a1", 1500, 200), # 1152 of the previous 1200 due, 952 lost + ("Z.AI", "a2", 800, 0), # another agent's first call + ("Z.AI", "a1", 600, 0), # prompt shrank: compaction, not a miss + ] + for provider, agent_id, input_tokens, cached_tokens in calls: + ledger.record_provider( + provider, + agent_id=agent_id, + input_tokens=input_tokens, + cached_tokens=cached_tokens, + cost=0.0, + cache_block_tokens=64, + ) + + providers = ledger.to_record()["providers"] + assert providers["DeepInfra"]["cache_misses"] == 1 + assert providers["DeepInfra"]["missed_tokens"] == 952 + assert providers["Z.AI"]["cache_misses"] == 0 + + +def test_openrouter_request_carries_agent_session_id() -> None: + _install_openrouter_stream_cost_capture() + config = ProviderConfigManager.get_provider_chat_config( + model="moonshotai/kimi-k3", provider=LlmProviders.OPENROUTER + ) + assert config is not None + messages: list[AllMessageValues] = [{"role": "user", "content": "hi"}] + + def body() -> dict[str, Any]: + return config.transform_request("moonshotai/kimi-k3", messages, {}, {}, {}) + + assert "session_id" not in body() + token = request_log.bind_call_context("a1b2c3d4", "root") + try: + session_id = body()["session_id"] + assert str(uuid.UUID(session_id)) == session_id + assert body()["session_id"] == session_id + finally: + request_log.reset_call_context(token)