Openrouter sticky sessions for caching, with telemetry

This commit is contained in:
Ian 2026-09-29 18:01:44 -04:00 • committed by Ahmed Allam
parent 463b149bdb
commit 95fbd8d687
6 changed files with 249 additions and 4 deletions

View file

@ -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.)

View file

@ -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",

View file

@ -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,

View file

@ -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:

View file

@ -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(),
},
)

View file

@ -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)