From 25fb7810c2668875530610861e82a2768e49a680 Mon Sep 17 00:00:00 2001 From: Mateo Wang <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 24 Sep 2026 18:21:47 -0700 Subject: [PATCH] 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-, 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- 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> --- .../proxy/common_utils/check_batch_cost.py | 7 +- .../proxy/hooks/managed_files.py | 3 +- litellm/integrations/prometheus.py | 5 +- .../websearch_interception/handler.py | 2 +- .../litellm_executed_batches.py | 2 +- .../proxy/common_utils/semantic_text_index.py | 2 +- .../proxy/hooks/parallel_request_limiter.py | 5 +- .../proxy/hooks/proxy_track_cost_callback.py | 4 +- litellm/proxy/litellm_pre_call_utils.py | 10 +- .../common_daily_activity.py | 10 +- .../pass_through_endpoints.py | 2 +- litellm/proxy/proxy_server.py | 7 +- .../spend_tracking/key_metadata_recovery.py | 105 ++++++++++++++---- .../spend_management_endpoints.py | 3 +- .../spend_tracking/spend_tracking_utils.py | 38 +++++-- .../integrations/test_prometheus_labels.py | 23 ++++ .../test_websearch_interception_handler.py | 33 ++++++ .../test_litellm_executed_batches.py | 17 ++- .../hooks/test_parallel_request_limiter.py | 23 ++++ .../test_common_daily_activity.py | 17 +++ .../test_pass_through_endpoints.py | 30 +++++ .../test_key_metadata_recovery.py | 76 +++++++++++++ .../test_spend_management_endpoints.py | 27 +++++ .../test_spend_tracking_utils.py | 58 ++++++++++ .../proxy/test_litellm_pre_call_utils.py | 34 ++++++ .../litellm_proxy/skills/test_skill_search.py | 23 ++++ .../common_utils/test_check_batch_cost.py | 21 ++++ 27 files changed, 533 insertions(+), 54 deletions(-) diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index 41974c26158..35fc3510414 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -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) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 5ac7c1e53c1..21bf7abdc2e 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -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 {}), diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 0396bca942c..c7bf291a887 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -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, diff --git a/litellm/integrations/websearch_interception/handler.py b/litellm/integrations/websearch_interception/handler.py index 90bbd5a00d8..6ebf485d717 100644 --- a/litellm/integrations/websearch_interception/handler.py +++ b/litellm/integrations/websearch_interception/handler.py @@ -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, } diff --git a/litellm/proxy/batches_endpoints/litellm_executed_batches.py b/litellm/proxy/batches_endpoints/litellm_executed_batches.py index 67201d99422..caf34404a7d 100644 --- a/litellm/proxy/batches_endpoints/litellm_executed_batches.py +++ b/litellm/proxy/batches_endpoints/litellm_executed_batches.py @@ -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, diff --git a/litellm/proxy/common_utils/semantic_text_index.py b/litellm/proxy/common_utils/semantic_text_index.py index b8d3595163e..d3fe68e65f7 100644 --- a/litellm/proxy/common_utils/semantic_text_index.py +++ b/litellm/proxy/common_utils/semantic_text_index.py @@ -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), } diff --git a/litellm/proxy/hooks/parallel_request_limiter.py b/litellm/proxy/hooks/parallel_request_limiter.py index d41acadc4dd..e3485ebf25d 100644 --- a/litellm/proxy/hooks/parallel_request_limiter.py +++ b/litellm/proxy/hooks/parallel_request_limiter.py @@ -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") diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 4dfea5cd472..e097debde77 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -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, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 8fc5faee2c9..56f647d5acc 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -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) diff --git a/litellm/proxy/management_endpoints/common_daily_activity.py b/litellm/proxy/management_endpoints/common_daily_activity.py index 1d797206ead..1c674d3640c 100644 --- a/litellm/proxy/management_endpoints/common_daily_activity.py +++ b/litellm/proxy/management_endpoints/common_daily_activity.py @@ -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( diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index a119335ba46..7d6db30e3e3 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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 diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 2a1ba1246d3..a5ae7ec0e44 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -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 diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index ee2e1cfeaf7..08e8fff8f1d 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -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) diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index e491175ea22..c4ed8713f95 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -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( diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 8f85ecdd480..1c51fb21d6e 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -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 "" ) diff --git a/tests/test_litellm/integrations/test_prometheus_labels.py b/tests/test_litellm/integrations/test_prometheus_labels.py index 41d0c44ff89..d598d59e183 100644 --- a/tests/test_litellm/integrations/test_prometheus_labels.py +++ b/tests/test_litellm/integrations/test_prometheus_labels.py @@ -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, diff --git a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py index d6450d0f1de..a5ab28ba72a 100644 --- a/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py +++ b/tests/test_litellm/integrations/websearch_interception/test_websearch_interception_handler.py @@ -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. diff --git a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py index 6f2341a578c..cae92ce18bb 100644 --- a/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py +++ b/tests/test_litellm/proxy/batches_endpoints/test_litellm_executed_batches.py @@ -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")} diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py index 0e2683dcbfd..a83ddc69863 100644 --- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py +++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter.py @@ -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", [ diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index cecdb937c40..a9604ef1296 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -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, diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py index 7a64d5f2218..2d929a832a5 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py @@ -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" diff --git a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py index 89be341c87b..1967d7b6aad 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py +++ b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py @@ -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" diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 0392c3f115b..347adc421a2 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -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) diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index f471e3f8fbb..2f284447dd6 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -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 = { diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index b2241191ced..ee42042bb8e 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -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"]) diff --git a/tests/unit/llms/litellm_proxy/skills/test_skill_search.py b/tests/unit/llms/litellm_proxy/skills/test_skill_search.py index 3f1fe0d5d68..e28e927cf61 100644 --- a/tests/unit/llms/litellm_proxy/skills/test_skill_search.py +++ b/tests/unit/llms/litellm_proxy/skills/test_skill_search.py @@ -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() diff --git a/tests/unit/proxy/common_utils/test_check_batch_cost.py b/tests/unit/proxy/common_utils/test_check_batch_cost.py index 6417c7c8aa6..4e7effad3bf 100644 --- a/tests/unit/proxy/common_utils/test_check_batch_cost.py +++ b/tests/unit/proxy/common_utils/test_check_batch_cost.py @@ -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