test(proxy): migrate DB and Redis backed proxy tests into tests/integration (#43996)

* test(proxy): migrate DB and Redis backed proxy tests into tests/integration

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): drop a type suppression comment from the key metadata integration test

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(proxy): scope integration test cleanup to owned rows and wait for backend stats flush

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(integration): seed NULL cache_hit and bound recovery reads from below

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yuneng <yuneng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-01 09:23:36 -07:00 • committed by GitHub
parent d96477abce
commit 6f123b7083
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 1124 additions and 1205 deletions

View file

@ -0,0 +1,58 @@
import os
from datetime import datetime, timedelta, timezone
from typing import Final
from uuid import uuid4
import pytest
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.utils import (
PrismaClient,
ProxyLogging,
_deprecated_key_cache,
_lookup_deprecated_key,
)
@pytest.mark.asyncio
async def test_deprecated_key_grace_period_cache_hit_path() -> None:
client: Final = PrismaClient(os.environ["DATABASE_URL"], ProxyLogging(UserApiKeyCache()))
old_token_hash: Final = f"old-{uuid4().hex}"
active_token_hash: Final = f"active-{uuid4().hex}"
_deprecated_key_cache.clear()
await client.connect()
try:
await client.db.litellm_verificationtoken.create(
data={
"token": active_token_hash,
"models": [],
}
)
await client.db.litellm_deprecatedverificationtoken.create(
data={
"token": old_token_hash,
"active_token_id": active_token_hash,
"revoke_at": datetime.now(timezone.utc) + timedelta(minutes=5),
}
)
first: Final = await _lookup_deprecated_key(db=client.db, hashed_token=old_token_hash)
assert first == active_token_hash
await client.db.litellm_deprecatedverificationtoken.delete_many(where={"token": old_token_hash})
second: Final = await _lookup_deprecated_key(db=client.db, hashed_token=old_token_hash)
third: Final = await _lookup_deprecated_key(db=client.db, hashed_token=old_token_hash)
assert second == active_token_hash
assert third == active_token_hash
cached: Final = _deprecated_key_cache.get(old_token_hash)
assert isinstance(cached, tuple)
assert len(cached) == 3
finally:
await client.db.litellm_deprecatedverificationtoken.delete_many(where={"token": old_token_hash})
await client.db.litellm_verificationtoken.delete_many(where={"token": active_token_hash})
_deprecated_key_cache.clear()
await client.disconnect()

View file

@ -0,0 +1,63 @@
import os
import uuid
from typing import Final
import pytest
from redis import Redis
from litellm.caching.caching import DualCache
from litellm.caching.redis_cache import RedisCache
from litellm.proxy.hooks.batch_enqueued_tokens import (
BatchEnqueuedTokenOverLimit,
BatchEnqueuedTokenReservation,
BatchEnqueuedTokenScope,
BatchEnqueuedTokenStore,
)
from litellm.proxy.utils import InternalUsageCache
@pytest.mark.asyncio
async def test_redis_lua_path_full_lifecycle() -> None:
redis_host: Final = os.environ["REDIS_HOST"]
redis_port: Final = int(os.environ["REDIS_PORT"])
redis_cache: Final = RedisCache(host=redis_host, port=redis_port)
store: Final = BatchEnqueuedTokenStore(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=redis_cache, default_in_memory_ttl=60))
)
suffix: Final = uuid.uuid4().hex[:8]
key_scope: Final = BatchEnqueuedTokenScope(key="api_key", value=f"api_key-{suffix}", limit=100)
team_scope: Final = BatchEnqueuedTokenScope(key="team", value=f"team-{suffix}", limit=50)
key_counter: Final = f"batch_enqueued_tokens:api_key:api_key-{suffix}"
team_counter: Final = f"batch_enqueued_tokens:team:team-{suffix}"
batch_id: Final = f"batch_{uuid.uuid4().hex}"
record_key: Final = f"batch_enqueued_token_reservation:{batch_id}"
try:
over: Final = await store.reserve(tokens=60, scopes=(key_scope, team_scope))
assert over == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0)
reservation: Final = await store.reserve(tokens=50, scopes=(key_scope, team_scope))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
assert reservation.backend == "redis"
with Redis(host=redis_host, port=redis_port) as raw:
assert int(raw.get(key_counter) or 0) == 50
assert int(raw.get(team_counter) or 0) == 50
assert isinstance(await store.reserve(tokens=1, scopes=(key_scope, team_scope)), BatchEnqueuedTokenOverLimit)
await store.save_reservation(batch_id, reservation)
popped: Final = await store.pop_reservation(batch_id)
assert popped == reservation
assert await store.pop_reservation(batch_id) is None
await store.refund(popped)
with Redis(host=redis_host, port=redis_port) as raw:
assert int(raw.get(key_counter) or 0) == 0
assert int(raw.get(team_counter) or 0) == 0
refill: Final = await store.reserve(tokens=50, scopes=(key_scope, team_scope))
assert isinstance(refill, BatchEnqueuedTokenReservation)
await store.refund(refill)
finally:
with Redis(host=redis_host, port=redis_port) as raw:
raw.delete(key_counter, team_counter, record_key)

View file

@ -0,0 +1,214 @@
import uuid
from collections.abc import Sequence
from datetime import datetime, timedelta
from typing import Final
import pytest
from pydantic import JsonValue, TypeAdapter
from litellm.constants import PTU_SENTINEL_API_KEY
from tests.integration._support.client import Gateway, object_value
from tests.integration._support.database import write_rows
_URL: Final = "/user/daily/activity/aggregated"
_RESULTS: Final = TypeAdapter(list[dict[str, JsonValue]])
def _unique_day() -> str:
return str((datetime(1900, 1, 1) + timedelta(days=uuid.uuid4().int % 200000)).date())
def _seed(day: str, rows: Sequence[tuple[object, ...]]) -> None:
for row in rows:
write_rows(
'INSERT INTO "LiteLLM_DailyUserSpend" (id, user_id, date, api_key, model, model_group,'
" custom_llm_provider, mcp_namespaced_tool_name, endpoint, prompt_tokens, spend, api_requests,"
" successful_requests, updated_at)"
" VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, now())",
tuple(str(value) if isinstance(value, (int, float)) else value for value in row),
)
def _clean(day: str) -> None:
write_rows('DELETE FROM "LiteLLM_DailyUserSpend" WHERE date = %s', (day,))
def _activity(gateway: Gateway, day: str, **params: str) -> dict[str, JsonValue]:
response: Final = gateway.request("GET", _URL, params={"start_date": day, "end_date": day, **params})
assert response.status_code == 200, response.text
return object_value(response.json())
def _row_id() -> str:
return f"agg-{uuid.uuid4().hex}"
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_returns_every_api_key(gateway: Gateway) -> None:
day: Final = _unique_day()
_seed(
day,
[
*[
(
_row_id(),
f"user-{i:03d}",
day,
f"key-{i:03d}",
"gpt-5",
"",
"openai",
None,
"/v1/chat/completions",
10,
6.0 if i == 4 else float(i + 1),
1,
1,
)
for i in range(105)
],
(_row_id(), None, day, PTU_SENTINEL_API_KEY, "gpt-5", "", "azure", None, None, 0, 1000.0, 0, 0),
],
)
try:
body: Final = _activity(gateway, day)
metadata: Final = object_value(body["metadata"])
assert metadata["total_spend"] == pytest.approx(6566.0)
assert metadata["total_api_requests"] == 105
results: Final = _RESULTS.validate_python(body["results"])
assert len(results) == 1
result_day: Final = object_value(results[0])
assert object_value(result_day["metrics"])["spend"] == pytest.approx(6566.0)
breakdown: Final = object_value(result_day["breakdown"])
expected_api_keys: Final = {f"key-{i:03d}" for i in range(105)}
api_keys: Final = object_value(breakdown["api_keys"])
assert set(api_keys) == expected_api_keys
assert PTU_SENTINEL_API_KEY not in api_keys
models: Final = object_value(breakdown["models"])
gpt5: Final = object_value(models["gpt-5"])
assert object_value(gpt5["metrics"])["spend"] == pytest.approx(6566.0)
assert set(object_value(gpt5["api_key_breakdown"])) == expected_api_keys
providers: Final = object_value(breakdown["providers"])
openai: Final = object_value(providers["openai"])
assert object_value(openai["metrics"])["spend"] == pytest.approx(5566.0)
assert set(object_value(openai["api_key_breakdown"])) == expected_api_keys
endpoints: Final = object_value(breakdown["endpoints"])
assert object_value(object_value(endpoints["/v1/chat/completions"])["metrics"])["api_requests"] == 105
finally:
_clean(day)
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_results(
gateway: Gateway,
) -> None:
day: Final = _unique_day()
_seed(
day,
[
(
_row_id(),
f"user-{i}",
day,
f"key-{i}",
"gpt-5",
"",
"openai",
None,
"/v1/chat/completions",
10,
float(i + 1),
1,
1,
)
for i in range(3)
],
)
try:
body: Final = _activity(gateway, day, api_key="key-1")
assert object_value(body["metadata"])["total_spend"] == 2.0
results: Final = _RESULTS.validate_python(body["results"])
assert len(results) == 1
breakdown: Final = object_value(object_value(results[0])["breakdown"])
api_keys: Final = object_value(breakdown["api_keys"])
assert set(api_keys) == {"key-1"}
assert object_value(object_value(api_keys["key-1"])["metrics"])["spend"] == 2.0
gpt5: Final = object_value(object_value(breakdown["models"])["gpt-5"])
assert object_value(gpt5["metrics"])["spend"] == 2.0
assert set(object_value(gpt5["api_key_breakdown"])) == {"key-1"}
finally:
_clean(day)
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_model_group_rollups_fall_back_to_model_name(
gateway: Gateway,
) -> None:
day: Final = _unique_day()
_seed(
day,
[
(
_row_id(),
"user-0",
day,
"key-0",
"gpt-5",
"gpt-5-eu",
"openai",
None,
"/v1/chat/completions",
10,
7.0,
1,
1,
),
(
_row_id(),
"user-1",
day,
"key-1",
"gpt-5",
"",
"openai",
None,
"/v1/chat/completions",
10,
3.0,
1,
1,
),
(
_row_id(),
"user-2",
day,
"key-2",
"claude-x",
None,
"anthropic",
None,
"/v1/messages",
10,
2.0,
1,
1,
),
],
)
try:
body: Final = _activity(gateway, day)
results: Final = _RESULTS.validate_python(body["results"])
assert len(results) == 1
breakdown: Final = object_value(object_value(results[0])["breakdown"])
model_groups: Final = object_value(breakdown["model_groups"])
assert set(model_groups) == {"gpt-5-eu", "gpt-5", "claude-x"}
assert object_value(object_value(model_groups["gpt-5-eu"])["metrics"])["spend"] == 7.0
gpt5_group: Final = object_value(model_groups["gpt-5"])
assert object_value(gpt5_group["metrics"])["spend"] == 3.0
assert object_value(object_value(model_groups["claude-x"])["metrics"])["spend"] == 2.0
assert set(object_value(gpt5_group["api_key_breakdown"])) == {"key-1"}
models: Final = object_value(breakdown["models"])
assert set(models) == {"gpt-5", "claude-x"}
assert object_value(object_value(models["gpt-5"])["metrics"])["spend"] == 10.0
finally:
_clean(day)

