mirror of
https://github.com/usestrix/strix.git
synced 2026-10-01 02:03:55 +00:00
Openrouter sticky sessions for caching, with telemetry
This commit is contained in:
parent
463b149bdb
commit
95fbd8d687
6 changed files with 249 additions and 4 deletions
|
|
@ -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.)
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
},
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue