From a1a42768c1cb2caba444750cec05a07742a45672 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Mon, 5 Oct 2026 22:36:25 +0000 Subject: [PATCH] fix(spend-tracking): stop caching failed spend-log metadata lookups as confirmed misses (#43560) * 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> * refactor(spend-tracking): share the short-lived miss cache write Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(spend): cover key alias recovery after spend log lookup failures across usage routes * test(spend): bound outage alias lookups per miss window instead of a fixed count * 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 * test(integration): audit spend-log metadata recovery under timeouts and dropped connections Cover the daily activity routes, the usage AI chat, the Vantage and CloudZero dry runs and exports under a locked spend-log table and under a database connection dropped mid-lookup, on a two-worker proxy, with the recovery after the outage asserted through the proxy's own miss TTL. Add a dropped_connection_relay that closes only the connection whose bytes carry a trigger, so a cell can drop the one connection the recovery query runs on while the rest of the pool keeps serving. Rewrite the sweep and JWT cells for the merged main: the export route reads metadata by SQL join and never calls the recovery, the search routes answer key rows and find deleted keys by alias, and the daily-spend owner recovery names the user while the alias stays blank. The sweep cell now times out a second lookup under the same lock, which pins the keys blank on the merge base and recovers on this branch. * test(integration): match a dropped-connection trigger split across two reads The dropped-connection relay checked each TCP read on its own, so a SQL marker that straddled two reads never tripped it and the outage cells would run without the outage they meant to exercise. Carry the tail of the previous read into the next check, as the held-statement relay already does, and pin that with a unit test that splits the trigger across two writes. * test(integration): scan relay triggers through an in-process helper The dropped-connection relay now matches its SQL trigger through a TriggerScanner that carries the previous read's tail, and the unit test exercises that scanner directly instead of opening loopback sockets, which tests/unit forbids. The relay's end to end behavior stays covered by the integration cells --------- Co-authored-by: gabriele Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../spend_tracking/key_metadata_recovery.py | 19 +- tests/integration/_support/database_relay.py | 80 ++ ...ily_activity_key_metadata_query_timeout.py | 1138 +++++++++++++++++ ...ey_metadata_recovery_dropped_connection.py | 626 +++++++++ .../test_database_relay.py | 28 + .../test_key_metadata_recovery.py | 59 + 6 files changed, 1943 insertions(+), 7 deletions(-) create mode 100644 tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py create mode 100644 tests/integration/spend/test_key_metadata_recovery_dropped_connection.py create mode 100644 tests/unit/integration_support/test_database_relay.py diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index 225e96179ff..900fea7e2cc 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -172,11 +172,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 @@ -441,6 +439,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: @@ -452,7 +454,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) @@ -469,10 +471,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/database_relay.py b/tests/integration/_support/database_relay.py index cb8b9098a64..20f23d0e010 100644 --- a/tests/integration/_support/database_relay.py +++ b/tests/integration/_support/database_relay.py @@ -146,6 +146,74 @@ class HeldStatementRelay: ) +class TriggerScanner: + def __init__(self, trigger: bytes) -> None: + self._trigger: Final = trigger + self._tail: bytes = b"" + + def feed(self, chunk: bytes) -> bool: + window: Final = self._tail + chunk + self._tail = window[-(len(self._trigger) - 1) :] + return self._trigger in window + + +class DroppedConnectionRelay: + def __init__(self, upstream_host: str, upstream_port: int, trigger: bytes) -> None: + self.port: Final = _free_port() + self._upstream_host: Final = upstream_host + self._upstream_port: Final = upstream_port + self._trigger: Final = trigger + self._loop: Final = asyncio.new_event_loop() + self._armed: Final = threading.Event() + self.dropped: Final = threading.Event() + self._ready: Final = threading.Event() + self._thread: Final = threading.Thread(target=self._run, daemon=True) + + def arm(self) -> None: + self._armed.set() + + def disarm(self) -> None: + self._armed.clear() + + def start(self) -> None: + self._thread.start() + assert self._ready.wait(10), "Database relay did not start" + + def stop(self) -> None: + self._loop.call_soon_threadsafe(self._loop.stop) + self._thread.join(10) + + def _run(self) -> None: + asyncio.set_event_loop(self._loop) + self._loop.run_until_complete(asyncio.start_server(self._serve, "127.0.0.1", self.port)) + self._ready.set() + self._loop.run_forever() + + async def _serve(self, client_reader: asyncio.StreamReader, client_writer: asyncio.StreamWriter) -> None: + server_reader, server_writer = await asyncio.open_connection(self._upstream_host, self._upstream_port) + + async def forward(reader: asyncio.StreamReader, writer: asyncio.StreamWriter, inspect: bool) -> None: + scanner: Final = TriggerScanner(self._trigger) + try: + while chunk := await reader.read(65536): + matched: Final = scanner.feed(chunk) + if inspect and self._armed.is_set() and matched: + self.dropped.set() + client_writer.close() + return + writer.write(chunk) + await writer.drain() + except (ConnectionError, asyncio.IncompleteReadError): + return + finally: + writer.close() + + await asyncio.gather( + forward(client_reader, server_writer, True), + forward(server_reader, client_writer, False), + ) + + def _relayed_url(database_url: str, port: int) -> str: parts: Final = urlsplit(database_url) credentials: Final = f"{parts.username}:{parts.password}@" if parts.username else "" @@ -174,3 +242,15 @@ def held_statement_relay(database_url: str, trigger: bytes) -> Generator[tuple[H yield relay, _relayed_url(database_url, relay.port) finally: relay.stop() + + +@contextmanager +def dropped_connection_relay(database_url: str, trigger: bytes) -> Generator[tuple[DroppedConnectionRelay, str]]: + parts: Final = urlsplit(database_url) + assert parts.hostname is not None and parts.port is not None, database_url + relay: Final = DroppedConnectionRelay(parts.hostname, parts.port, trigger) + relay.start() + try: + yield relay, _relayed_url(database_url, relay.port) + finally: + relay.stop() 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..747570a01ec --- /dev/null +++ b/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py @@ -0,0 +1,1138 @@ +import itertools +import json +import os +import signal +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 psutil +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 OwnedProxy, group_members, owned_proxy_process +from integration._support.wire import Reply, Request, wire_server +from jwt.algorithms import RSAAlgorithm +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter + +from litellm.constants import SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL + +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")) +KEYS_A_SEARCH_FINDS_BY_ALIAS: Final = frozenset(("live", "deleted")) +AGGREGATED: Final = "/user/daily/activity/aggregated" +BURST: Final = 16 + + +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 _owned_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[OwnedProxy]: + 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 + + +@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]: + with _owned_proxy( + gateway, directory, database_url, wire_url, general_settings=general_settings, environment=environment + ) 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 + + +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 _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 {"api_key": str() as digest, "metadata": dict() as metadata} if digest in digests: + meta: Final = _KeyMetadata.model_validate(metadata) + yield KeyRow(digest, meta.key_alias, meta.user_id, meta.user_email) + 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 _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 _sweep(pinned: Pinned, tenant: Tenant, digests: frozenset[str]) -> Mapping[str, frozenset[KeyRow]]: + return MappingProxyType( + {name: _walked(_get(pinned, path, params), digests) for name, (path, params) in _routes(tenant).items()} + ) + + +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, tenant.owner, tenant.email) + return _named_row(tenant, key) + + +def _expected(tenant: Tenant, keys: tuple[TenantKey, ...], *, outage: bool) -> Mapping[str, frozenset[KeyRow]]: + every: Final = frozenset(_outage_row(tenant, key) if outage else _named_row(tenant, key) for key in keys) + found: Final = frozenset(_named_row(tenant, key) for key in keys if key.name in KEYS_A_SEARCH_FINDS_BY_ALIAS) + searched: Final = frozenset(("user search", "team search")) + return MappingProxyType({name: found if name in searched else every for name in _routes(tenant)}) + + +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): + with _recording(database_url) as during_outage: + burst: Final = _burst(proxy, digests) + blank: Final = _sweep(pinned, tenant, digests) + with _recording(database_url) as during_retry: + retried: Final = eventually( + lambda: _get(pinned, AGGREGATED, {}), + lambda response: response.elapsed >= FAILED_LOOKUP_FLOOR, + seconds=MISS_TTL_BOUND, + ) + assert frozenset(burst) == {outage["user aggregated"]}, burst + assert dict(blank) == dict(outage) + 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) + assert _walked(retried, digests) == outage["user aggregated"], retried.text + assert len(_lookups(during_retry)) == 1 + eventually( + lambda: _walked(_get(pinned, AGGREGATED, {}), digests), + lambda rows: rows == healthy["user aggregated"], + seconds=MISS_TTL_BOUND, + ) + named: Final = _sweep(pinned, tenant, digests) + assert dict(named) == dict(healthy) + assert frozenset(_burst(proxy, digests)) == {healthy["user aggregated"]} + + +@pytest.mark.timeout(300) +def test_usage_page_retries_the_spend_log_lookup_for_a_rejected_jwt_caller_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(user_id=subject),), blank.text + assert len(_lookups(during_outage)) == 1 + retried: Final = eventually( + lambda: _probe(pinned, database_url, digest), + lambda probe: probe.ran_lookup, + seconds=MISS_TTL_BOUND, + ) + assert retried.aliases == (None,), retried + named: Final = _read(pinned, digest) + assert _metadata(named, digest) == (_KeyMetadata(user_id=subject),), named.text + + +def _full_metadata(spender: Spender) -> tuple[_KeyMetadata, ...]: + return (_KeyMetadata(key_alias=spender.alias, user_id=spender.user_id, user_email=spender.user_email),) + + +@pytest.mark.timeout(600) +def test_second_proxy_instance_names_the_key_while_the_first_recovers_from_its_own_misses( + gateway: Gateway, tmp_path: Path +) -> None: + first_home: Final = tmp_path / "first" + second_home: Final = tmp_path / "second" + first_home.mkdir() + second_home.mkdir() + with ( + scratch_database() as database_url, + wire_server(_respond) as wire, + _proxy(gateway, first_home, database_url, wire.url) as first, + _proxy(gateway, second_home, database_url, wire.url) as second, + ): + spender: Final = _spender(first, database_url, "second-instance") + with _pinned(first) as pinned_first, _pinned(second) as pinned_second: + with _locked_spend_logs(database_url): + timed_out: Final = _read(pinned_first, spender.digest) + assert timed_out.elapsed >= FAILED_LOOKUP_FLOOR, timed_out.elapsed + assert _aliases(timed_out, spender.digest) == (None,), timed_out.text + retried: Final = eventually( + lambda: _read(pinned_first, spender.digest), + lambda response: response.elapsed >= FAILED_LOOKUP_FLOOR, + seconds=MISS_TTL_BOUND, + ) + assert _aliases(retried, spender.digest) == (None,), retried.text + fresh: Final = _read(pinned_second, spender.digest) + assert fresh.elapsed < FAILED_LOOKUP_FLOOR, fresh.elapsed + assert _metadata(fresh, spender.digest) == _full_metadata(spender), fresh.text + recovered: Final = _named(pinned_first, spender) + assert _metadata(recovered, spender.digest) == _full_metadata(spender), recovered.text + + +@pytest.mark.timeout(600) +def test_usage_page_names_the_key_at_once_after_a_restart_ends_the_outage(gateway: Gateway, tmp_path: Path) -> None: + with scratch_database() as database_url, wire_server(_respond) as wire: + with _proxy(gateway, tmp_path, database_url, wire.url) as first: + spender: Final = _spender(first, database_url, "restart") + with _pinned(first) as pinned: + with _locked_spend_logs(database_url): + 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 + 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 _proxy(gateway, tmp_path, database_url, wire.url) as restarted, _pinned(restarted) as pinned_again: + named: Final = _read(pinned_again, spender.digest) + assert named.elapsed < FAILED_LOOKUP_FLOOR, named.elapsed + assert _metadata(named, spender.digest) == _full_metadata(spender), named.text + + +def _fresh_read(owned: OwnedProxy, digest: str) -> httpx.Response: + with httpx.Client(base_url=owned.gateway.client.base_url, timeout=60, trust_env=False) as client: + return client.get( + AGGREGATED, + params={"start_date": _day(-1), "end_date": _day(1), "api_key": digest}, + headers={"Authorization": f"Bearer {owned.gateway.key}", "Connection": "close"}, + ) + + +def _running_children(owned: OwnedProxy) -> tuple[int, ...]: + return tuple( + member.pid + for member in group_members(owned.process.pid) + if member.pid != owned.process.pid and member.is_running() and member.status() != psutil.STATUS_ZOMBIE + ) + + +def _worker_pids(owned: OwnedProxy) -> tuple[int, ...]: + return tuple( + member.pid + for member in group_members(owned.process.pid) + if member.pid != owned.process.pid and any("spawn_main" in part for part in member.cmdline()) + ) + + +@pytest.mark.timeout(600) +def test_usage_page_keeps_serving_when_a_worker_dies_mid_outage(gateway: Gateway, tmp_path: Path) -> None: + with ( + scratch_database() as database_url, + wire_server(_respond) as wire, + _owned_proxy(gateway, tmp_path, database_url, wire.url) as owned, + ): + spender: Final = _spender(owned.gateway, database_url, "worker-death") + children: Final = _running_children(owned) + workers: Final = _worker_pids(owned) + assert len(workers) == WORKERS, workers + with _locked_spend_logs(database_url): + timed_out: Final = _fresh_read(owned, spender.digest) + assert timed_out.status_code == 200, timed_out.text + assert timed_out.elapsed >= FAILED_LOOKUP_FLOOR, timed_out.elapsed + assert _aliases(timed_out, spender.digest) == (None,), timed_out.text + os.kill(workers[0], signal.SIGKILL) + after_kill: Final = _fresh_read(owned, spender.digest) + assert after_kill.status_code == 200, after_kill.text + assert _aliases(after_kill, spender.digest) == (None,), after_kill.text + respawned: Final = eventually( + lambda: _running_children(owned), + lambda pids: len(pids) >= len(children) and any(pid not in children for pid in pids), + seconds=30, + ) + assert workers[0] not in respawned, respawned + recovered: Final = eventually( + lambda: _fresh_read(owned, spender.digest), + lambda response: response.status_code == 200 and _aliases(response, spender.digest) == (spender.alias,), + seconds=MISS_TTL_BOUND, + ) + assert _metadata(recovered, spender.digest) == _full_metadata(spender), recovered.text + + +def _status(pinned: Pinned, path: str, params: Mapping[str, str]) -> int: + return pinned.request("GET", path, params={"start_date": _day(-1), "end_date": _day(1), **params}).status_code + + +@pytest.mark.timeout(360) +def test_usage_page_rejects_bad_key_filters_and_unrelated_routes_ignore_a_locked_spend_log_table( + gateway: Gateway, tmp_path: Path +) -> None: + with _rig(gateway, tmp_path) as (proxy, database_url): + spender: Final = _spender(proxy, database_url, "bad-filters") + with _pinned(proxy) as pinned, _locked_spend_logs(database_url): + odd_keys: Final = ("k" * 5000, "", "123", json.dumps([spender.digest])) + odd_statuses: Final = tuple(_status(pinned, AGGREGATED, {"api_key": api_key}) for api_key in odd_keys) + assert odd_statuses == (200, 200, 200, 200), odd_statuses + repeated: Final = _status(pinned, f"{AGGREGATED}?api_key={spender.digest}&api_key={spender.digest}", {}) + assert repeated == 200, repeated + page_sizes: Final = tuple( + _status(pinned, "/user/daily/activity", {"page_size": page_size}) for page_size in ("0", "abc") + ) + assert page_sizes == (422, 422), page_sizes + liveliness: Final = pinned.request("GET", "/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + readiness: Final = pinned.request("GET", "/health/readiness") + assert readiness.status_code == 200, readiness.text + gateway_activity: Final = _status(pinned, "/gateway/daily/activity", {}) + assert gateway_activity == 200, gateway_activity + chat: Final = proxy.chat(MODEL, text=f"locked {uuid.uuid4().hex}") + assert object_value(chat["usage"]) == dict(USAGE), chat + + +@pytest.mark.timeout(300) +def test_usage_ai_chat_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-ai-chat") + with _pinned(proxy) as pinned: + relay.arm() + 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 relay.tripped.is_set(), chat.text + 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 + recovered: Final = _named_once_reconnected(pinned, spender) + assert _metadata(recovered, spender.digest) == _full_metadata(spender), recovered.text diff --git a/tests/integration/spend/test_key_metadata_recovery_dropped_connection.py b/tests/integration/spend/test_key_metadata_recovery_dropped_connection.py new file mode 100644 index 00000000000..1c205c00162 --- /dev/null +++ b/tests/integration/spend/test_key_metadata_recovery_dropped_connection.py @@ -0,0 +1,626 @@ +import csv +import io +import json +import re +import socket +import socketserver +import threading +import uuid +from collections.abc import Generator, Mapping +from contextlib import contextmanager, suppress +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta +from email.message import Message +from email.parser import BytesParser +from email.policy import HTTP +from hashlib import sha256 +from pathlib import Path +from types import MappingProxyType +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, object_value, string_value +from integration._support.daily_activity import ( + DAY, + ROUTES, + TEAM_SPEND, + USER_SPEND, + Route, + activity_of_key, + assert_key_reported, + daily_rows, + key_metadata, + key_no_key_table_holds, + seeded_metrics, + seeded_row, + user_row, +) +from integration._support.database import scratch_database +from integration._support.database_relay import dropped_connection_relay +from integration._support.process import OwnedProxy, owned_proxy_process +from integration._support.tls import server_context, write_self_signed_cert +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +WORKERS: Final = 2 +PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml" +REVERSE_HASH_TRIGGER: Final = b"encode(sha256(convert_to(token, 'UTF8')), 'hex') = ANY(" +OWNER_RECOVERY_TRIGGER: Final = b"MIN(user_id) AS first_owner" +CLOUDZERO_HOST: Final = "api.cloudzero.com" +CLOUDZERO_AUTHORITY: Final = f"{CLOUDZERO_HOST}:443" +EXPORT_FILENAME: Final = re.compile(r"usage_\d{8}T\d{6}Z_\d{8}T\d{6}Z\.csv") +FILENAME_STAMP: Final = "%Y%m%dT%H%M%SZ" +VANTAGE_DONE: Final = "Vantage export completed successfully" +CLOUDZERO_DONE: Final = "CloudZero export completed successfully" +NO_ENVIRONMENT: Final[Mapping[str, str]] = MappingProxyType({}) +ENTITY_TABLES: Final = tuple( + dict.fromkeys((route.table, route.entity_column) for route in ROUTES if route.table != USER_SPEND) +) +FOCUS_ALIAS_FIELDS: Final = frozenset({"BillingAccountName", "Tags"}) +CBF_ALIAS_FIELDS: Final = frozenset({"resource/account", "resource/tag:api_key_alias"}) + + +@dataclass(frozen=True, slots=True) +class Owned: + owner: str + email: str + alias: str + key: str + token: str + double: str + + +@dataclass(frozen=True, slots=True) +class Sink: + forbidden: threading.Event + missing: threading.Event + + def respond(self, request: Request) -> Reply: + if self.forbidden.is_set(): + return Reply(status=403, body=b'{"error":"forbidden"}') + if self.missing.is_set(): + return Reply(status=404, body=b'{"error":"missing"}') + return Reply() + + +@dataclass(frozen=True, slots=True) +class Upload: + target: str + authorization: str + filename: str + row: Mapping[str, str] + + +def _identity(label: str) -> Owned: + stamp: Final = uuid.uuid4().hex[:8] + key: Final = f"sk-{uuid.uuid4().hex}" + token: Final = sha256(key.encode()).hexdigest() + return Owned( + owner=f"{label}-owner-{stamp}", + email=f"{label}-{uuid.uuid4().hex[:8]}@example.com", + alias=f"{label}-key-{stamp}", + key=key, + token=token, + double=sha256(token.encode()).hexdigest(), + ) + + +def _register(proxy: Gateway, owned: Owned) -> None: + proxy.post("/user/new", {"user_id": owned.owner, "user_email": owned.email, "auto_create_key": False}) + proxy.post("/key/generate", {"key": owned.key, "user_id": owned.owner, "key_alias": owned.alias}) + + +@contextmanager +def _relayed_proxy( + gateway: Gateway, + directory: Path, + relayed_url: str, + environment: Mapping[str, str] = NO_ENVIRONMENT, + config: Path | None = None, +) -> Generator[OwnedProxy]: + with owned_proxy_process( + gateway, + directory, + { + "DATABASE_URL": relayed_url, + "PRISMA_HEALTH_WATCHDOG_ENABLED": "false", + "LITELLM_DISABLE_NO_REDIS_WARNING": "true", + **environment, + }, + config=config, + remove_environment=("DATABASE_URL_READ_REPLICA",), + workers=WORKERS, + ) as owned: + yield owned + + +@contextmanager +def _reader(owned: OwnedProxy) -> Generator[Gateway]: + with httpx.Client(base_url=str(owned.gateway.client.base_url), timeout=90, trust_env=False) as client: + yield Gateway(client, owned.gateway.key, owned.gateway.upstream_url) + + +def _filters(route: Route, entity: str) -> dict[str, str]: + return {} if route.entity_filter is None else {route.entity_filter: entity} + + +def _config_allowing_a_base_url_in_the_body(directory: Path) -> Path: + config: Final = directory / "client_side_credentials.yaml" + config.write_text( + PROXY_CONFIG.read_text().replace( + "general_settings:\n", "general_settings:\n allow_client_side_credentials: true\n", 1 + ) + ) + return config + + +def _records(value: JsonValue) -> tuple[dict[str, JsonValue], ...]: + assert isinstance(value, list), value + return tuple(object_value(item) for item in value) + + +def _first(body: Mapping[str, JsonValue], name: str) -> dict[str, JsonValue]: + records: Final = _records(body[name]) + assert len(records) == 1, body + return records[0] + + +def _without(record: Mapping[str, JsonValue], names: frozenset[str]) -> dict[str, JsonValue]: + return {name: value for name, value in record.items() if name not in names} + + +def _tags(record: Mapping[str, JsonValue]) -> dict[str, JsonValue]: + return object_value(json.loads(string_value(record["Tags"]))) + + +def _dry_run(reader: Gateway, path: str) -> dict[str, JsonValue]: + return object_value(reader.post(path, {})["dry_run_data"]) + + +def _focus_tags(owned: Owned, *, alias: bool) -> dict[str, str]: + return { + **({"api_key_alias": owned.alias} if alias else {}), + "user_id": owned.owner, + "user_email": owned.email, + "model": "gpt-4o-mini", + "model_group": "gpt-4o-mini", + "custom_llm_provider": "openai", + } + + +def _assert_vantage_dry_run(body: Mapping[str, JsonValue], owned: Owned, *, alias: str | None) -> None: + usage: Final = _first(body, "usage_data") + assert {name: usage.get(name) for name in ("api_key", "api_key_alias", "user_id", "user_email", "spend")} == { + "api_key": owned.double, + "api_key_alias": alias, + "user_id": owned.owner, + "user_email": owned.email, + "spend": 0.25, + }, body + focus: Final = _first(body, "normalized_data") + assert { + name: focus.get(name) + for name in ("BillingAccountName", "BillingAccountId", "BilledCost", "ChargePeriodStart", "ChargePeriodEnd") + } == { + "BillingAccountName": alias, + "BillingAccountId": owned.double, + "BilledCost": 0.25, + "ChargePeriodStart": "2026-02-03T00:00:00Z", + "ChargePeriodEnd": "2026-02-04T00:00:00Z", + }, body + assert _tags(focus) == _focus_tags(owned, alias=alias is not None), body + + +def _assert_cloudzero_dry_run(body: Mapping[str, JsonValue], owned: Owned, *, alias: str | None) -> None: + usage: Final = _first(body, "usage_data") + assert {name: usage.get(name) for name in ("api_key", "api_key_alias", "user_id", "user_email")} == { + "api_key": owned.double, + "api_key_alias": alias, + "user_id": owned.owner, + "user_email": owned.email, + }, body + cbf: Final = _first(body, "cbf_data") + assert {name: cbf.get(name) for name in ("resource/account", "resource/tag:api_key_alias", "cost/cost")} == { + "resource/account": f"{alias}|{owned.double[:8]}" if alias else owned.double[:8], + "resource/tag:api_key_alias": str(alias), + "cost/cost": 0.25, + }, body + + +def _cbf_record(owned: Owned, *, alias: str | None) -> dict[str, str]: + prefix: Final = owned.double[:8] + return { + "time/usage_start": "2026-02-03T00:00:00Z", + "cost/cost": "0.25", + "resource/id": "czrn:litellm:openai:cross-region:unknown:llm-usage:gpt-4o-mini", + "usage/amount": "15", + "usage/units": "tokens", + "resource/service": "gpt-4o-mini", + "resource/account": f"{alias}|{prefix}" if alias else prefix, + "resource/region": "cross-region", + "resource/usage_family": "openai", + "action/operation": "", + "lineitem/type": "Usage", + "resource/tag:provider": "openai", + "resource/tag:model": "gpt-4o-mini", + "resource/tag:entity_type": "team", + "resource/tag:model_group": "gpt-4o-mini", + "resource/tag:api_key_prefix": prefix, + "resource/tag:api_key_alias": str(alias), + "resource/tag:user_email": owned.email, + "resource/tag:api_requests": "1", + "resource/tag:successful_requests": "1", + "resource/tag:failed_requests": "0", + "resource/tag:cache_creation_tokens": "0", + "resource/tag:cache_read_tokens": "0", + "resource/tag:prompt_tokens": "10", + "resource/tag:completion_tokens": "5", + } + + +def _window() -> tuple[datetime, datetime]: + now: Final = datetime.now(UTC).replace(microsecond=0) + return now - timedelta(hours=1), now + timedelta(hours=1) + + +def _window_body(window: tuple[datetime, datetime]) -> dict[str, JsonValue]: + return {"start_time_utc": window[0].isoformat(), "end_time_utc": window[1].isoformat()} + + +def _window_filename(window: tuple[datetime, datetime]) -> str: + return f"usage_{window[0].strftime(FILENAME_STAMP)}_{window[1].strftime(FILENAME_STAMP)}.csv" + + +def _last_received(wire: Wire) -> Request: + received: Final = wire.drain() + assert received, "The sink received nothing" + return received[-1] + + +def _csv_part(request: Request) -> Message: + message: Final = BytesParser(policy=HTTP).parsebytes( + b"content-type: " + request.headers["content-type"].encode() + b"\r\n\r\n" + request.body + ) + parts: Final = message.get_payload() + assert isinstance(parts, list) and len(parts) == 1, request.body + part: Final = parts[0] + assert isinstance(part, Message), request.body + assert part.get_param("name", header="content-disposition") == "csv", request.body + return part + + +def _upload(wire: Wire) -> Upload: + request: Final = _last_received(wire) + part: Final = _csv_part(request) + payload: Final = part.get_payload(decode=True) + assert isinstance(payload, bytes), request.body + rows: Final = tuple(csv.DictReader(io.StringIO(payload.decode()))) + assert len(rows) == 1, payload + filename: Final = part.get_filename() + assert filename is not None, request.body + return Upload(request.target, request.headers["authorization"], filename, rows[0]) + + +def _assert_csv_row(row: Mapping[str, str], owned: Owned, *, alias: str | None) -> None: + assert { + name: row.get(name) + for name in ( + "BilledCost", + "BillingAccountId", + "BillingAccountName", + "ChargeDescription", + "ChargePeriodStart", + "ChargePeriodEnd", + ) + } == { + "BilledCost": "0.25", + "BillingAccountId": owned.double, + "BillingAccountName": alias or "", + "ChargeDescription": "gpt-4o-mini", + "ChargePeriodStart": "2026-02-03T00:00:00Z", + "ChargePeriodEnd": "2026-02-04T00:00:00Z", + }, row + assert json.loads(row["Tags"]) == _focus_tags(owned, alias=alias is not None), row + + +def _pipe(source: socket.socket, sink: socket.socket) -> None: + with suppress(OSError): + for chunk in iter(lambda: source.recv(65536), b""): + sink.sendall(chunk) + with suppress(OSError): + sink.shutdown(socket.SHUT_WR) + + +@contextmanager +def _connect_tunnel(destination: Wire, authority: str) -> Generator[str]: + destination_port: Final = int(destination.url.rsplit(":", 1)[1]) + + class Tunnel(socketserver.StreamRequestHandler): + rbufsize = 0 + request: socket.socket + + def handle(self) -> None: + requested: Final = self.rfile.readline().decode().split()[1] + while self.rfile.readline() not in (b"\r\n", b""): + pass + if requested != authority: + self.wfile.write(b"HTTP/1.1 403 Forbidden\r\ncontent-length: 0\r\n\r\n") + return + self.wfile.write(b"HTTP/1.1 200 Connection established\r\n\r\n") + self.request.settimeout(10) + with socket.create_connection(("127.0.0.1", destination_port), timeout=10) as upstream: + outbound: Final = threading.Thread(target=_pipe, args=(self.request, upstream)) + outbound.start() + _pipe(upstream, self.request) + outbound.join(timeout=12) + + with socketserver.ThreadingTCPServer(("127.0.0.1", 0), Tunnel) as server: + thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05}) + thread.start() + try: + yield f"http://127.0.0.1:{server.server_address[1]}" + finally: + server.shutdown() + thread.join(timeout=6) + + +@pytest.mark.timeout(600) +def test_every_usage_route_survives_a_dropped_database_connection_during_reverse_hash_recovery( + gateway: Gateway, tmp_path: Path +) -> None: + owned_key: Final = _identity("dropped-reverse-hash") + entity: Final = f"dropped-reverse-hash-{uuid.uuid4().hex[:8]}" + rows: Final = ( + user_row(None, owned_key.double, DAY), + *(seeded_row(table, column, entity, owned_key.double, DAY) for table, column in ENTITY_TABLES), + ) + with ( + scratch_database() as database_url, + dropped_connection_relay(database_url, REVERSE_HASH_TRIGGER) as (relay, relayed_url), + _relayed_proxy(gateway, tmp_path, relayed_url) as owned, + _reader(owned) as reader, + daily_rows(rows, database_url=database_url), + ): + _register(owned.gateway, owned_key) + relay.arm() + for route in ROUTES: + relay.dropped.clear() + dropped: Final = activity_of_key(reader, route.path, owned_key.double, **_filters(route, entity)) + assert relay.dropped.is_set(), f"{route.path}: {dropped.text}" + assert_key_reported(dropped, owned_key.double, DAY, key_metadata(), seeded_metrics(1)) + relay.disarm() + named: Final = key_metadata(alias=owned_key.alias, user=owned_key.owner, email=owned_key.email) + for route in ROUTES: + recovered: Final = activity_of_key(reader, route.path, owned_key.double, **_filters(route, entity)) + assert_key_reported(recovered, owned_key.double, DAY, named, seeded_metrics(1)) + + +@pytest.mark.timeout(300) +def test_usage_page_survives_a_dropped_database_connection_during_user_detail_recovery( + gateway: Gateway, tmp_path: Path +) -> None: + owned_key: Final = _identity("dropped-user-detail") + with ( + scratch_database() as database_url, + dropped_connection_relay(database_url, owned_key.owner.encode()) as (relay, relayed_url), + _relayed_proxy(gateway, tmp_path, relayed_url) as owned, + _reader(owned) as reader, + daily_rows((user_row(None, owned_key.token, DAY),), database_url=database_url), + ): + _register(owned.gateway, owned_key) + relay.arm() + dropped: Final = activity_of_key(reader, "/user/daily/activity", owned_key.token) + assert relay.dropped.is_set(), dropped.text + assert_key_reported( + dropped, + owned_key.token, + DAY, + key_metadata(alias=owned_key.alias, user=owned_key.owner, exists=True), + seeded_metrics(1), + ) + relay.disarm() + recovered: Final = activity_of_key(reader, "/user/daily/activity", owned_key.token) + assert_key_reported( + recovered, + owned_key.token, + DAY, + key_metadata(alias=owned_key.alias, user=owned_key.owner, email=owned_key.email, exists=True), + seeded_metrics(1), + ) + + +@pytest.mark.timeout(300) +def test_team_usage_page_survives_a_dropped_database_connection_during_owner_recovery( + gateway: Gateway, tmp_path: Path +) -> None: + owner: Final = f"dropped-owner-{uuid.uuid4().hex[:8]}" + email: Final = f"owner-{uuid.uuid4().hex[:8]}@example.com" + team: Final = f"dropped-owner-team-{uuid.uuid4().hex[:8]}" + api_key: Final = key_no_key_table_holds() + rows: Final = (user_row(owner, api_key, DAY), seeded_row(TEAM_SPEND, "team_id", team, api_key, DAY)) + with ( + scratch_database() as database_url, + dropped_connection_relay(database_url, OWNER_RECOVERY_TRIGGER) as (relay, relayed_url), + _relayed_proxy(gateway, tmp_path, relayed_url) as owned, + _reader(owned) as reader, + daily_rows(rows, database_url=database_url), + ): + owned.gateway.post("/user/new", {"user_id": owner, "user_email": email, "auto_create_key": False}) + relay.arm() + dropped: Final = activity_of_key(reader, "/team/daily/activity", api_key, team_ids=team) + assert relay.dropped.is_set(), dropped.text + assert_key_reported(dropped, api_key, DAY, key_metadata(), seeded_metrics(1)) + relay.disarm() + recovered: Final = activity_of_key(reader, "/team/daily/activity", api_key, team_ids=team) + assert_key_reported(recovered, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1)) + + +@pytest.mark.timeout(300) +def test_vantage_and_cloudzero_dry_runs_survive_a_dropped_database_connection_during_reverse_hash_recovery( + gateway: Gateway, tmp_path: Path +) -> None: + owned_key: Final = _identity("dropped-dry-run") + with ( + scratch_database() as database_url, + dropped_connection_relay(database_url, owned_key.double.encode()) as (relay, relayed_url), + _relayed_proxy(gateway, tmp_path, relayed_url) as owned, + _reader(owned) as reader, + daily_rows((user_row(owned_key.owner, owned_key.double, DAY),), database_url=database_url), + ): + _register(owned.gateway, owned_key) + relay.arm() + vantage_dropped: Final = _dry_run(reader, "/vantage/dry-run") + assert relay.dropped.is_set(), vantage_dropped + _assert_vantage_dry_run(vantage_dropped, owned_key, alias=None) + relay.dropped.clear() + cloudzero_dropped: Final = _dry_run(reader, "/cloudzero/dry-run") + assert relay.dropped.is_set(), cloudzero_dropped + _assert_cloudzero_dry_run(cloudzero_dropped, owned_key, alias=None) + relay.disarm() + vantage: Final = _dry_run(reader, "/vantage/dry-run") + _assert_vantage_dry_run(vantage, owned_key, alias=owned_key.alias) + cloudzero: Final = _dry_run(reader, "/cloudzero/dry-run") + _assert_cloudzero_dry_run(cloudzero, owned_key, alias=owned_key.alias) + assert _without(_first(vantage_dropped, "normalized_data"), FOCUS_ALIAS_FIELDS) == _without( + _first(vantage, "normalized_data"), FOCUS_ALIAS_FIELDS + ) + assert _without(_first(cloudzero_dropped, "cbf_data"), CBF_ALIAS_FIELDS) == _without( + _first(cloudzero, "cbf_data"), CBF_ALIAS_FIELDS + ) + + +@pytest.mark.timeout(300) +def test_vantage_and_cloudzero_dry_runs_survive_a_dropped_database_connection_during_user_detail_recovery( + gateway: Gateway, tmp_path: Path +) -> None: + owned_key: Final = _identity("dropped-dry-run-detail") + with ( + scratch_database() as database_url, + dropped_connection_relay(database_url, owned_key.owner.encode()) as (relay, relayed_url), + _relayed_proxy(gateway, tmp_path, relayed_url) as owned, + _reader(owned) as reader, + daily_rows((user_row(owned_key.owner, owned_key.double, DAY),), database_url=database_url), + ): + _register(owned.gateway, owned_key) + relay.arm() + vantage_dropped: Final = _dry_run(reader, "/vantage/dry-run") + assert relay.dropped.is_set(), vantage_dropped + _assert_vantage_dry_run(vantage_dropped, owned_key, alias=owned_key.alias) + relay.dropped.clear() + cloudzero_dropped: Final = _dry_run(reader, "/cloudzero/dry-run") + assert relay.dropped.is_set(), cloudzero_dropped + _assert_cloudzero_dry_run(cloudzero_dropped, owned_key, alias=owned_key.alias) + relay.disarm() + vantage: Final = _dry_run(reader, "/vantage/dry-run") + cloudzero: Final = _dry_run(reader, "/cloudzero/dry-run") + assert _first(vantage_dropped, "normalized_data") == _first(vantage, "normalized_data") + assert _first(cloudzero_dropped, "cbf_data") == _first(cloudzero, "cbf_data") + + +@pytest.mark.timeout(420) +def test_vantage_export_delivers_a_blank_alias_under_a_dropped_database_connection_and_reports_sink_errors( + gateway: Gateway, tmp_path: Path +) -> None: + owned_key: Final = _identity("dropped-vantage-export") + sink: Final = Sink(threading.Event(), threading.Event()) + api_key: Final = f"vantage-api-key-{uuid.uuid4().hex}" + integration_token: Final = f"vantage-token-{uuid.uuid4().hex}" + costs_target: Final = f"/v2/integrations/{integration_token}/costs.csv" + with ( + scratch_database() as database_url, + dropped_connection_relay(database_url, owned_key.double.encode()) as (relay, relayed_url), + wire_server(sink.respond) as wire, + _relayed_proxy( + gateway, tmp_path, relayed_url, config=_config_allowing_a_base_url_in_the_body(tmp_path) + ) as owned, + _reader(owned) as reader, + daily_rows((user_row(owned_key.owner, owned_key.double, DAY),), database_url=database_url), + ): + _register(owned.gateway, owned_key) + owned.gateway.post( + "/vantage/init", {"api_key": api_key, "integration_token": integration_token, "base_url": wire.url} + ) + relay.arm() + dropped: Final = reader.post("/vantage/export", {}) + assert relay.dropped.is_set(), dropped + assert dropped["message"] == VANTAGE_DONE, dropped + blank: Final = _upload(wire) + assert (blank.target, blank.authorization) == (costs_target, f"Bearer {api_key}"), blank + assert EXPORT_FILENAME.fullmatch(blank.filename), blank + _assert_csv_row(blank.row, owned_key, alias=None) + relay.dropped.clear() + window: Final = _window() + windowed: Final = reader.post("/vantage/export", _window_body(window)) + assert relay.dropped.is_set(), windowed + assert windowed["message"] == VANTAGE_DONE, windowed + bounded: Final = _upload(wire) + assert bounded.filename == _window_filename(window), bounded + _assert_csv_row(bounded.row, owned_key, alias=None) + relay.disarm() + healthy: Final = reader.post("/vantage/export", {}) + assert healthy["message"] == VANTAGE_DONE, healthy + named: Final = _upload(wire) + _assert_csv_row(named.row, owned_key, alias=owned_key.alias) + assert _without(dict(blank.row), FOCUS_ALIAS_FIELDS) == _without(dict(named.row), FOCUS_ALIAS_FIELDS) + sink.forbidden.set() + forbidden: Final = reader.request("POST", "/vantage/export", {}) + assert forbidden.status_code == 500 and "Failed to perform Vantage export" in forbidden.text, forbidden.text + assert _last_received(wire).target == costs_target + sink.forbidden.clear() + sink.missing.set() + missing: Final = reader.request("POST", "/vantage/export", {}) + assert missing.status_code == 500 and "Failed to perform Vantage export" in missing.text, missing.text + assert _last_received(wire).target == costs_target + alive: Final = reader.request("GET", "/health/liveliness") + assert alive.status_code == 200, alive.text + + +@pytest.mark.timeout(420) +def test_cloudzero_export_delivers_a_blank_alias_under_a_dropped_database_connection( + gateway: Gateway, tmp_path: Path +) -> None: + owned_key: Final = _identity("dropped-cloudzero-export") + api_key: Final = f"cloudzero-api-key-{uuid.uuid4().hex}" + connection_id: Final = f"cloudzero-connection-{uuid.uuid4().hex[:8]}" + drops_target: Final = f"/v2/connections/billing/anycost/{connection_id}/billing_drops" + cert, key = write_self_signed_cert(tmp_path, (CLOUDZERO_HOST,)) + with ( + scratch_database() as database_url, + dropped_connection_relay(database_url, owned_key.double.encode()) as (relay, relayed_url), + wire_server(lambda request: Reply(), tls=server_context(cert, key)) as wire, + _connect_tunnel(wire, CLOUDZERO_AUTHORITY) as tunnel_url, + _relayed_proxy( + gateway, tmp_path, relayed_url, {"HTTPS_PROXY": tunnel_url, "SSL_CERT_FILE": str(cert)} + ) as owned, + _reader(owned) as reader, + daily_rows((user_row(owned_key.owner, owned_key.double, DAY),), database_url=database_url), + ): + _register(owned.gateway, owned_key) + owned.gateway.post("/cloudzero/init", {"api_key": api_key, "connection_id": connection_id, "timezone": "UTC"}) + relay.arm() + dropped: Final = reader.post("/cloudzero/export", {}) + assert relay.dropped.is_set(), dropped + assert dropped["message"] == CLOUDZERO_DONE, dropped + blank: Final = _last_received(wire) + assert (blank.method, blank.target) == ("POST", drops_target), blank + assert blank.headers["authorization"] == f"Bearer {api_key}", blank.headers + assert json.loads(blank.body) == { + "month": "2026-02", + "operation": "replace_hourly", + "data": [_cbf_record(owned_key, alias=None)], + }, blank.body + relay.dropped.clear() + windowed: Final = reader.post("/cloudzero/export", _window_body(_window())) + assert relay.dropped.is_set(), windowed + assert windowed["message"] == CLOUDZERO_DONE, windowed + assert json.loads(_last_received(wire).body) == json.loads(blank.body) + relay.disarm() + healthy: Final = reader.post("/cloudzero/export", {}) + assert healthy["message"] == CLOUDZERO_DONE, healthy + named: Final = _last_received(wire) + assert json.loads(named.body) == { + "month": "2026-02", + "operation": "replace_hourly", + "data": [_cbf_record(owned_key, alias=owned_key.alias)], + }, named.body diff --git a/tests/unit/integration_support/test_database_relay.py b/tests/unit/integration_support/test_database_relay.py new file mode 100644 index 00000000000..98501f84058 --- /dev/null +++ b/tests/unit/integration_support/test_database_relay.py @@ -0,0 +1,28 @@ +from __future__ import annotations + +from typing import Final + +import pytest + +from tests.integration._support.database_relay import TriggerScanner + +TRIGGER: Final = b'SELECT "startTime" FROM "LiteLLM_SpendLogs"' + + +@pytest.mark.parametrize("split_at", range(1, len(TRIGGER))) +def test_trigger_scanner_matches_a_trigger_split_across_two_reads(split_at: int) -> None: + scanner: Final = TriggerScanner(TRIGGER) + assert not scanner.feed(TRIGGER[:split_at]) + assert scanner.feed(TRIGGER[split_at:]) + + +def test_trigger_scanner_matches_a_trigger_arriving_one_byte_at_a_time() -> None: + scanner: Final = TriggerScanner(TRIGGER) + hits: Final = tuple(scanner.feed(TRIGGER[i : i + 1]) for i in range(len(TRIGGER))) + assert hits == (False,) * (len(TRIGGER) - 1) + (True,) + + +def test_trigger_scanner_reports_a_match_once() -> None: + scanner: Final = TriggerScanner(TRIGGER) + assert scanner.feed(b"x" + TRIGGER + b"y") + assert not scanner.feed(b"z") diff --git a/tests/unit/proxy/spend_tracking/test_key_metadata_recovery.py b/tests/unit/proxy/spend_tracking/test_key_metadata_recovery.py index 5d02b289360..2a5f03bcdf3 100644 --- a/tests/unit/proxy/spend_tracking/test_key_metadata_recovery.py +++ b/tests/unit/proxy/spend_tracking/test_key_metadata_recovery.py @@ -443,6 +443,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")