View file

@ -0,0 +1,380 @@
from collections.abc import Mapping, Sequence
from dataclasses import dataclass
from datetime import datetime, timedelta
from pathlib import Path
from typing import Final
import litellm_proxy_extras
import psycopg
import pytest
from psycopg.types.json import Jsonb
from pydantic import JsonValue
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.spend_tracking.key_metadata_recovery import recover_key_metadata_from_spend_logs
from litellm.proxy.utils import PrismaClient, ProxyLogging, hash_token
from tests.integration._support.client import eventually
from tests.integration._support.database import scratch_database, write_rows
_SPEND_LOGS_DDL: Final = """
CREATE TABLE "LiteLLM_SpendLogs" (
request_id TEXT PRIMARY KEY,
api_key TEXT NOT NULL DEFAULT '',
"startTime" TIMESTAMP(3) NOT NULL,
"user" TEXT DEFAULT '',
team_id TEXT,
metadata JSONB DEFAULT '{}'
)
"""
_API_KEY_START_TIME_INDEX_MIGRATION: Final = (
Path(litellm_proxy_extras.__file__).parent
/ "migrations"
/ "20260823000000_add_spend_logs_api_key_starttime_index"
/ "migration.sql"
)
_STATS_SQL: Final = """
SELECT seq_scan, idx_scan, seq_tup_read, idx_tup_fetch, n_tup_ins
FROM pg_stat_user_tables
WHERE relname = 'LiteLLM_SpendLogs'
"""
_OTHER_BACKENDS_SQL: Final = """
SELECT count(*) FROM pg_stat_activity
WHERE datname = current_database() AND pid <> pg_backend_pid() AND backend_type = 'client backend'
"""
@dataclass(frozen=True)
class _Settle:
previous: Mapping[str, int] | None
count: int
def _create_spend_logs_table(database_url: str) -> None:
write_rows(_SPEND_LOGS_DDL, (), database_url=database_url)
write_rows(_API_KEY_START_TIME_INDEX_MIGRATION.read_text(), (), database_url=database_url)
def _spend_log_stats(database_url: str) -> dict[str, int]:
with psycopg.connect(database_url) as connection:
row: Final = connection.execute(_STATS_SQL).fetchone()
if row is None:
return {"seq_scan": 0, "idx_scan": 0, "seq_tup_read": 0, "idx_tup_fetch": 0, "n_tup_ins": 0}
return {
"seq_scan": row[0],
"idx_scan": row[1],
"seq_tup_read": row[2],
"idx_tup_fetch": row[3],
"n_tup_ins": row[4],
}
def _other_client_backends(database_url: str) -> int:
with psycopg.connect(database_url) as connection:
row: Final = connection.execute(_OTHER_BACKENDS_SQL).fetchone()
return 0 if row is None else int(row[0])
def _settled_stats(database_url: str, seeded_rows: int | None = None) -> dict[str, int]:
eventually(
lambda: _other_client_backends(database_url),
lambda backends: backends == 0,
seconds=60,
)
settle = _Settle(previous=None, count=0)
def probe() -> dict[str, int]:
nonlocal settle
current: Final = _spend_log_stats(database_url)
if current == settle.previous:
settle = _Settle(previous=current, count=settle.count + 1)
else:
settle = _Settle(previous=current, count=0)
return current
settled: Final = eventually(
probe,
lambda stats: settle.count >= 5 and (seeded_rows is None or stats["n_tup_ins"] >= seeded_rows),
seconds=60,
)
return settled
def _rows_read_since(database_url: str, baseline: Mapping[str, int]) -> int:
settled: Final = _settled_stats(database_url)
return (settled["seq_tup_read"] + settled["idx_tup_fetch"]) - (baseline["seq_tup_read"] + baseline["idx_tup_fetch"])
def _insert_nameless_spend_logs(connection: psycopg.Connection[tuple[object, ...]], digest: str, rows: int) -> None:
connection.execute(
"""
INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime")
SELECT %(digest)s || '-' || g, %(digest)s, %(start)s + g * interval '1 minute'
FROM generate_series(1, %(rows)s) g
""",
{"digest": digest, "start": datetime(2026, 9, 7), "rows": rows},
)
def _named_spend_log(
digest: str, logged_at: datetime, alias: str | None, user: str | None, team: str | None = None
) -> tuple[str, str, datetime, str, str | None, Jsonb]:
return (
f"{digest}-{logged_at.isoformat()}",
digest,
logged_at,
user or "",
team,
Jsonb({"user_api_key_alias": alias} if alias else {}),
)
def _insert_spend_logs(database_url: str, rows: Sequence[tuple[str, str, datetime, str, str | None, Jsonb]]) -> None:
with psycopg.connect(database_url) as connection:
connection.cursor().executemany(
'INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime", "user", team_id, metadata)'
" VALUES (%s, %s, %s, %s, %s, %s)",
list(rows),
)
def _analyze(database_url: str, vacuum: bool) -> None:
with psycopg.connect(database_url, autocommit=True) as connection:
if vacuum:
connection.execute('VACUUM (ANALYZE) "LiteLLM_SpendLogs"')
else:
connection.execute('ANALYZE "LiteLLM_SpendLogs"')
async def _recover(
monkeypatch: pytest.MonkeyPatch,
database_url: str,
digests: set[str] | frozenset[str],
window: tuple[datetime, datetime],
) -> Mapping[str, JsonValue]:
monkeypatch.setenv("DATABASE_URL", database_url)
client: Final = PrismaClient(database_url, ProxyLogging(UserApiKeyCache()))
await client.connect()
try:
return await recover_key_metadata_from_spend_logs(client, digests, window, cache=InMemoryCache())
finally:
await client.disconnect()
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_names_a_key_by_its_oldest_and_newest_named_rows_in_the_window(
monkeypatch: pytest.MonkeyPatch,
) -> None:
with scratch_database() as database_url:
_create_spend_logs_table(database_url)
unnamed_edges, owner_logged_late, reowned, outside_window, never_named = (
hash_token(f"cli-session-{name}") for name in ("edges", "late", "reowned", "window", "never")
)
_insert_spend_logs(
database_url,
(
_named_spend_log(unnamed_edges, datetime(2026, 9, 7, 1), None, None),
_named_spend_log(unnamed_edges, datetime(2026, 9, 8), "cli-a", "alice", "team-a"),
_named_spend_log(unnamed_edges, datetime(2026, 9, 9), "cli-a", "alice", "team-a"),
_named_spend_log(unnamed_edges, datetime(2026, 9, 9, 23), None, None),
_named_spend_log(owner_logged_late, datetime(2026, 9, 7, 1), "cli-b", None),
_named_spend_log(owner_logged_late, datetime(2026, 9, 9), "cli-b", "bob"),
_named_spend_log(reowned, datetime(2026, 9, 7, 1), "cli-c", "carol"),
_named_spend_log(reowned, datetime(2026, 9, 9), "cli-c", "dave"),
_named_spend_log(outside_window, datetime(2026, 9, 6), "stale-alias", "erin"),
_named_spend_log(outside_window, datetime(2026, 9, 8), "cli-d", "erin"),
_named_spend_log(outside_window, datetime(2026, 9, 10), "later-alias", "erin"),
_named_spend_log(never_named, datetime(2026, 9, 8), None, None),
),
)
result: Final = await _recover(
monkeypatch,
database_url,
{unnamed_edges, owner_logged_late, reowned, outside_window, never_named},
(datetime(2026, 9, 7), datetime(2026, 9, 10)),
)
assert dict(result) == {
unnamed_edges: {"key_alias": "cli-a", "team_id": "team-a", "user_id": "alice"},
owner_logged_late: {"key_alias": "cli-b", "team_id": None, "user_id": "bob"},
reowned: {"key_alias": "cli-c", "team_id": None, "user_id": None},
outside_window: {"key_alias": "cli-d", "team_id": None, "user_id": "erin"},
}
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_reads_two_rows_per_key_however_many_the_key_logged(
monkeypatch: pytest.MonkeyPatch,
) -> None:
with scratch_database() as database_url:
_create_spend_logs_table(database_url)
owners: Final[Mapping[str, str]] = {hash_token(f"cli-session-busy-{i}"): f"user-{i}" for i in range(5)}
with psycopg.connect(database_url) as connection:
for digest, owner in owners.items():
connection.execute(
"""
INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime", "user", metadata)
SELECT %(digest)s || '-' || g, %(digest)s, %(start)s + g * interval '1 minute', %(owner)s,
jsonb_build_object('user_api_key_alias', 'cli-session-' || %(owner)s)
FROM generate_series(1, 2000) g
""",
{"digest": digest, "owner": owner, "start": datetime(2026, 9, 7)},
)
_analyze(database_url, vacuum=False)
baseline: Final = _settled_stats(database_url, seeded_rows=10000)
result: Final = await _recover(
monkeypatch, database_url, frozenset(owners), (datetime(2026, 9, 7), datetime(2026, 9, 10))
)
assert {digest: meta.get("user_id") for digest, meta in result.items()} == owners
rows_read: Final = _rows_read_since(database_url, baseline)
assert len(owners) <= rows_read <= 10
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_walks_a_bounded_number_of_nameless_rows_per_key(
monkeypatch: pytest.MonkeyPatch,
) -> None:
with scratch_database() as database_url:
_create_spend_logs_table(database_url)
named_late: Final[Mapping[str, str]] = {hash_token(f"cli-session-late-{i}"): f"user-{i}" for i in range(3)}
never_named: Final = frozenset(hash_token(f"cli-session-never-{i}") for i in range(3))
with psycopg.connect(database_url) as connection:
for digest in (*named_late, *never_named):
_insert_nameless_spend_logs(connection, digest, 3 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE)
for digest, owner in named_late.items():
connection.execute(
"""
INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime", "user", metadata)
VALUES (%(digest)s || '-newest', %(digest)s, %(logged_at)s, %(owner)s,
jsonb_build_object('user_api_key_alias', 'cli-session-' || %(owner)s))
""",
{"digest": digest, "owner": owner, "logged_at": datetime(2026, 9, 9)},
)
_analyze(database_url, vacuum=False)
baseline: Final = _settled_stats(database_url)
result: Final = await _recover(
monkeypatch,
database_url,
frozenset(named_late) | never_named,
(datetime(2026, 9, 7), datetime(2026, 9, 10)),
)
assert {digest: meta.get("user_id") for digest, meta in result.items()} == named_late
rows_read: Final = _rows_read_since(database_url, baseline)
assert len(frozenset(named_late) | never_named) <= rows_read <= 1800
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_reads_a_short_nameless_key_once(
monkeypatch: pytest.MonkeyPatch,
) -> None:
with scratch_database() as database_url:
_create_spend_logs_table(database_url)
rows_per_key: Final = SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE // 2
never_named: Final = frozenset(hash_token(f"cli-session-short-{i}") for i in range(20))
with psycopg.connect(database_url) as connection:
for digest in never_named:
_insert_nameless_spend_logs(connection, digest, rows_per_key)
_analyze(database_url, vacuum=True)
baseline: Final = _settled_stats(database_url)
result: Final = await _recover(
monkeypatch, database_url, never_named, (datetime(2026, 9, 7), datetime(2026, 9, 10))
)
assert dict(result) == {}
rows_read: Final = _rows_read_since(database_url, baseline)
assert len(never_named) <= rows_read <= 1000
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_bounds_a_busy_nameless_key_among_short_keys_before_any_vacuum(
monkeypatch: pytest.MonkeyPatch,
) -> None:
with scratch_database() as database_url:
_create_spend_logs_table(database_url)
busy: Final = frozenset(hash_token(f"cli-session-busy-nameless-{i}") for i in range(3))
with psycopg.connect(database_url) as connection:
for short_key in range(200):
_insert_nameless_spend_logs(
connection,
hash_token(f"cli-session-short-{short_key}"),
SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE // 5,
)
for digest in busy:
_insert_nameless_spend_logs(connection, digest, 30 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE)
_analyze(database_url, vacuum=False)
baseline: Final = _settled_stats(database_url)
result: Final = await _recover(monkeypatch, database_url, busy, (datetime(2026, 9, 7), datetime(2026, 9, 10)))
assert dict(result) == {}
rows_read: Final = _rows_read_since(database_url, baseline)
assert len(busy) <= rows_read <= 900
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_finds_a_name_logged_where_the_oldest_probe_stopped(
monkeypatch: pytest.MonkeyPatch,
) -> None:
with scratch_database() as database_url:
_create_spend_logs_table(database_url)
start: Final = datetime(2026, 9, 7)
past_the_stop, tied_with_the_stop = (hash_token(f"cli-session-{name}") for name in ("past", "tied"))
same_millisecond: Final = tuple(
start + timedelta(minutes=SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE, microseconds=n) for n in (100, 200, 300)
)
_insert_spend_logs(
database_url,
(
*(
_named_spend_log(past_the_stop, start + timedelta(minutes=minute), None, None)
for minute in range(1, SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 20)
),
_named_spend_log(
past_the_stop,
start + timedelta(minutes=SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 20),
"cli-p",
"pat",
),
*(
_named_spend_log(past_the_stop, start + timedelta(minutes=minute), None, None)
for minute in range(
SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 21,
SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 51,
)
),
*(
_named_spend_log(tied_with_the_stop, start + timedelta(minutes=minute), None, None)
for minute in range(1, SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE)
),
_named_spend_log(tied_with_the_stop, same_millisecond[0], None, None),
_named_spend_log(tied_with_the_stop, same_millisecond[1], None, None),
_named_spend_log(tied_with_the_stop, same_millisecond[2], "cli-t", "tess"),
),
)
_analyze(database_url, vacuum=False)
baseline: Final = _settled_stats(database_url)
digests: Final = {past_the_stop, tied_with_the_stop}
result: Final = await _recover(
monkeypatch,
database_url,
digests,
(start, datetime(2026, 9, 10)),
)
assert dict(result) == {
past_the_stop: {"key_alias": "cli-p", "team_id": None, "user_id": "pat"},
tied_with_the_stop: {"key_alias": "cli-t", "team_id": None, "user_id": "tess"},
}
assert len(digests) <= _rows_read_since(database_url, baseline)

