mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 390962bc61 into f8f05767da
This commit is contained in:
commit
890a5676a9
4 changed files with 1074 additions and 12 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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")
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue