test: add spend tracking tests

This commit is contained in:
mubashir1osmani 2026-06-18 20:07:35 -07:00
parent 8786e301bf
commit f616056214
No known key found for this signature in database
GPG key ID: AB055FF67D0B4D9A
5 changed files with 659 additions and 0 deletions

View file

@ -0,0 +1,77 @@
# Spend Tracking Test Coverage Matrix
Scope: every distinct spend-tracking code path, mapped to the test that exercises
it and the level it runs at. Highlights where a live e2e check is the only thing
that would catch a regression.
Companion: live suite `test_spend_tracking_e2e.py` + route breadth
`test_spend_routes.py` (this directory). Offline regression suite:
`tests/test_litellm/proxy/spend_tracking/`. Reference PR: BerriAI/litellm#29956.
Levels: `unit` mocked; `integration` real DB/cost-map; `live` real provider +
proxy + SpendLogs rows. Status: `covered` / `partial` / `gap`.
---
## SpendLogs row construction (`spend_tracking_utils.get_logging_payload`)
| Path | Existing | Level | Status | Live e2e |
|------|----------|-------|--------|----------|
| `_get_status_for_spend_log` | `test_spend_tracking_utils.py` | unit | covered | yes (status read off the row) |
| cache-hit `request_id` suffix | `test_spend_tracking_utils.py` | unit | covered | yes (`test_cache_hit_is_zero_cost_and_suffixed`) |
| failure status + zero spend | `test_spend_tracking_utils.py` | unit | partial | yes (`test_failure_call_writes_failure_status_row`) |
| field population (model/tokens/api_key/team/org) | `test_spend_tracking_utils.py` | unit | partial | yes (asserts real values) |
| `request_tags` propagation | `test_db_spend_update_writer.py` | unit | partial | yes (`test_request_tags_round_trip`) |
| `end_user` attribution | unit | unit | partial | yes (`test_end_user_spend_attributed_on_row`) |
## Cost calculation by modality
| Modality | Existing | Status | Live e2e |
|----------|----------|--------|----------|
| Chat (non-stream) | `test_cost_calculator.py`, `local_testing/test_completion_cost.py` | covered | yes (`test_chat_completion_writes_nonzero_spend_row`) |
| Chat (streaming) | `test_streaming_interrupt_spend_tracking.py` | partial | yes (`test_streaming_chat_completion_tracks_spend`) |
| Embedding | `test_cost_calculator.py` (#29956) | partial | yes (`test_embedding_writes_nonzero_spend_row`) |
| Pass-through (gemini/anthropic) | `pass_through_tests/*.test.js` + `llm_translation/` suite | covered | yes (llm_translation suite) |
| Image / audio / rerank / responses / realtime | per-provider unit cost tests | partial/gap | gap |
## Entity spend aggregation
| Entity | Existing | Status | Live e2e |
|--------|----------|--------|----------|
| API key | `test_db_spend_update_writer.py`, `test_spend_counters.py` | covered | yes (`test_key_spend_equals_sum_of_logs`) |
| Tag | `test_update_daily_tag_spend.py` | partial | yes (`test_tag_spend_matches_sum_of_tagged_logs`) |
| End-user | `test_proxy_update_spend.py` | covered | yes |
| Spend == sum(logs) consistency | none | gap | yes (key + tag aggregate == sum of rows) |
## Spend read endpoints (verification surface)
| Endpoint | Existing | Status | Live e2e |
|----------|----------|--------|----------|
| `/spend/logs` (request_id / api_key) | `test_spend_management_endpoints.py` | covered | yes (primary read path) |
| `/spend/calculate` | `local_testing/test_spend_calculate_endpoint.py` | covered | yes (`test_spend_calculate_returns_nonzero_cost`) |
| `/spend/tags` | `test_spend_management_endpoints.py` | partial | yes (tag accuracy test) |
| whole spend GET surface (22 routes) | unit per-handler | partial | yes (`test_spend_routes.py` probes each for 404/5xx) |
## What this suite pins
| Test | Invariant |
|------|-----------|
| `test_chat_completion_writes_nonzero_spend_row` | nonzero cost, token arithmetic, status, row findable by `response.id` |
| `test_streaming_chat_completion_tracks_spend` | streamed responses still costed |
| `test_embedding_writes_nonzero_spend_row` | embedding cost, `completion_tokens == 0` |
| `test_cache_hit_is_zero_cost_and_suffixed` | cache hits not double-charged; `_cache_hit` suffix |
| `test_key_spend_equals_sum_of_logs` | key aggregate == sum of rows |
| `test_request_tags_round_trip` | tags persist onto the row |
| `test_tag_spend_matches_sum_of_tagged_logs` | `/spend/tags` SUM/COUNT == tagged rows |
| `test_end_user_spend_attributed_on_row` | `end_user` attributed + costed |
| `test_failure_call_writes_failure_status_row` | failed call -> `status=failure`, `spend=0` |
| `test_spend_calculate_returns_nonzero_cost` | cost-map smoke (no batch wait) |
| `test_spend_routes.py` (23) | no spend route 404s or 5xxs |
## Design + timing
`proxy_batch_write_at` (~60s) means rows land late; every read polls to a deadline.
Fresh scoped key per test (isolation, xdist-safe, cleaned up). Assert invariants
(`spend > 0`, `total == prompt + completion`, aggregate == sum), not literal
$/token values, so pricing drift is not a failure. Skip on environment (no proxy /
no provider key), fail on behavior (a real 2xx call with a wrong/missing row).

View file

@ -0,0 +1,16 @@
"""Spend-tracking suite's `client` fixture.
The shared lifecycle (resources/scoped_key), proxy liveness skip, and e2e marker
live in the parent tests/e2e/conftest.py. SpendClient exposes the shared Gateway
(GatewayProvider), so the `resources` fixture cleans up keys and customers this
suite creates.
"""
import pytest
from spend_e2e_client import SpendClient, build_client
@pytest.fixture(scope="session")
def client() -> SpendClient:
return build_client()

View file

@ -0,0 +1,190 @@
"""Spend-tracking e2e client: a Gateway plus the spend-specific read endpoints.
Generic proxy operations (keys, customers, chat/embed, route probing, SpendLogs
polling) come from the shared Gateway, DI'd in (composition, not inheritance).
This client adds only the spend surface: /spend/calculate, /spend/tags,
key-spend polling, and the route probes the breadth test uses.
Re-exports unwrap / is_ok / unique_marker / SpendLogRow so the tests import their
helpers from one place.
"""
from __future__ import annotations
import os
import time
from collections.abc import Callable
from dataclasses import dataclass
from e2e_config import unique_marker
from e2e_http import NoBody, ProbeResult, Result, StreamingResponse, Success, is_ok, unwrap
from e2e_gateway import Gateway, build_gateway
from models import (
ChatBody,
ChatMessage,
ChatMetadata,
ChatResponse,
DateRangeParams,
EmbedBody,
EmbedResponse,
OpenAPISchema,
SpendCalculateBody,
SpendCalculateResponse,
SpendLogRow,
SpendTagsResponse,
TagSpend,
)
__all__ = [
"SpendClient",
"build_client",
"reset_spend_logs",
"unique_marker",
"unwrap",
"is_ok",
"SpendLogRow",
"ProbeResult",
]
def reset_spend_logs() -> None:
"""Truncate LiteLLM_SpendLogs for a clean slate. No proxy endpoint deletes
spend logs (/global/spend/reset keeps them), so go to the DB directly. Uses
DATABASE_URL (default: the local docker postgres on its mapped host port; note
the in-container `@db` host isn't resolvable from the host, so default to
localhost).
"""
import psycopg
url = os.environ.get(
"DATABASE_URL",
"postgresql://llmproxy:dbpassword9090@localhost:5432/litellm",
)
with psycopg.connect(url) as conn:
_ = conn.execute('TRUNCATE TABLE "LiteLLM_SpendLogs"')
def _chat_body(
model: str,
content: str,
*,
max_tokens: int | None = None,
tags: list[str] | None = None,
user: str | None = None,
stream: bool = False,
) -> ChatBody:
return ChatBody(
model=model,
messages=[ChatMessage(role="user", content=content)],
max_tokens=max_tokens,
stream=stream,
user=user,
metadata=ChatMetadata(tags=tags) if tags else None,
)
@dataclass(frozen=True, slots=True)
class SpendClient:
gateway: Gateway
def chat(
self,
key: str,
model: str,
content: str,
*,
max_tokens: int | None = None,
tags: list[str] | None = None,
user: str | None = None,
) -> Result[ChatResponse]:
return self.gateway.chat(
key, _chat_body(model, content, max_tokens=max_tokens, tags=tags, user=user)
)
def chat_stream(
self, key: str, model: str, content: str, *, max_tokens: int | None = None
) -> StreamingResponse:
return self.gateway.chat_stream(
key, _chat_body(model, content, max_tokens=max_tokens, stream=True)
)
def embed(self, key: str, model: str, content: str) -> Result[EmbedResponse]:
return self.gateway.embed(key, EmbedBody(model=model, input=content))
def poll_logs_for_key(
self,
key: str,
*,
min_rows: int = 1,
predicate: Callable[[list[SpendLogRow]], bool] | None = None,
) -> list[SpendLogRow]:
return self.gateway.poll_logs_for_key(
key, min_rows=min_rows, predicate=predicate
)
def calculate_spend(self, model: str, content: str) -> float:
return unwrap(
self.gateway.transport.post(
"/spend/calculate",
headers=self.gateway.transport.master,
json=SpendCalculateBody(
model=model, messages=[ChatMessage(role="user", content=content)]
),
response_type=SpendCalculateResponse,
)
).cost
def spend_by_tags(self) -> list[TagSpend]:
result = self.gateway.transport.get(
"/spend/tags",
headers=self.gateway.transport.master,
params=NoBody(),
response_type=SpendTagsResponse,
)
match result:
case Success(data=data):
return data.spend_per_tag or []
case _:
return []
def poll_tag_spend(self, tag: str, *, minimum: float = 0.0) -> TagSpend | None:
"""Poll /spend/tags until the tag's aggregate reaches `minimum`; last seen."""
deadline = time.monotonic() + self.gateway.poll_timeout
entry: TagSpend | None = None
while time.monotonic() < deadline:
matches = [
t for t in self.spend_by_tags() if t.individual_request_tag == tag
]
if matches:
entry = matches[0]
if (entry.total_spend or 0.0) >= minimum:
return entry
time.sleep(self.gateway.poll_interval)
return entry
def poll_key_spend(self, key: str, *, minimum: float = 0.0) -> float:
deadline = time.monotonic() + self.gateway.poll_timeout
spend = 0.0
while time.monotonic() < deadline:
spend = self.gateway.key_info(key).spend or 0.0
if spend > minimum:
return spend
time.sleep(self.gateway.poll_interval)
return spend
def probe(self, path: str, *, params: DateRangeParams) -> ProbeResult:
return self.gateway.transport.probe(path, params=params)
def openapi(self) -> OpenAPISchema:
return unwrap(
self.gateway.transport.get(
"/openapi.json",
headers=self.gateway.transport.master,
params=NoBody(),
response_type=OpenAPISchema,
)
)
def build_client() -> SpendClient:
return SpendClient(gateway=build_gateway())

View file

@ -0,0 +1,96 @@
"""Breadth check: query every route on the spend read surface and show what it
returns.
Spend tracking sprawls across many routes (model-cost / key / user / team / org /
customer aggregation, tags, and activity reports). Most are served with
`include_in_schema=False`, so they do NOT appear in `/openapi.json` - discovery
from the schema alone misses ~70% of the surface. So we probe a curated, verified
list directly, plus any spend route the schema does list (to auto-catch new ones).
Each probe captures status AND body, so a failure shows the proxy's actual error
(a 500 traceback, a 404 meaning the route was removed) rather than a bare code.
Run with `-rA` (or `-s`) to print every route's response, not just failures.
Healthy == route exists (not 404) and handler did not crash (not 5xx). A 4xx
(missing params / auth nuance) still means the route is wired and ran. Cheap and
fast: no batch-write wait, no provider calls.
"""
from datetime import datetime, timedelta, timezone
import pytest
from models import DateRangeParams
from spend_e2e_client import SpendClient
pytestmark = pytest.mark.e2e
# Verified present and responsive on a live proxy. One per row of the spend
# surface: key / user / team / org / customer aggregation, model-cost, tags,
# activity.
SPEND_ROUTES = (
"/spend/keys",
"/spend/users",
"/spend/tags",
"/spend/logs",
"/spend/logs/ui",
"/global/spend",
"/global/spend/keys",
"/global/spend/teams",
"/global/spend/models",
"/global/spend/provider",
"/global/spend/report",
"/global/spend/tags",
"/global/spend/logs",
"/global/spend/all_tag_names",
"/global/activity",
"/global/activity/model",
"/global/activity/exceptions",
"/key/list",
"/user/list",
"/team/list",
"/organization/list",
"/customer/list",
)
_SPEND_PREFIXES = ("/spend", "/global/spend", "/global/activity")
def _date_range() -> DateRangeParams:
# Satisfies date-required endpoints (report/activity/provider); ignored elsewhere.
end = datetime.now(timezone.utc).date()
start = end - timedelta(days=1)
return DateRangeParams(start_date=start.isoformat(), end_date=end.isoformat())
@pytest.mark.parametrize("route", SPEND_ROUTES)
def test_spend_route_responsive(client: SpendClient, route: str) -> None:
result = client.probe(route, params=_date_range())
print(f"{route} -> {result.status_code}\n{result.body[:600]}")
assert result.healthy, f"{route} -> {result.status_code}\n{result.body[:600]}"
def test_schema_listed_spend_routes_are_responsive(client: SpendClient) -> None:
"""Probe any spend GET route the schema lists that isn't in SPEND_ROUTES."""
schema = client.openapi()
assert schema.paths, "/openapi.json had no paths"
discovered = [
path
for path, spec in schema.paths.items()
if "get" in spec.methods
and "{" not in path
and any(path.startswith(prefix) for prefix in _SPEND_PREFIXES)
]
extras = [path for path in discovered if path not in SPEND_ROUTES]
params = _date_range()
results = [(path, client.probe(path, params=params)) for path in extras]
for path, result in results:
print(f"{path} -> {result.status_code}")
offenders = [
f"{path} -> {result.status_code}\n{result.body[:600]}"
for path, result in results
if not result.healthy
]
assert not offenders, "non-responsive schema spend routes:\n" + "\n".join(offenders)

View file

@ -0,0 +1,280 @@
"""Live end-to-end spend-tracking tests against a running proxy.
Run against a proxy started with the gateway config. Coverage rationale:
SPEND_TRACKING_COVERAGE_MATRIX.md.
Model names are literals from that config: chat tests hit "gemini-2.5-flash",
embedding tests hit "openai-text-embedding-3-small".
Every test: fresh scoped key (isolation) -> real provider call -> unwrap (hard
fail if the proxy couldn't make a call it should) -> poll /spend/logs to a
deadline (rows land ~60s later via proxy_batch_write_at) -> assert invariants on
the real row (spend, token arithmetic, status, cache).
Assertions target invariants, not literals: a regression in the spend pipeline
fails the test; a pricing or token-count drift does not.
"""
from collections.abc import Callable
import pytest
from lifecycle import ResourceManager
from spend_e2e_client import SpendClient, SpendLogRow, is_ok, unique_marker, unwrap
pytestmark = pytest.mark.e2e
def _approx_equal(actual: float, expected: float) -> bool:
"""Within 1% or 1e-9 absolute - spend math, not exact float identity."""
return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2)
def _summarize(rows: list[SpendLogRow]) -> list[dict[str, object]]:
fields = {
"request_id",
"model",
"spend",
"status",
"cache_hit",
"prompt_tokens",
"completion_tokens",
"total_tokens",
}
return [row.model_dump(include=fields) for row in rows]
def _require_row(
rows: list[SpendLogRow], predicate: Callable[[SpendLogRow], bool], what: str
) -> SpendLogRow:
matches = [r for r in rows if predicate(r)]
assert matches, (
f"no SpendLogs row {what} after polling; saw {len(rows)} row(s): "
f"{_summarize(rows)}"
)
return matches[0]
def test_chat_completion_writes_nonzero_spend_row(
client: SpendClient, scoped_key: str
) -> None:
chat = unwrap(
client.chat(
scoped_key,
"gemini-2.5-flash",
f"reply with one word {unique_marker()}",
max_tokens=16,
)
)
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(r.status == "success" for r in rs)
)
row = _require_row(rows, lambda r: r.status == "success", "for the chat call")
assert (row.spend or 0) > 0, f"chat row should cost > 0: {_summarize(rows)}"
assert row.status == "success"
assert row.cache_hit != "True", "fresh call must not be a cache hit"
assert "gemini-2.5-flash" in (row.model or "")
prompt = row.prompt_tokens or 0
completion = row.completion_tokens or 0
total = row.total_tokens or 0
assert prompt > 0 and completion > 0
assert total == prompt + completion, f"token arithmetic broken: {_summarize(rows)}"
if chat.id:
assert any(r.request_id == chat.id for r in rows), (
f"row request_id != client response.id ({chat.id})"
)
def test_streaming_chat_completion_tracks_spend(
client: SpendClient, scoped_key: str
) -> None:
result = client.chat_stream(
scoped_key, "gemini-2.5-flash", f"count to three {unique_marker()}", max_tokens=64
)
assert result.ok, f"stream failed (status {result.status_code}): {result.body[:300]}"
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any((r.spend or 0) > 0 for r in rs)
)
row = _require_row(
rows, lambda r: (r.spend or 0) > 0, "with nonzero spend for the stream"
)
prompt = row.prompt_tokens or 0
completion = row.completion_tokens or 0
assert prompt > 0 and completion > 0, f"streaming tokens not tracked: {_summarize(rows)}"
assert (row.total_tokens or 0) == prompt + completion
def test_embedding_writes_nonzero_spend_row(
client: SpendClient, scoped_key: str
) -> None:
_ = unwrap(
client.embed(
scoped_key,
"openai-text-embedding-3-small",
f"vectorize this sentence {unique_marker()}",
)
)
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any((r.spend or 0) > 0 for r in rs)
)
row = _require_row(
rows, lambda r: (r.spend or 0) > 0, "with nonzero spend for the embedding"
)
assert (row.prompt_tokens or 0) > 0
assert (row.completion_tokens or 0) == 0, "embeddings have no completion tokens"
assert "text-embedding-3-small" in (row.model or "")
def test_cache_hit_is_zero_cost_and_suffixed(
client: SpendClient, scoped_key: str
) -> None:
# Unique marker shared by both calls: call 1 is a guaranteed cache MISS (fresh
# content, paid), call 2 repeats the identical request and HITS the cache just
# populated. The marker keeps each run isolated - a fixed prompt would persist
# in the shared response cache across runs and make both calls hit (flaky).
prompt = f"What is the capital of France? Answer in one word. {unique_marker()}"
_ = unwrap(client.chat(scoped_key, "gemini-2.5-flash", prompt, max_tokens=16))
_ = unwrap(client.chat(scoped_key, "gemini-2.5-flash", prompt, max_tokens=16))
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(r.cache_hit == "True" for r in rs)
)
cache_rows = [r for r in rows if r.cache_hit == "True"]
if not cache_rows:
pytest.skip(
"no cache-hit row observed; caching may be disabled on this proxy. "
f"rows seen: {_summarize(rows)}"
)
cache_row = cache_rows[0]
assert (cache_row.spend or 0) == 0.0, (
f"cache hit was charged (double-charge regression): {_summarize(rows)}"
)
assert "_cache_hit" in (cache_row.request_id or ""), (
"cache-hit row missing the _cache_hit request_id suffix; "
"duplicate-key collisions will silently drop rows"
)
paid_rows = [r for r in rows if r.cache_hit != "True"]
assert any((r.spend or 0) > 0 for r in paid_rows), (
f"the non-cached call should still be charged: {_summarize(rows)}"
)
def test_key_spend_equals_sum_of_logs(client: SpendClient, scoped_key: str) -> None:
for _ in range(2):
_ = unwrap(
client.chat(
scoped_key, "gemini-2.5-flash", f"say hi {unique_marker()}", max_tokens=16
)
)
rows = client.poll_logs_for_key(
scoped_key,
min_rows=2,
predicate=lambda rs: sum((r.spend or 0) for r in rs) > 0,
)
assert len(rows) >= 2, f"expected >=2 rows for the key, saw {_summarize(rows)}"
logs_total = sum((r.spend or 0) for r in rows)
assert logs_total > 0
key_spend = client.poll_key_spend(scoped_key, minimum=logs_total * 0.999)
assert _approx_equal(key_spend, logs_total), (
f"key aggregate {key_spend} != sum of logs {logs_total}; rows: {_summarize(rows)}"
)
def test_request_tags_round_trip(client: SpendClient, scoped_key: str) -> None:
tag = f"e2e-spend-{unique_marker()}"
_ = unwrap(
client.chat(scoped_key, "gemini-2.5-flash", "tagged request", tags=[tag], max_tokens=16)
)
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(tag in (r.request_tags or []) for r in rs)
)
_require_row(
rows, lambda r: tag in (r.request_tags or []), f"carrying request tag {tag!r}"
)
def test_tag_spend_matches_sum_of_tagged_logs(
client: SpendClient, scoped_key: str
) -> None:
# Unique tag so /spend/tags can't be polluted by other rows; unique content
# per call so both are fresh misses (paid), not cache hits.
tag = f"e2e-tagspend-{unique_marker()}"
for _ in range(2):
_ = unwrap(
client.chat(
scoped_key, "gemini-2.5-flash", f"hi {unique_marker()}", tags=[tag], max_tokens=16
)
)
rows = client.poll_logs_for_key(
scoped_key,
min_rows=2,
predicate=lambda rs: sum((r.spend or 0) for r in rs) > 0,
)
tagged = [r for r in rows if tag in (r.request_tags or [])]
assert len(tagged) >= 2, f"expected 2 tagged rows, saw {_summarize(rows)}"
logs_total = sum((r.spend or 0) for r in tagged)
assert logs_total > 0
entry = client.poll_tag_spend(tag, minimum=logs_total * 0.999)
assert entry is not None, f"tag {tag!r} never appeared in /spend/tags"
assert _approx_equal(entry.total_spend or 0, logs_total), (
f"/spend/tags total_spend {entry} != sum of tagged rows {logs_total}"
)
assert (entry.log_count or 0) == len(tagged), (
f"/spend/tags log_count {entry.log_count} != tagged rows {len(tagged)}"
)
def test_end_user_spend_attributed_on_row(
client: SpendClient, scoped_key: str, resources: ResourceManager
) -> None:
customer = resources.customer(f"e2e-cust-{unique_marker()}")
_ = unwrap(
client.chat(scoped_key, "gemini-2.5-flash", "hi", user=customer, max_tokens=16)
)
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(r.end_user == customer for r in rs)
)
row = _require_row(
rows, lambda r: r.end_user == customer, f"attributed to end_user {customer!r}"
)
assert (row.spend or 0) > 0, f"end-user row should cost > 0: {_summarize(rows)}"
def test_failure_call_writes_failure_status_row(
client: SpendClient, scoped_key: str
) -> None:
result = client.chat(scoped_key, "gemini-2.5-flash", "", max_tokens=1)
if is_ok(result):
pytest.skip("call unexpectedly succeeded; could not induce a failure row")
rows = client.poll_logs_for_key(
scoped_key, predicate=lambda rs: any(r.status == "failure" for r in rs)
)
failure_rows = [r for r in rows if r.status == "failure"]
if not failure_rows:
pytest.skip(
"no failure-status row was logged for the rejected call; "
"failure logging is environment-specific"
)
assert (failure_rows[0].spend or 0) == 0.0, "failed call must not be charged"
def test_spend_calculate_returns_nonzero_cost(client: SpendClient) -> None:
cost = client.calculate_spend("gemini-2.5-flash", "estimate the cost of this request")
assert cost > 0, (
"/spend/calculate returned 0 for gemini-2.5-flash; "
"cost map may be missing this model"
)