View file

@ -0,0 +1,286 @@
import json
import uuid
from collections.abc import Mapping
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Final
import pytest
from pydantic import TypeAdapter
from litellm.types.management_endpoints.prompt_caching_requests import (
PromptCachingRequestFilter,
PromptCachingRequestsResponse,
)
from tests.integration._support.client import Gateway
from tests.integration._support.database import write_rows
_JSON_OBJECT: Final = TypeAdapter(dict[str, object])
_JSON_ROWS: Final = TypeAdapter(list[Mapping[str, object]])
_URL: Final = "/cost_optimization/prompt_caching/requests"
_MARKER: Final = "litellm_gateway_injected_cache"
_EXPECTED: Final = {
"injected": ("injected-empty", "injected-deployment"),
"hits": ("zero-fallback", "nested-read", "legacy-read", "boolean-number"),
"all": (
"zero-fallback",
"write",
"nested-write",
"nested-read",
"nested-creation",
"legacy-read",
"injected-empty",
"injected-deployment",
"boolean-number",
),
}
@dataclass(frozen=True)
class _Case:
request_id: str
metadata: Mapping[str, object]
cache_hit: str | None = None
start_time: datetime = datetime(2011, 9, 1, 12, 0, 0, 123456)
_CASES: Final = (
_Case("injected-empty", {_MARKER: ""}),
_Case("injected-deployment", {_MARKER: "dep-a"}),
_Case("wrong-deployment", {_MARKER: "dep-b"}),
_Case("legacy-read", {"usage_object": {"cache_read_input_tokens": 100}}),
_Case("nested-read", {"usage_object": {"prompt_tokens_details": {"cached_tokens": 100}}}),
_Case("write", {"usage_object": {"cache_creation_input_tokens": 100}}),
_Case("nested-write", {"usage_object": {"prompt_tokens_details": {"cache_write_tokens": 100}}}),
_Case("nested-creation", {"usage_object": {"prompt_tokens_details": {"cache_creation_tokens": 100}}}),
_Case(
"top-precedence",
{"usage_object": {"cache_read_input_tokens": -2, "prompt_tokens_details": {"cached_tokens": 100}}},
),
_Case(
"zero-fallback",
{"usage_object": {"cache_read_input_tokens": 0, "prompt_tokens_details": {"cached_tokens": 100}}},
),
_Case(
"fractional-precedence",
{"usage_object": {"cache_read_input_tokens": 0.5, "prompt_tokens_details": {"cached_tokens": 100}}},
),
_Case("malformed-number", {"usage_object": {"cache_read_input_tokens": "100"}}),
_Case("malformed-container", {"usage_object": [100]}),
_Case("boolean-number", {"usage_object": {"cache_read_input_tokens": True}}),
_Case("boolean-marker", {_MARKER: True}),
_Case("response-cache", {_MARKER: "", "usage_object": {"cache_read_input_tokens": 100}}, "True"),
_Case("outside-before", {_MARKER: ""}, start_time=datetime(2011, 8, 31, 23, 59, 59)),
_Case(
"outside-after", {"usage_object": {"cache_read_input_tokens": 100}}, start_time=datetime(2011, 9, 2, 0, 0, 1)
),
)
def _window(prefix: str) -> tuple[datetime, datetime]:
day: Final = datetime(1900, 1, 1) + timedelta(days=int(prefix[2:14], 16) % 200000)
return day, day + timedelta(days=1)
def _seed(prefix: str, cases: tuple[_Case, ...] = _CASES) -> None:
shift: Final = _window(prefix)[0] - datetime(2011, 9, 1)
for case in cases:
write_rows(
'INSERT INTO "LiteLLM_SpendLogs" (request_id, call_type, api_key, "startTime", "endTime", model,'
" model_id, custom_llm_provider, spend, metadata, cache_hit)"
" VALUES (%s, 'acompletion', %s, %s::timestamp, %s::timestamp, %s, %s, %s, %s, %s::jsonb, %s)",
(
f"{prefix}{case.request_id}",
"test-key",
(case.start_time + shift).isoformat(),
(datetime(2011, 9, 1, 12, 0, 1) + shift).isoformat(),
"claude-sonnet-5",
"dep-a",
"anthropic",
"0.01",
json.dumps(dict(case.metadata)),
case.cache_hit,
),
)
def _clean(prefix: str) -> None:
write_rows('DELETE FROM "LiteLLM_SpendLogs" WHERE request_id LIKE %s', (f"{prefix}%",))
def _strip(prefix: str, request_id: str) -> str:
assert request_id.startswith(prefix), request_id
return request_id[len(prefix) :]
def _run_filter_checks(
gateway: Gateway,
filter: PromptCachingRequestFilter,
prefix: str,
key: str | None,
window: tuple[datetime, datetime],
) -> None:
expected: Final = _EXPECTED[filter]
first: Final = gateway.request(
"GET",
_URL,
params={
"start_date": window[0].isoformat(),
"end_date": window[1].isoformat(),
"filter": filter,
"page_size": "2",
},
key=key,
)
assert first.status_code == 200, first.text
first_page: Final = PromptCachingRequestsResponse.model_validate_json(first.content)
assert tuple(_strip(prefix, row.request_id) for row in first_page.requests) == expected[:2]
assert first_page.has_more is (len(expected) > 2)
assert (first_page.next_cursor is not None) is first_page.has_more
if first_page.next_cursor is not None:
assert _strip(prefix, first_page.next_cursor.request_id) == expected[1]
assert first_page.next_cursor.start_time == first_page.requests[-1].start_time
next_response: Final = gateway.request(
"GET",
_URL,
params={
"start_date": window[0].isoformat(),
"end_date": window[1].isoformat(),
"filter": filter,
"page_size": "2",
"cursor_start_time": first_page.next_cursor.start_time.astimezone(
timezone(timedelta(hours=-7))
).isoformat(),
"cursor_request_id": first_page.next_cursor.request_id,
},
key=key,
)
assert next_response.status_code == 200, next_response.text
next_page: Final = PromptCachingRequestsResponse.model_validate_json(next_response.content)
assert tuple(_strip(prefix, row.request_id) for row in next_page.requests) == expected[2:4]
assert next_page.has_more is (len(expected) > 4)
assert (next_page.next_cursor is not None) is next_page.has_more
second: Final = gateway.request(
"GET",
_URL,
params={
"start_date": window[0].isoformat(),
"end_date": window[1].isoformat(),
"filter": filter,
"page_size": "100",
},
key=key,
)
assert second.status_code == 200, second.text
complete: Final = PromptCachingRequestsResponse.model_validate_json(second.content)
assert tuple(_strip(prefix, row.request_id) for row in complete.requests) == expected
assert complete.has_more is False
assert complete.next_cursor is None
assert all(row.start_time.tzinfo == timezone.utc for row in complete.requests)
payload: Final = _JSON_OBJECT.validate_json(second.content)
assert set(payload) == {"requests", "page_size", "has_more", "next_cursor"}
serialized_rows: Final = _JSON_ROWS.validate_python(payload["requests"])
assert set(serialized_rows[0]) == {
"request_id",
"start_time",
"model",
"gateway_injected",
"cache_read_tokens",
"cache_creation_tokens",
"spend",
"net_savings",
}
by_id: Final = {_strip(prefix, row.request_id): row for row in complete.requests}
if filter == "all":
assert by_id["injected-empty"].gateway_injected is True
assert by_id["injected-empty"].net_savings is None
assert by_id["legacy-read"].gateway_injected is False
assert by_id["legacy-read"].net_savings is not None and by_id["legacy-read"].net_savings > 0
assert by_id["write"].net_savings is not None and by_id["write"].net_savings < 0
@pytest.mark.asyncio
@pytest.mark.parametrize("filter", ["all", "injected", "hits"])
@pytest.mark.parametrize("role", ["admin", "view-only"])
async def test_request_filters_match_accounting_and_paginate_before_projection(
gateway: Gateway, filter: PromptCachingRequestFilter, role: str
) -> None:
prefix: Final = f"pc{uuid.uuid4().hex[:12]}:"
_seed(prefix)
try:
if role == "admin":
_run_filter_checks(gateway, filter, prefix, None, _window(prefix))
else:
with gateway.scenario() as scenario:
viewer: Final = scenario.user(user_role="proxy_admin_viewer")
_run_filter_checks(gateway, filter, prefix, scenario.key(user_id=viewer), _window(prefix))
finally:
_clean(prefix)
@pytest.mark.asyncio
@pytest.mark.parametrize("delete_before_cursor", [False, True])
async def test_cursor_keeps_remaining_requests_once_during_insertions_and_deletions(
gateway: Gateway, delete_before_cursor: bool
) -> None:
prefix: Final = f"pc{uuid.uuid4().hex[:12]}:"
cases: Final = (
*_CASES,
_Case(
"older-cache-read",
{"usage_object": {"cache_read_input_tokens": 100}},
start_time=datetime(2011, 9, 1, 11),
),
)
_seed(prefix, cases)
try:
window: Final = _window(prefix)
expected: Final = (*_EXPECTED["all"], "older-cache-read")
first: Final = gateway.request(
"GET",
_URL,
params={"start_date": window[0].isoformat(), "end_date": window[1].isoformat(), "page_size": "2"},
)
assert first.status_code == 200, first.text
first_page: Final = PromptCachingRequestsResponse.model_validate_json(first.content)
assert tuple(_strip(prefix, row.request_id) for row in first_page.requests) == expected[:2]
assert first_page.next_cursor is not None
write_rows(
'INSERT INTO "LiteLLM_SpendLogs" (request_id, call_type, api_key, "startTime", "endTime", model,'
" model_id, custom_llm_provider, spend, metadata, cache_hit)"
' SELECT %s, call_type, api_key, %s, "endTime", model, model_id, custom_llm_provider, spend,'
' metadata, cache_hit FROM "LiteLLM_SpendLogs" WHERE request_id = %s',
(f"{prefix}newer-request", (window[0] + timedelta(hours=13)).isoformat(), f"{prefix}{expected[0]}"),
)
write_rows(
'INSERT INTO "LiteLLM_SpendLogs" (request_id, call_type, api_key, "startTime", "endTime", model,'
" model_id, custom_llm_provider, spend, metadata, cache_hit)"
' SELECT %s, call_type, api_key, %s, "endTime", model, model_id, custom_llm_provider, spend,'
' metadata, cache_hit FROM "LiteLLM_SpendLogs" WHERE request_id = %s',
(
f"{prefix}zz-higher-id",
(cases[0].start_time + (window[0] - datetime(2011, 9, 1))).isoformat(),
f"{prefix}{expected[0]}",
),
)
if delete_before_cursor:
write_rows('DELETE FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (f"{prefix}{expected[0]}",))
following: Final = gateway.request(
"GET",
_URL,
params={
"start_date": window[0].isoformat(),
"end_date": window[1].isoformat(),
"page_size": "100",
"cursor_start_time": first_page.next_cursor.start_time.isoformat(),
"cursor_request_id": first_page.next_cursor.request_id,
},
)
assert following.status_code == 200, following.text
following_page: Final = PromptCachingRequestsResponse.model_validate_json(following.content)
assert tuple(_strip(prefix, row.request_id) for row in following_page.requests) == expected[2:]
assert following_page.has_more is False
assert following_page.next_cursor is None
finally:
_clean(prefix)

View file

@ -0,0 +1,65 @@
import os
import uuid
from typing import Final
import pytest
from redis import Redis
from litellm.caching.caching import DualCache
from litellm.caching.redis_cache import RedisCache
from litellm.proxy.hooks.parallel_request_limiter_v3 import (
_PROXY_MaxParallelRequestsHandler_v3 as _PROXY_MaxParallelRequestsHandler,
)
from litellm.proxy.utils import InternalUsageCache
from litellm.types.caching import RedisPipelineIncrementOperation
@pytest.mark.asyncio
async def test_async_increment_tokens_with_ttl_preservation() -> None:
redis_host: Final = os.environ["REDIS_HOST"]
redis_port: Final = int(os.environ["REDIS_PORT"])
redis_cache: Final = RedisCache(host=redis_host, port=redis_port)
handler: Final = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=redis_cache))
)
assert handler.token_increment_script is not None
suffix: Final = uuid.uuid4().hex[:8]
key_with_ttl: Final = f"{{test_ttl}}:with_ttl:{suffix}"
key_without_ttl: Final = f"{{test_ttl}}:without_ttl:{suffix}"
try:
await redis_cache.async_delete_cache(key_with_ttl)
await redis_cache.async_delete_cache(key_without_ttl)
await handler.async_increment_tokens_with_ttl_preservation(
pipeline_operations=[
RedisPipelineIncrementOperation(key=key_with_ttl, increment_value=10.0, ttl=60),
RedisPipelineIncrementOperation(key=key_without_ttl, increment_value=5.0, ttl=None),
]
)
assert await redis_cache.async_get_cache(key_with_ttl) == 10.0
assert await redis_cache.async_get_cache(key_without_ttl) == 5.0
first_ttl: Final = await redis_cache.async_get_ttl(key_with_ttl)
assert first_ttl is not None and 0 < first_ttl <= 60
assert await redis_cache.async_get_ttl(key_without_ttl) is None
with Redis(host=redis_host, port=redis_port) as raw:
assert raw.expire(key_with_ttl, 30, xx=True) == 1
await handler.async_increment_tokens_with_ttl_preservation(
pipeline_operations=[
RedisPipelineIncrementOperation(key=key_with_ttl, increment_value=15.0, ttl=60),
RedisPipelineIncrementOperation(key=key_without_ttl, increment_value=7.0, ttl=None),
]
)
assert await redis_cache.async_get_cache(key_with_ttl) == 25.0
assert await redis_cache.async_get_cache(key_without_ttl) == 12.0
second_ttl: Final = await redis_cache.async_get_ttl(key_with_ttl)
assert second_ttl is not None and 0 < second_ttl <= 30
assert await redis_cache.async_get_ttl(key_without_ttl) is None
finally:
await redis_cache.async_delete_cache(key_with_ttl)
await redis_cache.async_delete_cache(key_without_ttl)

