mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
fix(spend): attribute CLI session spend to the per-user cli-session alias instead of the hashed session token (#40541)
* fix(spend): attribute CLI session spend to the per-user cli-session alias instead of the hashed session token A CLI session token is a fresh random secret on every login, so since v1.99 each login's spend rows carried a different sha256 hash as api_key and the usage APIs could resolve neither key_alias nor user_email for them. Spend rows and logging callbacks now attribute a session request to its stable alias, cli-session-<user_id>, and the usage endpoints derive that alias and owner from the key itself instead of scanning for a matching digest * fix(spend): resolve the CLI session team from the user's first team in usage metadata A cli-session key carries no team of its own in the DB, so the usage breakdown showed team_id None for it and the export grouped it as Unassigned. The login attaches the user's first team to the session, so the recovery mirrors that rule for cli-session keys only. * fix(spend): claim the session team only for a single-team user The CLI login attaches a team on its own only when the user has exactly one; a user in several teams picks one per login, so usage metadata for the alias would otherwise name a team the login may not have used. * test(pass_through): mark the mocked auth object as a plain key The logged key follows the alias only for a session token; a bare MagicMock reads as one, so the test names the field it relies on. * fix(spend): attribute CLI session pass-through, queue, and managed batch spend to the cli-session alias Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(spend): only treat the exact cli-session-<created_by> value as a batch key alias A managed object row written by an older build can still carry the raw per-login session token, which shares the cli-session- prefix. Matching on the prefix alone would have surfaced that token as a trusted alias and persisted it verbatim in the batch cost spend log, so the alias check now requires the exact per-user value and every other prefixed value keeps going through redaction Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(spend): log proxy executed batch rows under the cli-session alias instead of the session token _row_metadata set user_api_key from the raw bearer token while user_api_key_hash carried the alias, so the spend log redaction rejected the alias as untrusted and hashed the random session token instead Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(spend): attribute semantic search embedding spend to the cli-session alias Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(spend): scope /key/spend/report for a CLI session to the cli-session alias Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(spend): use the cli-session alias for websearch spend, prometheus failure labels and the parallel limiter Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(spend): drop explanatory docstrings on get_logged_api_key and attach_user_details Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(spend): only recover cli-session usage keys whose suffix is a known user Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
e2302be068
commit
25fb7810c2
27 changed files with 533 additions and 54 deletions
|
|
@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Final, List, Literal, Optional, Protocol, Tupl
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm._uuid import uuid
|
||||
from litellm.constants import (
|
||||
CLI_SESSION_KEY_PREFIX,
|
||||
MANAGED_OBJECT_STALENESS_CUTOFF_DAYS,
|
||||
MAX_OBJECTS_PER_POLL_CYCLE,
|
||||
)
|
||||
|
|
@ -147,10 +148,12 @@ class CheckBatchCost:
|
|||
verbose_proxy_logger.error(f"CheckBatchCost: could not look up user {user_id} for batch {batch_id}: {e}")
|
||||
return {}
|
||||
|
||||
async def _get_key_alias(self, batch_id: str, api_key: str | None) -> str | None:
|
||||
async def _get_key_alias(self, batch_id: str, api_key: str | None, created_by: str | None) -> str | None:
|
||||
"""Resolve the creating virtual key's alias from its hashed token."""
|
||||
if not api_key:
|
||||
return None
|
||||
if created_by and api_key == f"{CLI_SESSION_KEY_PREFIX}-{created_by}":
|
||||
return api_key
|
||||
try:
|
||||
key_row: prisma_models.LiteLLM_VerificationToken | None = await _token_table(
|
||||
self.prisma_client
|
||||
|
|
@ -231,7 +234,7 @@ class CheckBatchCost:
|
|||
**(await self._get_user_info(batch_id, job.created_by)),
|
||||
}
|
||||
|
||||
key_alias = await self._get_key_alias(batch_id, api_key)
|
||||
key_alias = await self._get_key_alias(batch_id, api_key, job.created_by)
|
||||
if key_alias is not None:
|
||||
metadata["user_api_key_alias"] = key_alias
|
||||
team_alias = await self._get_team_alias(team_id)
|
||||
|
|
|
|||
|
|
@ -50,6 +50,7 @@ from litellm.proxy._types import (
|
|||
ProxyException,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.proxy.openai_files_endpoints.common_utils import (
|
||||
BATCH_CREATE_HIDDEN_PARAM,
|
||||
FILE_LIST_CONTINUATION_CHUNK_SIZE,
|
||||
|
|
@ -359,7 +360,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
|
|||
|
||||
from prisma import Json
|
||||
|
||||
api_key = user_api_key_dict.api_key or None
|
||||
api_key = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict) or None
|
||||
attribution_columns = (
|
||||
{
|
||||
**({"api_key": api_key} if api_key is not None else {}),
|
||||
|
|
|
|||
|
|
@ -2623,6 +2623,7 @@ class PrometheusLogger(CustomLogger):
|
|||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
StandardLoggingPayloadSetup,
|
||||
)
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
|
||||
status_code: Final = self._extract_status_code(exception=original_exception)
|
||||
|
||||
|
|
@ -2641,7 +2642,9 @@ class PrometheusLogger(CustomLogger):
|
|||
end_user=user_api_key_dict.end_user_id,
|
||||
user=user_api_key_dict.user_id,
|
||||
user_email=user_api_key_dict.user_email,
|
||||
hashed_api_key=None if status_code == 401 else user_api_key_dict.api_key,
|
||||
hashed_api_key=None
|
||||
if status_code == 401
|
||||
else LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
|
||||
api_key_alias=user_api_key_dict.key_alias,
|
||||
team=user_api_key_dict.team_id,
|
||||
team_alias=user_api_key_dict.team_alias,
|
||||
|
|
|
|||
|
|
@ -1687,7 +1687,7 @@ class WebSearchInterceptionLogger(CustomLogger):
|
|||
**user_api_key_metadata,
|
||||
**parent_correlation.as_search_metadata(),
|
||||
"model_group": search_tool_name,
|
||||
"user_api_key": user_api_key_auth.api_key,
|
||||
"user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_auth),
|
||||
"user_api_key_auth": user_api_key_auth,
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -656,7 +656,7 @@ class LiteLLMExecutedBatchRunner:
|
|||
def _row_metadata(self, run: _BatchRun) -> dict[str, object]: # mutable-ok: router updates metadata in place
|
||||
return { # mutable-ok: the router updates request metadata in place
|
||||
**LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(run.user_api_key_dict),
|
||||
"user_api_key": run.user_api_key_dict.api_key,
|
||||
"user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(run.user_api_key_dict),
|
||||
"user_api_end_user_max_budget": run.user_api_key_dict.end_user_max_budget,
|
||||
"tags": list(run.request_tags), # mutable-ok: litellm types request tags as a list
|
||||
"batch_id": run.unified_batch_id,
|
||||
|
|
|
|||
|
|
@ -68,7 +68,7 @@ def embedding_spend_metadata(user_api_key_dict: UserAPIKeyAuth) -> dict[str, obj
|
|||
|
||||
return { # mutable-ok: the router mutates the metadata dict it is handed
|
||||
**LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict),
|
||||
"user_api_key": user_api_key_dict.api_key,
|
||||
"user_api_key": LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
|
||||
}
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ from litellm.proxy.auth.auth_utils import (
|
|||
from litellm.proxy.auth.budget_throttle import throttled_limit
|
||||
from litellm.proxy.common_utils.proxy_rate_limit_error import ProxyRateLimitError
|
||||
from litellm.proxy.hooks.rate_limiter_utils import resolve_llm_provider_for_rate_limit
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.types.utils import Usage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -250,7 +251,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
call_type: str,
|
||||
):
|
||||
self.print_verbose("Inside Max Parallel Request Pre-Call Hook")
|
||||
api_key: Final = user_api_key_dict.api_key
|
||||
api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
|
||||
max_parallel_requests = user_api_key_dict.max_parallel_requests
|
||||
if max_parallel_requests is None:
|
||||
max_parallel_requests = sys.maxsize
|
||||
|
|
@ -803,7 +804,7 @@ class _PROXY_MaxParallelRequestsHandler(CustomLogger):
|
|||
"""
|
||||
Retrieve the key's remaining rate limits.
|
||||
"""
|
||||
api_key: Final = user_api_key_dict.api_key
|
||||
api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
|
||||
current_date: Final = datetime.now().strftime("%Y-%m-%d")
|
||||
current_hour: Final = datetime.now().strftime("%H")
|
||||
current_minute: Final = datetime.now().strftime("%M")
|
||||
|
|
|
|||
|
|
@ -165,7 +165,7 @@ class _ProxyDBLogger(CustomLogger):
|
|||
_metadata = dict(
|
||||
LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict)
|
||||
)
|
||||
_metadata["user_api_key"] = user_api_key_dict.api_key
|
||||
_metadata["user_api_key"] = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
|
||||
_metadata["status"] = "failure"
|
||||
_error_information = StandardLoggingPayloadSetup.get_error_information(
|
||||
original_exception=original_exception,
|
||||
|
|
@ -259,7 +259,7 @@ class _ProxyDBLogger(CustomLogger):
|
|||
)
|
||||
|
||||
await self._spend_writer().update_database(
|
||||
token=user_api_key_dict.api_key,
|
||||
token=LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
|
||||
response_cost=recovered_response_cost,
|
||||
user_id=user_api_key_dict.user_id,
|
||||
end_user_id=user_api_key_dict.end_user_id,
|
||||
|
|
|
|||
|
|
@ -1602,6 +1602,12 @@ class LiteLLMProxyRequestSetup:
|
|||
data[_metadata_variable_name].update(metadata_from_headers)
|
||||
return data
|
||||
|
||||
@staticmethod
|
||||
def get_logged_api_key(user_api_key_dict: UserAPIKeyAuth) -> str | None:
|
||||
if user_api_key_dict.is_session_token and user_api_key_dict.key_alias:
|
||||
return user_api_key_dict.key_alias
|
||||
return user_api_key_dict.api_key
|
||||
|
||||
@staticmethod
|
||||
def get_sanitized_user_information_from_key(
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
|
|
@ -1609,7 +1615,7 @@ class LiteLLMProxyRequestSetup:
|
|||
stripped_metadata: Final = strip_callback_config(user_api_key_dict.metadata)
|
||||
auth_metadata: Final = cast("dict[str, str] | None", stripped_metadata) # cast-ok: metadata is free-form JSON
|
||||
user_api_key_logged_metadata: Final = StandardLoggingUserAPIKeyMetadata(
|
||||
user_api_key_hash=user_api_key_dict.api_key, # just the hashed token
|
||||
user_api_key_hash=LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
|
||||
user_api_key_alias=user_api_key_dict.key_alias,
|
||||
user_api_key_spend=user_api_key_dict.spend,
|
||||
user_api_key_max_budget=user_api_key_dict.max_budget,
|
||||
|
|
@ -1647,7 +1653,7 @@ class LiteLLMProxyRequestSetup:
|
|||
user_api_key_dict=user_api_key_dict
|
||||
)
|
||||
data[_metadata_variable_name].update(user_api_key_logged_metadata)
|
||||
data[_metadata_variable_name]["user_api_key"] = user_api_key_dict.api_key # this is just the hashed token
|
||||
data[_metadata_variable_name]["user_api_key"] = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
|
||||
|
||||
# Key-owned agent_id for spend attribution; keep existing (e.g. from header) if key has none
|
||||
_key_agent_id: Final = getattr(user_api_key_dict, "agent_id", None)
|
||||
|
|
|
|||
|
|
@ -15,7 +15,8 @@ from litellm.constants import PTU_SENTINEL_API_KEY, USAGE_TOP_API_KEYS_LIMIT
|
|||
from litellm.proxy._types import CommonProxyErrors
|
||||
from litellm.proxy.spend_tracking.daily_global_spend_rollup import GLOBAL_SPEND_TABLE_NAME, reconciled_through
|
||||
from litellm.proxy.spend_tracking.key_metadata_recovery import (
|
||||
attach_user_emails,
|
||||
attach_user_details,
|
||||
recover_cli_session_key_metadata,
|
||||
recover_double_hashed_key_metadata,
|
||||
recover_key_metadata_from_spend_logs,
|
||||
)
|
||||
|
|
@ -551,11 +552,12 @@ async def get_api_key_metadata(
|
|||
e,
|
||||
)
|
||||
|
||||
still_missing: Final = api_keys - frozenset(result)
|
||||
from_session_keys: Final = await recover_cli_session_key_metadata(prisma_client, api_keys - frozenset(result))
|
||||
still_missing: Final = api_keys - frozenset(result) - frozenset(from_session_keys)
|
||||
from_reverse_hash: Final = (
|
||||
await recover_double_hashed_key_metadata(prisma_client, still_missing) if still_missing else _EMPTY_KEY_METADATA
|
||||
)
|
||||
after_token_recovery: Final = MappingProxyType({**result, **from_reverse_hash})
|
||||
after_token_recovery: Final = MappingProxyType({**result, **from_session_keys, **from_reverse_hash})
|
||||
unresolved: Final = api_keys - frozenset(after_token_recovery)
|
||||
from_spend_logs: Final = (
|
||||
await recover_key_metadata_from_spend_logs(prisma_client, unresolved, spend_logs_window)
|
||||
|
|
@ -563,7 +565,7 @@ async def get_api_key_metadata(
|
|||
else _EMPTY_KEY_METADATA
|
||||
)
|
||||
combined: Final = MappingProxyType({**after_token_recovery, **from_spend_logs})
|
||||
return await attach_user_emails(prisma_client, combined)
|
||||
return await attach_user_details(prisma_client, combined)
|
||||
|
||||
|
||||
def _adjust_dates_for_timezone(
|
||||
|
|
|
|||
|
|
@ -607,7 +607,7 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
# body that mirrors them cannot clobber the authenticated key, the real
|
||||
# parent span, or the proxy's own session-id decision.
|
||||
_metadata.pop(SESSION_ID_OMITTED_METADATA_KEY, None)
|
||||
_metadata["user_api_key"] = user_api_key_dict.api_key
|
||||
_metadata["user_api_key"] = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
|
||||
_metadata["litellm_parent_otel_span"] = user_api_key_dict.parent_otel_span
|
||||
_metadata["user_api_key_budget_reservation"] = user_api_key_dict.budget_reservation
|
||||
_metadata[MODEL_ACCESS_GROUP_METADATA_KEY] = user_api_key_dict.matched_model_access_groups
|
||||
|
|
|
|||
|
|
@ -554,7 +554,7 @@ from litellm.proxy.list_api.common import (
|
|||
problem_response,
|
||||
request_validation_problem,
|
||||
)
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, add_litellm_data_to_request
|
||||
from litellm.proxy.logging_endpoints.callback_logs_endpoints import (
|
||||
rust_control_plane_router,
|
||||
)
|
||||
|
|
@ -16491,8 +16491,9 @@ async def async_queue_request(
|
|||
# Covers both missing and JSON-string metadata (multipart /
|
||||
# extra_body); see above for the same guard upstream.
|
||||
data["metadata"] = {}
|
||||
data["metadata"]["user_api_key"] = user_api_key_dict.api_key
|
||||
data["metadata"]["user_api_key_hash"] = user_api_key_dict.api_key
|
||||
logged_api_key: Final = LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict)
|
||||
data["metadata"]["user_api_key"] = logged_api_key
|
||||
data["metadata"]["user_api_key_hash"] = logged_api_key
|
||||
data["metadata"]["user_api_key_metadata"] = strip_callback_config(user_api_key_dict.metadata)
|
||||
_headers: Final = _safe_get_request_headers(request).copy()
|
||||
_headers.pop("authorization", None) # do not store the original `sk-..` api key in the db
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import asyncio
|
||||
from collections.abc import Awaitable, Callable, Mapping, Sequence
|
||||
from collections.abc import Set as AbstractSet
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta
|
||||
from types import MappingProxyType
|
||||
from typing import Final, TypeVar
|
||||
|
|
@ -11,6 +12,7 @@ from typing_extensions import ReadOnly, TypedDict
|
|||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.caching.in_memory_cache import InMemoryCache
|
||||
from litellm.constants import (
|
||||
CLI_SESSION_KEY_PREFIX,
|
||||
SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS,
|
||||
SPEND_LOG_KEY_METADATA_CACHE_TTL,
|
||||
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL,
|
||||
|
|
@ -62,6 +64,7 @@ _SPEND_LOG_STATEMENT_TIMEOUT_SQL: Final = f"SET LOCAL statement_timeout = {SPEND
|
|||
_SPEND_LOG_TRANSACTION_TIMEOUT: Final = timedelta(milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS)
|
||||
|
||||
_HASHED_JWT_PREFIX: Final = "hashed-jwt-"
|
||||
_CLI_SESSION_KEY_PREFIX: Final = f"{CLI_SESSION_KEY_PREFIX}-"
|
||||
|
||||
|
||||
class KeyMetadataDict(TypedDict, total=False):
|
||||
|
|
@ -109,7 +112,6 @@ _SPEND_LOG_METADATA_CACHE: Final = InMemoryCache(
|
|||
)
|
||||
_SPEND_LOG_QUERY_LOCK: Final = asyncio.Lock()
|
||||
_EMPTY_KEY_METADATA: Final[Mapping[str, KeyMetadataDict]] = MappingProxyType({})
|
||||
_EMPTY_EMAILS: Final[Mapping[str, str]] = MappingProxyType({})
|
||||
|
||||
|
||||
async def _db_or_empty(
|
||||
|
|
@ -149,54 +151,107 @@ async def _reverse_hash_key_metadata(
|
|||
)
|
||||
|
||||
|
||||
async def _emails_for_user_ids(
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class _UserDetails:
|
||||
email: str | None
|
||||
only_team: str | None
|
||||
|
||||
|
||||
_EMPTY_USER_DETAILS: Final[Mapping[str, _UserDetails]] = MappingProxyType({})
|
||||
|
||||
|
||||
async def _details_for_user_ids(
|
||||
prisma_client: PrismaClient,
|
||||
user_ids: AbstractSet[str],
|
||||
) -> Mapping[str, str]:
|
||||
) -> Mapping[str, _UserDetails]:
|
||||
if not user_ids:
|
||||
return _EMPTY_EMAILS
|
||||
return _EMPTY_USER_DETAILS
|
||||
users: Final = await _db_or_empty(
|
||||
lambda: UserRepository(prisma_client).table.find_many(
|
||||
where={"user_id": {"in": list(user_ids)}}, # mutable-ok: Prisma find_many where= is a dict
|
||||
),
|
||||
"Failed user_email recovery for %d user ids: %s",
|
||||
"Failed user detail recovery for %d user ids: %s",
|
||||
len(user_ids),
|
||||
)
|
||||
if users is None:
|
||||
return _EMPTY_EMAILS
|
||||
return _EMPTY_USER_DETAILS
|
||||
return MappingProxyType(
|
||||
{
|
||||
user.user_id: user.user_email
|
||||
user.user_id: _UserDetails(
|
||||
email=getattr(user, "user_email", None) or None,
|
||||
only_team=_only_team(getattr(user, "teams", None)),
|
||||
)
|
||||
for user in users
|
||||
if getattr(user, "user_id", None) and getattr(user, "user_email", None)
|
||||
if getattr(user, "user_id", None)
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
def _meta_with_email(meta: KeyMetadataDict, emails: Mapping[str, str]) -> KeyMetadataDict:
|
||||
if meta.get("user_email"):
|
||||
return meta
|
||||
def _only_team(teams: object) -> str | None:
|
||||
if not isinstance(teams, list) or len(teams) != 1:
|
||||
return None
|
||||
team: Final = teams[0]
|
||||
return team if isinstance(team, str) and team else None
|
||||
|
||||
|
||||
def _is_cli_session_key(api_key: str) -> bool:
|
||||
return api_key.startswith(_CLI_SESSION_KEY_PREFIX) and len(api_key) > len(_CLI_SESSION_KEY_PREFIX)
|
||||
|
||||
|
||||
def _meta_with_user_details(
|
||||
api_key: str, meta: KeyMetadataDict, details: Mapping[str, _UserDetails]
|
||||
) -> KeyMetadataDict:
|
||||
user_id: Final = meta.get("user_id")
|
||||
if not isinstance(user_id, str) or user_id not in emails:
|
||||
if not isinstance(user_id, str) or user_id not in details:
|
||||
return meta
|
||||
updated: Final[KeyMetadataDict] = {**meta, "user_email": emails[user_id]}
|
||||
user: Final = details[user_id]
|
||||
email: Final = meta.get("user_email") or user.email
|
||||
team_id: Final = meta.get("team_id") or (user.only_team if _is_cli_session_key(api_key) else None)
|
||||
updated: Final[KeyMetadataDict] = {
|
||||
**meta,
|
||||
**({"user_email": email} if email else {}),
|
||||
**({"team_id": team_id} if team_id else {}),
|
||||
}
|
||||
return updated
|
||||
|
||||
|
||||
async def attach_user_emails(
|
||||
async def attach_user_details(
|
||||
prisma_client: PrismaClient,
|
||||
recovered: Mapping[str, KeyMetadataDict],
|
||||
) -> Mapping[str, KeyMetadataDict]:
|
||||
needing_email: Final = frozenset(
|
||||
needing_details: Final = frozenset(
|
||||
user_id
|
||||
for meta in recovered.values()
|
||||
for api_key, meta in recovered.items()
|
||||
for user_id in (meta.get("user_id"),)
|
||||
if isinstance(user_id, str) and user_id and not meta.get("user_email")
|
||||
if isinstance(user_id, str)
|
||||
and user_id
|
||||
and (not meta.get("user_email") or (_is_cli_session_key(api_key) and not meta.get("team_id")))
|
||||
)
|
||||
emails: Final = await _emails_for_user_ids(prisma_client, needing_email)
|
||||
if not emails:
|
||||
details: Final = await _details_for_user_ids(prisma_client, needing_details)
|
||||
if not details:
|
||||
return recovered
|
||||
return MappingProxyType({api_key: _meta_with_email(meta, emails) for api_key, meta in recovered.items()})
|
||||
return MappingProxyType(
|
||||
{api_key: _meta_with_user_details(api_key, meta, details) for api_key, meta in recovered.items()}
|
||||
)
|
||||
|
||||
|
||||
async def recover_cli_session_key_metadata(
|
||||
prisma_client: PrismaClient,
|
||||
missing_keys: AbstractSet[str],
|
||||
) -> Mapping[str, KeyMetadataDict]:
|
||||
candidates: Final = MappingProxyType(
|
||||
{key: key.removeprefix(_CLI_SESSION_KEY_PREFIX) for key in missing_keys if _is_cli_session_key(key)}
|
||||
)
|
||||
if not candidates:
|
||||
return _EMPTY_KEY_METADATA
|
||||
known_users: Final = await _details_for_user_ids(prisma_client, frozenset(candidates.values()))
|
||||
return MappingProxyType(
|
||||
{
|
||||
key: KeyMetadataDict(key_alias=key, user_id=user_id)
|
||||
for key, user_id in candidates.items()
|
||||
if user_id in known_users
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
async def recover_double_hashed_key_metadata(
|
||||
|
|
@ -384,9 +439,15 @@ async def fill_missing_api_key_aliases(
|
|||
if not missing_keys:
|
||||
return tuple(rows)
|
||||
|
||||
recovered: Final = await attach_user_emails(
|
||||
from_session_keys: Final = await recover_cli_session_key_metadata(prisma_client, missing_keys)
|
||||
recovered: Final = await attach_user_details(
|
||||
prisma_client,
|
||||
await recover_double_hashed_key_metadata(prisma_client, missing_keys),
|
||||
MappingProxyType(
|
||||
{
|
||||
**from_session_keys,
|
||||
**await recover_double_hashed_key_metadata(prisma_client, missing_keys - frozenset(from_session_keys)),
|
||||
}
|
||||
),
|
||||
)
|
||||
if not recovered:
|
||||
return tuple(rows)
|
||||
|
|
|
|||
|
|
@ -38,6 +38,7 @@ from litellm.litellm_core_utils.classifier_logging import classifier_audit_field
|
|||
from litellm.proxy._types import *
|
||||
from litellm.proxy._types import ProviderBudgetResponse, ProviderBudgetResponseObject
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
|
||||
from litellm.proxy.spend_tracking.spend_capture_rate import (
|
||||
ProviderBillingCredentialMissing,
|
||||
ProviderBillingRequestFailed,
|
||||
|
|
@ -1925,7 +1926,7 @@ async def get_key_spend_report(
|
|||
scoped_api_key = _resolve_spend_report_scope(
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
requested=requested,
|
||||
caller_value=user_api_key_dict.api_key,
|
||||
caller_value=LiteLLMProxyRequestSetup.get_logged_api_key(user_api_key_dict),
|
||||
scope_name="api_key",
|
||||
)
|
||||
db_response: Sequence[Mapping[str, object]] | None = await _query_raw_or_none(
|
||||
|
|
|
|||
|
|
@ -14,6 +14,7 @@ from pydantic import BaseModel, JsonValue
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.constants import (
|
||||
CLI_SESSION_KEY_PREFIX,
|
||||
EMPTY_MAPPING,
|
||||
LITELLM_PROXY_MASTER_KEY_ALIAS,
|
||||
LITELLM_TRUNCATED_PAYLOAD_FIELD,
|
||||
|
|
@ -102,19 +103,28 @@ _NON_SECRET_KEY_ALIASES: Final = frozenset(
|
|||
)
|
||||
|
||||
|
||||
def _is_non_secret_key_value(value: str) -> bool:
|
||||
def _is_cli_session_alias(value: str, key_alias: object) -> bool:
|
||||
return value.startswith(f"{CLI_SESSION_KEY_PREFIX}-") and value == key_alias
|
||||
|
||||
|
||||
def _is_non_secret_key_value(value: str, *, key_alias: object = None) -> bool:
|
||||
return (
|
||||
value in _NON_SECRET_KEY_ALIASES or is_valid_sha256_hash(value) or _HASHED_JWT_RE.fullmatch(value) is not None
|
||||
value in _NON_SECRET_KEY_ALIASES
|
||||
or is_valid_sha256_hash(value)
|
||||
or _HASHED_JWT_RE.fullmatch(value) is not None
|
||||
or _is_cli_session_alias(value, key_alias)
|
||||
)
|
||||
|
||||
|
||||
def _redact_logged_api_key(value: str | None, *, already_redacted: bool = False) -> str | None:
|
||||
def _redact_logged_api_key(
|
||||
value: str | None, *, already_redacted: bool = False, key_alias: object = None
|
||||
) -> str | None:
|
||||
if not isinstance(value, str) or not value:
|
||||
return None
|
||||
stripped: Final = re.sub(r"(?i)^bearer ", "", value)
|
||||
if not stripped:
|
||||
return None
|
||||
if already_redacted and _is_non_secret_key_value(stripped):
|
||||
if already_redacted and _is_non_secret_key_value(stripped, key_alias=key_alias):
|
||||
return stripped
|
||||
return hash_token(stripped)
|
||||
|
||||
|
|
@ -230,10 +240,15 @@ def _get_spend_logs_metadata(
|
|||
)
|
||||
_raw_key: Final = clean_metadata.get("user_api_key")
|
||||
_trusted_hash: Final = metadata.get("user_api_key_hash")
|
||||
_key_alias: Final = metadata.get("user_api_key_alias")
|
||||
_already_redacted: Final = (
|
||||
isinstance(_trusted_hash, str) and _is_non_secret_key_value(_trusted_hash) and _trusted_hash == _raw_key
|
||||
isinstance(_trusted_hash, str)
|
||||
and _is_non_secret_key_value(_trusted_hash, key_alias=_key_alias)
|
||||
and _trusted_hash == _raw_key
|
||||
)
|
||||
clean_metadata["user_api_key"] = _redact_logged_api_key(
|
||||
_raw_key, already_redacted=_already_redacted, key_alias=_key_alias
|
||||
)
|
||||
clean_metadata["user_api_key"] = _redact_logged_api_key(_raw_key, already_redacted=_already_redacted)
|
||||
clean_metadata["applied_guardrails"] = applied_guardrails
|
||||
clean_metadata["batch_models"] = batch_models
|
||||
clean_metadata["batch_successful_requests"] = batch_successful_requests
|
||||
|
|
@ -537,10 +552,13 @@ def get_logging_payload(
|
|||
standard_logging_completion_tokens = standard_logging_payload.get("completion_tokens", 0)
|
||||
standard_logging_total_tokens = standard_logging_payload.get("total_tokens", 0)
|
||||
_trusted_hash = metadata.get("user_api_key_hash")
|
||||
_key_alias = metadata.get("user_api_key_alias")
|
||||
_key_already_redacted = (
|
||||
isinstance(_trusted_hash, str) and _is_non_secret_key_value(_trusted_hash) and _trusted_hash == api_key
|
||||
isinstance(_trusted_hash, str)
|
||||
and _is_non_secret_key_value(_trusted_hash, key_alias=_key_alias)
|
||||
and _trusted_hash == api_key
|
||||
)
|
||||
api_key = _redact_logged_api_key(api_key, already_redacted=_key_already_redacted) or ""
|
||||
api_key = _redact_logged_api_key(api_key, already_redacted=_key_already_redacted, key_alias=_key_alias) or ""
|
||||
|
||||
if (
|
||||
standard_logging_payload is not None
|
||||
|
|
@ -548,7 +566,9 @@ def get_logging_payload(
|
|||
api_key = (
|
||||
api_key
|
||||
or _redact_logged_api_key(
|
||||
standard_logging_payload["metadata"].get("user_api_key_hash"), already_redacted=True
|
||||
standard_logging_payload["metadata"].get("user_api_key_hash"),
|
||||
already_redacted=True,
|
||||
key_alias=standard_logging_payload["metadata"].get("user_api_key_alias"),
|
||||
)
|
||||
or ""
|
||||
)
|
||||
|
|
|
|||
|
|
@ -716,6 +716,29 @@ async def test_failure_hook_emits_api_provider_value_on_failed_requests_metric()
|
|||
_clear_prometheus_registry()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failure_hook_labels_a_cli_session_with_the_per_user_alias_not_the_login_token():
|
||||
from litellm.integrations.prometheus import PrometheusLogger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
_clear_prometheus_registry()
|
||||
try:
|
||||
await PrometheusLogger().async_post_call_failure_hook(
|
||||
request_data={"model": "gpt-4o-mini", "metadata": {}},
|
||||
original_exception=Exception("boom"),
|
||||
user_api_key_dict=UserAPIKeyAuth(
|
||||
api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa",
|
||||
user_id="alice",
|
||||
key_alias="cli-session-alice",
|
||||
is_session_token=True,
|
||||
),
|
||||
)
|
||||
hashed_keys = {s.labels.get("hashed_api_key") for s in _collected_samples("litellm_proxy_failed_requests_metric_total")}
|
||||
assert hashed_keys == {"cli-session-alice"}, hashed_keys
|
||||
finally:
|
||||
_clear_prometheus_registry()
|
||||
|
||||
|
||||
async def _failed_requests_api_provider_labels(
|
||||
request_data: dict[str, object],
|
||||
original_exception: Exception,
|
||||
|
|
|
|||
|
|
@ -230,6 +230,39 @@ async def test_execute_search_passes_selected_search_tool_litellm_params(monkeyp
|
|||
assert forwarded_kwargs["max_retries"] == 2
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_search_attributes_cli_session_spend_to_the_per_user_alias_not_the_login_token(monkeypatch):
|
||||
import litellm
|
||||
from litellm.proxy import proxy_server
|
||||
|
||||
logger = WebSearchInterceptionLogger(enabled_providers=["bedrock"], search_tool_name="perplexity-sonar-pro")
|
||||
router = MagicMock()
|
||||
router.search_tools = [
|
||||
{
|
||||
"search_tool_name": "perplexity-sonar-pro",
|
||||
"litellm_params": {"search_provider": "perplexity", "api_key": "fake-key"},
|
||||
}
|
||||
]
|
||||
mock_asearch = AsyncMock(return_value=SearchResponse(object="search", results=[]))
|
||||
session = UserAPIKeyAuth(
|
||||
api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa",
|
||||
user_id="alice",
|
||||
key_alias="cli-session-alice",
|
||||
is_session_token=True,
|
||||
)
|
||||
monkeypatch.setattr(proxy_server, "llm_router", router)
|
||||
monkeypatch.setattr(litellm, "asearch", mock_asearch)
|
||||
|
||||
await logger._execute_search(
|
||||
"what is litellm",
|
||||
kwargs={"litellm_params": {"metadata": {"user_api_key_auth": session}}},
|
||||
)
|
||||
|
||||
forwarded_metadata = mock_asearch.await_args.kwargs["litellm_metadata"]
|
||||
assert forwarded_metadata["user_api_key"] == "cli-session-alice"
|
||||
assert forwarded_metadata["user_api_key_hash"] == "cli-session-alice"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_search_attributes_spend_to_the_calling_key(monkeypatch):
|
||||
"""An intercepted search is billed and logged against the key that made the LLM request.
|
||||
|
|
|
|||
|
|
@ -378,6 +378,7 @@ def make_runner(
|
|||
general_settings: Mapping[str, object] = MappingProxyType({}),
|
||||
heartbeat_seconds: float = 30.0,
|
||||
completion_window_seconds: float = 24 * 60 * 60,
|
||||
user: UserAPIKeyAuth | None = None,
|
||||
) -> Harness:
|
||||
store = store_factory({INPUT_FILE_ID: managed_input_file()} if files is None else files)
|
||||
router = FakeRouter()
|
||||
|
|
@ -385,7 +386,7 @@ def make_runner(
|
|||
storage = FakeStorageBackend({STORAGE_URL: content})
|
||||
storage_factory = FakeStorageBackendFactory(storage, storage_error)
|
||||
prisma = FakePrismaClient(store.objects)
|
||||
user = UserAPIKeyAuth(
|
||||
user = user or UserAPIKeyAuth(
|
||||
api_key="sk-batch-key", user_id="user-1", team_id="team-1", key_alias="alias-1", user_email="user@example.com"
|
||||
)
|
||||
runner = LiteLLMExecutedBatchRunner(
|
||||
|
|
@ -668,6 +669,20 @@ async def test_create_dispatches_each_row_with_the_batch_model_and_the_key_metad
|
|||
assert metadata["user_api_key_user_email"] == "user@example.com"
|
||||
|
||||
|
||||
async def test_create_dispatches_cli_session_rows_under_the_per_user_alias_not_the_login_token() -> None:
|
||||
session = UserAPIKeyAuth(
|
||||
api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa",
|
||||
user_id="alice",
|
||||
key_alias="cli-session-alice",
|
||||
is_session_token=True,
|
||||
)
|
||||
harness = make_runner(user=session)
|
||||
await harness.create_and_finish()
|
||||
|
||||
logged_keys = {call.kwargs["metadata"]["user_api_key"] for call in harness.router.acompletion.await_args_list}
|
||||
assert logged_keys == {"cli-session-alice"}
|
||||
|
||||
|
||||
async def test_create_uploads_one_output_line_per_row_with_the_router_response() -> None:
|
||||
harness = make_runner()
|
||||
replies = {"hi 1": chat_response("hi 1"), "hi 2": chat_response("hi 2")}
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from datetime import datetime
|
|||
import pytest
|
||||
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.hooks.parallel_request_limiter import (
|
||||
_PROXY_MaxParallelRequestsHandler,
|
||||
)
|
||||
|
|
@ -14,6 +15,28 @@ from litellm.proxy.utils import InternalUsageCache, hash_token
|
|||
from litellm.types.utils import EmbeddingResponse, TextCompletionResponse, Usage
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_pre_call_hook_counts_a_cli_session_under_the_per_user_alias_not_the_login_token():
|
||||
handler = _PROXY_MaxParallelRequestsHandler(internal_usage_cache=InternalUsageCache(DualCache()))
|
||||
session = UserAPIKeyAuth(
|
||||
api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa",
|
||||
user_id="alice",
|
||||
key_alias="cli-session-alice",
|
||||
is_session_token=True,
|
||||
max_parallel_requests=5,
|
||||
)
|
||||
|
||||
await handler.async_pre_call_hook(
|
||||
user_api_key_dict=session, cache=DualCache(), data={"model": "gpt-4o-mini"}, call_type="completion"
|
||||
)
|
||||
|
||||
precise_minute = datetime.now().strftime("%Y-%m-%d-%H-%M")
|
||||
counted = await handler.internal_usage_cache.async_get_cache(
|
||||
key=f"cli-session-alice::{precise_minute}::request_count", litellm_parent_otel_span=None
|
||||
)
|
||||
assert counted == {"current_requests": 1, "current_tpm": 0, "current_rpm": 1}, counted
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"response_obj",
|
||||
[
|
||||
|
|
|
|||
|
|
@ -2871,6 +2871,23 @@ def test_spend_logs_window_is_none_when_no_date_parses():
|
|||
assert _spend_logs_window({"garbage", ""}) is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_api_key_metadata_resolves_cli_session_keys_from_the_key_itself():
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[])
|
||||
mock_prisma.db.query_raw = AsyncMock(side_effect=AssertionError("no reverse-hash or spend-log scan expected"))
|
||||
mock_prisma.db.litellm_usertable.find_many = AsyncMock(
|
||||
return_value=[SimpleNamespace(user_id="alice", user_email="alice@example.com", teams=["team-a"])]
|
||||
)
|
||||
|
||||
result = await get_api_key_metadata(prisma_client=mock_prisma, api_keys={"cli-session-alice"})
|
||||
|
||||
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"
|
||||
|
||||
|
||||
_DAILY_TEAM_SPEND_DDL: Final = """
|
||||
CREATE TABLE "LiteLLM_DailyTeamSpend" (
|
||||
id TEXT PRIMARY KEY,
|
||||
|
|
|
|||
|
|
@ -1451,6 +1451,7 @@ async def test_pass_through_request_contains_proxy_server_request_in_kwargs():
|
|||
mock_user_api_key_dict = MagicMock()
|
||||
mock_user_api_key_dict.api_key = "test-api-key"
|
||||
mock_user_api_key_dict.key_alias = "test-alias"
|
||||
mock_user_api_key_dict.is_session_token = False
|
||||
mock_user_api_key_dict.user_email = "test@example.com"
|
||||
mock_user_api_key_dict.user_id = "test-user-id"
|
||||
mock_user_api_key_dict.team_id = "test-team-id"
|
||||
|
|
@ -7416,3 +7417,32 @@ async def test_chat_completion_pass_through_endpoint_failure_carries_the_callers
|
|||
record = next(r for r in caplog.records if "Exception occured" in r.getMessage())
|
||||
assert record.litellm_call_id == call_id
|
||||
assert call_id in record.getMessage()
|
||||
|
||||
|
||||
def test_passthrough_attributes_a_cli_session_to_its_alias_not_the_login_token():
|
||||
from litellm.proxy.spend_tracking.spend_tracking_utils import _get_spend_logs_metadata
|
||||
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = "POST"
|
||||
mock_request.url = "http://0.0.0.0:4000/anthropic/v1/messages"
|
||||
mock_request.headers = Headers({})
|
||||
mock_request.scope = {}
|
||||
session = UserAPIKeyAuth(
|
||||
api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa",
|
||||
key_alias="cli-session-alice",
|
||||
user_id="alice",
|
||||
is_session_token=True,
|
||||
)
|
||||
|
||||
kwargs = HttpPassThroughEndpointHelpers._init_kwargs_for_pass_through_endpoint(
|
||||
request=mock_request,
|
||||
user_api_key_dict=session,
|
||||
passthrough_logging_payload=MagicMock(),
|
||||
logging_obj=MagicMock(),
|
||||
_parsed_body={},
|
||||
litellm_call_id="lit-6852-passthrough-call-id",
|
||||
)
|
||||
|
||||
metadata = kwargs["litellm_params"]["metadata"]
|
||||
assert metadata["user_api_key"] == "cli-session-alice"
|
||||
assert _get_spend_logs_metadata(metadata)["user_api_key"] == "cli-session-alice"
|
||||
|
|
|
|||
|
|
@ -15,7 +15,9 @@ from litellm.constants import (
|
|||
SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS,
|
||||
)
|
||||
from litellm.proxy.spend_tracking.key_metadata_recovery import (
|
||||
attach_user_details,
|
||||
fill_missing_api_key_aliases,
|
||||
recover_cli_session_key_metadata,
|
||||
recover_double_hashed_key_metadata,
|
||||
recover_key_metadata_from_spend_logs,
|
||||
)
|
||||
|
|
@ -588,3 +590,77 @@ async def test_recover_key_metadata_from_spend_logs_bounds_the_scan_with_a_state
|
|||
assert mock_prisma.db.tx.call_args.kwargs["timeout"] == timedelta(
|
||||
milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS
|
||||
)
|
||||
|
||||
|
||||
@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()
|
||||
mock_prisma.db.litellm_usertable.find_many = AsyncMock(
|
||||
return_value=[SimpleNamespace(user_id="alice", user_email="alice@example.com", teams=[])]
|
||||
)
|
||||
raw_login_token = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa"
|
||||
|
||||
result = await recover_cli_session_key_metadata(
|
||||
mock_prisma, {"cli-session-alice", raw_login_token, hash_token("sk-other"), "cli-session-", "sk-raw"}
|
||||
)
|
||||
|
||||
assert dict(result) == {"cli-session-alice": {"key_alias": "cli-session-alice", "user_id": "alice"}}
|
||||
assert sorted(mock_prisma.db.litellm_usertable.find_many.call_args.kwargs["where"]["user_id"]["in"]) == [
|
||||
"Qm7xJ2kP9sLw4vT1nR8yAa",
|
||||
"alice",
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fill_missing_api_key_aliases_resolves_cli_session_keys_without_a_digest_lookup():
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.query_raw = AsyncMock(side_effect=AssertionError("no reverse-hash lookup expected"))
|
||||
mock_prisma.db.litellm_usertable.find_many = AsyncMock(
|
||||
return_value=[SimpleNamespace(user_id="alice", user_email="alice@example.com", teams=["team-a"])]
|
||||
)
|
||||
rows = ({"api_key": "cli-session-alice", "api_key_alias": None, "team_id": None, "user_email": None, "spend": 2.0},)
|
||||
|
||||
filled = await fill_missing_api_key_aliases(mock_prisma, rows)
|
||||
|
||||
assert filled[0]["api_key_alias"] == "cli-session-alice"
|
||||
assert filled[0]["user_email"] == "alice@example.com"
|
||||
assert filled[0]["team_id"] == "team-a"
|
||||
assert filled[0]["spend"] == 2.0
|
||||
assert mock_prisma.db.litellm_usertable.find_many.call_args.kwargs["where"] == {"user_id": {"in": ["alice"]}}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attach_user_details_gives_the_login_team_only_to_cli_session_keys():
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_usertable.find_many = AsyncMock(
|
||||
return_value=[SimpleNamespace(user_id="alice", user_email="alice@example.com", teams=["team-a"])]
|
||||
)
|
||||
personal_key = hash_token("sk-personal")
|
||||
|
||||
attached = await attach_user_details(
|
||||
mock_prisma,
|
||||
{
|
||||
"cli-session-alice": {"key_alias": "cli-session-alice", "user_id": "alice"},
|
||||
personal_key: {"key_alias": "personal", "user_id": "alice"},
|
||||
},
|
||||
)
|
||||
|
||||
assert attached["cli-session-alice"]["team_id"] == "team-a"
|
||||
assert attached["cli-session-alice"]["user_email"] == "alice@example.com"
|
||||
assert "team_id" not in attached[personal_key]
|
||||
assert attached[personal_key]["user_email"] == "alice@example.com"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attach_user_details_claims_no_team_for_a_multi_team_user_session_key():
|
||||
mock_prisma = MagicMock()
|
||||
mock_prisma.db.litellm_usertable.find_many = AsyncMock(
|
||||
return_value=[SimpleNamespace(user_id="bob", user_email="bob@example.com", teams=["team-a", "team-b"])]
|
||||
)
|
||||
|
||||
attached = await attach_user_details(
|
||||
mock_prisma, {"cli-session-bob": {"key_alias": "cli-session-bob", "user_id": "bob"}}
|
||||
)
|
||||
|
||||
assert "team_id" not in attached["cli-session-bob"]
|
||||
assert attached["cli-session-bob"]["user_email"] == "bob@example.com"
|
||||
|
|
|
|||
|
|
@ -6585,6 +6585,33 @@ def test_key_spend_report_scopes_to_caller_key(client, monkeypatch):
|
|||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_key_spend_report_scopes_a_cli_session_to_the_per_user_alias_not_the_login_token(client, monkeypatch):
|
||||
mock_prisma = _spend_report_mock_prisma(
|
||||
query_raw_returns=[{"api_key": "cli-session-alice", "total_cost": 1.5}]
|
||||
)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", True)
|
||||
app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(
|
||||
user_role=LitellmUserRoles.INTERNAL_USER,
|
||||
user_id="alice",
|
||||
api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa",
|
||||
key_alias="cli-session-alice",
|
||||
is_session_token=True,
|
||||
)
|
||||
try:
|
||||
response = client.get(
|
||||
"/key/spend/report",
|
||||
params={"start_date": "2026-07-01", "end_date": "2026-07-31", "api_key": "cli-session-alice"},
|
||||
headers={"Authorization": "Bearer sk-test"},
|
||||
)
|
||||
assert response.status_code == 200, response.text
|
||||
assert response.json() == [{"api_key": "cli-session-alice", "total_cost": 1.5}]
|
||||
args, _ = mock_prisma.db.query_raw.await_args
|
||||
assert args[3] == "cli-session-alice"
|
||||
finally:
|
||||
app.dependency_overrides.pop(ps.user_api_key_auth, None)
|
||||
|
||||
|
||||
def test_key_spend_report_non_admin_override_403(client, monkeypatch):
|
||||
mock_prisma = _spend_report_mock_prisma()
|
||||
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
|
||||
|
|
|
|||
|
|
@ -5156,6 +5156,64 @@ def test_spend_log_request_id_is_the_response_id_a_bridged_messages_caller_recei
|
|||
)
|
||||
|
||||
|
||||
_CLI_SESSION_ALIAS: Final = "cli-session-alice"
|
||||
_CLI_SESSION_TOKEN: Final = "cli-session-Qm7xJ2kP9sLw4vT1nR8yAa"
|
||||
|
||||
|
||||
def _cli_session_request_metadata(logged_key: str) -> dict[str, str]:
|
||||
return {
|
||||
"user_api_key": logged_key,
|
||||
"user_api_key_hash": logged_key,
|
||||
"user_api_key_alias": _CLI_SESSION_ALIAS,
|
||||
"user_api_key_user_id": "alice",
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("call_type", ["acompletion", "anthropic_messages", "aresponses"])
|
||||
def test_get_logging_payload_attributes_a_cli_session_to_its_alias(call_type: str):
|
||||
payload = get_logging_payload(
|
||||
kwargs={
|
||||
"call_type": call_type,
|
||||
"model": "gpt-5.4-nano",
|
||||
"response_cost": 0.00001,
|
||||
"litellm_params": {"metadata": _cli_session_request_metadata(_CLI_SESSION_ALIAS)},
|
||||
},
|
||||
response_obj=litellm.ModelResponse(id=f"{call_type}-1", choices=[], usage=litellm.Usage()),
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
end_time=datetime.datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
assert payload["api_key"] == _CLI_SESSION_ALIAS
|
||||
assert json.loads(payload["metadata"])["user_api_key"] == _CLI_SESSION_ALIAS
|
||||
assert json.loads(payload["metadata"])["user_api_key_alias"] == _CLI_SESSION_ALIAS
|
||||
|
||||
|
||||
@pytest.mark.parametrize("call_type", ["acompletion", "anthropic_messages", "aresponses"])
|
||||
def test_get_logging_payload_never_lands_a_raw_cli_session_token(call_type: str):
|
||||
payload = get_logging_payload(
|
||||
kwargs={
|
||||
"call_type": call_type,
|
||||
"model": "gpt-5.4-nano",
|
||||
"litellm_params": {"metadata": _cli_session_request_metadata(_CLI_SESSION_TOKEN)},
|
||||
},
|
||||
response_obj=litellm.ModelResponse(id=f"{call_type}-2", choices=[], usage=litellm.Usage()),
|
||||
start_time=datetime.datetime.now(timezone.utc),
|
||||
end_time=datetime.datetime.now(timezone.utc),
|
||||
)
|
||||
|
||||
assert payload["api_key"] == hash_token(_CLI_SESSION_TOKEN)
|
||||
assert json.loads(payload["metadata"])["user_api_key"] == hash_token(_CLI_SESSION_TOKEN)
|
||||
|
||||
|
||||
def test_redact_logged_api_key_cli_session_alias_needs_alias_provenance():
|
||||
assert _redact_logged_api_key(_CLI_SESSION_ALIAS, already_redacted=True, key_alias=_CLI_SESSION_ALIAS) == (
|
||||
_CLI_SESSION_ALIAS
|
||||
)
|
||||
assert _redact_logged_api_key(_CLI_SESSION_ALIAS, already_redacted=True) == hash_token(_CLI_SESSION_ALIAS)
|
||||
assert _redact_logged_api_key(_CLI_SESSION_ALIAS, key_alias=_CLI_SESSION_ALIAS) == hash_token(_CLI_SESSION_ALIAS)
|
||||
assert _redact_logged_api_key("alice", already_redacted=True, key_alias="alice") == hash_token("alice")
|
||||
|
||||
|
||||
def test_azure_spillover_stamped_from_response_headers():
|
||||
"""Raw provider response headers on the logging kwargs mark the request as spilled."""
|
||||
kwargs: Final = {
|
||||
|
|
|
|||
|
|
@ -8396,6 +8396,40 @@ def test_default_team_settings_bool_turn_off_message_logging_redacts():
|
|||
)
|
||||
|
||||
|
||||
def test_add_user_api_key_auth_to_request_metadata_attributes_a_cli_session_to_its_alias():
|
||||
data = {"model": "gpt-5.4-nano", "messages": [{"role": "user", "content": "hi"}], "litellm_metadata": {}}
|
||||
session = UserAPIKeyAuth(
|
||||
api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa",
|
||||
key_alias="cli-session-alice",
|
||||
user_id="alice",
|
||||
is_session_token=True,
|
||||
)
|
||||
|
||||
metadata = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
|
||||
data=data, user_api_key_dict=session, _metadata_variable_name="litellm_metadata"
|
||||
)["litellm_metadata"]
|
||||
|
||||
assert metadata["user_api_key"] == "cli-session-alice"
|
||||
assert metadata["user_api_key_hash"] == "cli-session-alice"
|
||||
assert metadata["user_api_key_alias"] == "cli-session-alice"
|
||||
assert "Qm7xJ2kP9sLw4vT1nR8yAa" not in (metadata["user_api_key"], metadata["user_api_key_hash"])
|
||||
|
||||
|
||||
def test_add_user_api_key_auth_to_request_metadata_keeps_the_hashed_token_for_virtual_keys():
|
||||
from litellm.proxy._types import hash_token
|
||||
|
||||
data = {"model": "gpt-5.4-nano", "messages": [], "litellm_metadata": {}}
|
||||
hashed = hash_token("sk-virtual-key")
|
||||
virtual_key = UserAPIKeyAuth(api_key=hashed, key_alias="cli-session-alice", user_id="alice")
|
||||
|
||||
metadata = LiteLLMProxyRequestSetup.add_user_api_key_auth_to_request_metadata(
|
||||
data=data, user_api_key_dict=virtual_key, _metadata_variable_name="litellm_metadata"
|
||||
)["litellm_metadata"]
|
||||
|
||||
assert metadata["user_api_key"] == hashed
|
||||
assert metadata["user_api_key_hash"] == hashed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("path", ["/mcp-rest/tools/call", "/v1/responses", "/v1/chat/completions"])
|
||||
@pytest.mark.parametrize("custom_auth", ["x-mcp-auth", "x-private-mcp-token"])
|
||||
|
|
|
|||
|
|
@ -304,6 +304,29 @@ class TestSearchSkills:
|
|||
assert metadata["user_api_key_team_id"] == "team-1"
|
||||
assert metadata["user_api_key_user_id"] == "user-1"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cli_session_embedding_spend_is_attributed_to_the_per_user_alias_not_the_login_token(self) -> None:
|
||||
router = _embedding_router()
|
||||
session = UserAPIKeyAuth(
|
||||
api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa",
|
||||
user_id="alice",
|
||||
key_alias="cli-session-alice",
|
||||
is_session_token=True,
|
||||
)
|
||||
await search_skills(
|
||||
"language translation",
|
||||
SKILLS,
|
||||
1,
|
||||
router=router,
|
||||
embedding_model="text-embedding-3-small",
|
||||
index=SkillSearchIndex(),
|
||||
user_api_key_dict=session,
|
||||
proxy_logging_obj=_pass_through_key_limits(),
|
||||
)
|
||||
metadata = router.aembedding.await_args.kwargs["metadata"]
|
||||
assert metadata["user_api_key"] == "cli-session-alice"
|
||||
assert metadata["user_api_key_hash"] == "cli-session-alice"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_key_limits_are_checked_against_the_real_embedding_call_before_it_runs(self) -> None:
|
||||
router = _embedding_router()
|
||||
|
|
|
|||
|
|
@ -2620,6 +2620,27 @@ class TestBatchCostAttribution:
|
|||
assert metadata["user_api_key"] == "hash-alice"
|
||||
assert metadata.get("user_api_key_alias") is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cli_session_batch_keeps_its_alias_without_a_key_row(self):
|
||||
instance = self._instance(key_row=None)
|
||||
|
||||
metadata = await instance._build_creator_attribution_metadata(
|
||||
self._job(api_key="cli-session-alice"), "batch-1"
|
||||
)
|
||||
|
||||
assert metadata["user_api_key"] == "cli-session-alice"
|
||||
assert metadata["user_api_key_alias"] == "cli-session-alice"
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raw_cli_session_token_on_a_legacy_batch_row_is_not_treated_as_the_alias(self):
|
||||
instance = self._instance(key_row=None)
|
||||
|
||||
metadata = await instance._build_creator_attribution_metadata(
|
||||
self._job(api_key="cli-session-Qm7xJ2kP9sLw4vT1nR8yAa"), "batch-1"
|
||||
)
|
||||
|
||||
assert metadata.get("user_api_key_alias") is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_unnamed_key_keeps_the_creating_user_alias(self):
|
||||
"""Regression: a key generated without key_alias resolves to no alias, and the
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue