diff --git a/litellm/litellm_core_utils/core_helpers.py b/litellm/litellm_core_utils/core_helpers.py index b095b4b12c6..64f94ed3799 100644 --- a/litellm/litellm_core_utils/core_helpers.py +++ b/litellm/litellm_core_utils/core_helpers.py @@ -339,6 +339,13 @@ def get_or_create_metadata_bucket( return metadata_key, metadata_bucket +def proxy_stamped_used_client_oauth_token(metadata: object, litellm_params: Mapping[str, object] | None) -> object: + litellm_metadata: Final = litellm_params.get("litellm_metadata") if litellm_params is not None else None + if isinstance(litellm_metadata, Mapping) and "used_client_oauth_token" in litellm_metadata: + return litellm_metadata["used_client_oauth_token"] + return metadata.get("used_client_oauth_token") if isinstance(metadata, Mapping) else None + + def get_litellm_metadata_from_kwargs(kwargs: dict): """ Helper to get litellm metadata from all litellm request kwargs diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 7784addb215..3958cda3efe 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -71,6 +71,7 @@ from litellm.litellm_core_utils.classifier_logging import ( from litellm.litellm_core_utils.core_helpers import ( get_provider_response_headers_from_hidden_params, is_expected_client_error, + proxy_stamped_used_client_oauth_token, reconstruct_model_name, set_response_cost_in_hidden_params, ) @@ -288,13 +289,6 @@ _STANDARD_LOGGING_METADATA_KEYS: Final[frozenset[str]] = ( ) -def _proxy_stamped_used_client_oauth_token(metadata: object, litellm_params: Mapping[str, object] | None) -> object: - litellm_metadata: Final = litellm_params.get("litellm_metadata") if litellm_params is not None else None - if isinstance(litellm_metadata, Mapping) and "used_client_oauth_token" in litellm_metadata: - return litellm_metadata["used_client_oauth_token"] - return metadata.get("used_client_oauth_token") if isinstance(metadata, Mapping) else None - - def _get_provider_request_id(original_exception: Exception) -> str | None: try: error_response: Final = getattr(original_exception, "response", None) @@ -5793,7 +5787,7 @@ class StandardLoggingPayloadSetup: team_alias=None, team_id=None, used_client_oauth_token=resolve_used_client_oauth_token( - _proxy_stamped_used_client_oauth_token(metadata, litellm_params), + proxy_stamped_used_client_oauth_token(metadata, litellm_params), custom_llm_provider, ), ) diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index bd53adeec41..76c9ea48948 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -30,7 +30,7 @@ from litellm.proxy.db.db_spend_update_writer import ( debitable_model_access_groups, get_llm_router, ) -from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup +from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup, metadata_variable_name_for_route from litellm.proxy.spend_tracking.spend_event import ( ObjectMapping, SpendEventBuildError, @@ -86,8 +86,15 @@ _CAPTURED_IDENTITY_CALL_TYPES: Final[frozenset[str]] = frozenset( ) -def _proxy_stamped_used_client_oauth_token(request_data: Mapping[str, object]) -> bool | None: - proxy_metadata: Final = request_data.get(get_metadata_variable_name_from_kwargs(request_data)) +def _proxy_stamped_used_client_oauth_token( + request_data: Mapping[str, object], request_route: str | None +) -> bool | None: + proxy_bucket: Final = ( + get_metadata_variable_name_from_kwargs(request_data) + if request_route is None + else metadata_variable_name_for_route(request_route) + ) + proxy_metadata: Final = request_data.get(proxy_bucket) stamped: Final = proxy_metadata.get("used_client_oauth_token") if isinstance(proxy_metadata, dict) else None return stamped if isinstance(stamped, bool) else None @@ -198,7 +205,7 @@ class _ProxyDBLogger(CustomLogger): metadata=_metadata, original_exception=original_exception ) - _metadata["used_client_oauth_token"] = _proxy_stamped_used_client_oauth_token(request_data) + _metadata["used_client_oauth_token"] = _proxy_stamped_used_client_oauth_token(request_data, request_route) existing_metadata: Final[dict] = request_data.get("metadata", None) or {} existing_metadata.update(_metadata) diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 4a501eb3d09..bbd3ed66ff4 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -7,7 +7,7 @@ from collections import OrderedDict from collections.abc import Mapping, MutableMapping, Sequence from datetime import datetime from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, cast +from typing import TYPE_CHECKING, Any, Final, Literal, cast from fastapi import HTTPException, Request from pydantic import TypeAdapter @@ -649,11 +649,14 @@ def _get_metadata_variable_name(request: Request) -> str: # Inline imports — auth_utils/route_checks participate in a proxy import cycle. from litellm.proxy.auth.auth_utils import get_request_route # noqa: PLC0415 - path: Final = get_request_route(request) - if "thread" in path or "assistant" in path: + return metadata_variable_name_for_route(get_request_route(request)) + + +def metadata_variable_name_for_route(route: str) -> Literal["metadata", "litellm_metadata"]: + if "thread" in route or "assistant" in route: return "litellm_metadata" - if any(route in path for route in LITELLM_METADATA_ROUTES): + if any(metadata_route in route for metadata_route in LITELLM_METADATA_ROUTES): return "litellm_metadata" return "metadata" diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 79f8e72d2aa..d736739c57d 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -33,6 +33,7 @@ from litellm.constants import ( from litellm.litellm_core_utils.classifier_logging import classifier_audit_fields, without_classifier_audit from litellm.litellm_core_utils.core_helpers import ( get_litellm_metadata_from_kwargs, + proxy_stamped_used_client_oauth_token, reconstruct_model_name, ) from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider @@ -720,7 +721,7 @@ def get_logging_payload( router_correlation_id=litellm_call_id, ), used_client_oauth_token=resolve_used_client_oauth_token( - metadata.get("used_client_oauth_token") if metadata is not None else None, custom_llm_provider + proxy_stamped_used_client_oauth_token(litellm_params.get("metadata"), litellm_params), custom_llm_provider ), azure_spillover=azure_spillover( response_headers=kwargs.get("response_headers") diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 06fdcba7b70..44c5f7a9d1b 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -245,6 +245,53 @@ async def test_async_post_call_failure_hook_never_lets_caller_metadata_set_used_ assert json.loads(payload["metadata"])["used_client_oauth_token"] is expected +@pytest.mark.asyncio +@pytest.mark.parametrize( + "request_route, metadata_buckets, expected", + [ + ( + "/v1/chat/completions", + {"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"user_api_key_hash": "guardrail"}}, + True, + ), + ( + "/v1/messages", + {"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"user_api_key_hash": "proxy"}}, + None, + ), + ], +) +async def test_async_post_call_failure_hook_reads_used_client_oauth_token_from_the_routes_stamped_bucket( + request_route: str, metadata_buckets: dict, expected: bool | None +): + logger = _ProxyDBLogger() + request_data = { + "model": "claude-sonnet-5", + "custom_llm_provider": "anthropic", + "messages": [{"role": "user", "content": "Hello"}], + "proxy_server_request": {"request_id": "test_request_id"}, + **metadata_buckets, + } + + with patch( + "litellm.proxy.db.db_spend_update_writer.DBSpendUpdateWriter.update_database", + new_callable=AsyncMock, + ) as mock_update_database: + await logger.async_post_call_failure_hook( + request_data=request_data, + original_exception=Exception("rate limited"), + user_api_key_dict=UserAPIKeyAuth(api_key="test_api_key", request_route=request_route), + ) + + payload = get_logging_payload( + kwargs=mock_update_database.call_args[1]["kwargs"], + response_obj={}, + start_time=datetime.now(), + end_time=datetime.now(), + ) + assert json.loads(payload["metadata"])["used_client_oauth_token"] is expected + + @pytest.mark.asyncio async def test_async_post_call_failure_hook_bills_guardrail_cost_on_blocked_request(): """LIT-5651: a request blocked by a guardrail never reaches the LLM, but the 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 05c6a556da1..fb351d75cc4 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 @@ -3305,6 +3305,31 @@ def test_get_logging_payload_records_used_client_oauth_token_for_the_selected_pr assert _get_spend_logs_metadata(None)["used_client_oauth_token"] is None +@pytest.mark.parametrize( + "litellm_params, expected", + [ + ( + {"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"user_api_key_hash": "guardrail"}}, + True, + ), + ( + {"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"used_client_oauth_token": False}}, + False, + ), + ], +) +def test_get_logging_payload_reads_used_client_oauth_token_from_the_bucket_the_proxy_stamped( + litellm_params: dict, expected: bool +): + payload = get_logging_payload( + kwargs={"model": "claude-sonnet-5", "custom_llm_provider": "anthropic", "litellm_params": litellm_params}, + response_obj={}, + start_time=datetime.datetime.now(timezone.utc), + end_time=datetime.datetime.now(timezone.utc), + ) + assert json.loads(payload["metadata"])["used_client_oauth_token"] is expected + + def test_redact_logged_api_key_bearer_only_returns_none(): # "bearer " with nothing after stripping is equivalent to no key assert _redact_logged_api_key("bearer ") is None