diff --git a/strix/report/pricing.py b/strix/report/pricing.py new file mode 100644 index 00000000..57c89959 --- /dev/null +++ b/strix/report/pricing.py @@ -0,0 +1,54 @@ +"""LiteLLM model-name resolution for local cost estimates.""" + +from __future__ import annotations + +from functools import lru_cache +from typing import Any, cast + + +@lru_cache(maxsize=512) +def resolve_litellm_model(model: str) -> str | None: + """Return a provider-qualified model name that LiteLLM can price.""" + try: + import litellm + + normalized = model.strip() + for prefix in ("litellm/", "any-llm/", "openai/"): + if normalized.startswith(prefix): + normalized = normalized.removeprefix(prefix) + break + if not normalized: + return None + + model_cost = cast( + "dict[str, dict[str, Any]]", + getattr(litellm, "model_cost"), # noqa: B009 + ) + bare_entry = model_cost.get(normalized) + if "/" not in normalized and isinstance(bare_entry, dict): + provider = bare_entry.get("litellm_provider") + if isinstance(provider, str) and provider: + return f"{provider}/{normalized}" + if "/" in normalized and isinstance(bare_entry, dict): + return normalized + + names = [normalized] + if "/" in normalized: + names.append(normalized.rsplit("/", 1)[-1]) + for name in names: + matches = sorted(key for key in model_cost if key.endswith(f"/{name}")) + if not matches: + continue + prices = { + ( + model_cost[key].get("input_cost_per_token"), + model_cost[key].get("output_cost_per_token"), + ) + for key in matches + if isinstance(model_cost.get(key), dict) + } + if len(matches) == 1 or len(prices) == 1: + return matches[0] + return None # noqa: TRY300 + except Exception: # noqa: BLE001 + return None diff --git a/strix/report/state.py b/strix/report/state.py index 490afa96..ce7662b5 100644 --- a/strix/report/state.py +++ b/strix/report/state.py @@ -14,6 +14,7 @@ from agents.usage import Usage from strix.config import codex from strix.config.loader import load_settings from strix.core.paths import run_dir_for +from strix.report.pricing import resolve_litellm_model from strix.report.sarif import write_sarif from strix.report.usage import LLMUsageLedger from strix.report.writer import ( @@ -696,10 +697,13 @@ def _estimate_response_cost(kwargs: Any, completion_response: Any) -> float | No candidates.append(model.rsplit("/", 1)[-1]) for candidate in candidates: + resolved = resolve_litellm_model(candidate) + if not resolved: + continue try: value = completion_cost( - completion_response={"model": candidate, "usage": usage_payload}, - model=candidate, + completion_response={"model": resolved, "usage": usage_payload}, + model=resolved, ) except Exception: # nosec B112 # noqa: BLE001, S112 continue diff --git a/strix/report/usage.py b/strix/report/usage.py index e3ddf494..3d6be050 100644 --- a/strix/report/usage.py +++ b/strix/report/usage.py @@ -7,6 +7,8 @@ from typing import Any from agents.usage import Usage, deserialize_usage, serialize_usage +from strix.report.pricing import resolve_litellm_model + logger = logging.getLogger(__name__) @@ -18,7 +20,9 @@ class LLMUsageLedger: self._total_usage = Usage() self._agent_usage: dict[str, Usage] = {} self._agent_metadata: dict[str, dict[str, str]] = {} - self._total_cost = 0.0 + self._observed_cost = 0.0 + self._estimated_cost = 0.0 + self._has_observed_cost = False # 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 @@ -44,10 +48,10 @@ class LLMUsageLedger: if model: metadata["model"] = model - if not self.zero_cost and not _is_litellm_routed(model): + if not self.zero_cost: estimated = _estimate_litellm_cost(usage, model) if estimated: - self._total_cost += estimated + self._estimated_cost += estimated return True @@ -55,15 +59,18 @@ class LLMUsageLedger: if self.zero_cost: return if isinstance(cost, int | float) and cost > 0: - self._total_cost += float(cost) + self._observed_cost += float(cost) + self._has_observed_cost = True @property def total_cost(self) -> float: - return _round_cost(self._total_cost) + if self.zero_cost: + return 0.0 + return _round_cost(self._observed_cost if self._has_observed_cost else self._estimated_cost) def to_record(self) -> dict[str, Any]: record = serialize_usage(self._total_usage) - record["cost"] = _round_cost(self._total_cost) + record["cost"] = self.total_cost record["agents"] = [] agent_tokens = {aid: _resolve_total_tokens(u) for aid, u in self._agent_usage.items()} @@ -72,7 +79,7 @@ class LLMUsageLedger: usage = self._agent_usage[agent_id] metadata = self._agent_metadata.get(agent_id, {}) agent_cost = ( - self._total_cost * (agent_tokens[agent_id] / total_tokens) if total_tokens else 0.0 + self.total_cost * (agent_tokens[agent_id] / total_tokens) if total_tokens else 0.0 ) agent_record = serialize_usage(usage) @@ -92,7 +99,9 @@ class LLMUsageLedger: self._total_usage = Usage() self._agent_usage.clear() self._agent_metadata.clear() - self._total_cost = 0.0 + self._observed_cost = 0.0 + self._estimated_cost = 0.0 + self._has_observed_cost = False if not isinstance(raw_usage, dict): return @@ -103,7 +112,9 @@ class LLMUsageLedger: logger.exception("Failed to hydrate aggregate llm_usage from run.json") self._total_usage = Usage() - self._total_cost = _float_or_zero(raw_usage.get("cost")) + persisted_cost = _float_or_zero(raw_usage.get("cost")) + self._observed_cost = persisted_cost + self._estimated_cost = persisted_cost for raw_agent in raw_usage.get("agents") or []: if not isinstance(raw_agent, dict): @@ -136,15 +147,6 @@ def _resolve_total_tokens(usage: Usage) -> int: return prompt + completion -def _is_litellm_routed(model: str | None) -> bool: - if not model: - return False - name = model.strip().lower() - if "/" not in name: - return False - return not name.startswith("openai/") - - def _usage_has_activity(usage: Usage) -> bool: return bool( usage.requests @@ -201,24 +203,23 @@ def _estimate_litellm_entry_cost(entry: Any, model: str) -> float | None: candidates = [model] if "/" in model: - candidates.append(model.split("/", 1)[-1]) + candidates.append(model.rsplit("/", 1)[-1]) - cost: Any = None for candidate in candidates: + resolved = resolve_litellm_model(candidate) + if not resolved: + continue try: cost = completion_cost( - completion_response={"model": candidate, "usage": usage_payload}, - model=model, + completion_response={"model": resolved, "usage": usage_payload}, + model=resolved, ) - break except Exception: # nosec B112 # noqa: BLE001, S112 continue - - if cost is None: - logger.debug("LiteLLM cost estimate unavailable for model %s", model) - return None - - return cost if isinstance(cost, int | float) and cost >= 0 else None + if cost > 0: + return float(cost) + logger.debug("LiteLLM cost estimate unavailable for model %s", model) + return None def _litellm_model_name(model: str | None) -> str | None: diff --git a/tests/test_cost_tracking.py b/tests/test_cost_tracking.py index 30d4db44..6db31145 100644 --- a/tests/test_cost_tracking.py +++ b/tests/test_cost_tracking.py @@ -143,7 +143,7 @@ def test_cost_callback_estimates_cost_with_bare_model_fallback() -> None: } def fake_completion_cost(**kwargs: object) -> float: - if kwargs["model"] == "gpt-4o-mini": + if kwargs["model"] == "openai/gpt-4o-mini": return 0.025 raise ValueError(kwargs["model"]) diff --git a/tests/test_pricing.py b/tests/test_pricing.py new file mode 100644 index 00000000..abff873d --- /dev/null +++ b/tests/test_pricing.py @@ -0,0 +1,120 @@ +from __future__ import annotations + +from unittest.mock import patch + +import litellm +from agents.usage import Usage + +from strix.report.pricing import resolve_litellm_model +from strix.report.usage import LLMUsageLedger + + +def test_resolves_common_bare_model_names() -> None: + resolve_litellm_model.cache_clear() + assert resolve_litellm_model("deepseek-v4-flash") == "deepseek/deepseek-v4-flash" + assert resolve_litellm_model("openai/deepseek-v4-flash") == "deepseek/deepseek-v4-flash" + assert resolve_litellm_model("grok-4.5") == "xai/grok-4.5" + assert resolve_litellm_model("MiniMax-M3") == "minimax/MiniMax-M3" + + +def test_resolver_returns_none_for_unresolvable_model() -> None: + resolve_litellm_model.cache_clear() + assert resolve_litellm_model("provider/not-a-real-model") is None + + +def test_ledger_uses_estimate_when_routed_provider_reports_no_cost() -> None: + usage = Usage() + usage.requests = 1 + usage.input_tokens = 1000 + usage.output_tokens = 200 + usage.total_tokens = 1200 + ledger = LLMUsageLedger() + + with patch("litellm.completion_cost", return_value=0.42): + ledger.record(agent_id="a", usage=usage, model="openai/deepseek-v4-flash") + + assert ledger.total_cost == 0.42 + + +def test_ledger_prefers_observed_cost_over_estimate() -> None: + usage = Usage() + usage.requests = 1 + usage.input_tokens = 1000 + usage.output_tokens = 200 + usage.total_tokens = 1200 + ledger = LLMUsageLedger() + + with patch("litellm.completion_cost", return_value=0.42): + ledger.record(agent_id="a", usage=usage, model="openai/deepseek-v4-flash") + ledger.record_observed_cost(0.17) + + assert ledger.total_cost == 0.17 + + +def test_hydrated_estimate_continues_accumulating_new_estimates() -> None: + usage = Usage() + usage.requests = 1 + usage.input_tokens = 1000 + usage.output_tokens = 200 + usage.total_tokens = 1200 + ledger = LLMUsageLedger() + ledger.hydrate({"cost": 0.42}) + + with patch("litellm.completion_cost", return_value=0.17): + ledger.record(agent_id="a", usage=usage, model="openai/deepseek-v4-flash") + + assert ledger.total_cost == 0.59 + + +def test_zero_cost_disables_both_observed_and_estimated_costs() -> None: + usage = Usage() + usage.requests = 1 + usage.input_tokens = 1000 + usage.output_tokens = 200 + usage.total_tokens = 1200 + ledger = LLMUsageLedger() + ledger.zero_cost = True + + with patch("litellm.completion_cost", return_value=0.42) as estimate: + ledger.record(agent_id="a", usage=usage, model="deepseek-v4-flash") + ledger.record_observed_cost(1.0) + + estimate.assert_not_called() + assert ledger.total_cost == 0.0 + + +def test_resolver_uses_provider_when_bare_entry_has_one() -> None: + original = litellm.model_cost + litellm.model_cost = { + "example": { + "litellm_provider": "example-provider", + "input_cost_per_token": 1.0, + "output_cost_per_token": 2.0, + } + } + try: + resolve_litellm_model.cache_clear() + assert resolve_litellm_model("example") == "example-provider/example" + finally: + litellm.model_cost = original + resolve_litellm_model.cache_clear() + + +def test_resolver_does_not_guess_between_differently_priced_providers() -> None: + original = litellm.model_cost + litellm.model_cost = { + "provider-a/example": { + "input_cost_per_token": 1.0, + "output_cost_per_token": 2.0, + }, + "provider-b/example": { + "input_cost_per_token": 3.0, + "output_cost_per_token": 4.0, + }, + } + try: + resolve_litellm_model.cache_clear() + assert resolve_litellm_model("example") is None + finally: + litellm.model_cost = original + resolve_litellm_model.cache_clear()