strix/tests/test_cost_tracking.py
2026-09-30 04:46:01 +03:00

391 lines
14 KiB
Python

"""Tests for provider-reported LLM cost capture."""
from __future__ import annotations
import uuid
from types import SimpleNamespace
from typing import TYPE_CHECKING, Any
from unittest.mock import MagicMock, call, patch
import httpx
import litellm
import pytest
from litellm.types.utils import LlmProviders
from litellm.utils import ProviderConfigManager
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,
openrouter_stream_cost,
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)
def _clear_streamed_costs() -> None:
streamed_openrouter_costs.clear()
def test_streaming_logging_stays_enabled_for_cost_callback() -> None:
with (
patch.object(litellm, "disable_streaming_logging", new=True),
patch("strix.config.models._register_litellm_cost_callback") as register,
):
_configure_litellm_compatibility()
assert litellm.disable_streaming_logging is False
register.assert_called_once_with()
def test_cost_callback_reads_openrouter_stream_usage_cost() -> None:
report_state = MagicMock()
response = SimpleNamespace(
usage=SimpleNamespace(cost=1.2345),
_hidden_params={},
)
with patch("strix.report.state.get_global_report_state", return_value=report_state):
litellm_cost_callback({"response_cost": None}, response)
report_state.record_observed_llm_cost.assert_called_once_with(1.2345)
def test_cost_callback_reads_usage_cost_from_mapping_response() -> None:
report_state = MagicMock()
response = {"usage": {"cost": 0.125}}
with patch("strix.report.state.get_global_report_state", return_value=report_state):
litellm_cost_callback({}, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.125)
def test_cost_callback_reads_byok_upstream_inference_cost() -> None:
report_state = MagicMock()
response = SimpleNamespace(
usage=SimpleNamespace(
cost=0,
is_byok=True,
cost_details=SimpleNamespace(upstream_inference_cost=6.75e-06),
),
_hidden_params={},
)
with patch("strix.report.state.get_global_report_state", return_value=report_state):
litellm_cost_callback({"response_cost": None}, response)
report_state.record_observed_llm_cost.assert_called_once_with(6.75e-06)
def test_cost_callback_sums_usage_cost_and_upstream_inference_cost() -> None:
report_state = MagicMock()
response = {
"usage": {
"cost": 0.01,
"is_byok": True,
"cost_details": {"upstream_inference_cost": 0.2},
}
}
with patch("strix.report.state.get_global_report_state", return_value=report_state):
litellm_cost_callback({}, response)
report_state.record_observed_llm_cost.assert_called_once_with(pytest.approx(0.21))
def test_cost_callback_ignores_upstream_cost_for_non_byok_responses() -> None:
report_state = MagicMock()
response = {
"usage": {
"cost": 0.05,
"is_byok": False,
"cost_details": {"upstream_inference_cost": 0.04},
}
}
with patch("strix.report.state.get_global_report_state", return_value=report_state):
litellm_cost_callback({}, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.05)
def test_cost_callback_estimates_cost_with_provider_prefixed_model() -> None:
report_state = MagicMock()
response = {"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}
kwargs = {
"response_cost": None,
"model": "anthropic/claude-sonnet-4.5",
"litellm_params": {"custom_llm_provider": "openrouter"},
}
def fake_completion_cost(**kwargs: object) -> float:
if kwargs["model"] == "openrouter/anthropic/claude-sonnet-4.5":
return 0.5
raise ValueError(kwargs["model"])
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", side_effect=fake_completion_cost),
):
litellm_cost_callback(kwargs, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.5)
def test_cost_callback_estimates_cost_with_bare_model_fallback() -> None:
report_state = MagicMock()
response = {"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}
kwargs = {
"response_cost": None,
"model": "openai/gpt-4o-mini",
"litellm_params": {"custom_llm_provider": "openrouter"},
}
def fake_completion_cost(**kwargs: object) -> float:
if kwargs["model"] == "openai/gpt-4o-mini":
return 0.025
raise ValueError(kwargs["model"])
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", side_effect=fake_completion_cost),
):
litellm_cost_callback(kwargs, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.025)
def test_cost_callback_records_nothing_when_no_cost_available() -> None:
report_state = MagicMock()
response = {"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}}
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", side_effect=ValueError("unknown model")),
):
litellm_cost_callback({"response_cost": None, "model": "x/y"}, response)
report_state.record_observed_llm_cost.assert_not_called()
def test_openrouter_stream_cost_extracts_plain_and_byok_totals() -> None:
assert openrouter_stream_cost({"cost": 0.003168}) == pytest.approx(0.003168)
assert openrouter_stream_cost(
{"cost": 0.01, "is_byok": True, "cost_details": {"upstream_inference_cost": 0.2}}
) == pytest.approx(0.21)
# Upstream cost is only added for BYOK responses.
assert openrouter_stream_cost(
{"cost": 0.05, "is_byok": False, "cost_details": {"upstream_inference_cost": 0.04}}
) == pytest.approx(0.05)
assert openrouter_stream_cost({"prompt_tokens": 10}) is None
assert openrouter_stream_cost(None) is None
def test_cost_callback_recovers_streamed_openrouter_cost_by_response_id() -> None:
report_state = MagicMock()
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={})
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", side_effect=ValueError("unknown model")),
):
litellm_cost_callback({"response_cost": None, "model": "moonshotai/kimi-k3"}, response)
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 streamed_openrouter_costs.take(response) is None
def test_streamed_openrouter_cost_prefers_provider_report_over_estimate() -> None:
report_state = MagicMock()
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),
_hidden_params={},
)
with (
patch("strix.report.state.get_global_report_state", return_value=report_state),
patch("litellm.completion_cost", return_value=0.1) as estimate,
):
litellm_cost_callback({"response_cost": None, "model": "moonshotai/kimi-k3"}, response)
report_state.record_observed_llm_cost.assert_called_once_with(0.9)
estimate.assert_not_called()
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
def test_streamed_openrouter_costs_cleared_on_new_run() -> None:
streamed_openrouter_costs.remember("gen-stale", {"cost": 0.7})
try:
set_global_report_state(ReportState.__new__(ReportState))
assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stale")) is None
finally:
set_global_report_state(None)
def test_openrouter_stream_handler_records_cost() -> None:
_install_openrouter_stream_cost_capture()
# Resolve the config the way LiteLLM does in production so we prove the
# override is actually reachable through provider resolution, not just as a
# directly-constructed class.
config = ProviderConfigManager.get_provider_chat_config(
model="moonshotai/kimi-k3", provider=LlmProviders.OPENROUTER
)
assert config is not None
assert type(config).__name__ == "_StrixOpenrouterConfig"
handler = config.get_model_response_iterator(streaming_response=iter([]), sync_stream=True)
chunk = {
"id": "gen-stream",
"created": 1,
"model": "moonshotai/kimi-k3",
"choices": [{"index": 0, "delta": {"content": None}}],
"usage": {"prompt_tokens": 89, "completion_tokens": 138, "cost": 0.0035055},
}
handler.chunk_parser(chunk)
assert streamed_openrouter_costs.take(SimpleNamespace(id="gen-stream")) == pytest.approx(
0.0035055
)
def test_openrouter_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,
}
)
# Non-streamed replies (LLM_DISABLE_STREAMING) carry the same fields.
reply = {
"choices": [{"message": {"role": "assistant"}}],
"provider": "Together",
"usage": usage,
}
config.transform_response(
"z-ai/glm-5.3",
httpx.Response(200, json=reply),
litellm.ModelResponse(),
MagicMock(),
{},
[],
{},
{},
None,
)
tally = call("Together", agent_id=None, input_tokens=1000, cached_tokens=900, cost=0.002)
assert report_state.record_llm_provider.call_args_list == [tally, tally]
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=128,
)
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), # 40 short of the previous 1000: within a block
("DeepInfra", "a1", 1500, 200), # 1000 of the previous 1200 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=128,
)
providers = ledger.to_record()["providers"]
assert providers["DeepInfra"]["cache_misses"] == 1
assert providers["DeepInfra"]["missed_tokens"] == 1000
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:
assert "session_id" not in body()
with patch("strix.config.models.load_settings") as settings:
settings.return_value.llm.openrouter_sticky_sessions = True
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)