From 01903cb3ae39092badfb89a6772ca9b8b42ee94a Mon Sep 17 00:00:00 2001 From: gabriele Date: Mon, 28 Sep 2026 11:28:00 +0000 Subject: [PATCH 1/5] fix(spend-tracking): stop caching failed spend-log metadata lookups as confirmed misses Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../spend_tracking/key_metadata_recovery.py | 13 ++++++-- .../test_key_metadata_recovery.py | 32 +++++++++++++++++++ 2 files changed, 42 insertions(+), 3 deletions(-) diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index ce96dc62780..120ff0a680f 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -370,10 +370,17 @@ async def _spend_log_metadata_one_query_at_a_time( fresh: Final = ( await _query_spend_log_metadata(prisma_client, pending, window) if pending else _EMPTY_KEY_METADATA ) - found: Final = fresh if fresh is not None else _EMPTY_KEY_METADATA + if fresh is None: + for digest in pending: + cache.set_cache( + _spend_log_cache_key(digest, window), + KeyMetadataDict(), + ttl=SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL, + ) + return settled for digest in pending: - _remember_spend_log_metadata(cache, digest, window, found.get(digest)) - return MappingProxyType({**settled, **found}) + _remember_spend_log_metadata(cache, digest, window, fresh.get(digest)) + return MappingProxyType({**settled, **fresh}) async def recover_key_metadata_from_spend_logs( diff --git a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py index acd03964bf3..d96156a9d10 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py +++ b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py @@ -3,6 +3,7 @@ import time from collections.abc import Sequence from datetime import datetime, timedelta from types import SimpleNamespace +from typing import Final from unittest.mock import AsyncMock, MagicMock import pytest @@ -439,6 +440,37 @@ async def test_recover_key_metadata_from_spend_logs_retries_a_failed_query_only_ assert result[digest]["key_alias"] == "back-online" +@pytest.mark.asyncio +async def test_recover_key_metadata_from_spend_logs_never_caches_repeated_query_failures_as_long_as_a_hit(): + digest: Final = hash_token("cli-session-repeated-failure") + window: Final = (datetime(2026, 9, 7), datetime(2026, 9, 10)) + cache: Final = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL) + mock_prisma: Final = MagicMock() + _spend_log_transaction( + mock_prisma, + AsyncMock( + side_effect=[ + PrismaError("statement timeout"), + PrismaError("statement timeout"), + [_spend_log_row(digest, "back-online", None, None)], + ] + ), + ) + await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) + miss_key: Final = next(key for key in cache.ttl_dict if digest in key and not key.endswith(":missed-before")) + cache.ttl_dict[miss_key] = time.time() - 1 + second_query_started: Final = time.time() + + await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) + + assert cache.ttl_dict[miss_key] - second_query_started <= SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL + 1 + cache.ttl_dict[miss_key] = time.time() - 1 + + result: Final = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) + + assert result[digest]["key_alias"] == "back-online" + + @pytest.mark.asyncio async def test_recover_key_metadata_from_spend_logs_drops_the_owner_of_a_digest_shared_by_several_users(): shared_ui_digest = hash_token("ui-token") From 7c288cd7ed29e1dee335e475227e591cc983fa8b Mon Sep 17 00:00:00 2001 From: gabriele Date: Mon, 28 Sep 2026 11:29:26 +0000 Subject: [PATCH 2/5] refactor(spend-tracking): share the short-lived miss cache write Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../proxy/spend_tracking/key_metadata_recovery.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index 120ff0a680f..6299f576c37 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -342,6 +342,10 @@ async def _query_spend_log_metadata( ) +def _remember_short_lived_miss(cache: InMemoryCache, key: str) -> None: + cache.set_cache(key, KeyMetadataDict(), ttl=SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL) + + def _remember_spend_log_metadata( cache: InMemoryCache, digest: str, window: tuple[datetime, datetime], meta: KeyMetadataDict | None ) -> None: @@ -353,7 +357,7 @@ def _remember_spend_log_metadata( if cache.get_cache(missed_before) is not None: cache.set_cache(key, KeyMetadataDict()) return - cache.set_cache(key, KeyMetadataDict(), ttl=SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL) + _remember_short_lived_miss(cache, key) cache.set_cache(missed_before, True) @@ -372,11 +376,7 @@ async def _spend_log_metadata_one_query_at_a_time( ) if fresh is None: for digest in pending: - cache.set_cache( - _spend_log_cache_key(digest, window), - KeyMetadataDict(), - ttl=SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL, - ) + _remember_short_lived_miss(cache, _spend_log_cache_key(digest, window)) return settled for digest in pending: _remember_spend_log_metadata(cache, digest, window, fresh.get(digest)) From 6eb7d9e8bc9007a0feec4c846aa2d20d09d21e74 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 21:44:06 -0700 Subject: [PATCH 3/5] test(spend): cover key alias recovery after spend log lookup failures across usage routes --- ...ily_activity_key_metadata_query_timeout.py | 974 ++++++++++++++++++ 1 file changed, 974 insertions(+) create mode 100644 tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py diff --git a/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py b/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py new file mode 100644 index 00000000000..b06114ada87 --- /dev/null +++ b/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py @@ -0,0 +1,974 @@ +import csv +import io +import itertools +import json +import threading +import uuid +from collections.abc import Generator, Iterator, Mapping +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta +from hashlib import sha256 +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final + +import httpcore +import httpx +import jwt +import psycopg +import pytest +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway, eventually, object_value, string_value +from integration._support.database import read_rows, scratch_database, write_rows +from integration._support.database_relay import database_relay +from integration._support.process import owned_proxy_process +from integration._support.wire import Reply, Request, wire_server +from jwt.algorithms import RSAAlgorithm +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter + +MODEL: Final = "key-metadata-recovery-audit" +LOOKUP_MARKER: Final = "first_alias" +FAILED_LOOKUP_FLOOR: Final = timedelta(seconds=4) +MISS_TTL_BOUND: Final = 60 +USAGE: Final = MappingProxyType({"prompt_tokens": 10, "completion_tokens": 30, "total_tokens": 40}) +REPLY_TEXT: Final = "You spent a little this week." +WAITING_LOOKUPS: Final = ( + "SELECT pid, query_start::text AS started FROM pg_stat_activity " + "WHERE datname = current_database() AND pid <> pg_backend_pid() AND state = 'active' " + "AND wait_event_type = 'Lock' AND position(%s in query) > 0" +) +JSON_VALUE: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) +WAITING_ROWS: Final[TypeAdapter[tuple[tuple[int, str], ...]]] = TypeAdapter(tuple[tuple[int, str], ...]) +SOCKET_ADDRESS: Final[TypeAdapter[tuple[str, int]]] = TypeAdapter(tuple[str, int]) +NO_SETTINGS: Final[Mapping[str, JsonValue]] = MappingProxyType({}) +NO_ENVIRONMENT: Final[Mapping[str, str]] = MappingProxyType({}) +JWT_SETTINGS: Final[Mapping[str, JsonValue]] = MappingProxyType( + {"enable_jwt_auth": True, "litellm_jwtauth": {"user_id_jwt_field": "sub", "user_id_upsert": True}} +) +JWKS_KEY_ID: Final = "key-metadata-recovery-audit" +SPEND_TABLES: Final = ( + "LiteLLM_SpendLogs", + "LiteLLM_DailyUserSpend", + "LiteLLM_DailyTeamSpend", + "LiteLLM_DailyOrganizationSpend", + "LiteLLM_DailyEndUserSpend", + "LiteLLM_DailyTagSpend", + "LiteLLM_DailyAgentSpend", +) +LANDED_KEYS: Final = " UNION ".join( + f"SELECT DISTINCT '{table}' AS source, api_key FROM \"{table}\"" for table in SPEND_TABLES +) +TENANT_KEYS: Final = ("chat", "stream", "messages", "responses", "live", "deleted") +KEYS_ONLY_IN_SPEND_LOGS: Final = frozenset(("chat", "stream", "messages", "responses")) +AGGREGATED: Final = "/user/daily/activity/aggregated" +EXPORT: Final = "/team/daily/activity/export" +BURST: Final = 16 +EXPORT_ROUTES: Final = ("team export json", "team export csv") + + +class _KeyMetadata(BaseModel): + model_config = ConfigDict(frozen=True) + key_alias: str | None = None + user_id: str | None = None + user_email: str | None = None + + +class _KeyBreakdown(BaseModel): + model_config = ConfigDict(frozen=True) + metadata: _KeyMetadata + + +class _Breakdown(BaseModel): + model_config = ConfigDict(frozen=True) + api_keys: Mapping[str, _KeyBreakdown] + + +class _Day(BaseModel): + model_config = ConfigDict(frozen=True) + breakdown: _Breakdown + + +class _Activity(BaseModel): + model_config = ConfigDict(frozen=True) + results: tuple[_Day, ...] + + +class _UpstreamMessage(BaseModel): + model_config = ConfigDict(frozen=True) + role: str | None = None + + +class _UpstreamRequest(BaseModel): + model_config = ConfigDict(frozen=True) + stream: bool | None = None + tools: tuple[JsonValue, ...] | None = None + messages: tuple[_UpstreamMessage, ...] = () + + +@dataclass(frozen=True, slots=True) +class Spender: + alias: str + user_id: str + user_email: str + digest: str + + +@dataclass(frozen=True, slots=True) +class Pinned: + client: httpx.Client + port: int + + def request( + self, + method: str, + path: str, + *, + params: Mapping[str, str] | None = None, + body: Mapping[str, JsonValue] | None = None, + ) -> httpx.Response: + response: Final = self.client.request(method, path, params=params, json=body) + assert _local_port(response) == self.port, "Pinned connection moved to another worker socket" + return response + + +@dataclass(frozen=True, slots=True) +class Probe: + ran_lookup: bool + aliases: tuple[str | None, ...] + + +def _day(offset: int) -> str: + return (datetime.now(UTC) + timedelta(days=offset)).date().isoformat() + + +def _completion(message: Mapping[str, JsonValue], finish_reason: str) -> bytes: + return json.dumps( + { + "id": f"chatcmpl-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": dict(message), "finish_reason": finish_reason}], + "usage": dict(USAGE), + } + ).encode() + + +def _stream_frames() -> tuple[bytes, ...]: + identity: Final = f"chatcmpl-{uuid.uuid4().hex}" + frames: Final[tuple[Mapping[str, JsonValue], ...]] = ( + {"choices": [{"index": 0, "delta": {"role": "assistant", "content": REPLY_TEXT}, "finish_reason": None}]}, + {"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + {"choices": [], "usage": dict(USAGE)}, + ) + envelope: Final = {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"} + return (*(f"data: {json.dumps({**envelope, **frame})}\n\n".encode() for frame in frames), b"data: [DONE]\n\n") + + +def _usage_tool_call() -> Mapping[str, JsonValue]: + arguments: Final = json.dumps({"start_date": _day(-1), "end_date": _day(1)}) + return { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": "call_usage", "type": "function", "function": {"name": "get_usage_data", "arguments": arguments}} + ], + } + + +def _response_object() -> bytes: + return json.dumps( + { + "id": f"resp_{uuid.uuid4().hex}", + "object": "response", + "status": "completed", + "created_at": 1, + "model": "gpt-4o-mini", + "output": [ + { + "type": "message", + "id": f"msg_{uuid.uuid4().hex}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": REPLY_TEXT, "annotations": []}], + } + ], + "usage": { + "input_tokens": USAGE["prompt_tokens"], + "output_tokens": USAGE["completion_tokens"], + "total_tokens": USAGE["total_tokens"], + "input_tokens_details": {"cached_tokens": 0}, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + } + ).encode() + + +def _respond(request: Request) -> Reply: + if request.target.endswith("/responses"): + return Reply(body=_response_object()) + body: Final = _UpstreamRequest.model_validate_json(request.body or b"{}") + if body.stream: + return Reply(content_type="text/event-stream", chunks=_stream_frames()) + roles: Final = frozenset(message.role for message in body.messages) + if body.tools and "tool" not in roles: + return Reply(body=_completion(_usage_tool_call(), "tool_calls")) + return Reply(body=_completion({"role": "assistant", "content": REPLY_TEXT}, "stop")) + + +def _config(wire_url: str, general_settings: Mapping[str, JsonValue]) -> str: + return json.dumps( + { + "model_list": [ + { + "model_name": MODEL, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_base": f"{wire_url}/v1", + "api_key": "sk-upstream", + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + }, + } + ], + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + "store_model_in_db": True, + "disable_spend_logs": False, + "proxy_batch_write_at": 1, + "proxy_batch_polling_interval": 1, + **general_settings, + }, + "router_settings": {"disable_cooldowns": True}, + } + ) + + +@contextmanager +def _proxy( + gateway: Gateway, + directory: Path, + database_url: str, + wire_url: str, + *, + general_settings: Mapping[str, JsonValue] = NO_SETTINGS, + environment: Mapping[str, str] = NO_ENVIRONMENT, +) -> Generator[Gateway]: + config: Final = directory / "key_metadata_recovery.yaml" + config.write_text(_config(wire_url, general_settings)) + with owned_proxy_process( + gateway, + directory, + { + "DATABASE_URL": database_url, + "KEEPALIVE_TIMEOUT": "600", + "PRISMA_HEALTH_WATCHDOG_ENABLED": "false", + "OPENAI_API_KEY": "sk-upstream", + "OPENAI_BASE_URL": f"{wire_url}/v1", + "LITELLM_DISABLE_NO_REDIS_WARNING": "true", + **environment, + }, + config=config, + remove_environment=("DATABASE_URL_READ_REPLICA",), + workers=2, + ) as owned: + yield owned.gateway + + +def _local_port(response: httpx.Response) -> int: + match response.extensions: + case {"network_stream": httpcore.NetworkStream() as stream}: + return SOCKET_ADDRESS.validate_python(stream.get_extra_info("client_addr"))[1] + case _: + raise AssertionError(f"No network stream on {response.request.url}") + + +@contextmanager +def _pinned(proxy: Gateway) -> Generator[Pinned]: + with httpx.Client( + base_url=proxy.client.base_url, + headers={"Authorization": f"Bearer {proxy.key}"}, + limits=httpx.Limits(max_connections=1, max_keepalive_connections=1, keepalive_expiry=600), + timeout=60, + trust_env=False, + ) as client: + opened: Final = client.get("/health/liveliness") + assert opened.status_code == 200, opened.text + yield Pinned(client, _local_port(opened)) + + +def _landed(database_url: str, digest: str) -> bool: + daily: Final = read_rows( + 'SELECT 1 FROM "LiteLLM_DailyUserSpend" WHERE api_key=%s', (digest,), database_url=database_url + ) + logged: Final = read_rows( + 'SELECT 1 FROM "LiteLLM_SpendLogs" WHERE api_key=%s', (digest,), database_url=database_url + ) + return bool(daily) and bool(logged) + + +def _spender(proxy: Gateway, database_url: str, label: str) -> Spender: + user_id: Final = f"{label}-{uuid.uuid4().hex[:8]}" + user_email: Final = f"{user_id}@example.com" + alias: Final = f"{label}-laptop-key" + proxy.post("/user/new", {"user_id": user_id, "user_email": user_email, "auto_create_key": False}) + key: Final = string_value( + proxy.post("/key/generate", {"user_id": user_id, "key_alias": alias, "models": [MODEL]})["key"] + ) + digest: Final = sha256(key.encode()).hexdigest() + usage: Final = object_value(proxy.chat(MODEL, key=key, text=f"spend {uuid.uuid4().hex}")["usage"]) + assert {name: usage.get(name) for name in USAGE} == dict(USAGE), usage + eventually(lambda: _landed(database_url, digest), bool, seconds=70) + write_rows('DELETE FROM "LiteLLM_VerificationToken" WHERE token=%s', (digest,), database_url=database_url) + return Spender(alias, user_id, user_email, digest) + + +def _read(pinned: Pinned, digest: str) -> httpx.Response: + response: Final = pinned.request( + "GET", + "/user/daily/activity/aggregated", + params={"start_date": _day(-1), "end_date": _day(1), "api_key": digest}, + ) + assert response.status_code == 200, response.text + return response + + +def _metadata(response: httpx.Response, digest: str) -> tuple[_KeyMetadata, ...]: + days: Final = _Activity.model_validate_json(response.content).results + return tuple(day.breakdown.api_keys[digest].metadata for day in days) + + +def _aliases(response: httpx.Response, digest: str) -> tuple[str | None, ...]: + return tuple(meta.key_alias for meta in _metadata(response, digest)) + + +def _named(pinned: Pinned, spender: Spender) -> httpx.Response: + return eventually( + lambda: _read(pinned, spender.digest), + lambda response: _aliases(response, spender.digest) == (spender.alias,), + seconds=MISS_TTL_BOUND, + ) + + +@contextmanager +def _locked_spend_logs(database_url: str) -> Generator[None]: + with psycopg.connect(database_url) as connection: + connection.execute('LOCK TABLE "LiteLLM_SpendLogs" IN ACCESS EXCLUSIVE MODE') + try: + yield + finally: + connection.rollback() + + +def _poll_waiting_lookups(database_url: str, stop: threading.Event, seen: SimpleQueue[tuple[int, str]]) -> None: + with psycopg.connect(database_url, autocommit=True) as connection: + while not stop.wait(0.05): + for lookup in WAITING_ROWS.validate_python( + connection.execute(WAITING_LOOKUPS, (LOOKUP_MARKER,)).fetchall() + ): + seen.put(lookup) + + +@contextmanager +def _recording(database_url: str) -> Generator[SimpleQueue[tuple[int, str]]]: + seen: Final[SimpleQueue[tuple[int, str]]] = SimpleQueue() + stop: Final = threading.Event() + poller: Final = threading.Thread(target=_poll_waiting_lookups, args=(database_url, stop, seen)) + poller.start() + try: + yield seen + finally: + stop.set() + poller.join(timeout=5) + assert not poller.is_alive(), "Lookup recorder survived its recording window" + + +def _lookups(seen: SimpleQueue[tuple[int, str]]) -> frozenset[tuple[int, str]]: + return frozenset(seen.get_nowait() for _ in range(seen.qsize())) + + +def _waiting_lookup_count(database_url: str) -> int: + return len(read_rows(WAITING_LOOKUPS, (LOOKUP_MARKER,), database_url=database_url)) + + +def _probe(pinned: Pinned, database_url: str, digest: str) -> Probe: + with ThreadPoolExecutor(max_workers=1) as pool: + with _locked_spend_logs(database_url): + pending: Final = pool.submit(_read, pinned, digest) + ran_lookup: Final = eventually( + lambda: (pending.done(), _waiting_lookup_count(database_url) > 0), + lambda state: state[0] or state[1], + seconds=15, + )[1] + return Probe(ran_lookup, _aliases(pending.result(), digest)) + + +def _park(database_url: str, digest: str) -> None: + write_rows( + """UPDATE "LiteLLM_SpendLogs" SET api_key = 'parked-' || api_key WHERE api_key = %s""", + (digest,), + database_url=database_url, + ) + + +def _restore(database_url: str, digest: str) -> None: + write_rows( + """UPDATE "LiteLLM_SpendLogs" SET api_key = substr(api_key, 8) WHERE api_key = 'parked-' || %s""", + (digest,), + database_url=database_url, + ) + + +def _events(response: httpx.Response) -> tuple[Mapping[str, JsonValue], ...]: + return tuple( + json.loads(line.removeprefix("data: ")) for line in response.text.splitlines() if line.startswith("data: ") + ) + + +@dataclass(frozen=True, slots=True) +class Tenant: + label: str + team_id: str + organization_id: str + + @property + def owner(self) -> str: + return f"{self.label}-owner" + + @property + def email(self) -> str: + return f"{self.owner}@example.com" + + @property + def customer(self) -> str: + return f"{self.label}-customer" + + @property + def tag(self) -> str: + return f"{self.label}-tag" + + @property + def agent(self) -> str: + return f"{self.label}-agent" + + @property + def headers(self) -> Mapping[str, str]: + return MappingProxyType( + {"x-litellm-end-user-id": self.customer, "x-litellm-tags": self.tag, "x-litellm-agent-id": self.agent} + ) + + +@dataclass(frozen=True, slots=True) +class TenantKey: + name: str + key: str + digest: str + alias: str + + +@dataclass(frozen=True, slots=True) +class KeyRow: + digest: str + key_alias: str | None + user_id: str | None + user_email: str | None + + +@dataclass(frozen=True, slots=True) +class UserRow: + user_id: str | None + keys: int + + +@dataclass(frozen=True, slots=True) +class Sweep: + keys: Mapping[str, frozenset[KeyRow]] + users: frozenset[UserRow] + + +class _KeyExportRow(BaseModel): + model_config = ConfigDict(frozen=True) + api_key: str + key_alias: str | None = None + user_id: str | None = None + user_email: str | None = None + + +class _KeyExport(BaseModel): + model_config = ConfigDict(frozen=True) + data: tuple[_KeyExportRow, ...] + + +class _UserExportRow(BaseModel): + model_config = ConfigDict(frozen=True) + user_id: str | None = None + keys: int + + +class _UserExport(BaseModel): + model_config = ConfigDict(frozen=True) + data: tuple[_UserExportRow, ...] + + +def _tenant(proxy: Gateway) -> Tenant: + label: Final = f"audit-{uuid.uuid4().hex[:6]}" + organization: Final = proxy.post("/organization/new", {"organization_alias": f"{label}-org", "models": [MODEL]}) + organization_id: Final = string_value(organization["organization_id"]) + owner: Final = f"{label}-owner" + proxy.post("/user/new", {"user_id": owner, "user_email": f"{owner}@example.com", "auto_create_key": False}) + team: Final = proxy.post( + "/team/new", + { + "team_alias": f"{label}-team", + "organization_id": organization_id, + "models": [MODEL], + "members_with_roles": [{"role": "user", "user_id": owner}], + }, + ) + return Tenant(label, string_value(team["team_id"]), organization_id) + + +def _tenant_key(proxy: Gateway, tenant: Tenant, name: str) -> TenantKey: + alias: Final = f"{tenant.label}-{name}" + generated: Final = proxy.post( + "/key/generate", {"user_id": tenant.owner, "team_id": tenant.team_id, "key_alias": alias, "models": [MODEL]} + ) + key: Final = string_value(generated["key"]) + return TenantKey(name, key, sha256(key.encode()).hexdigest(), alias) + + +def _spend_request(name: str) -> tuple[str, Mapping[str, JsonValue]]: + prompt: Final = f"spend {uuid.uuid4().hex}" + match name: + case "stream": + return "/v1/chat/completions", { + "model": MODEL, + "messages": [{"role": "user", "content": prompt}], + "stream": True, + "stream_options": {"include_usage": True}, + } + case "messages": + return "/v1/messages", {"model": MODEL, "max_tokens": 64, "messages": [{"role": "user", "content": prompt}]} + case "responses": + return "/v1/responses", {"model": MODEL, "input": prompt} + case _: + return "/v1/chat/completions", {"model": MODEL, "messages": [{"role": "user", "content": prompt}]} + + +def _chunk_text(frame: JsonValue) -> str: + match frame: + case {"choices": [{"delta": {"content": str() as text}}]}: + return text + case _: + return "" + + +def _spent_text(response: httpx.Response) -> str: + if response.headers["content-type"].startswith("text/event-stream"): + return "".join( + _chunk_text(JSON_VALUE.validate_json(line.removeprefix("data: "))) + for line in response.text.splitlines() + if line.startswith("data: ") and line != "data: [DONE]" + ) + match JSON_VALUE.validate_json(response.content): + case {"choices": [{"message": {"content": str() as text}}]}: + return text + case {"content": [{"text": str() as text}]}: + return text + case {"output": [{"content": [{"text": str() as text}]}]}: + return text + case _: + return response.text + + +def _spend(proxy: Gateway, tenant: Tenant, key: TenantKey) -> None: + path, body = _spend_request(key.name) + response: Final = proxy.request("POST", path, body, key=key.key, headers=tenant.headers) + assert response.status_code == 200, f"{key.name}: {response.status_code} {response.text}" + assert _spent_text(response) == REPLY_TEXT, f"{key.name}: {response.text}" + + +def _landed_rows(database_url: str) -> frozenset[tuple[str, str]]: + return frozenset( + (str(row["source"]), str(row["api_key"])) for row in read_rows(LANDED_KEYS, (), database_url=database_url) + ) + + +def _dash(value: str) -> str | None: + return None if value == "-" else value + + +def _named_key_row(name: str, child: JsonValue, digests: frozenset[str]) -> tuple[KeyRow, ...]: + match child: + case {"metadata": dict() as metadata} if name in digests: + meta: Final = _KeyMetadata.model_validate(metadata) + return (KeyRow(name, meta.key_alias, meta.user_id, meta.user_email),) + case _: + return () + + +def _key_rows(value: JsonValue, digests: frozenset[str]) -> Iterator[KeyRow]: + match value: + case list(): + for item in value: + yield from _key_rows(item, digests) + case dict(): + for name, child in value.items(): + yield from _named_key_row(name, child, digests) + yield from _key_rows(child, digests) + case _: + return + + +def _walked(response: httpx.Response, digests: frozenset[str]) -> frozenset[KeyRow]: + return frozenset(_key_rows(JSON_VALUE.validate_json(response.content), digests)) + + +def _csv_key_row(record: Mapping[str, str]) -> KeyRow: + return KeyRow(record["Key ID"], _dash(record["Key Alias"]), _dash(record["User ID"]), _dash(record["User Email"])) + + +def _csv_key_rows(text: str) -> frozenset[KeyRow]: + header, *records = tuple(csv.reader(io.StringIO(text))) + return frozenset(_csv_key_row(dict(zip(header, record, strict=True))) for record in records) + + +def _export_key_rows(response: httpx.Response) -> frozenset[KeyRow]: + return frozenset( + KeyRow(row.api_key, row.key_alias, row.user_id, row.user_email) + for row in _KeyExport.model_validate_json(response.content).data + ) + + +def _routes(tenant: Tenant) -> Mapping[str, tuple[str, Mapping[str, str]]]: + return MappingProxyType( + { + "user": ("/user/daily/activity", {}), + "user aggregated": (AGGREGATED, {}), + "user search": ("/user/daily/activity/aggregated/search", {"search": tenant.label}), + "team": ("/team/daily/activity", {"team_ids": tenant.team_id}), + "team aggregated": ("/team/daily/activity/aggregated", {"team_ids": tenant.team_id}), + "team search": ( + "/team/daily/activity/aggregated/search", + {"search": tenant.label, "team_ids": tenant.team_id}, + ), + "organization": ("/organization/daily/activity", {"organization_ids": tenant.organization_id}), + "customer": ("/customer/daily/activity", {"end_user_ids": tenant.customer}), + "end user": ("/end_user/daily/activity", {"end_user_ids": tenant.customer}), + "tag": ("/tag/daily/activity", {"tags": tenant.tag}), + "agent": ("/agent/daily/activity", {"agent_ids": tenant.agent}), + } + ) + + +def _get(pinned: Pinned, path: str, params: Mapping[str, str]) -> httpx.Response: + response: Final = pinned.request( + "GET", path, params={"start_date": _day(-1), "end_date": _day(1), "page_size": "100", **params} + ) + assert response.status_code == 200, f"GET {path}: {response.status_code} {response.text}" + return response + + +def _export(pinned: Pinned, tenant: Tenant, export_type: str, export_format: str) -> httpx.Response: + return _get(pinned, EXPORT, {"export_type": export_type, "format": export_format, "team_id": tenant.team_id}) + + +def _sweep(pinned: Pinned, tenant: Tenant, digests: frozenset[str]) -> Sweep: + routed: Final = MappingProxyType( + {name: _walked(_get(pinned, path, params), digests) for name, (path, params) in _routes(tenant).items()} + ) + return Sweep( + MappingProxyType( + { + **routed, + EXPORT_ROUTES[0]: _export_key_rows(_export(pinned, tenant, "daily_with_keys", "json")), + EXPORT_ROUTES[1]: _csv_key_rows(_export(pinned, tenant, "daily_with_keys", "csv").text), + } + ), + frozenset( + UserRow(row.user_id, row.keys) + for row in _UserExport.model_validate_json(_export(pinned, tenant, "daily_with_users", "json").content).data + ), + ) + + +def _named_row(tenant: Tenant, key: TenantKey) -> KeyRow: + return KeyRow(key.digest, key.alias, tenant.owner, tenant.email) + + +def _outage_row(tenant: Tenant, key: TenantKey) -> KeyRow: + if key.name in KEYS_ONLY_IN_SPEND_LOGS: + return KeyRow(key.digest, None, None, None) + return _named_row(tenant, key) + + +def _expected(tenant: Tenant, keys: tuple[TenantKey, ...], *, outage: bool) -> Sweep: + every: Final = frozenset(_outage_row(tenant, key) if outage else _named_row(tenant, key) for key in keys) + live: Final = frozenset(_named_row(tenant, key) for key in keys if key.name == "live") + searched: Final = frozenset(("user search", "team search")) + owned: Final = sum(1 for row in every if row.user_id is not None) + return Sweep( + MappingProxyType({name: live if name in searched else every for name in (*_routes(tenant), *EXPORT_ROUTES)}), + frozenset(row for row in (UserRow(tenant.owner, owned), UserRow(None, len(every) - owned)) if row.keys > 0), + ) + + +def _burst(proxy: Gateway, digests: frozenset[str]) -> tuple[frozenset[KeyRow], ...]: + params: Final = {"start_date": _day(-1), "end_date": _day(1)} + with ( + httpx.Client( + base_url=proxy.client.base_url, + headers={"Authorization": f"Bearer {proxy.key}"}, + timeout=60, + trust_env=False, + ) as client, + ThreadPoolExecutor(max_workers=BURST) as pool, + ): + + def read_aggregated(_: int) -> httpx.Response: + return client.get(AGGREGATED, params=params) + + responses: Final = tuple(pool.map(read_aggregated, range(BURST))) + assert all(response.status_code == 200 for response in responses), tuple(r.text for r in responses) + return tuple(_walked(response, digests) for response in responses) + + +def _jwks_reply(public_jwk: str) -> Reply: + return Reply(body=json.dumps({"keys": [{**json.loads(public_jwk), "kid": JWKS_KEY_ID}]}).encode()) + + +@contextmanager +def _rig(gateway: Gateway, directory: Path) -> Generator[tuple[Gateway, str]]: + with ( + scratch_database() as database_url, + wire_server(_respond) as wire, + _proxy(gateway, directory, database_url, wire.url) as proxy, + ): + yield proxy, database_url + + +@pytest.mark.timeout(360) +def test_usage_ai_chat_timeout_then_second_timeout_still_recovers_key_alias_after_database_frees( + gateway: Gateway, tmp_path: Path +) -> None: + with _rig(gateway, tmp_path) as (proxy, database_url): + spender: Final = _spender(proxy, database_url, "ai-chat") + with _pinned(proxy) as pinned: + with _locked_spend_logs(database_url): + with _recording(database_url) as during_chat: + chat: Final = pinned.request( + "POST", + "/usage/ai/chat", + body={ + "messages": [{"role": "user", "content": "What did we spend?"}], + "model": "openai/gpt-4o-mini", + }, + ) + assert chat.status_code == 200, chat.text + assert chat.elapsed >= FAILED_LOOKUP_FLOOR, chat.elapsed + tool_call: Final = { + "type": "tool_call", + "tool_name": "get_usage_data", + "tool_label": "global usage data", + "arguments": {"start_date": _day(-1), "end_date": _day(1)}, + } + assert _events(chat) == ( + {"type": "status", "message": "Thinking..."}, + {**tool_call, "status": "running"}, + {**tool_call, "status": "complete"}, + {"type": "status", "message": "Analyzing results..."}, + {"type": "chunk", "content": REPLY_TEXT}, + {"type": "done"}, + ), chat.text + assert len(_lookups(during_chat)) == 1 + cached_miss: Final = _read(pinned, spender.digest) + assert cached_miss.elapsed < FAILED_LOOKUP_FLOOR, cached_miss.elapsed + assert _aliases(cached_miss, spender.digest) == (None,), cached_miss.text + with _recording(database_url) as during_retry: + retried: Final = eventually( + lambda: _read(pinned, spender.digest), + lambda response: response.elapsed >= FAILED_LOOKUP_FLOOR, + seconds=MISS_TTL_BOUND, + ) + assert _aliases(retried, spender.digest) == (None,), retried.text + assert len(_lookups(during_retry)) == 1 + recovered: Final = _named(pinned, spender) + assert _metadata(recovered, spender.digest) == ( + _KeyMetadata(key_alias=spender.alias, user_id=spender.user_id, user_email=spender.user_email), + ), recovered.text + + +@pytest.mark.timeout(300) +def test_usage_page_timeout_then_genuine_miss_still_recovers_key_alias_once_spend_logs_return( + gateway: Gateway, tmp_path: Path +) -> None: + with _rig(gateway, tmp_path) as (proxy, database_url): + spender: Final = _spender(proxy, database_url, "miss-after-timeout") + with _pinned(proxy) as pinned: + with _locked_spend_logs(database_url), _recording(database_url) as during_timeout: + timed_out: Final = _read(pinned, spender.digest) + assert timed_out.elapsed >= FAILED_LOOKUP_FLOOR, timed_out.elapsed + assert _aliases(timed_out, spender.digest) == (None,), timed_out.text + assert len(_lookups(during_timeout)) == 1 + _park(database_url, spender.digest) + missed: Final = eventually( + lambda: _probe(pinned, database_url, spender.digest), + lambda probe: probe.ran_lookup, + seconds=MISS_TTL_BOUND, + ) + assert missed.aliases == (None,), missed + _restore(database_url, spender.digest) + recovered: Final = _named(pinned, spender) + assert _metadata(recovered, spender.digest) == ( + _KeyMetadata(key_alias=spender.alias, user_id=spender.user_id, user_email=spender.user_email), + ), recovered.text + + +@pytest.mark.timeout(360) +def test_usage_page_pins_a_key_blank_only_after_two_genuine_misses(gateway: Gateway, tmp_path: Path) -> None: + with _rig(gateway, tmp_path) as (proxy, database_url): + missed_twice: Final = _spender(proxy, database_url, "missed-twice") + missed_once: Final = _spender(proxy, database_url, "missed-once") + with _pinned(proxy) as pinned: + _park(database_url, missed_twice.digest) + _park(database_url, missed_once.digest) + first_misses: Final = ( + _probe(pinned, database_url, missed_twice.digest), + _probe(pinned, database_url, missed_once.digest), + ) + assert first_misses == (Probe(True, (None,)), Probe(True, (None,))), first_misses + _restore(database_url, missed_once.digest) + second_miss: Final = eventually( + lambda: _probe(pinned, database_url, missed_twice.digest), + lambda probe: probe.ran_lookup, + seconds=MISS_TTL_BOUND, + ) + assert second_miss.aliases == (None,), second_miss + _restore(database_url, missed_twice.digest) + _named(pinned, missed_once) + pinned_blank: Final = eventually( + lambda: _probe(pinned, database_url, missed_twice.digest), + lambda probe: probe.ran_lookup, + seconds=45, + return_last_on_timeout=True, + ) + assert pinned_blank == Probe(False, (None,)), pinned_blank + + +@pytest.mark.timeout(300) +def test_usage_page_survives_a_dropped_database_connection_during_alias_recovery( + gateway: Gateway, tmp_path: Path +) -> None: + with ( + scratch_database() as database_url, + database_relay(database_url, b"AS " + LOOKUP_MARKER.encode()) as (relay, relayed_url), + wire_server(_respond) as wire, + _proxy(gateway, tmp_path, relayed_url, wire.url) as proxy, + ): + spender: Final = _spender(proxy, database_url, "dropped-connection") + with _pinned(proxy) as pinned: + relay.arm() + dropped: Final = _read(pinned, spender.digest) + assert relay.tripped.is_set(), dropped.text + assert _aliases(dropped, spender.digest) == (None,), dropped.text + recovered: Final = _named(pinned, spender) + assert _metadata(recovered, spender.digest) == ( + _KeyMetadata(key_alias=spender.alias, user_id=spender.user_id, user_email=spender.user_email), + ), recovered.text + + +def _drop_token_rows(database_url: str, keys: tuple[TenantKey, ...]) -> None: + write_rows( + """DELETE FROM "LiteLLM_VerificationToken" WHERE token = ANY(string_to_array(%s, ','))""", + (",".join(key.digest for key in keys if key.name in KEYS_ONLY_IN_SPEND_LOGS),), + database_url=database_url, + ) + + +@pytest.mark.timeout(420) +def test_every_usage_route_names_spend_log_only_keys_again_once_an_outage_ends( + gateway: Gateway, tmp_path: Path +) -> None: + with _rig(gateway, tmp_path) as (proxy, database_url): + tenant: Final = _tenant(proxy) + keys: Final = tuple(_tenant_key(proxy, tenant, name) for name in TENANT_KEYS) + for key in keys: + _spend(proxy, tenant, key) + digests: Final = frozenset(key.digest for key in keys) + every_table: Final = frozenset(itertools.product(SPEND_TABLES, digests)) + eventually(lambda: every_table - _landed_rows(database_url), lambda missing: not missing, seconds=90) + _drop_token_rows(database_url, keys) + deleted: Final = next(key for key in keys if key.name == "deleted") + proxy.post("/key/delete", {"keys": [deleted.key]}) + outage: Final = _expected(tenant, keys, outage=True) + healthy: Final = _expected(tenant, keys, outage=False) + with _pinned(proxy) as pinned: + with _locked_spend_logs(database_url), _recording(database_url) as during_outage: + burst: Final = _burst(proxy, digests) + blank: Final = _sweep(pinned, tenant, digests) + assert frozenset(burst) == {outage.keys["user aggregated"]}, burst + assert dict(blank.keys) == dict(outage.keys) + assert blank.users == outage.users + assert 1 <= len(_lookups(during_outage)) <= 2, "A worker looked up again inside its 30s miss window" + eventually( + lambda: _walked(_get(pinned, AGGREGATED, {}), digests), + lambda rows: rows == healthy.keys["user aggregated"], + seconds=MISS_TTL_BOUND, + ) + named: Final = _sweep(pinned, tenant, digests) + assert dict(named.keys) == dict(healthy.keys) + assert named.users == healthy.users + assert frozenset(_burst(proxy, digests)) == {healthy.keys["user aggregated"]} + + +@pytest.mark.timeout(300) +def test_usage_page_names_a_rejected_jwt_caller_again_once_spend_logs_free_up(gateway: Gateway, tmp_path: Path) -> None: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_jwk: Final = RSAAlgorithm.to_jwk(private_key.public_key()) + subject: Final = f"jwt-{uuid.uuid4().hex[:12]}" + token: Final = jwt.encode( + {"sub": subject, "exp": int((datetime.now(UTC) + timedelta(minutes=10)).timestamp())}, + private_key, + algorithm="RS256", + headers={"kid": JWKS_KEY_ID}, + ) + digest: Final = f"hashed-jwt-{sha256(token.encode()).hexdigest()}" + with ( + scratch_database() as database_url, + wire_server(lambda _: _jwks_reply(public_jwk)) as jwks, + wire_server(_respond) as wire, + _proxy( + gateway, + tmp_path, + database_url, + wire.url, + general_settings=JWT_SETTINGS, + environment={"JWT_PUBLIC_KEY_URL": jwks.url}, + ) as proxy, + ): + broke: Final = proxy.request("POST", "/user/new", {"user_id": subject, "max_budget": 0}) + assert broke.status_code == 200, broke.text + rejected: Final = proxy.request( + "POST", + "/v1/chat/completions", + {"model": MODEL, "messages": [{"role": "user", "content": f"spend {uuid.uuid4().hex}"}]}, + key=token, + ) + assert rejected.status_code == 422, rejected.text + assert f"User={subject} over budget" in rejected.text, rejected.text + eventually(lambda: _landed(database_url, digest), bool, seconds=70) + with _pinned(proxy) as pinned: + with _locked_spend_logs(database_url), _recording(database_url) as during_outage: + blank: Final = _read(pinned, digest) + assert blank.elapsed >= FAILED_LOOKUP_FLOOR, blank.elapsed + assert _metadata(blank, digest) == (_KeyMetadata(),), blank.text + assert len(_lookups(during_outage)) == 1 + eventually( + lambda: _metadata(_read(pinned, digest), digest), + lambda metadata: metadata == (_KeyMetadata(user_id=subject),), + seconds=MISS_TTL_BOUND, + ) From 036b551f4ec7a21566ee52a0dfbbfdd89ed41684 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Mon, 28 Sep 2026 22:23:04 -0700 Subject: [PATCH 4/5] test(spend): bound outage alias lookups per miss window instead of a fixed count --- ...t_daily_activity_key_metadata_query_timeout.py | 15 +++++++++++++-- 1 file changed, 13 insertions(+), 2 deletions(-) diff --git a/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py b/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py index b06114ada87..9e684ebc317 100644 --- a/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py +++ b/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py @@ -4,6 +4,7 @@ import itertools import json import threading import uuid +from bisect import bisect_left from collections.abc import Generator, Iterator, Mapping from concurrent.futures import ThreadPoolExecutor from contextlib import contextmanager @@ -27,12 +28,15 @@ from integration._support.database_relay import database_relay from integration._support.process import owned_proxy_process from integration._support.wire import Reply, Request, wire_server from jwt.algorithms import RSAAlgorithm +from litellm.constants import SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter MODEL: Final = "key-metadata-recovery-audit" LOOKUP_MARKER: Final = "first_alias" FAILED_LOOKUP_FLOOR: Final = timedelta(seconds=4) MISS_TTL_BOUND: Final = 60 +MISS_WINDOW: Final = timedelta(seconds=SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL) +WORKERS: Final = 2 USAGE: Final = MappingProxyType({"prompt_tokens": 10, "completion_tokens": 30, "total_tokens": 40}) REPLY_TEXT: Final = "You spent a little this week." WAITING_LOOKUPS: Final = ( @@ -274,7 +278,7 @@ def _proxy( }, config=config, remove_environment=("DATABASE_URL_READ_REPLICA",), - workers=2, + workers=WORKERS, ) as owned: yield owned.gateway @@ -391,6 +395,11 @@ def _lookups(seen: SimpleQueue[tuple[int, str]]) -> frozenset[tuple[int, str]]: return frozenset(seen.get_nowait() for _ in range(seen.qsize())) +def _busiest_miss_window(lookups: frozenset[tuple[int, str]]) -> int: + starts: Final = sorted(datetime.fromisoformat(started) for _, started in lookups) + return max((bisect_left(starts, start + MISS_WINDOW) - index for index, start in enumerate(starts)), default=0) + + def _waiting_lookup_count(database_url: str) -> int: return len(read_rows(WAITING_LOOKUPS, (LOOKUP_MARKER,), database_url=database_url)) @@ -913,7 +922,9 @@ def test_every_usage_route_names_spend_log_only_keys_again_once_an_outage_ends( assert frozenset(burst) == {outage.keys["user aggregated"]}, burst assert dict(blank.keys) == dict(outage.keys) assert blank.users == outage.users - assert 1 <= len(_lookups(during_outage)) <= 2, "A worker looked up again inside its 30s miss window" + outage_lookups: Final = _lookups(during_outage) + assert outage_lookups, "No alias lookup reached the locked spend logs" + assert _busiest_miss_window(outage_lookups) <= WORKERS, sorted(outage_lookups) eventually( lambda: _walked(_get(pinned, AGGREGATED, {}), digests), lambda rows: rows == healthy.keys["user aggregated"], From 390962bc61e2b0c936679a88dbc649da3242ae6d Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Wed, 30 Sep 2026 16:26:46 -0700 Subject: [PATCH 5/5] fix(spend-tracking): treat any spend-log lookup failure as a short-lived miss The Prisma client raises a plain AttributeError when the database drops the connection mid-query, so the PrismaError catch let it through and the whole usage call answered 500. Any failure now keeps the 30 second backoff only, and the integration proxy patches its test entitlement at import so uvicorn's spawned workers inherit it --- .../spend_tracking/key_metadata_recovery.py | 4 +-- tests/integration/_support/proxy.py | 11 ++++---- ...ily_activity_key_metadata_query_timeout.py | 18 ++++++++++-- .../test_key_metadata_recovery.py | 28 +++++++++++++++++++ 4 files changed, 50 insertions(+), 11 deletions(-) diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index 6299f576c37..796f6ffd995 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -120,11 +120,9 @@ async def _db_or_empty( warning: str, count: int, ) -> _T | None: - from prisma.errors import PrismaError - try: return await load() - except PrismaError as e: + except Exception as e: verbose_proxy_logger.warning(warning, count, e) return None diff --git a/tests/integration/_support/proxy.py b/tests/integration/_support/proxy.py index a444b93757d..193254b9e7f 100644 --- a/tests/integration/_support/proxy.py +++ b/tests/integration/_support/proxy.py @@ -1,4 +1,4 @@ -"""Run the normal single-process CLI with the existing behavior-suite test entitlement.""" +"""Run the normal CLI with the behavior-suite test entitlement, patched at import so spawned workers inherit it.""" import signal import sys @@ -7,6 +7,10 @@ from unittest.mock import patch from litellm import run_server +patch( # test-quality-ok: route entitlement only; license validation is outside these HTTP/DB contracts + "litellm.proxy.auth.litellm_license.LicenseCheck.is_premium", return_value=True +).start() + def _exit_on_reraised_term(signum: int, frame: FrameType | None) -> None: sys.exit(0) @@ -14,10 +18,7 @@ def _exit_on_reraised_term(signum: int, frame: FrameType | None) -> None: def main() -> None: signal.signal(signal.SIGTERM, _exit_on_reraised_term) - with patch( # test-quality-ok: route entitlement only; license validation is outside these HTTP/DB contracts - "litellm.proxy.auth.litellm_license.LicenseCheck.is_premium", return_value=True - ): - run_server() + run_server() if __name__ == "__main__": diff --git a/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py b/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py index 9e684ebc317..d037909545f 100644 --- a/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py +++ b/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py @@ -331,12 +331,16 @@ def _spender(proxy: Gateway, database_url: str, label: str) -> Spender: return Spender(alias, user_id, user_email, digest) -def _read(pinned: Pinned, digest: str) -> httpx.Response: - response: Final = pinned.request( +def _fetch(pinned: Pinned, digest: str) -> httpx.Response: + return pinned.request( "GET", "/user/daily/activity/aggregated", params={"start_date": _day(-1), "end_date": _day(1), "api_key": digest}, ) + + +def _read(pinned: Pinned, digest: str) -> httpx.Response: + response: Final = _fetch(pinned, digest) assert response.status_code == 200, response.text return response @@ -358,6 +362,14 @@ def _named(pinned: Pinned, spender: Spender) -> httpx.Response: ) +def _named_once_reconnected(pinned: Pinned, spender: Spender) -> httpx.Response: + return eventually( + lambda: _fetch(pinned, spender.digest), + lambda response: response.status_code == 200 and _aliases(response, spender.digest) == (spender.alias,), + seconds=MISS_TTL_BOUND, + ) + + @contextmanager def _locked_spend_logs(database_url: str) -> Generator[None]: with psycopg.connect(database_url) as connection: @@ -884,7 +896,7 @@ def test_usage_page_survives_a_dropped_database_connection_during_alias_recovery dropped: Final = _read(pinned, spender.digest) assert relay.tripped.is_set(), dropped.text assert _aliases(dropped, spender.digest) == (None,), dropped.text - recovered: Final = _named(pinned, spender) + recovered: Final = _named_once_reconnected(pinned, spender) assert _metadata(recovered, spender.digest) == ( _KeyMetadata(key_alias=spender.alias, user_id=spender.user_id, user_email=spender.user_email), ), recovered.text diff --git a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py index d96156a9d10..c52d11e4bb0 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py +++ b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py @@ -471,6 +471,34 @@ async def test_recover_key_metadata_from_spend_logs_never_caches_repeated_query_ assert result[digest]["key_alias"] == "back-online" +@pytest.mark.asyncio +async def test_recover_key_metadata_from_spend_logs_treats_a_dropped_connection_error_as_a_short_lived_miss(): + digest: Final = hash_token("cli-session-dropped-connection") + window: Final = (datetime(2026, 9, 7), datetime(2026, 9, 10)) + cache: Final = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL) + mock_prisma: Final = MagicMock() + _spend_log_transaction( + mock_prisma, + AsyncMock( + side_effect=[ + AttributeError("'NoneType' object has no attribute 'get'"), + [_spend_log_row(digest, "back-online", None, None)], + ] + ), + ) + started: Final = time.time() + + assert await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) == {} + + miss_key: Final = next(key for key in cache.ttl_dict if digest in key and not key.endswith(":missed-before")) + assert cache.ttl_dict[miss_key] - started <= SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL + 1 + assert f"{miss_key}:missed-before" not in cache.ttl_dict + cache.ttl_dict[miss_key] = time.time() - 1 + + result: Final = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) + + assert result[digest]["key_alias"] == "back-online" + @pytest.mark.asyncio async def test_recover_key_metadata_from_spend_logs_drops_the_owner_of_a_digest_shared_by_several_users(): shared_ui_digest = hash_token("ui-token")