mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
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:
parent
d96477abce
commit
6f123b7083
14 changed files with 1124 additions and 1205 deletions
|
|
@ -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()
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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}
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue