diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index 560363ca7d7..672a7aab622 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -171,11 +171,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 @@ -426,6 +424,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: @@ -437,7 +439,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) @@ -454,10 +456,13 @@ 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: + _remember_short_lived_miss(cache, _spend_log_cache_key(digest, window)) + 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/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 new file mode 100644 index 00000000000..d037909545f --- /dev/null +++ b/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py @@ -0,0 +1,997 @@ +import csv +import io +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 +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 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 = ( + "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=WORKERS, + ) 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 _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 + + +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, + ) + + +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: + 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 _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)) + + +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_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 + + +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 + 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"], + 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, + ) 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 bc0bbd4dd38..36540e407fb 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 @@ -449,6 +449,65 @@ 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_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")