View file

@ -0,0 +1,56 @@
from datetime import date
from typing import Final
import pytest
from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache
from litellm.proxy.spend_tracking.spend_capture_rate import captured_spend_by_day
from litellm.proxy.utils import PrismaClient, ProxyLogging
from tests.integration._support.database import scratch_database, write_rows
_DAILY_USER_SPEND_DDL: Final = """
CREATE TABLE "LiteLLM_DailyUserSpend" (
id TEXT PRIMARY KEY,
date TEXT NOT NULL,
custom_llm_provider TEXT,
spend DOUBLE PRECISION DEFAULT 0
)
"""
@pytest.mark.asyncio
async def test_captured_spend_sums_only_the_openai_billed_providers_inside_the_window(
monkeypatch: pytest.MonkeyPatch,
) -> None:
with scratch_database() as database_url:
monkeypatch.setenv("DATABASE_URL", database_url)
write_rows(_DAILY_USER_SPEND_DDL, (), database_url=database_url)
for index, (day, provider, spend) in enumerate(
(
("2026-09-19", "openai", 1.0),
("2026-09-20", "openai", 2.0),
("2026-09-20", "openai", 3.0),
("2026-09-20", "text-completion-openai", 0.5),
("2026-09-20", "anthropic", 100.0),
("2026-09-21", "azure", 100.0),
("2026-09-22", "openai", 4.0),
)
):
write_rows(
'INSERT INTO "LiteLLM_DailyUserSpend" (id, date, custom_llm_provider, spend) VALUES (%s, %s, %s, %s)',
(f"row-{index}", day, provider, str(spend)),
database_url=database_url,
)
client: Final = PrismaClient(database_url, ProxyLogging(UserApiKeyCache()))
await client.connect()
try:
captured: Final = await captured_spend_by_day(
client,
litellm_providers=("openai", "text-completion-openai"),
start_date=date(2026, 9, 20),
end_date=date(2026, 9, 21),
)
finally:
await client.disconnect()
assert dict(captured) == {"2026-09-20": 5.5}

View file

@ -10,11 +10,8 @@ Covers the critical gaps:
6. Rotation count increments correctly over multiple rotations
"""
import os
from datetime import datetime, timedelta, timezone
from typing import cast
from unittest.mock import AsyncMock, MagicMock, patch
from uuid import uuid4
import pytest
@ -24,11 +21,6 @@ from litellm.proxy._types import (
LiteLLM_VerificationToken,
)
from litellm.proxy.common_utils.key_rotation_manager import KeyRotationManager
from litellm.proxy.utils import (
PrismaClient,
_deprecated_key_cache,
_lookup_deprecated_key,
)
class TestMultiPodKeyRotation:
@ -562,85 +554,3 @@ class TestKeyRotationInitialization:
assert acquire_call.kwargs.get("cronjob_id") == KEY_ROTATION_JOB_NAME
assert release_call.kwargs.get("cronjob_id") == KEY_ROTATION_JOB_NAME
class TestDeprecatedKeyLookupDbE2E:
"""DB-backed integration tests for deprecated key lookup behavior."""
@pytest.mark.asyncio
async def test_deprecated_key_grace_period_cache_hit_path(self):
"""
End-to-end validation against a real Prisma-backed DB:
- old key hash resolves through LiteLLM_DeprecatedVerificationToken
- repeated lookups hit the in-memory deprecated-key cache
- no ValueError/401 regression on subsequent requests
"""
database_url = os.getenv("DATABASE_URL")
if not database_url:
pytest.skip("DATABASE_URL not set; skipping DB-backed key-rotation E2E test.")
db_url = cast(str, database_url)
proxy_logging_obj = MagicMock()
proxy_logging_obj.failure_handler = AsyncMock()
prisma_client = PrismaClient(
database_url=db_url, proxy_logging_obj=proxy_logging_obj
)
old_token_hash = f"old-{uuid4().hex}"
active_token_hash = f"active-{uuid4().hex}"
_deprecated_key_cache.clear()
await prisma_client.connect()
try:
await prisma_client.db.litellm_verificationtoken.create(
data={
"token": active_token_hash,
"models": [],
}
)
await prisma_client.db.litellm_deprecatedverificationtoken.create(
data={
"token": old_token_hash,
"active_token_id": active_token_hash,
"revoke_at": datetime.now(timezone.utc) + timedelta(minutes=5),
}
)
# Request 1 (DB path) + Request 2/3 (cache-hit path)
r1 = await _lookup_deprecated_key(
db=prisma_client.db,
hashed_token=old_token_hash,
)
r2 = await _lookup_deprecated_key(
db=prisma_client.db,
hashed_token=old_token_hash,
)
r3 = await _lookup_deprecated_key(
db=prisma_client.db,
hashed_token=old_token_hash,
)
assert r1 == active_token_hash
assert r2 == active_token_hash
assert r3 == active_token_hash
cached = _deprecated_key_cache.get(old_token_hash)
assert isinstance(cached, tuple)
assert len(cached) == 3
finally:
# Best-effort cleanup for idempotent reruns.
try:
await prisma_client.db.litellm_deprecatedverificationtoken.delete_many(
where={"token": old_token_hash}
)
except Exception:
pass
try:
await prisma_client.db.litellm_verificationtoken.delete_many(
where={"token": active_token_hash}
)
except Exception:
pass
_deprecated_key_cache.clear()
await prisma_client.disconnect()

View file

@ -8,7 +8,6 @@ response-shape helpers the v3 limiter's post-call hooks rely on.
import base64
import logging
import socket
import uuid
from collections.abc import Mapping, Sequence
from types import MappingProxyType, SimpleNamespace
@ -402,47 +401,6 @@ def test_batch_response_view_accepts_batch_objects_only():
assert batch_response_view("batch_1") is None
def _local_redis_port() -> int | None:
for port in (6379,):
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.settimeout(0.2)
if sock.connect_ex(("127.0.0.1", port)) == 0:
return port
return None
@pytest.mark.asyncio
@pytest.mark.skipif(_local_redis_port() is None, reason="requires a local Redis on 6379 for the Lua script path")
async def test_redis_lua_path_full_lifecycle():
from litellm.caching.redis_cache import RedisCache
port = _local_redis_port()
redis_cache = RedisCache(host="127.0.0.1", port=port)
store = BatchEnqueuedTokenStore(
internal_usage_cache=InternalUsageCache(DualCache(redis_cache=redis_cache, default_in_memory_ttl=60))
)
key_scope = _scope(limit=100, key="api_key")
team_scope = _scope(limit=50, key="team")
over = await store.reserve(tokens=60, scopes=(key_scope, team_scope))
assert over == BatchEnqueuedTokenOverLimit(scope=team_scope, enqueued=0)
reservation = await store.reserve(tokens=50, scopes=(key_scope, team_scope))
assert isinstance(reservation, BatchEnqueuedTokenReservation)
assert isinstance(await store.reserve(tokens=1, scopes=(key_scope, team_scope)), BatchEnqueuedTokenOverLimit)
batch_id = f"batch_{uuid.uuid4().hex}"
await store.save_reservation(batch_id, reservation)
popped = await store.pop_reservation(batch_id)
assert popped == reservation
assert await store.pop_reservation(batch_id) is None
await store.refund(popped)
refill = await store.reserve(tokens=50, scopes=(key_scope, team_scope))
assert isinstance(refill, BatchEnqueuedTokenReservation)
await store.refund(refill)
class _OpenBreakerRedis:
def async_register_script(self, script: str):
async def refused(keys: Sequence[str], args: Sequence[str | bytes | int | float]) -> object:

View file

@ -4,7 +4,6 @@ Unit Tests for the max parallel request limiter v3 for the proxy
import asyncio
import logging
import os
import sys
import time
from collections.abc import Iterator, Sequence
@ -1561,200 +1560,6 @@ async def test_dynamic_rate_limiting_v3():
), "RPM limit should be enforced when dynamic mode and failures detected"
@pytest.mark.flaky(retries=3, delay=2)
@pytest.mark.asyncio
async def test_async_increment_tokens_with_ttl_preservation():
"""
Test TTL preservation functionality for token increment operations.
This test verifies that:
1. Keys are created with proper TTL on first increment
2. TTL is preserved on subsequent increments (not reset)
3. Both TTL and non-TTL operations work correctly in the same call
Environment variables required:
- REDIS_HOST: Redis server hostname
- REDIS_PORT: Redis server port
- REDIS_PASSWORD: Redis password (optional)
Test scenario:
1. First call: Create keys with TTL=60s and TTL=None
2. Wait 2 seconds
3. Second call: Increment same keys
4. Verify TTL decreased but wasn't reset to 60s
"""
import time
from litellm.caching.redis_cache import RedisCache
from litellm.types.caching import RedisPipelineIncrementOperation
# Skip test if Redis environment variables are not set
redis_host = os.getenv("REDIS_HOST")
redis_port = os.getenv("REDIS_PORT")
redis_password = os.getenv("REDIS_PASSWORD")
if not redis_host or not redis_port:
pytest.skip("Redis environment variables (REDIS_HOST, REDIS_PORT) not set")
# Setup Redis cache
redis_cache = RedisCache(
host=redis_host,
port=int(redis_port),
password=redis_password,
)
local_cache = DualCache(redis_cache=redis_cache)
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
internal_usage_cache=InternalUsageCache(local_cache)
)
# Verify Redis connection is working
try:
await redis_cache.ping()
except Exception as e:
pytest.skip(f"Redis connection failed: {str(e)}")
# Verify the TTL preservation script is registered
if parallel_request_handler.token_increment_script is None:
pytest.skip(
"Token increment script not available - Redis Lua scripting may not be supported"
)
# Test keys - use hash tags to ensure they map to same Redis cluster slot
# Use a unique suffix per test run to avoid stale state from prior runs
import uuid
unique_suffix = str(uuid.uuid4())[:8]
test_key_with_ttl = f"{{test_ttl}}:with_ttl:{unique_suffix}"
test_key_without_ttl = f"{{test_ttl}}:without_ttl:{unique_suffix}"
try:
# Clean up any existing test keys
try:
await redis_cache.async_delete_cache(test_key_with_ttl)
await redis_cache.async_delete_cache(test_key_without_ttl)
except Exception:
# Keys might not exist, ignore cleanup errors
pass
# First increment: Create operations with mixed TTL scenarios
pipeline_operations_first = [
RedisPipelineIncrementOperation(
key=test_key_with_ttl, increment_value=10.0, ttl=60
),
RedisPipelineIncrementOperation(
key=test_key_without_ttl, increment_value=5.0, ttl=None # No TTL
),
]
# Execute first increment
await parallel_request_handler.async_increment_tokens_with_ttl_preservation(
pipeline_operations=pipeline_operations_first
)
# Small delay to ensure Redis has processed the commands
await asyncio.sleep(0.1)
# Verify keys exist and check initial TTL
ttl_after_first = await redis_cache.async_get_ttl(test_key_with_ttl)
value_after_first_with_ttl = await redis_cache.async_get_cache(
test_key_with_ttl
)
value_after_first_without_ttl = await redis_cache.async_get_cache(
test_key_without_ttl
)
assert (
value_after_first_with_ttl == 10.0
), f"First increment should set value to 10.0, got {value_after_first_with_ttl}"
assert (
value_after_first_without_ttl == 5.0
), "First increment should set value to 5.0"
assert (
ttl_after_first is not None and ttl_after_first > 0
), "Key with TTL should have positive TTL after first increment"
assert ttl_after_first <= 60, "TTL should not exceed the set value"
# Check TTL for key without TTL (should be None, meaning no expiry)
ttl_no_ttl_key = await redis_cache.async_get_ttl(test_key_without_ttl)
assert (
ttl_no_ttl_key is None
), "Key without TTL should have no expiry (None from async_get_ttl)"
# Wait a moment to ensure TTL decreases
await asyncio.sleep(2)
# Second increment: Same operations to test TTL preservation
pipeline_operations_second = [
RedisPipelineIncrementOperation(
key=test_key_with_ttl, increment_value=15.0, ttl=60 # Same TTL value
),
RedisPipelineIncrementOperation(
key=test_key_without_ttl, increment_value=7.0, ttl=None # No TTL
),
]
# Execute second increment
await parallel_request_handler.async_increment_tokens_with_ttl_preservation(
pipeline_operations=pipeline_operations_second
)
# Small delay to ensure Redis has processed the commands
await asyncio.sleep(0.1)
# Verify TTL preservation and value updates
ttl_after_second = await redis_cache.async_get_ttl(test_key_with_ttl)
value_after_second_with_ttl = await redis_cache.async_get_cache(
test_key_with_ttl
)
value_after_second_without_ttl = await redis_cache.async_get_cache(
test_key_without_ttl
)
assert (
value_after_second_with_ttl == 25.0
), "Second increment should update value to 25.0"
assert (
value_after_second_without_ttl == 12.0
), "Second increment should update value to 12.0"
# Critical test: TTL should be preserved (not reset to 60)
assert ttl_after_second is not None, "TTL should still exist"
assert (
ttl_after_second < ttl_after_first
), "TTL should have decreased (not been reset)"
assert ttl_after_second > 0, "TTL should still be positive"
# TTL should not be close to the original 60 seconds (proving it wasn't reset)
assert (
ttl_after_second < 59
), "TTL should be significantly less than original, proving preservation"
# Key without TTL should still have no expiry
ttl_no_ttl_key_after_second = await redis_cache.async_get_ttl(
test_key_without_ttl
)
assert (
ttl_no_ttl_key_after_second is None
), "Key without TTL should still have no expiry"
finally:
# Clean up test keys
try:
await redis_cache.async_delete_cache(test_key_with_ttl)
await redis_cache.async_delete_cache(test_key_without_ttl)
except Exception:
# Ignore cleanup errors
pass
# Properly close Redis connections to prevent warnings
try:
await redis_cache.disconnect()
except Exception:
# Ignore disconnect errors
pass
@pytest.mark.asyncio
async def test_async_increment_tokens_fallback_behavior():
"""

View file

@ -1,16 +1,11 @@
import re
from collections.abc import Sequence
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import psycopg
import pytest
from psycopg.rows import dict_row
from pytest_postgresql import factories
from litellm.constants import PTU_SENTINEL_API_KEY
from litellm.proxy.management_endpoints.common_daily_activity import (
_adjust_dates_for_timezone,
_build_aggregated_sql_query,
@ -1454,252 +1449,6 @@ async def test_get_daily_activity_aggregated_empty_result_set():
assert result.metadata.total_compression_saved_tokens == 0
_aggregated_postgresql_proc: Final = factories.postgresql_proc()
_aggregated_postgresql: Final = factories.postgresql("_aggregated_postgresql_proc")
_DAILY_USER_SPEND_DDL: Final = """
CREATE TABLE "LiteLLM_DailyUserSpend" (
id TEXT PRIMARY KEY,
user_id TEXT,
date TEXT NOT NULL,
api_key TEXT NOT NULL,
model TEXT,
model_group TEXT,
custom_llm_provider TEXT,
mcp_namespaced_tool_name TEXT,
endpoint TEXT,
prompt_tokens BIGINT DEFAULT 0,
completion_tokens BIGINT DEFAULT 0,
cache_read_input_tokens BIGINT DEFAULT 0,
cache_creation_input_tokens BIGINT DEFAULT 0,
compression_saved_tokens BIGINT DEFAULT 0,
compression_savings_spend DOUBLE PRECISION DEFAULT 0,
prompt_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
gateway_injected_caching_savings_spend DOUBLE PRECISION DEFAULT 0,
autorouter_savings_spend DOUBLE PRECISION DEFAULT 0,
spend DOUBLE PRECISION DEFAULT 0,
api_requests BIGINT DEFAULT 0,
successful_requests BIGINT DEFAULT 0,
failed_requests BIGINT DEFAULT 0,
total_response_time_ms BIGINT DEFAULT 0,
timed_requests BIGINT DEFAULT 0
)
"""
def _seed_daily_user_spend(conn: psycopg.Connection, rows: Sequence[tuple[object, ...]]) -> None:
with conn.cursor() as cur:
cur.execute(_DAILY_USER_SPEND_DDL)
cur.executemany(
"""
INSERT INTO "LiteLLM_DailyUserSpend"
(id, user_id, date, api_key, model, model_group, custom_llm_provider,
endpoint, prompt_tokens, spend, api_requests, successful_requests)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
""",
rows,
)
conn.commit()
def _psycopg_query_raw(conn: psycopg.Connection, row_counts: list[int]):
"""Run the proxy's $N-parameterized SQL through psycopg, recording each result size."""
async def query_raw(sql: str, *params: str) -> list[dict[str, object]]:
converted: Final = re.sub(r"\$(\d+)", r"%(p\1)s", sql)
with conn.cursor(row_factory=dict_row) as cur:
cur.execute(
converted, # pyright: ignore[reportArgumentType] # psycopg stubs want a literal-typed query
{f"p{i}": v for i, v in enumerate(params, start=1)},
)
rows: Final = cur.fetchall()
row_counts.append(len(rows))
return rows
return query_raw
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_returns_every_api_key(
_aggregated_postgresql: psycopg.Connection,
):
n_keys: Final = 105
key_rows: Final = [
(
f"row-{i:03d}",
f"user-{i:03d}",
"2026-06-01",
f"key-{i:03d}",
"gpt-5",
"",
"openai",
"/v1/chat/completions",
10,
6.0 if i == 4 else float(i + 1),
1,
1,
)
for i in range(n_keys)
]
sentinel_row: Final = (
"row-ptu",
None,
"2026-06-01",
PTU_SENTINEL_API_KEY,
"gpt-5",
"",
"azure",
None,
0,
1000.0,
0,
0,
)
_seed_daily_user_spend(_aggregated_postgresql, [*key_rows, sentinel_row])
key_spend: Final = sum(6.0 if i == 4 else float(i + 1) for i in range(n_keys))
expected_api_keys: Final = {f"key-{i:03d}" for i in range(n_keys)}
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, [])
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_usertable = MagicMock()
mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[])
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
model=None,
api_key=None,
)
assert result.metadata.total_spend == pytest.approx(key_spend + 1000.0)
assert result.metadata.total_api_requests == 105
day: Final = result.results[0]
assert day.metrics.spend == pytest.approx(key_spend + 1000.0)
assert set(day.breakdown.api_keys) == expected_api_keys
assert PTU_SENTINEL_API_KEY not in day.breakdown.api_keys
assert day.breakdown.models["gpt-5"].metrics.spend == pytest.approx(key_spend + 1000.0)
assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == expected_api_keys
assert day.breakdown.providers["openai"].metrics.spend == pytest.approx(key_spend)
assert set(day.breakdown.providers["openai"].api_key_breakdown) == expected_api_keys
assert day.breakdown.endpoints["/v1/chat/completions"].metrics.api_requests == 105
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_results(
_aggregated_postgresql: psycopg.Connection,
):
"""An explicit api_key filter must scope the results to that key alone."""
rows: Final = [
(
f"row-{i}",
f"user-{i}",
"2026-06-01",
f"key-{i}",
"gpt-5",
"",
"openai",
"/v1/chat/completions",
10,
float(i + 1),
1,
1,
)
for i in range(3)
]
_seed_daily_user_spend(_aggregated_postgresql, rows)
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, [])
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_usertable = MagicMock()
mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[])
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
model=None,
api_key="key-1",
)
assert result.metadata.total_spend == 2.0
day: Final = result.results[0]
assert set(day.breakdown.api_keys) == {"key-1"}
assert day.breakdown.api_keys["key-1"].metrics.spend == 2.0
assert day.breakdown.models["gpt-5"].metrics.spend == 2.0
assert set(day.breakdown.models["gpt-5"].api_key_breakdown) == {"key-1"}
@pytest.mark.asyncio
async def test_get_daily_activity_aggregated_model_group_rollups_fall_back_to_model_name(
_aggregated_postgresql: psycopg.Connection,
):
"""Rows stored with an empty or NULL model_group must land in the model_groups
breakdown under their model name instead of vanishing from the usage UI."""
rows: Final = [
(
"row-0",
"user-0",
"2026-06-01",
"key-0",
"gpt-5",
"gpt-5-eu",
"openai",
"/v1/chat/completions",
10,
7.0,
1,
1,
),
("row-1", "user-1", "2026-06-01", "key-1", "gpt-5", "", "openai", "/v1/chat/completions", 10, 3.0, 1, 1),
("row-2", "user-2", "2026-06-01", "key-2", "claude-x", None, "anthropic", "/v1/messages", 10, 2.0, 1, 1),
]
_seed_daily_user_spend(_aggregated_postgresql, rows)
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_prisma.db.query_raw = _psycopg_query_raw(_aggregated_postgresql, [])
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_usertable = MagicMock()
mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[])
result = await get_daily_activity_aggregated(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2026-06-01",
end_date="2026-06-01",
model=None,
api_key=None,
)
breakdown: Final = result.results[0].breakdown
assert set(breakdown.model_groups) == {"gpt-5-eu", "gpt-5", "claude-x"}
assert breakdown.model_groups["gpt-5-eu"].metrics.spend == 7.0
assert breakdown.model_groups["gpt-5"].metrics.spend == 3.0
assert breakdown.model_groups["claude-x"].metrics.spend == 2.0
assert set(breakdown.model_groups["gpt-5"].api_key_breakdown) == {"key-1"}
assert set(breakdown.models) == {"gpt-5", "claude-x"}
assert breakdown.models["gpt-5"].metrics.spend == 10.0
def _no_spend_record():
"""A rollup row for a key with no spend, where SUM() returns NULL (None)."""
return SimpleNamespace(

View file

@ -1,148 +1,19 @@
import json
from collections.abc import AsyncIterator, Mapping
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from types import SimpleNamespace
from collections.abc import Mapping
from typing import Final
import httpx
import psycopg
import pytest
import pytest_asyncio
from fastapi import FastAPI
from prisma import Prisma
from pydantic import TypeAdapter
from pytest_postgresql import factories
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.management_endpoints.prompt_caching_requests import router
from litellm.proxy.spend_tracking.savings import (
extract_cache_creation_tokens,
extract_cache_read_tokens,
marks_gateway_injection,
)
from litellm.types.management_endpoints.prompt_caching_requests import (
PromptCachingRequestFilter,
PromptCachingRequestsResponse,
)
pytestmark = pytest.mark.usefixtures("local_model_cost_map")
_cache_postgresql_proc: Final = factories.postgresql_proc() # pyright: ignore[reportUnknownMemberType] # third-party fixture factory has incomplete callable types
_cache_postgresql: Final = factories.postgresql("_cache_postgresql_proc")
_JSON_OBJECT: Final = TypeAdapter(Mapping[str, object])
_JSON_ROWS: Final = TypeAdapter(tuple[Mapping[str, object], ...])
_START: Final = "2026-09-01T00:00:00Z"
_END: Final = "2026-09-02T00:00:00Z"
_URL: Final = "/cost_optimization/prompt_caching/requests"
_MODEL: Final = "claude-sonnet-5"
_MARKER: Final = "litellm_gateway_injected_cache"
_DDL: Final = """
CREATE TABLE "LiteLLM_SpendLogs" (
request_id TEXT PRIMARY KEY, "startTime" TIMESTAMP, "endTime" TIMESTAMP,
model TEXT, model_id TEXT, custom_llm_provider TEXT, spend DOUBLE PRECISION,
metadata JSONB, cache_hit TEXT
)
"""
@dataclass(frozen=True)
class _Case:
request_id: str
metadata: Mapping[str, object]
cache_hit: str | None = None
start_time: datetime = datetime(2026, 9, 1, 12, 0, 0, 123456)
def matches(self, filter: PromptCachingRequestFilter) -> bool:
if self.cache_hit is not None and self.cache_hit.lower() == "true":
return False
if not datetime(2026, 9, 1) <= self.start_time <= datetime(2026, 9, 2):
return False
usage: Final = self.metadata.get("usage_object")
normalized: Final = _JSON_OBJECT.validate_python(usage) if isinstance(usage, Mapping) else None
injected: Final = marks_gateway_injection(self.metadata, "dep-a")
reads: Final = extract_cache_read_tokens(normalized)
writes: Final = extract_cache_creation_tokens(normalized)
match filter:
case "injected":
return injected
case "hits":
return reads > 0
case "all":
return injected or reads > 0 or writes > 0
_CASES: Final = (
_Case("injected-empty", {_MARKER: ""}),
_Case("injected-deployment", {_MARKER: "dep-a"}),
_Case("wrong-deployment", {_MARKER: "dep-b"}),
_Case("legacy-read", {"usage_object": {"cache_read_input_tokens": 100}}),
_Case("nested-read", {"usage_object": {"prompt_tokens_details": {"cached_tokens": 100}}}),
_Case("write", {"usage_object": {"cache_creation_input_tokens": 100}}),
_Case("nested-write", {"usage_object": {"prompt_tokens_details": {"cache_write_tokens": 100}}}),
_Case("nested-creation", {"usage_object": {"prompt_tokens_details": {"cache_creation_tokens": 100}}}),
_Case(
"top-precedence",
{"usage_object": {"cache_read_input_tokens": -2, "prompt_tokens_details": {"cached_tokens": 100}}},
),
_Case(
"zero-fallback",
{"usage_object": {"cache_read_input_tokens": 0, "prompt_tokens_details": {"cached_tokens": 100}}},
),
_Case(
"fractional-precedence",
{"usage_object": {"cache_read_input_tokens": 0.5, "prompt_tokens_details": {"cached_tokens": 100}}},
),
_Case("malformed-number", {"usage_object": {"cache_read_input_tokens": "100"}}),
_Case("malformed-container", {"usage_object": [100]}),
_Case("boolean-number", {"usage_object": {"cache_read_input_tokens": True}}),
_Case("boolean-marker", {_MARKER: True}),
_Case("response-cache", {_MARKER: "", "usage_object": {"cache_read_input_tokens": 100}}, "True"),
_Case("outside-before", {_MARKER: ""}, start_time=datetime(2026, 8, 31, 23, 59, 59)),
_Case(
"outside-after", {"usage_object": {"cache_read_input_tokens": 100}}, start_time=datetime(2026, 9, 2, 0, 0, 1)
),
)
@pytest_asyncio.fixture(loop_scope="function")
async def _cache_prisma(
_cache_postgresql: psycopg.Connection[tuple[object, ...]],
) -> AsyncIterator[Prisma]:
info: Final = _cache_postgresql.info
database: Final = Prisma(datasource={
"url": f"postgresql://{info.user}@{info.host}:{info.port}/{info.dbname}?connection_limit=1",
})
await database.connect()
try:
yield database
finally:
await database.disconnect()
def _seed(connection: psycopg.Connection[tuple[object, ...]], cases: tuple[_Case, ...] = _CASES) -> None:
with connection.cursor() as cursor:
cursor.execute(_DDL)
cursor.executemany(
"""INSERT INTO "LiteLLM_SpendLogs"
VALUES (%s, %s, %s, %s, %s, %s, %s, %s::jsonb, %s)""",
tuple(
(
case.request_id,
case.start_time,
datetime(2026, 9, 1, 12, 0, 1),
_MODEL,
"dep-a",
"anthropic",
0.01,
json.dumps(dict(case.metadata)),
case.cache_hit,
)
for case in cases
),
)
connection.commit()
def _app(role: LitellmUserRoles | None) -> FastAPI:
@ -156,79 +27,6 @@ def _app(role: LitellmUserRoles | None) -> FastAPI:
return application
@pytest.mark.asyncio
@pytest.mark.parametrize("filter", ["all", "injected", "hits"])
@pytest.mark.parametrize("role", [LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY])
async def test_request_filters_match_accounting_and_paginate_before_projection(
_cache_postgresql: psycopg.Connection[tuple[object, ...]],
_cache_prisma: Prisma,
monkeypatch: pytest.MonkeyPatch,
filter: PromptCachingRequestFilter,
role: LitellmUserRoles,
) -> None:
from litellm.proxy import proxy_server
_seed(_cache_postgresql)
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=_cache_prisma))
monkeypatch.setattr(proxy_server, "llm_router", None)
expected: Final = tuple(sorted((case.request_id for case in _CASES if case.matches(filter)), reverse=True))
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=_app(role)), base_url="http://test") as client:
first: Final = await client.get(
_URL, params={"start_date": _START, "end_date": _END, "filter": filter, "page_size": 2}
)
assert first.status_code == 200
first_page: Final = PromptCachingRequestsResponse.model_validate_json(first.content)
assert tuple(row.request_id for row in first_page.requests) == expected[:2]
assert first_page.has_more is (len(expected) > 2)
assert (first_page.next_cursor is not None) is first_page.has_more
if first_page.next_cursor is not None:
assert first_page.next_cursor.request_id == expected[1]
assert first_page.next_cursor.start_time == first_page.requests[-1].start_time
next_response: Final = await client.get(
_URL, params={
"start_date": _START, "end_date": _END, "filter": filter, "page_size": 2,
"cursor_start_time": first_page.next_cursor.start_time.astimezone(
timezone(timedelta(hours=-7))
).isoformat(),
"cursor_request_id": first_page.next_cursor.request_id,
}
)
assert next_response.status_code == 200
next_page: Final = PromptCachingRequestsResponse.model_validate_json(next_response.content)
assert tuple(row.request_id for row in next_page.requests) == expected[2:4]
assert next_page.has_more is (len(expected) > 4)
assert (next_page.next_cursor is not None) is next_page.has_more
second: Final = await client.get(
_URL, params={"start_date": _START, "end_date": _END, "filter": filter, "page_size": 100}
)
assert second.status_code == 200
complete: Final = PromptCachingRequestsResponse.model_validate_json(second.content)
assert tuple(row.request_id for row in complete.requests) == expected
assert complete.has_more is False
assert complete.next_cursor is None
assert all(row.start_time.tzinfo == timezone.utc for row in complete.requests)
payload: Final = _JSON_OBJECT.validate_json(second.content)
assert set(payload) == {"requests", "page_size", "has_more", "next_cursor"}
serialized_rows: Final = _JSON_ROWS.validate_python(payload["requests"])
assert set(serialized_rows[0]) == {
"request_id",
"start_time",
"model",
"gateway_injected",
"cache_read_tokens",
"cache_creation_tokens",
"spend",
"net_savings",
}
by_id: Final = {row.request_id: row for row in complete.requests}
if filter == "all":
assert by_id["injected-empty"].gateway_injected is True
assert by_id["injected-empty"].net_savings is None
assert by_id["legacy-read"].gateway_injected is False
assert by_id["legacy-read"].net_savings is not None and by_id["legacy-read"].net_savings > 0
assert by_id["write"].net_savings is not None and by_id["write"].net_savings < 0
@pytest.mark.asyncio
@pytest.mark.parametrize("role", [None, LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY])
async def test_non_admin_is_denied_before_database_access(
@ -269,53 +67,3 @@ async def test_incomplete_cursor_is_rejected(
) as client:
response: Final = await client.get(_URL, params={"start_date": _START, "end_date": _END, **params})
assert response.status_code == 400
@pytest.mark.asyncio
@pytest.mark.parametrize("delete_before_cursor", [False, True])
async def test_cursor_keeps_remaining_requests_once_during_insertions_and_deletions(
_cache_postgresql: psycopg.Connection[tuple[object, ...]],
_cache_prisma: Prisma,
monkeypatch: pytest.MonkeyPatch,
delete_before_cursor: bool,
) -> None:
from litellm.proxy import proxy_server
cases: Final = (*_CASES, _Case(
"older-cache-read", {"usage_object": {"cache_read_input_tokens": 100}}, start_time=datetime(2026, 9, 1, 11),
))
_seed(_cache_postgresql, cases)
monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=_cache_prisma))
monkeypatch.setattr(proxy_server, "llm_router", None)
expected: Final = (*sorted((case.request_id for case in _CASES if case.matches("all")), reverse=True), "older-cache-read")
async with httpx.AsyncClient(
transport=httpx.ASGITransport(app=_app(LitellmUserRoles.PROXY_ADMIN)), base_url="http://test"
) as client:
first: Final = await client.get(_URL, params={"start_date": _START, "end_date": _END, "page_size": 2})
assert first.status_code == 200
first_page: Final = PromptCachingRequestsResponse.model_validate_json(first.content)
assert tuple(row.request_id for row in first_page.requests) == expected[:2]
assert first_page.next_cursor is not None
with _cache_postgresql.cursor() as cursor:
cursor.executemany(
"""INSERT INTO "LiteLLM_SpendLogs"
SELECT %s, %s, "endTime", model, model_id, custom_llm_provider, spend, metadata, cache_hit
FROM "LiteLLM_SpendLogs" WHERE request_id = %s""",
(
("newer-request", datetime(2026, 9, 1, 13), expected[0]),
("zz-higher-id", cases[0].start_time, expected[0]),
),
)
if delete_before_cursor:
cursor.execute('DELETE FROM "LiteLLM_SpendLogs" WHERE request_id = %s', (expected[0],))
_cache_postgresql.commit()
following: Final = await client.get(_URL, params={
"start_date": _START, "end_date": _END, "page_size": 100,
"cursor_start_time": first_page.next_cursor.start_time.isoformat(),
"cursor_request_id": first_page.next_cursor.request_id,
})
assert following.status_code == 200
following_page: Final = PromptCachingRequestsResponse.model_validate_json(following.content)
assert tuple(row.request_id for row in following_page.requests) == expected[2:]
assert following_page.has_more is False
assert following_page.next_cursor is None

View file

@ -1,27 +1,19 @@
import asyncio
import re
import time
from collections.abc import Mapping, Sequence
from collections.abc import Sequence
from datetime import datetime, timedelta
from pathlib import Path
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import litellm_proxy_extras
import psycopg
import pytest
from prisma.errors import PrismaError
from psycopg.rows import dict_row
from psycopg.types.json import Jsonb
from pytest_postgresql import factories
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import (
SPEND_LOG_KEY_METADATA_CACHE_TTL,
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL,
SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS,
SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE,
)
from litellm.proxy.spend_tracking.key_metadata_recovery import (
attach_user_details,
@ -606,304 +598,6 @@ async def test_recover_key_metadata_from_spend_logs_bounds_the_scan_with_a_state
)
_spend_logs_postgresql_proc: Final = factories.postgresql_proc()
_spend_logs_postgresql: Final = factories.postgresql("_spend_logs_postgresql_proc")
_SPEND_LOGS_DDL: Final = """
CREATE TABLE "LiteLLM_SpendLogs" (
request_id TEXT PRIMARY KEY,
api_key TEXT NOT NULL DEFAULT '',
"startTime" TIMESTAMP(3) NOT NULL,
"user" TEXT DEFAULT '',
team_id TEXT,
metadata JSONB DEFAULT '{}'
)
"""
_API_KEY_START_TIME_INDEX_MIGRATION: Final = (
Path(litellm_proxy_extras.__file__).parent
/ "migrations"
/ "20260823000000_add_spend_logs_api_key_starttime_index"
/ "migration.sql"
)
_SPEND_LOG_ROWS_READ_IN_THIS_TRANSACTION_SQL: Final = """
SELECT COALESCE(seq_tup_read, 0) + COALESCE(idx_tup_fetch, 0) AS rows_read
FROM pg_stat_xact_user_tables
WHERE relname = 'LiteLLM_SpendLogs'
"""
def _create_spend_logs_table(conn: psycopg.Connection) -> None:
conn.execute(_SPEND_LOGS_DDL) # pyright: ignore[reportArgumentType] # DDL literal
conn.execute(_API_KEY_START_TIME_INDEX_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # migration file
def _psycopg_prisma(conn: psycopg.Connection) -> MagicMock:
async def query_raw(sql: str, *params: object) -> list[dict[str, object]]:
with conn.cursor(row_factory=dict_row) as cur:
cur.execute(
re.sub(r"\$(\d+)", r"%(p\1)s", sql), # pyright: ignore[reportArgumentType] # proxy SQL is not a literal
{f"p{i}": v for i, v in enumerate(params, start=1)},
)
return cur.fetchall()
async def execute_raw(sql: str) -> int:
conn.execute(sql) # pyright: ignore[reportArgumentType] # proxy SQL is not a literal
return 0
mock_prisma: Final = MagicMock()
transaction: Final = MagicMock()
transaction.query_raw = AsyncMock(side_effect=query_raw)
transaction.execute_raw = AsyncMock(side_effect=execute_raw)
mock_prisma.db.tx.return_value.__aenter__.return_value = transaction
return mock_prisma
def _commit_and_vacuum(conn: psycopg.Connection) -> None:
conn.commit()
conn.set_autocommit(True)
conn.execute('VACUUM (ANALYZE) "LiteLLM_SpendLogs"')
conn.set_autocommit(False)
def _insert_nameless_spend_logs(conn: psycopg.Connection, digest: str, rows: int) -> None:
conn.execute(
"""
INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime")
SELECT %(digest)s || '-' || g, %(digest)s, %(start)s + g * interval '1 minute'
FROM generate_series(1, %(rows)s) g
""",
{"digest": digest, "start": datetime(2026, 9, 7), "rows": rows},
)
def _named_spend_log(
digest: str, logged_at: datetime, alias: str | None, user: str | None, team: str | None = None
) -> tuple[str, str, datetime, str, str | None, Jsonb]:
return (
f"{digest}-{logged_at.isoformat()}",
digest,
logged_at,
user or "",
team,
Jsonb({"user_api_key_alias": alias} if alias else {}),
)
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_names_a_key_by_its_oldest_and_newest_named_rows_in_the_window(
_spend_logs_postgresql: psycopg.Connection,
):
conn: Final = _spend_logs_postgresql
_create_spend_logs_table(conn)
unnamed_edges, owner_logged_late, reowned, outside_window, never_named = (
hash_token(f"cli-session-{name}") for name in ("edges", "late", "reowned", "window", "never")
)
with conn.cursor() as cur:
cur.executemany(
'INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime", "user", team_id, metadata)'
" VALUES (%s, %s, %s, %s, %s, %s)",
(
_named_spend_log(unnamed_edges, datetime(2026, 9, 7, 1), None, None),
_named_spend_log(unnamed_edges, datetime(2026, 9, 8), "cli-a", "alice", "team-a"),
_named_spend_log(unnamed_edges, datetime(2026, 9, 9), "cli-a", "alice", "team-a"),
_named_spend_log(unnamed_edges, datetime(2026, 9, 9, 23), None, None),
_named_spend_log(owner_logged_late, datetime(2026, 9, 7, 1), "cli-b", None),
_named_spend_log(owner_logged_late, datetime(2026, 9, 9), "cli-b", "bob"),
_named_spend_log(reowned, datetime(2026, 9, 7, 1), "cli-c", "carol"),
_named_spend_log(reowned, datetime(2026, 9, 9), "cli-c", "dave"),
_named_spend_log(outside_window, datetime(2026, 9, 6), "stale-alias", "erin"),
_named_spend_log(outside_window, datetime(2026, 9, 8), "cli-d", "erin"),
_named_spend_log(outside_window, datetime(2026, 9, 10), "later-alias", "erin"),
_named_spend_log(never_named, datetime(2026, 9, 8), None, None),
),
)
conn.commit()
result = await recover_key_metadata_from_spend_logs(
_psycopg_prisma(conn),
{unnamed_edges, owner_logged_late, reowned, outside_window, never_named},
(datetime(2026, 9, 7), datetime(2026, 9, 10)),
cache=InMemoryCache(),
)
assert dict(result) == {
unnamed_edges: {"key_alias": "cli-a", "team_id": "team-a", "user_id": "alice"},
owner_logged_late: {"key_alias": "cli-b", "team_id": None, "user_id": "bob"},
reowned: {"key_alias": "cli-c", "team_id": None, "user_id": None},
outside_window: {"key_alias": "cli-d", "team_id": None, "user_id": "erin"},
}
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_reads_two_rows_per_key_however_many_the_key_logged(
_spend_logs_postgresql: psycopg.Connection,
):
conn: Final = _spend_logs_postgresql
_create_spend_logs_table(conn)
owners: Final[Mapping[str, str]] = {hash_token(f"cli-session-busy-{i}"): f"user-{i}" for i in range(5)}
for digest, owner in owners.items():
conn.execute(
"""
INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime", "user", metadata)
SELECT %(digest)s || '-' || g, %(digest)s, %(start)s + g * interval '1 minute', %(owner)s,
jsonb_build_object('user_api_key_alias', 'cli-session-' || %(owner)s)
FROM generate_series(1, 2000) g
""",
{"digest": digest, "owner": owner, "start": datetime(2026, 9, 7)},
)
conn.execute('ANALYZE "LiteLLM_SpendLogs"')
conn.commit()
result = await recover_key_metadata_from_spend_logs(
_psycopg_prisma(conn), frozenset(owners), (datetime(2026, 9, 7), datetime(2026, 9, 10)), cache=InMemoryCache()
)
assert {digest: meta.get("user_id") for digest, meta in result.items()} == owners
rows_read: Final = conn.execute(_SPEND_LOG_ROWS_READ_IN_THIS_TRANSACTION_SQL).fetchone() # pyright: ignore[reportArgumentType] # SQL literal
assert rows_read is not None and rows_read[0] <= 2 * len(owners)
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_walks_a_bounded_number_of_nameless_rows_per_key(
_spend_logs_postgresql: psycopg.Connection,
):
conn: Final = _spend_logs_postgresql
_create_spend_logs_table(conn)
named_late: Final[Mapping[str, str]] = {hash_token(f"cli-session-late-{i}"): f"user-{i}" for i in range(3)}
never_named: Final = frozenset(hash_token(f"cli-session-never-{i}") for i in range(3))
for digest in (*named_late, *never_named):
_insert_nameless_spend_logs(conn, digest, 3 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE)
for digest, owner in named_late.items():
conn.execute(
"""
INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime", "user", metadata)
VALUES (%(digest)s || '-newest', %(digest)s, %(logged_at)s, %(owner)s,
jsonb_build_object('user_api_key_alias', 'cli-session-' || %(owner)s))
""",
{"digest": digest, "owner": owner, "logged_at": datetime(2026, 9, 9)},
)
conn.execute('ANALYZE "LiteLLM_SpendLogs"')
conn.commit()
result = await recover_key_metadata_from_spend_logs(
_psycopg_prisma(conn),
frozenset(named_late) | never_named,
(datetime(2026, 9, 7), datetime(2026, 9, 10)),
cache=InMemoryCache(),
)
assert {digest: meta.get("user_id") for digest, meta in result.items()} == named_late
rows_read: Final = conn.execute(_SPEND_LOG_ROWS_READ_IN_THIS_TRANSACTION_SQL).fetchone() # pyright: ignore[reportArgumentType] # SQL literal
assert rows_read is not None
assert rows_read[0] <= 3 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE * (len(named_late) + len(never_named))
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_reads_a_short_nameless_key_once(
_spend_logs_postgresql: psycopg.Connection,
):
conn: Final = _spend_logs_postgresql
_create_spend_logs_table(conn)
rows_per_key: Final = SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE // 2
never_named: Final = frozenset(hash_token(f"cli-session-short-{i}") for i in range(20))
for digest in never_named:
_insert_nameless_spend_logs(conn, digest, rows_per_key)
_commit_and_vacuum(conn)
result = await recover_key_metadata_from_spend_logs(
_psycopg_prisma(conn), never_named, (datetime(2026, 9, 7), datetime(2026, 9, 10)), cache=InMemoryCache()
)
assert dict(result) == {}
rows_read: Final = conn.execute(_SPEND_LOG_ROWS_READ_IN_THIS_TRANSACTION_SQL).fetchone() # pyright: ignore[reportArgumentType] # SQL literal
assert rows_read is not None
assert rows_read[0] <= rows_per_key * len(never_named)
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_bounds_a_busy_nameless_key_among_short_keys_before_any_vacuum(
_spend_logs_postgresql: psycopg.Connection,
):
conn: Final = _spend_logs_postgresql
_create_spend_logs_table(conn)
busy: Final = frozenset(hash_token(f"cli-session-busy-nameless-{i}") for i in range(3))
for short_key in range(200):
_insert_nameless_spend_logs(
conn, hash_token(f"cli-session-short-{short_key}"), SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE // 5
)
for digest in busy:
_insert_nameless_spend_logs(conn, digest, 30 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE)
conn.execute('ANALYZE "LiteLLM_SpendLogs"')
conn.commit()
result = await recover_key_metadata_from_spend_logs(
_psycopg_prisma(conn), busy, (datetime(2026, 9, 7), datetime(2026, 9, 10)), cache=InMemoryCache()
)
assert dict(result) == {}
rows_read: Final = conn.execute(_SPEND_LOG_ROWS_READ_IN_THIS_TRANSACTION_SQL).fetchone() # pyright: ignore[reportArgumentType] # SQL literal
assert rows_read is not None
assert rows_read[0] <= 3 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE * len(busy)
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_finds_a_name_logged_where_the_oldest_probe_stopped(
_spend_logs_postgresql: psycopg.Connection,
):
conn: Final = _spend_logs_postgresql
_create_spend_logs_table(conn)
start: Final = datetime(2026, 9, 7)
past_the_stop, tied_with_the_stop = (hash_token(f"cli-session-{name}") for name in ("past", "tied"))
same_millisecond: Final = tuple(
start + timedelta(minutes=SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE, microseconds=n) for n in (100, 200, 300)
)
with conn.cursor() as cur:
cur.executemany(
'INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime", "user", team_id, metadata)'
" VALUES (%s, %s, %s, %s, %s, %s)",
(
*(
_named_spend_log(past_the_stop, start + timedelta(minutes=minute), None, None)
for minute in range(1, SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 20)
),
_named_spend_log(
past_the_stop,
start + timedelta(minutes=SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 20),
"cli-p",
"pat",
),
*(
_named_spend_log(past_the_stop, start + timedelta(minutes=minute), None, None)
for minute in range(
SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 21, SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 51
)
),
*(
_named_spend_log(tied_with_the_stop, start + timedelta(minutes=minute), None, None)
for minute in range(1, SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE)
),
_named_spend_log(tied_with_the_stop, same_millisecond[0], None, None),
_named_spend_log(tied_with_the_stop, same_millisecond[1], None, None),
_named_spend_log(tied_with_the_stop, same_millisecond[2], "cli-t", "tess"),
),
)
conn.commit()
result = await recover_key_metadata_from_spend_logs(
_psycopg_prisma(conn),
{past_the_stop, tied_with_the_stop},
(start, datetime(2026, 9, 10)),
cache=InMemoryCache(),
)
assert dict(result) == {
past_the_stop: {"key_alias": "cli-p", "team_id": None, "user_id": "pat"},
tied_with_the_stop: {"key_alias": "cli-t", "team_id": None, "user_id": "tess"},
}
@pytest.mark.asyncio
async def test_recover_cli_session_key_metadata_names_the_owner_only_when_the_suffix_is_a_real_user():
mock_prisma = MagicMock()

View file

@ -1,16 +1,12 @@
import json
import re
from collections.abc import Mapping
from datetime import date, datetime, timezone
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import httpx
import psycopg
import pytest
from psycopg.rows import dict_row
from pydantic import ValidationError
from pytest_postgresql import factories
from litellm.constants import (
SPEND_CAPTURE_RATE_CHECK_JOB_ID,
@ -22,7 +18,6 @@ from litellm.proxy.spend_tracking.spend_capture_rate import (
ProviderBillingCredentialMissing,
ProviderBillingRequestFailed,
alert_message,
captured_spend_by_day,
compute_capture_rate,
run_scheduled_spend_capture_rate_check,
run_spend_capture_rate_check,
@ -375,65 +370,3 @@ def test_settings_reject_typos_and_out_of_range_values():
json.loads('{"providers": ["openai"], "threshold": 0.8, "lookback_days": 3, "openai_project_ids": ["p"]}')
)
assert (parsed.threshold, parsed.lookback_days, parsed.openai_project_ids) == (0.8, 3, ("p",))
_capture_postgresql_proc: Final = factories.postgresql_proc()
_capture_postgresql: Final = factories.postgresql("_capture_postgresql_proc")
_DAILY_USER_SPEND_DDL: Final = """
CREATE TABLE "LiteLLM_DailyUserSpend" (
id TEXT PRIMARY KEY,
date TEXT NOT NULL,
custom_llm_provider TEXT,
spend DOUBLE PRECISION DEFAULT 0
)
"""
class _PsycopgPrisma:
"""``prisma_client.db.query_raw`` on a real connection, with ``$n`` placeholders converted for psycopg."""
def __init__(self, conn: psycopg.Connection) -> None:
self.db = self
self._conn = conn
async def query_raw(self, sql: str, *params: object) -> list[dict[str, object]]:
converted: Final = re.sub(r"\$(\d+)", r"%(p\1)s", sql)
with self._conn.cursor(row_factory=dict_row) as cur:
cur.execute(
converted, # pyright: ignore[reportArgumentType] # psycopg stubs want a literal-typed query
{f"p{i}": list(v) if isinstance(v, tuple) else v for i, v in enumerate(params, start=1)},
)
return cur.fetchall()
@pytest.mark.asyncio
async def test_captured_spend_sums_only_the_openai_billed_providers_inside_the_window(
_capture_postgresql: psycopg.Connection,
):
conn: Final = _capture_postgresql
conn.execute(_DAILY_USER_SPEND_DDL) # pyright: ignore[reportArgumentType] # DDL literal
rows: Final = (
("2026-09-19", "openai", 1.0),
("2026-09-20", "openai", 2.0),
("2026-09-20", "openai", 3.0),
("2026-09-20", "text-completion-openai", 0.5),
("2026-09-20", "anthropic", 100.0),
("2026-09-21", "azure", 100.0),
("2026-09-22", "openai", 4.0),
)
for index, (day, provider, spend) in enumerate(rows):
conn.execute(
'INSERT INTO "LiteLLM_DailyUserSpend" (id, date, custom_llm_provider, spend) VALUES (%s, %s, %s, %s)',
(f"row-{index}", day, provider, spend),
)
conn.commit()
captured = await captured_spend_by_day(
_PsycopgPrisma(conn), # pyright: ignore[reportArgumentType] # duck-typed prisma for the raw query
litellm_providers=("openai", "text-completion-openai"),
start_date=date(2026, 9, 20),
end_date=date(2026, 9, 21),
)
assert dict(captured) == {"2026-09-20": 5.5}