Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 27 additions & 3 deletions litellm/proxy/management_endpoints/common_daily_activity.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
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
Expand Down Expand Up @@ -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],
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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))
Expand Down Expand Up @@ -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))
Expand Down
62 changes: 50 additions & 12 deletions litellm/proxy/spend_tracking/key_metadata_recovery.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,6 +61,13 @@
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 <> ''
Comment thread
greptile-apps[bot] marked this conversation as resolved.
GROUP BY api_key
Comment thread
greptile-apps[bot] marked this conversation as resolved.
"""

_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)

Expand Down Expand Up @@ -104,15 +111,23 @@ def metadata(self) -> KeyMetadataDict:
)


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,
default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL,
)
_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(
Expand All @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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),
)
Expand Down
Loading
Loading