fix(proxy): recover session key owners from daily spend for usage attribution (#43642)

* fix(proxy): recover daily spend key owners

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

* fix(proxy): simplify daily spend owner recovery

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

* style(proxy): format daily activity metadata

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

* test(proxy): cover recovered owner metadata merge

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

* fix(proxy): bound the daily spend owner lookup with the statement timeout

* test(integration): audit the daily activity key owner fallback on every usage route

Thirty five integration cells under tests/integration/spend cover the daily
spend owner fallback on all nine daily activity routes and /usage/ai/chat:
the happy path per route, the unanimity rules (two users, blank and null
rows, an owner the user table lacks, live and deleted keys with and without
their own user, a spend log alias), a non admin reader, an invalid key, a 5 KB
key, a locked LiteLLM_DailyUserSpend, 300 keys of one team, repeated reads, a
second user landing between reads, a concurrent burst across the unified
endpoints, a killed worker, and a proxy restart

The traffic cells ignore the GET /v1/models call the proxy's five minute token
limit refresh makes to every registered OpenAI compatible deployment, since it
lands on a test's provider wire whenever the refresh instant falls inside the
test

---------

Co-authored-by: jesus <jesus@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-29 15:36:09 -07:00 • committed by GitHub
parent 56a63b4b29
commit f5a1c9f1f1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
11 changed files with 1426 additions and 37 deletions

View file

@ -16,6 +16,7 @@ from litellm.proxy.spend_tracking.key_metadata_recovery import (
recover_cli_session_key_metadata,
recover_double_hashed_key_metadata,
recover_key_metadata_from_spend_logs,
recover_key_owner_from_daily_spend,
)
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
from litellm.proxy.utils import PrismaClient
@ -468,6 +469,17 @@ def _parse_spend_date(raw: str | None) -> datetime | None:
_EMPTY_KEY_METADATA: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType({})
def _metadata_with_recovered_owner(
metadata: Mapping[str, _KeyMetadataDict],
key: str,
owner: str,
) -> _KeyMetadataDict:
current: Final = metadata.get(key)
if current is None:
return {"user_id": owner}
return {**current, "user_id": owner}
async def get_api_key_metadata(
prisma_client: PrismaClient,
api_keys: AbstractSet[str],
@ -530,7 +542,19 @@ async def get_api_key_metadata(
else _EMPTY_KEY_METADATA
)
combined: Final = MappingProxyType({**after_token_recovery, **from_spend_logs})
return await attach_user_details(prisma_client, combined)
ownerless: Final = frozenset(
key
for key in api_keys
if not combined.get(key, {}).get("user_id") and not combined.get(key, {}).get("key_exists")
)
owners: Final = await recover_key_owner_from_daily_spend(prisma_client, ownerless)
metadata_with_owners: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType(
{
**combined,
**{key: _metadata_with_recovered_owner(combined, key, owner) for key, owner in owners.items()},
}
)
return await attach_user_details(prisma_client, metadata_with_owners)
def _adjust_dates_for_timezone(
@ -944,7 +968,7 @@ async def _aggregate_spend_records(
record.api_key for record in records if record.api_key and record.api_key != PTU_SENTINEL_API_KEY
}
api_key_metadata: dict[str, _KeyMetadataDict] = {}
api_key_metadata: Mapping[str, _KeyMetadataDict] = MappingProxyType({})
if api_keys:
api_key_metadata = await get_api_key_metadata(
prisma_client, api_keys, _spend_logs_window(frozenset(record.date for record in records))
@ -1144,7 +1168,7 @@ async def _aggregate_grouping_sets_records(
"""Async wrapper: fetch api_key_metadata, then dispatch on a worker thread."""
api_keys: Final[set[str]] = {r.api_key for r in records if r.api_key and r.api_key != PTU_SENTINEL_API_KEY}
api_key_metadata: dict[str, _KeyMetadataDict] = {}
api_key_metadata: Mapping[str, _KeyMetadataDict] = MappingProxyType({})
if api_keys:
api_key_metadata = await get_api_key_metadata(
prisma_client, api_keys, _spend_logs_window(frozenset(r.date for r in records))

View file

@ -61,6 +61,13 @@ WHERE COALESCE(key_alias, user_id, team_id) IS NOT NULL
GROUP BY api_key
"""
_DAILY_USER_SPEND_OWNER_SQL: Final = """
SELECT api_key, MIN(user_id) AS first_owner, MAX(user_id) AS last_owner
FROM "LiteLLM_DailyUserSpend"
WHERE api_key = ANY($1::text[]) AND user_id IS NOT NULL AND user_id <> ''
GROUP BY api_key
"""
_SPEND_LOG_STATEMENT_TIMEOUT_SQL: Final = f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}"
_SPEND_LOG_TRANSACTION_TIMEOUT: Final = timedelta(milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS)
@ -104,8 +111,15 @@ class _SpendLogDigestRow(BaseModel):
)
class _DailyUserSpendOwnerRow(BaseModel):
api_key: str
first_owner: str | None = None
last_owner: str | None = None
_TOKEN_DIGEST_ROWS: Final = TypeAdapter(tuple[_TokenDigestRow, ...])
_SPEND_LOG_DIGEST_ROWS: Final = TypeAdapter(tuple[_SpendLogDigestRow, ...])
_DAILY_USER_SPEND_OWNER_ROWS: Final = TypeAdapter(tuple[_DailyUserSpendOwnerRow, ...])
_CACHED_KEY_METADATA: Final = TypeAdapter(KeyMetadataDict)
_SPEND_LOG_METADATA_CACHE: Final = InMemoryCache(
max_size_in_memory=SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS,
@ -113,6 +127,7 @@ _SPEND_LOG_METADATA_CACHE: Final = InMemoryCache(
)
_SPEND_LOG_QUERY_LOCK: Final = asyncio.Lock()
_EMPTY_KEY_METADATA: Final[Mapping[str, KeyMetadataDict]] = MappingProxyType({})
_EMPTY_KEY_OWNERS: Final[Mapping[str, str]] = MappingProxyType({})
async def _db_or_empty(
@ -129,6 +144,16 @@ async def _db_or_empty(
return None
async def _rows_within_the_statement_timeout(
prisma_client: PrismaClient,
sql: str,
*params: object,
) -> Sequence[Mapping[str, object]]:
async with prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction:
await transaction.execute_raw(_SPEND_LOG_STATEMENT_TIMEOUT_SQL)
return await transaction.query_raw(sql, *params)
async def _reverse_hash_key_metadata(
prisma_client: PrismaClient,
sql: str,
@ -152,6 +177,29 @@ async def _reverse_hash_key_metadata(
)
async def recover_key_owner_from_daily_spend(
prisma_client: PrismaClient,
keys: AbstractSet[str],
) -> Mapping[str, str]:
if not keys:
return _EMPTY_KEY_OWNERS
rows: Final = await _db_or_empty(
lambda: _rows_within_the_statement_timeout(prisma_client, _DAILY_USER_SPEND_OWNER_SQL, sorted(keys)),
"Failed daily-spend key owner recovery for %d keys: %s",
len(keys),
)
if rows is None:
return _EMPTY_KEY_OWNERS
return MappingProxyType(
{
row.api_key: owner
for row in _DAILY_USER_SPEND_OWNER_ROWS.validate_python(rows)
for owner in (_unanimous(row.first_owner, row.last_owner),)
if row.api_key in keys and owner is not None
}
)
@dataclass(frozen=True, slots=True)
class _UserDetails:
email: str | None
@ -309,24 +357,14 @@ def _cached_spend_log_metadata(
)
async def _spend_log_rows_within_the_statement_timeout(
prisma_client: PrismaClient,
digests: AbstractSet[str],
window: tuple[datetime, datetime],
) -> Sequence[Mapping[str, object]]:
start, end = window
async with prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction:
await transaction.execute_raw(_SPEND_LOG_STATEMENT_TIMEOUT_SQL)
return await transaction.query_raw(_SPEND_LOG_ALIAS_SQL, sorted(digests), start, end)
async def _query_spend_log_metadata(
prisma_client: PrismaClient,
digests: AbstractSet[str],
window: tuple[datetime, datetime],
) -> Mapping[str, KeyMetadataDict] | None:
start, end = window
rows: Final = await _db_or_empty(
lambda: _spend_log_rows_within_the_statement_timeout(prisma_client, digests, window),
lambda: _rows_within_the_statement_timeout(prisma_client, _SPEND_LOG_ALIAS_SQL, sorted(digests), start, end),
"Failed spend-log alias recovery for %d missing keys: %s",
len(digests),
)

View file

@ -0,0 +1,237 @@
import os
import uuid
from collections.abc import Iterator, Mapping, Sequence
from contextlib import contextmanager
from dataclasses import dataclass
from itertools import chain
from typing import Final
import httpx
import psycopg
import pytest
from integration._support.client import Gateway, Scenario, object_value
from psycopg import sql
from psycopg.types.json import Jsonb
from pydantic import JsonValue
USER_SPEND: Final = "LiteLLM_DailyUserSpend"
TEAM_SPEND: Final = "LiteLLM_DailyTeamSpend"
TAG_SPEND: Final = "LiteLLM_DailyTagSpend"
ORGANIZATION_SPEND: Final = "LiteLLM_DailyOrganizationSpend"
END_USER_SPEND: Final = "LiteLLM_DailyEndUserSpend"
AGENT_SPEND: Final = "LiteLLM_DailyAgentSpend"
DAY: Final = "2026-02-03"
AGGREGATED_USER_ACTIVITY: Final = "/user/daily/activity/aggregated"
INSERT_DAILY_ROW: Final = sql.SQL(
"INSERT INTO {table} (id, {entity}, date, api_key, model, model_group, custom_llm_provider, prompt_tokens,"
" completion_tokens, spend, api_requests, successful_requests, failed_requests, updated_at)"
" VALUES (gen_random_uuid()::text, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, now())"
)
DELETE_DAILY_ROWS: Final = sql.SQL("DELETE FROM {table} WHERE api_key = ANY(%s)")
INSERT_SPEND_LOG: Final = (
'INSERT INTO "LiteLLM_SpendLogs" (request_id, call_type, api_key, "startTime", "endTime", metadata)'
" VALUES (%s, 'acompletion', %s, %s::timestamp, %s::timestamp, %s)"
)
DELETE_SPEND_LOG: Final = 'DELETE FROM "LiteLLM_SpendLogs" WHERE request_id = %s'
LOCK_TABLE: Final = sql.SQL("LOCK TABLE {table} IN ACCESS EXCLUSIVE MODE")
@dataclass(frozen=True, slots=True)
class Route:
path: str
table: str
entity_column: str
entity_filter: str | None
ROUTES: Final = (
Route("/user/daily/activity", USER_SPEND, "user_id", None),
Route(AGGREGATED_USER_ACTIVITY, USER_SPEND, "user_id", None),
Route("/team/daily/activity", TEAM_SPEND, "team_id", "team_ids"),
Route("/team/daily/activity/aggregated", TEAM_SPEND, "team_id", "team_ids"),
Route("/tag/daily/activity", TAG_SPEND, "tag", "tags"),
Route("/organization/daily/activity", ORGANIZATION_SPEND, "organization_id", "organization_ids"),
Route("/customer/daily/activity", END_USER_SPEND, "end_user_id", "end_user_ids"),
Route("/end_user/daily/activity", END_USER_SPEND, "end_user_id", "end_user_ids"),
Route("/agent/daily/activity", AGENT_SPEND, "agent_id", "agent_ids"),
)
def user_with_an_email(scenario: Scenario) -> tuple[str, str]:
email: Final = f"integration-{uuid.uuid4().hex}@example.com"
return scenario.user(user_email=email), email
def key_no_key_table_holds() -> str:
return f"integration-ownerless-{uuid.uuid4().hex}"
def activity_of_key(
gateway: Gateway, path: str, api_key: str, *, reader: str | None = None, **filters: str
) -> httpx.Response:
return gateway.request(
"GET", path, params={"start_date": DAY, "end_date": DAY, "api_key": api_key, **filters}, key=reader
)
@dataclass(frozen=True, slots=True)
class DailyRow:
table: str
entity_column: str
entity: str | None
api_key: str
date: str
model: str
provider: str
prompt_tokens: int
completion_tokens: int
spend: float
successful_requests: int
failed_requests: int
def _insert(connection: psycopg.Connection[tuple[object, ...]], row: DailyRow) -> None:
connection.execute(
INSERT_DAILY_ROW.format(table=sql.Identifier(row.table), entity=sql.Identifier(row.entity_column)),
(
row.entity,
row.date,
row.api_key,
row.model,
row.model,
row.provider,
row.prompt_tokens,
row.completion_tokens,
row.spend,
row.successful_requests + row.failed_requests,
row.successful_requests,
row.failed_requests,
),
)
def insert_daily_rows(rows: Sequence[DailyRow], *, database_url: str | None = None) -> None:
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:
for row in rows:
_insert(connection, row)
def delete_daily_rows(rows: Sequence[DailyRow], *, database_url: str | None = None) -> None:
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:
for table in sorted({row.table for row in rows}):
connection.execute(
DELETE_DAILY_ROWS.format(table=sql.Identifier(table)),
(sorted({row.api_key for row in rows if row.table == table}),),
)
@contextmanager
def daily_rows(rows: Sequence[DailyRow], *, database_url: str | None = None) -> Iterator[None]:
insert_daily_rows(rows, database_url=database_url)
try:
yield
finally:
delete_daily_rows(rows, database_url=database_url)
@contextmanager
def spend_log_naming_only_an_alias(request_id: str, api_key: str, started: str, alias: str) -> Iterator[None]:
with psycopg.connect(os.environ["DATABASE_URL"]) as connection:
connection.execute(
INSERT_SPEND_LOG, (request_id, api_key, started, started, Jsonb({"user_api_key_alias": alias}))
)
try:
yield
finally:
with psycopg.connect(os.environ["DATABASE_URL"]) as connection:
connection.execute(DELETE_SPEND_LOG, (request_id,))
@contextmanager
def locked_table(table: str, *, database_url: str | None = None) -> Iterator[None]:
with psycopg.connect(database_url or os.environ["DATABASE_URL"]) as connection:
connection.execute(LOCK_TABLE.format(table=sql.Identifier(table)))
try:
yield
finally:
connection.rollback()
def records_of_key(node: JsonValue, api_key: str) -> tuple[JsonValue, ...]:
if isinstance(node, list):
return tuple(chain.from_iterable(records_of_key(item, api_key) for item in node))
if not isinstance(node, dict):
return ()
nested: Final = tuple(chain.from_iterable(records_of_key(value, api_key) for value in node.values()))
return (node[api_key], *nested) if api_key in node else nested
def seeded_row(table: str, entity_column: str, entity: str | None, api_key: str, date: str) -> DailyRow:
return DailyRow(table, entity_column, entity, api_key, date, "gpt-4o-mini", "openai", 10, 5, 0.25, 1, 0)
def user_row(user: str | None, api_key: str, date: str) -> DailyRow:
return seeded_row(USER_SPEND, "user_id", user, api_key, date)
def seeded_metrics(rows: int) -> dict[str, float]:
return {
"spend": 0.25 * rows,
"prompt_tokens": 10 * rows,
"completion_tokens": 5 * rows,
"total_tokens": 15 * rows,
"api_requests": rows,
"successful_requests": rows,
}
def key_metadata(
*,
alias: str | None = None,
team: str | None = None,
user: str | None = None,
email: str | None = None,
exists: bool = False,
) -> dict[str, JsonValue]:
return {"key_alias": alias, "team_id": team, "user_id": user, "user_email": email, "key_exists": exists}
def counted(metrics: JsonValue) -> dict[str, JsonValue]:
return {name: value for name, value in object_value(metrics).items() if value}
def assert_key_reported(
response: httpx.Response,
api_key: str,
date: str,
metadata: Mapping[str, JsonValue],
metrics: Mapping[str, float],
) -> None:
assert response.status_code == 200, response.text
body: Final = object_value(response.json())
records: Final = tuple(object_value(record) for record in records_of_key(body, api_key))
assert records, response.text
assert all(record["metadata"] == metadata for record in records), response.text
assert all(counted(record["metrics"]) == pytest.approx(metrics) for record in records), response.text
days: Final = body["results"]
assert isinstance(days, list) and len(days) == 1, response.text
day: Final = object_value(days[0])
assert day["date"] == date, response.text
assert counted(day["metrics"]) == pytest.approx(metrics), response.text
assert object_value(body["metadata"])["total_spend"] == pytest.approx(metrics["spend"]), response.text
def assert_key_owner_and_totals(
response: httpx.Response,
api_key: str,
metadata: Mapping[str, JsonValue],
totals: Mapping[str, float],
) -> None:
assert response.status_code == 200, response.text
body: Final = object_value(response.json())
records: Final = tuple(object_value(record) for record in records_of_key(body, api_key))
assert records, response.text
assert all(record["metadata"] == metadata for record in records), response.text
reported: Final = object_value(body["metadata"])
assert {name: reported[name] for name in totals} == pytest.approx(totals), response.text

View file

@ -0,0 +1,196 @@
import uuid
from hashlib import sha256
from typing import Final
import pytest
from integration._support.client import Gateway, Scenario, string_value
from integration._support.daily_activity import (
AGGREGATED_USER_ACTIVITY,
DAY,
ROUTES,
USER_SPEND,
Route,
activity_of_key,
assert_key_reported,
daily_rows,
key_metadata,
key_no_key_table_holds,
seeded_metrics,
seeded_row,
spend_log_naming_only_an_alias,
user_row,
user_with_an_email,
)
@pytest.mark.parametrize("route", ROUTES, ids=lambda route: route.path.strip("/").replace("/", "_"))
def test_key_missing_from_the_key_tables_is_reported_with_the_one_user_its_daily_spend_names(
gateway: Gateway, route: Route
) -> None:
api_key: Final = key_no_key_table_holds()
entity: Final = f"integration-entity-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
entity_rows: Final = (
() if route.table == USER_SPEND else (seeded_row(route.table, route.entity_column, entity, api_key, DAY),)
)
filters: Final = {} if route.entity_filter is None else {route.entity_filter: entity}
with daily_rows((user_row(owner, api_key, DAY), *entity_rows)):
assert_key_reported(
activity_of_key(gateway, route.path, api_key, **filters),
api_key,
DAY,
key_metadata(user=owner, email=email),
seeded_metrics(1),
)
def test_key_whose_daily_spend_names_two_users_is_reported_with_no_owner(gateway: Gateway) -> None:
api_key: Final = key_no_key_table_holds()
with gateway.scenario() as scenario:
first, _ = user_with_an_email(scenario)
second, _ = user_with_an_email(scenario)
with daily_rows((user_row(first, api_key, DAY), user_row(second, api_key, DAY))):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(),
seeded_metrics(2),
)
@pytest.mark.parametrize("unnamed", ["", None], ids=["blank_user", "null_user"])
def test_daily_spend_rows_naming_no_user_do_not_hide_the_one_user_the_others_name(
gateway: Gateway, unnamed: str | None
) -> None:
api_key: Final = key_no_key_table_holds()
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with daily_rows((user_row(owner, api_key, DAY), user_row(unnamed, api_key, DAY))):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(user=owner, email=email),
seeded_metrics(2),
)
def test_key_whose_daily_spend_names_no_user_at_all_is_reported_with_no_owner(gateway: Gateway) -> None:
api_key: Final = key_no_key_table_holds()
with daily_rows((user_row("", api_key, DAY), user_row(None, api_key, DAY))):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(),
seeded_metrics(2),
)
def test_owner_the_user_table_does_not_hold_is_reported_by_id_with_no_email(gateway: Gateway) -> None:
api_key: Final = key_no_key_table_holds()
owner: Final = f"integration-departed-{uuid.uuid4().hex}"
with daily_rows((user_row(owner, api_key, DAY),)):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(user=owner),
seeded_metrics(1),
)
def _stored_form(token: str) -> str:
return sha256(token.encode()).hexdigest()
def _deleted_key(gateway: Gateway, scenario: Scenario, alias: str, **fields: str) -> str:
token: Final = string_value(gateway.post("/key/generate", {"key_alias": alias, **fields})["key"])
scenario.delete_key(token)
return _stored_form(token)
def test_live_key_keeps_its_own_user_when_its_daily_spend_names_another(gateway: Gateway) -> None:
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
other, _ = user_with_an_email(scenario)
api_key: Final = _stored_form(scenario.key(user_id=owner, key_alias=alias))
with daily_rows((user_row(other, api_key, DAY),)):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(alias=alias, user=owner, email=email, exists=True),
seeded_metrics(1),
)
def test_live_key_with_no_user_is_not_given_the_user_its_daily_spend_names(gateway: Gateway) -> None:
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
spender, _ = user_with_an_email(scenario)
api_key: Final = _stored_form(scenario.key(key_alias=alias))
with daily_rows((user_row(spender, api_key, DAY),)):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(alias=alias, exists=True),
seeded_metrics(1),
)
def test_deleted_key_keeps_its_own_user_when_its_daily_spend_names_another(gateway: Gateway) -> None:
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
other, _ = user_with_an_email(scenario)
api_key: Final = _deleted_key(gateway, scenario, alias, user_id=owner)
with daily_rows((user_row(other, api_key, DAY),)):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(alias=alias, user=owner, email=email),
seeded_metrics(1),
)
def test_deleted_key_with_no_user_keeps_its_alias_and_gains_the_one_user_its_daily_spend_names(
gateway: Gateway,
) -> None:
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
api_key: Final = _deleted_key(gateway, scenario, alias)
with daily_rows((user_row(owner, api_key, DAY),)):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(alias=alias, user=owner, email=email),
seeded_metrics(1),
)
def test_key_named_only_by_a_spend_log_alias_keeps_that_alias_and_gains_the_one_user_its_daily_spend_names(
gateway: Gateway,
) -> None:
api_key: Final = sha256(uuid.uuid4().bytes).hexdigest()
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with (
spend_log_naming_only_an_alias(f"integration-{uuid.uuid4().hex}", api_key, f"{DAY} 12:00:00", alias),
daily_rows((user_row(owner, api_key, DAY),)),
):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(alias=alias, user=owner, email=email),
seeded_metrics(1),
)

View file

@ -0,0 +1,264 @@
import os
import signal
import time
import uuid
from collections.abc import Iterator
from contextlib import contextmanager
from itertools import chain
from pathlib import Path
from typing import Final
import httpx
import psutil
import pytest
from integration._support.client import Gateway, eventually, object_value
from integration._support.daily_activity import (
AGGREGATED_USER_ACTIVITY,
DAY,
TEAM_SPEND,
USER_SPEND,
activity_of_key,
assert_key_reported,
daily_rows,
insert_daily_rows,
key_metadata,
key_no_key_table_holds,
locked_table,
records_of_key,
seeded_metrics,
seeded_row,
user_row,
user_with_an_email,
)
from integration._support.database import scratch_database
from integration._support.process import OwnedProxy, group_members, owned_proxy_process
USER_ACTIVITY: Final = "/user/daily/activity"
TEAM_ACTIVITY: Final = "/team/daily/activity"
AGGREGATED_TEAM_ACTIVITY: Final = "/team/daily/activity/aggregated"
KEYS_OF_ONE_TEAM: Final = 300
GIVES_UP_WITHIN_SECONDS: Final = 10
READS_AFTER_THE_WORKER_IS_REPLACED: Final = 6
@contextmanager
def _proxy_on(gateway: Gateway, directory: Path, database_url: str, *, workers: int = 1) -> Iterator[OwnedProxy]:
with owned_proxy_process(
gateway,
directory,
{"DATABASE_URL": database_url},
remove_environment=("DATABASE_URL_READ_REPLICA",),
workers=workers,
) as owned:
yield owned
def _owner_on(candidate: Gateway) -> tuple[str, str]:
owner: Final = f"integration-{uuid.uuid4().hex}"
email: Final = f"{owner}@example.com"
candidate.post("/user/new", {"user_id": owner, "user_email": email, "auto_create_key": False})
return owner, email
def _read_on_a_new_connection(candidate: Gateway, api_key: str) -> httpx.Response:
return candidate.request(
"GET",
AGGREGATED_USER_ACTIVITY,
params={"start_date": DAY, "end_date": DAY, "api_key": api_key},
headers={"Connection": "close"},
)
def _running_children(owned: OwnedProxy) -> tuple[int, ...]:
return tuple(
member.pid
for member in group_members(owned.process.pid)
if member.pid != owned.process.pid and member.is_running() and member.status() != psutil.STATUS_ZOMBIE
)
def test_user_reading_a_key_shared_with_another_user_is_shown_no_owner_and_nothing_of_the_other_user(
gateway: Gateway,
) -> None:
api_key: Final = key_no_key_table_holds()
with gateway.scenario() as scenario:
reader, _ = user_with_an_email(scenario)
other, other_email = user_with_an_email(scenario)
reader_key: Final = scenario.key(user_id=reader)
with daily_rows((user_row(reader, api_key, DAY), user_row(other, api_key, DAY))):
response: Final = activity_of_key(gateway, USER_ACTIVITY, api_key, reader=reader_key)
assert_key_reported(response, api_key, DAY, key_metadata(), seeded_metrics(1))
assert other not in response.text
assert other_email not in response.text
def test_user_reading_a_key_only_they_spent_with_is_shown_themselves_as_its_owner(gateway: Gateway) -> None:
api_key: Final = key_no_key_table_holds()
with gateway.scenario() as scenario:
reader, email = user_with_an_email(scenario)
reader_key: Final = scenario.key(user_id=reader)
with daily_rows((user_row(reader, api_key, DAY),)):
response: Final = activity_of_key(gateway, USER_ACTIVITY, api_key, reader=reader_key)
assert_key_reported(response, api_key, DAY, key_metadata(user=reader, email=email), seeded_metrics(1))
def test_user_reading_a_key_only_another_user_spent_with_is_shown_nothing_of_it(gateway: Gateway) -> None:
api_key: Final = key_no_key_table_holds()
with gateway.scenario() as scenario:
reader, _ = user_with_an_email(scenario)
other, other_email = user_with_an_email(scenario)
reader_key: Final = scenario.key(user_id=reader)
with daily_rows((user_row(other, api_key, DAY),)):
response: Final = activity_of_key(gateway, USER_ACTIVITY, api_key, reader=reader_key)
assert response.status_code == 200, response.text
assert object_value(response.json())["results"] == [], response.text
assert other not in response.text
assert other_email not in response.text
def test_invalid_key_is_refused_without_naming_the_owner(gateway: Gateway) -> None:
api_key: Final = key_no_key_table_holds()
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with daily_rows((user_row(owner, api_key, DAY),)):
response: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key, reader="sk-not-a-key")
assert response.status_code == 401, response.text
assert owner not in response.text
assert email not in response.text
def test_five_kilobyte_key_is_reported_with_the_one_user_its_daily_spend_names(gateway: Gateway) -> None:
api_key: Final = f"integration-5kb-{uuid.uuid4().hex}-{'k' * 5000}"
with gateway.scenario() as scenario:
owner, email = user_with_an_email(scenario)
with daily_rows((user_row(owner, api_key, DAY),)):
assert_key_reported(
activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key),
api_key,
DAY,
key_metadata(user=owner, email=email),
seeded_metrics(1),
)
def test_key_with_no_daily_spend_is_reported_as_no_activity(gateway: Gateway) -> None:
response: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, key_no_key_table_holds())
assert response.status_code == 200, response.text
body: Final = object_value(response.json())
assert body["results"] == [], response.text
totals: Final = object_value(body["metadata"])
assert [totals["total_spend"], totals["total_api_requests"]] == [0.0, 0], response.text
def test_every_key_of_a_team_is_reported_with_its_own_user(gateway: Gateway) -> None:
team: Final = f"integration-entity-{uuid.uuid4().hex}"
owners: Final = {key_no_key_table_holds(): f"integration-owner-{uuid.uuid4().hex}" for _ in range(KEYS_OF_ONE_TEAM)}
rows: Final = tuple(
chain.from_iterable(
(user_row(owner, api_key, DAY), seeded_row(TEAM_SPEND, "team_id", team, api_key, DAY))
for api_key, owner in owners.items()
)
)
with daily_rows(rows):
response: Final = gateway.request(
"GET", AGGREGATED_TEAM_ACTIVITY, params={"start_date": DAY, "end_date": DAY, "team_ids": team}
)
assert response.status_code == 200, response.text
body: Final = object_value(response.json())
days: Final = body["results"]
assert isinstance(days, list) and len(days) == 1, response.text
reported: Final = object_value(object_value(object_value(days[0])["breakdown"])["api_keys"])
assert {api_key: object_value(record)["metadata"] for api_key, record in reported.items()} == {
api_key: key_metadata(user=owner) for api_key, owner in owners.items()
}, response.text
totals: Final = object_value(body["metadata"])
assert totals["total_api_requests"] == KEYS_OF_ONE_TEAM, response.text
assert totals["total_spend"] == pytest.approx(0.25 * KEYS_OF_ONE_TEAM), response.text
def test_reading_the_same_activity_twice_gives_the_same_answer(gateway: Gateway) -> None:
api_key: Final = key_no_key_table_holds()
with gateway.scenario() as scenario:
owner, _ = user_with_an_email(scenario)
with daily_rows((user_row(owner, api_key, DAY),)):
first: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key)
second: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key)
assert [first.status_code, second.status_code] == [200, 200], [first.text, second.text]
assert records_of_key(first.json(), api_key), first.text
assert first.json() == second.json(), [first.text, second.text]
def test_key_stops_being_reported_with_an_owner_once_a_second_user_spends_with_it(gateway: Gateway) -> None:
api_key: Final = key_no_key_table_holds()
with gateway.scenario() as scenario:
first, email = user_with_an_email(scenario)
second, _ = user_with_an_email(scenario)
with daily_rows((user_row(first, api_key, DAY),)):
alone: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key)
with daily_rows((user_row(second, api_key, DAY),)):
shared: Final = activity_of_key(gateway, AGGREGATED_USER_ACTIVITY, api_key)
assert_key_reported(alone, api_key, DAY, key_metadata(user=first, email=email), seeded_metrics(1))
assert_key_reported(shared, api_key, DAY, key_metadata(), seeded_metrics(2))
@pytest.mark.timeout(300)
def test_owner_lookup_gives_up_while_daily_user_spend_is_locked_and_answers_once_it_is_not(
gateway: Gateway, tmp_path: Path
) -> None:
api_key: Final = key_no_key_table_holds()
team: Final = f"integration-entity-{uuid.uuid4().hex}"
with scratch_database() as database_url, _proxy_on(gateway, tmp_path, database_url) as owned:
owner, email = _owner_on(owned.gateway)
rows: Final = (user_row(owner, api_key, DAY), seeded_row(TEAM_SPEND, "team_id", team, api_key, DAY))
with daily_rows(rows, database_url=database_url):
with locked_table(USER_SPEND, database_url=database_url):
started: Final = time.monotonic()
locked: Final = activity_of_key(owned.gateway, TEAM_ACTIVITY, api_key, team_ids=team)
waited: Final = time.monotonic() - started
unlocked: Final = activity_of_key(owned.gateway, TEAM_ACTIVITY, api_key, team_ids=team)
assert waited < GIVES_UP_WITHIN_SECONDS, waited
assert_key_reported(locked, api_key, DAY, key_metadata(), seeded_metrics(1))
assert_key_reported(unlocked, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1))
@pytest.mark.timeout(300)
def test_owner_is_reported_while_a_worker_is_killed_and_after_it_is_replaced(gateway: Gateway, tmp_path: Path) -> None:
api_key: Final = key_no_key_table_holds()
with scratch_database() as database_url, _proxy_on(gateway, tmp_path, database_url, workers=2) as owned:
owner, email = _owner_on(owned.gateway)
with daily_rows((user_row(owner, api_key, DAY),), database_url=database_url):
before: Final = _read_on_a_new_connection(owned.gateway, api_key)
members: Final = tuple(
member for member in group_members(owned.process.pid) if member.pid != owned.process.pid
)
children: Final = tuple(member.pid for member in members)
workers: Final = tuple(
member.pid for member in members if any("spawn_main" in part for part in member.cmdline())
)
assert len(workers) >= 2, workers
os.kill(workers[0], signal.SIGKILL)
during: Final = _read_on_a_new_connection(owned.gateway, api_key)
eventually(
lambda: _running_children(owned),
lambda pids: len(pids) >= len(children) and any(pid not in children for pid in pids),
seconds=30,
)
after: Final = tuple(
_read_on_a_new_connection(owned.gateway, api_key) for _ in range(READS_AFTER_THE_WORKER_IS_REPLACED)
)
for response in (before, during, *after):
assert_key_reported(response, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1))
@pytest.mark.timeout(300)
def test_owner_is_reported_again_after_the_proxy_restarts(gateway: Gateway, tmp_path: Path) -> None:
api_key: Final = key_no_key_table_holds()
with scratch_database() as database_url:
with _proxy_on(gateway, tmp_path, database_url) as first:
owner, email = _owner_on(first.gateway)
insert_daily_rows((user_row(owner, api_key, DAY),), database_url=database_url)
before: Final = activity_of_key(first.gateway, AGGREGATED_USER_ACTIVITY, api_key)
with _proxy_on(gateway, tmp_path, database_url) as second:
after: Final = activity_of_key(second.gateway, AGGREGATED_USER_ACTIVITY, api_key)
for response in (before, after):
assert_key_reported(response, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1))

View file

@ -0,0 +1,397 @@
import json
import os
import threading
import uuid
from collections.abc import Iterable
from concurrent.futures import ThreadPoolExecutor
from datetime import UTC, datetime, timedelta
from hashlib import sha256
from pathlib import Path
from queue import SimpleQueue
from typing import Final
import httpx
import pytest
from integration._support.client import Gateway, Scenario, eventually
from integration._support.daily_activity import (
AGGREGATED_USER_ACTIVITY,
DAY,
ROUTES,
USER_SPEND,
Route,
activity_of_key,
assert_key_owner_and_totals,
assert_key_reported,
daily_rows,
key_metadata,
key_no_key_table_holds,
seeded_metrics,
seeded_row,
user_row,
user_with_an_email,
)
from integration._support.database import read_rows, scratch_database
from integration._support.process import owned_proxy
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue
from litellm.proxy._types import LiteLLM_UserTable
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken
REQUESTS_OF_KEY: Final = (
'SELECT COALESCE(SUM(api_requests), 0)::int AS requests FROM "LiteLLM_DailyUserSpend" '
"WHERE api_key=%s AND user_id=%s"
)
UNIFIED_ENDPOINTS: Final = ("/v1/chat/completions", "/v1/messages", "/v1/responses")
REQUESTS_OF_A_BURST: Final = 21
READS_DURING_A_BURST: Final = 30
TOKEN_LIMIT_DISCOVERY: Final = ("GET", "/v1/models")
TOOL_CALL: Final = "call_integration_usage"
ANSWER: Final = "One request cost $0.25"
SUMMARY_OF_ONE_SEEDED_ROW: Final = "\n".join(
(
"Total Spend: $0.2500",
"Total Requests: 1",
"Successful: 1 | Failed: 0",
"Total Tokens: 15",
"",
"Top Models by Spend:",
" - gpt-4o-mini: $0.2500 (1 reqs, 15 tokens)",
"",
"Top Providers by Spend:",
" - openai: $0.2500 (1 reqs)",
)
)
def _chat_completion() -> dict[str, JsonValue]:
return {
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
}
def _response() -> dict[str, JsonValue]:
return {
"id": f"resp_{uuid.uuid4().hex}",
"object": "response",
"created_at": 1,
"status": "completed",
"model": "gpt-4o-mini",
"output": [
{
"type": "message",
"id": f"msg_{uuid.uuid4().hex}",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "ok", "annotations": []}],
}
],
"usage": {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15},
}
def _provider(request: Request) -> Reply:
body: Final = _response() if request.target.endswith("/responses") else _chat_completion()
return Reply(body=json.dumps(body).encode())
def _usage_tool_call() -> dict[str, JsonValue]:
call: Final[dict[str, JsonValue]] = {
"id": TOOL_CALL,
"type": "function",
"function": {
"name": "get_usage_data",
"arguments": json.dumps({"start_date": DAY, "end_date": DAY}),
},
}
return {
"id": f"chatcmpl-{uuid.uuid4().hex}",
"object": "chat.completion",
"created": 1,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {"role": "assistant", "content": None, "tool_calls": [call]},
"finish_reason": "tool_calls",
}
],
"usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15},
}
def _streamed_chunk(delta: dict[str, JsonValue], finish_reason: str | None) -> bytes:
chunk: Final = {
"id": "chatcmpl-integration-usage",
"object": "chat.completion.chunk",
"created": 1,
"model": "gpt-4o-mini",
"choices": [{"index": 0, "delta": delta, "finish_reason": finish_reason}],
}
return f"data: {json.dumps(chunk)}\n\n".encode()
def _usage_analyst(request: Request) -> Reply:
if json.loads(request.body).get("stream"):
return Reply(
chunks=(
_streamed_chunk({"role": "assistant", "content": ANSWER}, None),
_streamed_chunk({}, "stop"),
b"data: [DONE]\n\n",
),
content_type="text/event-stream",
)
return Reply(body=json.dumps(_usage_tool_call()).encode())
def _sent_for_callers(requests: Iterable[Request]) -> tuple[Request, ...]:
return tuple(request for request in requests if (request.method, request.target) != TOKEN_LIMIT_DISCOVERY)
def _priced_model(scenario: Scenario, provider_url: str) -> str:
return scenario.model(
api_base=f"{provider_url}/v1", input_cost_per_token=0.001, output_cost_per_token=0.002, num_retries=0
)
def _request_body(endpoint: str, model: str, prompt: str) -> dict[str, JsonValue]:
if endpoint == "/v1/chat/completions":
return {"model": model, "messages": [{"role": "user", "content": prompt}]}
if endpoint == "/v1/messages":
return {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": prompt}]}
return {"model": model, "input": prompt}
def _activity_on_route(gateway: Gateway, route: Route, api_key: str, entity: str) -> httpx.Response:
filters: Final = {} if route.entity_filter is None else {route.entity_filter: entity}
return activity_of_key(gateway, route.path, api_key, **filters)
def _prompt() -> str:
return f"daily activity owner {uuid.uuid4().hex}"
def _totals_of_requests(requests: int) -> dict[str, float]:
return {
"total_spend": 0.02 * requests,
"total_prompt_tokens": 10 * requests,
"total_completion_tokens": 5 * requests,
"total_tokens": 15 * requests,
"total_api_requests": requests,
"total_successful_requests": requests,
"total_failed_requests": 0,
}
def _activity_around_today(gateway: Gateway, api_key: str) -> httpx.Response:
today: Final = datetime.now(UTC).date()
return gateway.request(
"GET",
AGGREGATED_USER_ACTIVITY,
params={
"start_date": str(today - timedelta(days=1)),
"end_date": str(today + timedelta(days=1)),
"timezone": "0",
"api_key": api_key,
},
)
def _wait_for_requests(api_key: str, user: str, requests: int) -> None:
eventually(
lambda: read_rows(REQUESTS_OF_KEY, (api_key, user)),
lambda rows: rows[0]["requests"] == requests,
seconds=70,
)
def _cli_session_token(user: str, team: str) -> str:
cli_user: Final = LiteLLM_UserTable(user_id=user, user_role="internal_user", teams=[team], models=[])
return ExperimentalUIJWTToken.get_cli_jwt_auth_token(user_info=cli_user, team_id=team, team_alias="cli-team")
def test_key_used_on_every_unified_endpoint_is_reported_with_its_own_alias_and_user(gateway: Gateway) -> None:
chat_prompt, messages_prompt, responses_prompt = _prompt(), _prompt(), _prompt()
with wire_server(_provider) as wire, gateway.scenario() as scenario:
model: Final = _priced_model(scenario, wire.url)
owner, email = user_with_an_email(scenario)
alias: Final = f"integration-alias-{uuid.uuid4().hex}"
key: Final = scenario.key(user_id=owner, key_alias=alias, models=[model])
stored: Final = sha256(key.encode()).hexdigest()
prompts: Final = (chat_prompt, messages_prompt, responses_prompt)
answers: Final = tuple(
gateway.request("POST", endpoint, _request_body(endpoint, model, prompt), key=key)
for endpoint, prompt in zip(UNIFIED_ENDPOINTS, prompts, strict=True)
)
assert [answer.status_code for answer in answers] == [200, 200, 200], [answer.text for answer in answers]
received: Final = _sent_for_callers(wire.drain())
assert [request.target for request in received] == ["/v1/chat/completions", "/v1/responses", "/v1/responses"]
assert [json.loads(request.body)["model"] for request in received] == ["gpt-4o-mini"] * 3
assert json.loads(received[0].body)["messages"] == [{"role": "user", "content": chat_prompt}]
assert messages_prompt in received[1].body.decode()
assert json.loads(received[2].body)["input"] == responses_prompt
_wait_for_requests(stored, owner, 3)
assert_key_owner_and_totals(
_activity_around_today(gateway, stored),
stored,
key_metadata(alias=alias, user=owner, email=email, exists=True),
_totals_of_requests(3),
)
def test_cli_session_spend_is_reported_with_the_user_and_team_of_the_session(
gateway: Gateway, monkeypatch: pytest.MonkeyPatch
) -> None:
monkeypatch.setenv("LITELLM_SALT_KEY", os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt"))
prompt: Final = _prompt()
with wire_server(_provider) as wire, gateway.scenario() as scenario:
model: Final = _priced_model(scenario, wire.url)
owner, email = user_with_an_email(scenario)
team: Final = scenario.team(models=[model], members_with_roles=[{"role": "user", "user_id": owner}])
answer: Final = gateway.request(
"POST",
"/v1/chat/completions",
{"model": model, "messages": [{"role": "user", "content": prompt}]},
key=_cli_session_token(owner, team),
)
assert answer.status_code == 200, answer.text
received: Final = _sent_for_callers(wire.drain())
assert [request.target for request in received] == ["/v1/chat/completions"]
assert json.loads(received[0].body) == {
"model": "gpt-4o-mini",
"messages": [{"role": "user", "content": prompt}],
}
stored: Final = f"cli-session-{owner}"
_wait_for_requests(stored, owner, 1)
assert_key_owner_and_totals(
_activity_around_today(gateway, stored),
stored,
key_metadata(alias=stored, team=team, user=owner, email=email),
_totals_of_requests(1),
)
@pytest.mark.timeout(300)
def test_usage_ai_chat_hands_the_model_the_usage_summary_without_any_key_owner(
gateway: Gateway, tmp_path: Path
) -> None:
question: Final = f"what did we spend {uuid.uuid4().hex}"
owner: Final = f"integration-{uuid.uuid4().hex}"
ownerless_key: Final = f"integration-ownerless-{uuid.uuid4().hex}"
with (
scratch_database() as scratch_url,
wire_server(_usage_analyst) as wire,
owned_proxy(
gateway,
tmp_path,
{
"DATABASE_URL": scratch_url,
"OPENAI_API_BASE": f"{wire.url}/v1",
"OPENAI_BASE_URL": f"{wire.url}/v1",
"OPENAI_API_KEY": "integration-provider-key",
},
remove_environment=("DATABASE_URL_READ_REPLICA",),
) as candidate,
):
candidate.post("/user/new", {"user_id": owner, "user_email": f"{owner}@example.com", "auto_create_key": False})
with daily_rows((user_row(owner, ownerless_key, DAY),), database_url=scratch_url):
answer: Final = candidate.request(
"POST",
"/usage/ai/chat",
{"messages": [{"role": "user", "content": question}], "model": "openai/gpt-4o-mini"},
)
assert answer.status_code == 200, answer.text
tool_call: Final = {
"type": "tool_call",
"tool_name": "get_usage_data",
"tool_label": "global usage data",
"arguments": {"start_date": DAY, "end_date": DAY},
}
events: Final = [
json.loads(line.removeprefix("data: ")) for line in answer.text.splitlines() if line.startswith("data: ")
]
assert events == [
{"type": "status", "message": "Thinking..."},
{**tool_call, "status": "running"},
{**tool_call, "status": "complete"},
{"type": "status", "message": "Analyzing results..."},
{"type": "chunk", "content": ANSWER},
{"type": "done"},
], answer.text
asked, analysed = wire.drain()
assert [asked.target, analysed.target] == ["/v1/chat/completions", "/v1/chat/completions"]
assert json.loads(asked.body)["messages"][-1] == {"role": "user", "content": question}
assert json.loads(analysed.body)["messages"][-1] == {
"role": "tool",
"tool_call_id": TOOL_CALL,
"content": SUMMARY_OF_ONE_SEEDED_ROW,
}
assert owner not in analysed.body.decode()
assert ownerless_key not in analysed.body.decode()
@pytest.mark.timeout(300)
def test_owner_is_reported_on_every_route_while_a_burst_of_requests_waits_on_the_provider(gateway: Gateway) -> None:
released: Final = threading.Event()
held: Final[SimpleQueue[str]] = SimpleQueue()
def held_provider(request: Request) -> Reply:
if (request.method, request.target) == TOKEN_LIMIT_DISCOVERY:
return _provider(request)
held.put(request.target)
assert released.wait(timeout=120), "The burst was never released"
return _provider(request)
api_key: Final = key_no_key_table_holds()
entity: Final = f"integration-entity-{uuid.uuid4().hex}"
prompts: Final = tuple(_prompt() for _ in range(REQUESTS_OF_A_BURST))
entity_columns: Final = {route.table: route.entity_column for route in ROUTES if route.table != USER_SPEND}
with (
wire_server(held_provider) as wire,
gateway.scenario() as scenario,
httpx.Client(base_url=gateway.client.base_url, timeout=180, trust_env=False) as patient,
ThreadPoolExecutor(max_workers=REQUESTS_OF_A_BURST) as traffic,
ThreadPoolExecutor(max_workers=READS_DURING_A_BURST) as readers,
):
model: Final = _priced_model(scenario, wire.url)
owner, email = user_with_an_email(scenario)
key: Final = scenario.key(models=[model])
rows: Final = (
user_row(owner, api_key, DAY),
*(seeded_row(table, column, entity, api_key, DAY) for table, column in entity_columns.items()),
)
try:
with daily_rows(rows):
burst: Final = tuple(
traffic.submit(
patient.post,
UNIFIED_ENDPOINTS[index % len(UNIFIED_ENDPOINTS)],
json=_request_body(UNIFIED_ENDPOINTS[index % len(UNIFIED_ENDPOINTS)], model, prompt),
headers={"Authorization": f"Bearer {key}"},
)
for index, prompt in enumerate(prompts)
)
eventually(held.qsize, lambda waiting: waiting >= REQUESTS_OF_A_BURST, seconds=60)
reads: Final = tuple(
readers.submit(_activity_on_route, gateway, ROUTES[index % len(ROUTES)], api_key, entity)
for index in range(READS_DURING_A_BURST)
)
activity: Final = tuple(read.result() for read in reads)
still_waiting: Final = [call.done() for call in burst]
finally:
released.set()
answers: Final = tuple(call.result() for call in burst)
received: Final = tuple(request.body.decode() for request in _sent_for_callers(wire.drain()))
assert still_waiting == [False] * REQUESTS_OF_A_BURST
assert [answer.status_code for answer in answers] == [200] * REQUESTS_OF_A_BURST, [
answer.text for answer in answers
]
assert [sum(prompt in body for body in received) for prompt in prompts] == [1] * REQUESTS_OF_A_BURST
assert len(received) == REQUESTS_OF_A_BURST, len(received)
for response in activity:
assert_key_reported(response, api_key, DAY, key_metadata(user=owner, email=email), seeded_metrics(1))

View file

@ -23,6 +23,7 @@ from litellm.proxy.management_endpoints.common_daily_activity import (
update_metrics,
)
from litellm.proxy.spend_tracking.ptu_feature_flag import PTU_COST_ATTRIBUTION_ENV_VAR
from litellm.proxy.utils import hash_token
from litellm.types.proxy.management_endpoints.common_daily_activity import (
DailySpendMetadata,
SpendMetrics,
@ -505,6 +506,7 @@ async def test_get_api_key_metadata_permanent_miss_never_pages_tokens_or_reads_s
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[])
mock_prisma.db.query_raw = AsyncMock(return_value=[])
recovery_query_raw = _recovery_transaction(mock_prisma)
result = await get_api_key_metadata(
prisma_client=mock_prisma,
@ -512,9 +514,10 @@ async def test_get_api_key_metadata_permanent_miss_never_pages_tokens_or_reads_s
)
assert double_hashed not in result
issued_sql = [call.args[0] for call in mock_prisma.db.query_raw.call_args_list]
assert len(issued_sql) == 2
assert not any("LiteLLM_SpendLogs" in sql for sql in issued_sql)
assert mock_prisma.db.query_raw.await_count == 2
((owner_sql, owner_keys),) = [call.args for call in recovery_query_raw.call_args_list]
assert _DAILY_USER_SPEND in owner_sql
assert owner_keys == [double_hashed]
token_lookups = (
mock_prisma.db.litellm_verificationtoken.find_many.call_args_list
+ mock_prisma.db.litellm_deletedverificationtoken.find_many.call_args_list
@ -522,14 +525,29 @@ async def test_get_api_key_metadata_permanent_miss_never_pages_tokens_or_reads_s
assert all("take" not in call.kwargs and "skip" not in call.kwargs for call in token_lookups)
def _spend_log_transaction(mock_prisma: MagicMock, rows: list[dict[str, str | None]]) -> AsyncMock:
_DAILY_USER_SPEND: Final = '"LiteLLM_DailyUserSpend"'
_SPEND_LOGS: Final = '"LiteLLM_SpendLogs"'
def _recovery_transaction(
mock_prisma: MagicMock,
spend_log_rows: Sequence[dict[str, str | None]] = (),
daily_spend_owner_rows: Sequence[dict[str, str | None]] = (),
) -> AsyncMock:
async def query_raw(sql: str, *_: object) -> Sequence[dict[str, str | None]]:
return daily_spend_owner_rows if _DAILY_USER_SPEND in sql else spend_log_rows
transaction = MagicMock()
transaction.execute_raw = AsyncMock(return_value=0)
transaction.query_raw = AsyncMock(return_value=rows)
transaction.query_raw = AsyncMock(side_effect=query_raw)
mock_prisma.db.tx.return_value.__aenter__.return_value = transaction
return transaction.query_raw
def _calls_reading(query_raw: AsyncMock, table: str) -> tuple[tuple[object, ...], ...]:
return tuple(call.args for call in query_raw.call_args_list if table in call.args[0])
def _spend_log_row(digest: str, key_alias: str, user_id: str) -> dict[str, str | None]:
return {
"digest": digest,
@ -553,15 +571,17 @@ async def test_get_api_key_metadata_permanent_miss_with_a_window_reads_spend_log
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[])
mock_prisma.db.query_raw = AsyncMock(return_value=[])
spend_log_query_raw = _spend_log_transaction(mock_prisma, [])
recovery_query_raw = _recovery_transaction(mock_prisma)
result = await get_api_key_metadata(prisma_client=mock_prisma, api_keys={double_hashed}, spend_logs_window=window)
assert double_hashed not in result
assert mock_prisma.db.query_raw.await_count == 2
((_, digests, start, end),) = [call.args for call in spend_log_query_raw.call_args_list]
((_, digests, start, end),) = _calls_reading(recovery_query_raw, _SPEND_LOGS)
assert digests == [double_hashed]
assert (start, end) == window
((_, owner_keys),) = _calls_reading(recovery_query_raw, _DAILY_USER_SPEND)
assert owner_keys == [double_hashed]
@pytest.mark.asyncio
@ -583,7 +603,7 @@ async def test_get_daily_activity_recovers_a_session_key_alias_from_spend_logs_a
)
mock_prisma.db.query_raw = AsyncMock(return_value=[])
spend_log_query_raw = _spend_log_transaction(
spend_log_query_raw = _recovery_transaction(
mock_prisma, [_spend_log_row(session_digest, "cli-session-alias", "session-user")]
)
@ -1544,6 +1564,8 @@ async def test_get_daily_activity_aggregated_returns_every_api_key(
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,
@ -1599,6 +1621,8 @@ async def test_get_daily_activity_aggregated_explicit_api_key_filter_scopes_resu
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,
@ -1651,6 +1675,8 @@ async def test_get_daily_activity_aggregated_model_group_rollups_fall_back_to_mo
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,
@ -2580,7 +2606,7 @@ async def test_get_api_key_metadata_resolves_session_key_via_spend_log_window():
)
mock_prisma.db.query_raw = AsyncMock(return_value=[])
spend_log_query_raw = _spend_log_transaction(
spend_log_query_raw = _recovery_transaction(
mock_prisma, [_spend_log_row(session_digest, "cli-session-user-42", "user-42")]
)
@ -2627,3 +2653,90 @@ async def test_get_api_key_metadata_resolves_cli_session_keys_from_the_key_itsel
assert result["cli-session-alice"]["key_alias"] == "cli-session-alice"
assert result["cli-session-alice"]["user_email"] == "alice@example.com"
assert result["cli-session-alice"]["team_id"] == "team-a"
@pytest.mark.asyncio
async def test_get_api_key_metadata_recovers_legacy_hashed_jwt_owner_from_daily_spend():
api_key: Final = f"hashed-jwt-{hash_token('legacy-cli-session-daily-spend-owner')}"
user_id: Final = "legacy-owner"
mock_prisma: Final = MagicMock()
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.find_many = AsyncMock(
return_value=[SimpleNamespace(user_id=user_id, user_email="legacy-owner@example.com", teams=[])]
)
recovery_query_raw: Final = _recovery_transaction(
mock_prisma,
daily_spend_owner_rows=[{"api_key": api_key, "first_owner": user_id, "last_owner": user_id}],
)
result: Final = await get_api_key_metadata(
prisma_client=mock_prisma,
api_keys={api_key},
spend_logs_window=(datetime(2026, 9, 7), datetime(2026, 9, 10)),
)
assert result.get(api_key, {}).get("user_id") == user_id
assert result.get(api_key, {}).get("user_email") == "legacy-owner@example.com"
assert len(_calls_reading(recovery_query_raw, _SPEND_LOGS)) == 1
assert len(_calls_reading(recovery_query_raw, _DAILY_USER_SPEND)) == 1
@pytest.mark.asyncio
async def test_get_api_key_metadata_preserves_deleted_key_metadata_when_recovering_daily_spend_owner():
api_key: Final = f"hashed-jwt-{hash_token('legacy-cli-session-daily-spend-metadata')}"
user_id: Final = "legacy-owner"
mock_prisma: Final = MagicMock()
deleted_key: Final = MagicMock()
deleted_key.token = api_key
deleted_key.key_alias = "legacy-cli-key"
deleted_key.team_id = "team-legacy"
deleted_key.user_id = None
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[deleted_key])
mock_prisma.db.litellm_usertable.find_many = AsyncMock(
return_value=[SimpleNamespace(user_id=user_id, user_email="legacy-owner@example.com", teams=[])]
)
recovery_query_raw: Final = _recovery_transaction(
mock_prisma,
daily_spend_owner_rows=[{"api_key": api_key, "first_owner": user_id, "last_owner": user_id}],
)
result: Final = await get_api_key_metadata(prisma_client=mock_prisma, api_keys={api_key})
recovered_metadata: Final = result[api_key]
assert recovered_metadata.get("key_alias") == "legacy-cli-key"
assert recovered_metadata.get("team_id") == "team-legacy"
assert recovered_metadata.get("user_id") == user_id
assert recovered_metadata.get("user_email") == "legacy-owner@example.com"
recovery_query_raw.assert_awaited_once()
@pytest.mark.asyncio
async def test_get_api_key_metadata_does_not_recover_daily_spend_owner_for_active_keys():
api_key: Final = "active-token-value"
mock_prisma: Final = MagicMock()
active_key: Final = MagicMock()
active_key.token = api_key
active_key.key_alias = "active-key-alias"
active_key.team_id = "active-team"
active_key.user_id = "active-owner"
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[active_key])
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
mock_prisma.db.litellm_usertable.find_many = AsyncMock(
return_value=[SimpleNamespace(user_id="active-owner", user_email="active-owner@example.com", teams=[])]
)
recovery_query_raw: Final = _recovery_transaction(
mock_prisma,
daily_spend_owner_rows=[{"api_key": api_key, "first_owner": "other-owner", "last_owner": "other-owner"}],
)
result: Final = await get_api_key_metadata(prisma_client=mock_prisma, api_keys={api_key})
active_metadata: Final = result[api_key]
assert active_metadata.get("key_alias") == "active-key-alias"
assert active_metadata.get("team_id") == "active-team"
assert active_metadata.get("user_id") == "active-owner"
assert active_metadata.get("user_email") == "active-owner@example.com"
assert active_metadata.get("key_exists") is True
recovery_query_raw.assert_not_awaited()

View file

@ -3,6 +3,7 @@ import time
from collections.abc import Sequence
from datetime import datetime, timedelta
from types import SimpleNamespace
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -20,6 +21,7 @@ from litellm.proxy.spend_tracking.key_metadata_recovery import (
recover_cli_session_key_metadata,
recover_double_hashed_key_metadata,
recover_key_metadata_from_spend_logs,
recover_key_owner_from_daily_spend,
)
from litellm.proxy.utils import hash_token
@ -702,3 +704,91 @@ async def test_attach_user_details_leaves_metadata_unchanged_when_a_later_chunk_
assert mock_prisma.db.litellm_usertable.find_many.call_count == 2
assert attached == recovered
def _daily_spend_owner_row(api_key: str, first_owner: str, last_owner: str) -> dict[str, str]:
return {"api_key": api_key, "first_owner": first_owner, "last_owner": last_owner}
def _daily_spend_transaction(mock_prisma: MagicMock, query_raw: AsyncMock) -> MagicMock:
transaction: Final = MagicMock()
transaction.execute_raw = AsyncMock(return_value=0)
transaction.query_raw = query_raw
mock_prisma.db.tx.return_value.__aenter__.return_value = transaction
return transaction
@pytest.mark.asyncio
async def test_recover_key_owner_from_daily_spend_keeps_a_unanimous_owner():
key: Final = "hashed-jwt-digest-a"
mock_prisma: Final = MagicMock()
_daily_spend_transaction(mock_prisma, AsyncMock(return_value=[_daily_spend_owner_row(key, "owner-a", "owner-a")]))
result: Final = await recover_key_owner_from_daily_spend(mock_prisma, {key})
assert dict(result) == {key: "owner-a"}
@pytest.mark.asyncio
async def test_recover_key_owner_from_daily_spend_drops_conflicting_owners():
key: Final = "hashed-jwt-digest-b"
mock_prisma: Final = MagicMock()
_daily_spend_transaction(mock_prisma, AsyncMock(return_value=[_daily_spend_owner_row(key, "owner-a", "owner-b")]))
result: Final = await recover_key_owner_from_daily_spend(mock_prisma, {key})
assert dict(result) == {}
@pytest.mark.asyncio
async def test_recover_key_owner_from_daily_spend_skips_empty_input():
mock_prisma: Final = MagicMock()
transaction: Final = _daily_spend_transaction(mock_prisma, AsyncMock(return_value=[]))
result: Final = await recover_key_owner_from_daily_spend(mock_prisma, frozenset())
assert dict(result) == {}
transaction.query_raw.assert_not_awaited()
mock_prisma.db.tx.assert_not_called()
@pytest.mark.asyncio
async def test_recover_key_owner_from_daily_spend_returns_empty_on_prisma_error():
mock_prisma: Final = MagicMock()
_daily_spend_transaction(mock_prisma, AsyncMock(side_effect=PrismaError("db down")))
result: Final = await recover_key_owner_from_daily_spend(mock_prisma, {"hashed-jwt-digest-c"})
assert dict(result) == {}
@pytest.mark.asyncio
async def test_recover_key_owner_from_daily_spend_names_no_owner_when_the_lookup_hits_the_statement_timeout():
mock_prisma: Final = MagicMock()
_daily_spend_transaction(
mock_prisma, AsyncMock(side_effect=PrismaError("canceling statement due to statement timeout"))
)
result: Final = await recover_key_owner_from_daily_spend(mock_prisma, {"hashed-jwt-digest-d"})
assert dict(result) == {}
@pytest.mark.asyncio
async def test_recover_key_owner_from_daily_spend_bounds_the_lookup_with_a_statement_timeout():
key: Final = "hashed-jwt-digest-e"
mock_prisma: Final = MagicMock()
transaction: Final = _daily_spend_transaction(
mock_prisma, AsyncMock(return_value=[_daily_spend_owner_row(key, "owner-a", "owner-a")])
)
result: Final = await recover_key_owner_from_daily_spend(mock_prisma, {key})
assert dict(result) == {key: "owner-a"}
assert [name for name, _, _ in transaction.mock_calls] == ["execute_raw", "query_raw"]
transaction.execute_raw.assert_awaited_once_with(
f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}"
)
assert mock_prisma.db.tx.call_args.kwargs["timeout"] == timedelta(
milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS
)

View file

@ -58,6 +58,7 @@ export interface TopApiKeyData {
api_key: string;
key_alias: string | null;
team_id: string | null;
user: string | null;
spend: number;
requests: number;
tokens: number;

View file

@ -258,10 +258,20 @@ describe("ActivityMetrics", () => {
api_key: "key-123",
key_alias: "Test Key",
team_id: "team1",
user: "owner@example.com",
spend: 50.25,
requests: 25,
tokens: 12500,
},
{
api_key: "key-456",
key_alias: "Owner Alias",
team_id: null,
user: "Owner Alias",
spend: 40.25,
requests: 20,
tokens: 10000,
},
],
},
};
@ -269,6 +279,9 @@ describe("ActivityMetrics", () => {
render(<ActivityMetrics modelMetrics={modelWithTopKeys} />);
expect(screen.getByText("Top Virtual Keys by Spend")).toBeInTheDocument();
expect(screen.getByText("Test Key")).toBeInTheDocument();
expect(screen.getByText("User: owner@example.com")).toBeInTheDocument();
expect(screen.getByText("Owner Alias")).toBeInTheDocument();
expect(screen.queryByText("User: Owner Alias")).not.toBeInTheDocument();
});
it("should display API key hash when alias is missing", () => {
@ -280,6 +293,7 @@ describe("ActivityMetrics", () => {
api_key: "key-1234567890",
key_alias: null,
team_id: null,
user: null,
spend: 50.25,
requests: 25,
tokens: 12500,
@ -301,6 +315,7 @@ describe("ActivityMetrics", () => {
api_key: "key-123",
key_alias: "Test Key",
team_id: "team1",
user: null,
spend: 50.25,
requests: 25,
tokens: 12500,
@ -1056,6 +1071,8 @@ describe("processActivityData", () => {
metadata: {
key_alias: "test-key-1",
team_id: "team1",
user_id: "owner-id-1",
user_email: "owner-1@example.com",
},
},
"key-2": {
@ -1073,6 +1090,7 @@ describe("processActivityData", () => {
metadata: {
key_alias: "test-key-2",
team_id: "team2",
user_id: "owner-id-2",
},
},
},
@ -1094,6 +1112,10 @@ describe("processActivityData", () => {
expect(result["gpt-4"].top_api_keys[0].spend).toBe(60.0);
expect(result["gpt-4"].top_api_keys[0].api_key).toBe("key-1");
expect(result["gpt-4"].top_api_keys[1].spend).toBe(40.5);
expect(result["gpt-4"].top_api_keys.map(({ api_key, user }) => [api_key, user])).toEqual([
["key-1", "owner-1@example.com"],
["key-2", "owner-id-2"],
]);
});
it("should limit top_api_keys to 5 entries", () => {

View file

@ -102,20 +102,26 @@ const ModelSection = ({
<h3 className="text-lg font-medium text-foreground">Top Virtual Keys by Spend</h3>
<div className="mt-3">
<div className="grid grid-cols-1 gap-2">
{metrics.top_api_keys.map((keyData) => (
<div key={keyData.api_key} className="flex justify-between items-center p-3 bg-muted rounded-lg">
<div>
<p className="font-medium">{keyData.key_alias || `${keyData.api_key.substring(0, 10)}...`}</p>
{keyData.team_id && <p className="text-xs text-muted-foreground">Team: {keyData.team_id}</p>}
{metrics.top_api_keys.map((keyData) => {
const keyLabel = keyData.key_alias || `${keyData.api_key.substring(0, 10)}...`;
return (
<div key={keyData.api_key} className="flex justify-between items-center p-3 bg-muted rounded-lg">
<div>
<p className="font-medium">{keyLabel}</p>
{keyData.team_id && <p className="text-xs text-muted-foreground">Team: {keyData.team_id}</p>}
{keyData.user && keyData.user !== keyLabel && (
<p className="text-xs text-muted-foreground">User: {keyData.user}</p>
)}
</div>
<div className="text-right">
<p className="font-medium">${formatNumberWithCommas(keyData.spend, 2)}</p>
<p className="text-xs text-muted-foreground">
{keyData.requests.toLocaleString()} requests | {keyData.tokens.toLocaleString()} tokens
</p>
</div>
</div>
<div className="text-right">
<p className="font-medium">${formatNumberWithCommas(keyData.spend, 2)}</p>
<p className="text-xs text-muted-foreground">
{keyData.requests.toLocaleString()} requests | {keyData.tokens.toLocaleString()} tokens
</p>
</div>
</div>
))}
);
})}
</div>
</div>
</CardContent>
@ -585,6 +591,7 @@ export const processActivityData = (
api_key: apiKey,
key_alias: keyActivityLabel(keyData.metadata, "") || null,
team_id: keyData.metadata.team_id,
user: keyData.metadata.user_email ?? keyData.metadata.user_id ?? null,
spend: 0,
requests: 0,
tokens: 0,