From 05b35191d3ffe9fbba17cbcbb2ff427bade679d8 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 24 Sep 2026 16:55:11 -0700 Subject: [PATCH] fix(proxy): resolve used_client_oauth_token against the provider the call was sent to --- .../get_provider_specific_headers.py | 10 +++++- litellm/litellm_core_utils/litellm_logging.py | 14 ++++++-- litellm/proxy/litellm_pre_call_utils.py | 3 +- .../spend_tracking/spend_tracking_utils.py | 9 ++++- .../test_litellm_logging.py | 33 +++++++++++++++++ .../hooks/test_proxy_track_cost_callback.py | 10 ++++-- .../test_spend_tracking_utils.py | 36 ++++++++++++++++--- .../proxy/test_litellm_pre_call_utils.py | 15 ++++++-- 8 files changed, 114 insertions(+), 16 deletions(-) diff --git a/litellm/litellm_core_utils/get_provider_specific_headers.py b/litellm/litellm_core_utils/get_provider_specific_headers.py index 2618aee9afa..0d5f77b71a0 100644 --- a/litellm/litellm_core_utils/get_provider_specific_headers.py +++ b/litellm/litellm_core_utils/get_provider_specific_headers.py @@ -1,7 +1,15 @@ from collections.abc import Sequence from typing import Final -from litellm.types.utils import ProviderSpecificHeader +from litellm.types.utils import LlmProviders, ProviderSpecificHeader + +ANTHROPIC_OAUTH_FORWARD_PROVIDERS: Final[frozenset[str]] = frozenset((LlmProviders.ANTHROPIC.value,)) + + +def resolve_used_client_oauth_token(client_sent_oauth_token: object, custom_llm_provider: str | None) -> bool | None: + if not isinstance(client_sent_oauth_token, bool): + return None + return client_sent_oauth_token and custom_llm_provider in ANTHROPIC_OAUTH_FORWARD_PROVIDERS class ProviderSpecificHeaderUtils: diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index e35ec13e023..319bec87a06 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -79,6 +79,7 @@ from litellm.litellm_core_utils.core_helpers import ( ) from litellm.litellm_core_utils.error_normalization import normalize_error from litellm.litellm_core_utils.get_litellm_params import get_litellm_params +from litellm.litellm_core_utils.get_provider_specific_headers import resolve_used_client_oauth_token from litellm.litellm_core_utils.internal_call_metadata import ( MODEL_ACCESS_GROUP_METADATA_KEY, is_unbilled_non_inference_call, @@ -283,7 +284,10 @@ else: _PAGERDUTY_ALERTING_FACTORY: Final = PagerDutyAlerting _in_memory_loggers: Final[list[CustomLogger]] = [] -_STANDARD_LOGGING_METADATA_KEYS: Final[frozenset[str]] = frozenset(StandardLoggingMetadata.__annotations__.keys()) +_STANDARD_LOGGING_METADATA_RESOLVED_KEYS: Final[frozenset[str]] = frozenset(("used_client_oauth_token",)) +_STANDARD_LOGGING_METADATA_KEYS: Final[frozenset[str]] = ( + frozenset(StandardLoggingMetadata.__annotations__.keys()) - _STANDARD_LOGGING_METADATA_RESOLVED_KEYS +) def _get_provider_request_id(original_exception: Exception) -> str | None: @@ -5726,6 +5730,7 @@ class StandardLoggingPayloadSetup: proxy_server_request: dict | None = None, start_time: dt_object | None = None, response_id: str | None = None, + custom_llm_provider: str | None = None, ) -> StandardLoggingMetadata: """ Clean and filter the metadata dictionary to include only the specified keys in StandardLoggingMetadata. @@ -5789,7 +5794,10 @@ class StandardLoggingPayloadSetup: user_api_key_auth_metadata=None, team_alias=None, team_id=None, - used_client_oauth_token=None, + used_client_oauth_token=resolve_used_client_oauth_token( + metadata.get("used_client_oauth_token") if isinstance(metadata, dict) else None, + custom_llm_provider, + ), ) if isinstance(metadata, dict): for key in metadata.keys() & _STANDARD_LOGGING_METADATA_KEYS: @@ -6513,6 +6521,7 @@ def get_standard_logging_object_payload( stream=kwargs.get("stream", False), ) # clean up litellm metadata + selected_provider: Final = kwargs.get("custom_llm_provider") clean_metadata: Final = StandardLoggingPayloadSetup.get_standard_logging_metadata( metadata=metadata, litellm_params=litellm_params, @@ -6524,6 +6533,7 @@ def get_standard_logging_object_payload( proxy_server_request=proxy_server_request, start_time=start_time, response_id=id, + custom_llm_provider=selected_provider if isinstance(selected_provider, str) else None, ) _request_body: Final = proxy_server_request.get("body", {}) end_user_id: Final = clean_metadata["user_api_key_end_user_id"] or _request_body.get( diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 907dbc58ace..96dac16bd55 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -34,6 +34,7 @@ from litellm.constants import ( ) from litellm.litellm_core_utils.core_helpers import is_codex_user_agent from litellm.litellm_core_utils.credential_accessor import CredentialAccessor +from litellm.litellm_core_utils.get_provider_specific_headers import ANTHROPIC_OAUTH_FORWARD_PROVIDERS from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( TRUSTED_CALLBACK_VARS_FIELD, _request_blocked_callback_params, @@ -3456,7 +3457,7 @@ _ANTHROPIC_API_HEADER_PROVIDERS: Final = ",".join( LlmProviders.VERTEX_AI.value, ) ) -_ANTHROPIC_OAUTH_CREDENTIAL_PROVIDERS: Final = LlmProviders.ANTHROPIC.value +_ANTHROPIC_OAUTH_CREDENTIAL_PROVIDERS: Final = ",".join(sorted(ANTHROPIC_OAUTH_FORWARD_PROVIDERS)) def add_provider_specific_headers_to_request( diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index aca72a954b1..3129d867824 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -35,6 +35,7 @@ from litellm.litellm_core_utils.core_helpers import ( reconstruct_model_name, ) from litellm.litellm_core_utils.get_llm_provider_logic import declared_authenticating_provider +from litellm.litellm_core_utils.get_provider_specific_headers import resolve_used_client_oauth_token from litellm.litellm_core_utils.internal_call_metadata import is_unbilled_non_inference_call from litellm.litellm_core_utils.litellm_logging import ( coerce_model_access_groups, @@ -144,6 +145,7 @@ _STAMPED_METADATA_KEYS: Final = frozenset( "autorouter_savings", "autorouter_savings_estimate", "autorouter_baseline_observation", + "used_client_oauth_token", ) ) @@ -168,6 +170,7 @@ def _get_spend_logs_metadata( autorouter_baseline_observation: str | None = None, router_metadata: SpendLogsRouterMetadata | None = None, azure_spillover: AzureSpillover | None = None, + used_client_oauth_token: bool | None = None, ) -> SpendLogsMetadata: if metadata is None: return SpendLogsMetadata( @@ -212,7 +215,7 @@ def _get_spend_logs_metadata( litellm_call_id=litellm_call_id, router_metadata=router_metadata, azure_spillover=azure_spillover, - used_client_oauth_token=None, + used_client_oauth_token=used_client_oauth_token, ) verbose_proxy_logger.debug( "getting payload for SpendLogs, available keys in metadata: " + str(list(metadata.keys())) @@ -228,6 +231,7 @@ def _get_spend_logs_metadata( autorouter_baseline_observation=autorouter_baseline_observation, router_metadata=router_metadata, azure_spillover=azure_spillover, + used_client_oauth_token=used_client_oauth_token, ) _raw_key: Final = clean_metadata.get("user_api_key") _trusted_hash: Final = metadata.get("user_api_key_hash") @@ -695,6 +699,9 @@ def get_logging_payload( selected_provider=custom_llm_provider, 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 + ), azure_spillover=azure_spillover( response_headers=kwargs.get("response_headers") if isinstance(kwargs.get("response_headers"), Mapping) diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index bb8098e6e23..f75233d8a37 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -4846,6 +4846,39 @@ def test_get_standard_logging_object_payload_includes_litellm_call_id(logging_ob assert payload["litellm_call_id"] == call_id +@pytest.mark.parametrize( + "client_sent_oauth_token, custom_llm_provider, expected", + [(True, "anthropic", True), (True, "bedrock", False), (False, "anthropic", False), (None, "anthropic", None)], +) +def test_get_standard_logging_object_payload_resolves_used_client_oauth_token_against_the_selected_provider( + logging_obj, client_sent_oauth_token: bool | None, custom_llm_provider: str, expected: bool | None +): + """The proxy stamps whether the client presented an Anthropic OAuth bearer before routing, but the + bearer only reaches an Anthropic deployment, so the logged flag must follow the provider that was called.""" + from datetime import datetime + + from litellm.litellm_core_utils.litellm_logging import get_standard_logging_object_payload + + request_metadata = {} if client_sent_oauth_token is None else {"used_client_oauth_token": client_sent_oauth_token} + now = datetime.now() + payload = get_standard_logging_object_payload( + kwargs={ + "model": "claude-sonnet-5", + "messages": [], + "custom_llm_provider": custom_llm_provider, + "litellm_params": {"metadata": request_metadata}, + }, + init_response_obj={}, + start_time=now, + end_time=now, + logging_obj=logging_obj, + status="success", + ) + + assert payload is not None + assert payload["metadata"]["used_client_oauth_token"] is expected + + def test_get_standard_logging_object_payload_carries_matched_access_groups(logging_obj): """Access groups stamped at auth time reach the logging payload, so integrations see what a request billed.""" from datetime import datetime 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 5f274e609ee..cf305115bf0 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 @@ -161,9 +161,12 @@ async def test_async_post_call_failure_hook_does_not_clobber_guardrail_info_in_m @pytest.mark.asyncio -@pytest.mark.parametrize("used_client_oauth_token", [True, False]) +@pytest.mark.parametrize( + "used_client_oauth_token, custom_llm_provider, expected", + [(True, "anthropic", True), (True, "bedrock", False), (False, "anthropic", False)], +) async def test_async_post_call_failure_hook_carries_used_client_oauth_token_from_litellm_metadata( - used_client_oauth_token: bool, + used_client_oauth_token: bool, custom_llm_provider: str, expected: bool ): """ /v1/messages and /v1/responses stamp the proxy's own fields into request_data["litellm_metadata"] @@ -173,6 +176,7 @@ async def test_async_post_call_failure_hook_carries_used_client_oauth_token_from logger = _ProxyDBLogger() request_data = { "model": "claude-sonnet-5", + "custom_llm_provider": custom_llm_provider, "messages": [{"role": "user", "content": "Hello"}], "metadata": {"user_id": "anthropic-native-metadata"}, "litellm_metadata": {"used_client_oauth_token": used_client_oauth_token}, @@ -194,7 +198,7 @@ async def test_async_post_call_failure_hook_carries_used_client_oauth_token_from payload = get_logging_payload( kwargs=call_kwargs, response_obj={}, start_time=datetime.now(), end_time=datetime.now() ) - assert json.loads(payload["metadata"])["used_client_oauth_token"] is used_client_oauth_token + assert json.loads(payload["metadata"])["used_client_oauth_token"] is expected @pytest.mark.asyncio 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 f9a415afd63..b0e83f7fbe7 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 @@ -3271,12 +3271,38 @@ def test_get_spend_logs_metadata_keeps_user_agent(): assert _get_spend_logs_metadata(None)["user_agent"] is None -@pytest.mark.parametrize("used_client_oauth_token", [True, False]) -def test_get_spend_logs_metadata_keeps_used_client_oauth_token(used_client_oauth_token: bool): - meta = _get_spend_logs_metadata({"used_client_oauth_token": used_client_oauth_token}) - assert meta["used_client_oauth_token"] is used_client_oauth_token +@pytest.mark.parametrize( + "client_sent_oauth_token, custom_llm_provider, expected", + [ + (True, "anthropic", True), + (True, "bedrock", False), + (True, "vertex_ai", False), + (False, "anthropic", False), + (None, "anthropic", None), + ], +) +def test_get_logging_payload_records_used_client_oauth_token_for_the_selected_provider( + client_sent_oauth_token: bool | None, custom_llm_provider: str, expected: bool | None +): + """The client's OAuth bearer is only forwarded to an Anthropic deployment, so a request that + the router sent to Bedrock or Vertex paid with the configured key and must not read true.""" + request_metadata = ( + {"user_agent": "claude-cli/2.1.0"} + if client_sent_oauth_token is None + else {"user_agent": "claude-cli/2.1.0", "used_client_oauth_token": client_sent_oauth_token} + ) + payload = get_logging_payload( + kwargs={ + "model": "claude-sonnet-5", + "custom_llm_provider": custom_llm_provider, + "litellm_params": {"metadata": request_metadata}, + }, + 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 assert _get_spend_logs_metadata(None)["used_client_oauth_token"] is None - assert _get_spend_logs_metadata({"user_agent": "curl/8.7.1"})["used_client_oauth_token"] is None def test_redact_logged_api_key_bearer_only_returns_none(): 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 d7692f9f3ff..1c7484383e4 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -38,8 +38,7 @@ from litellm.proxy.litellm_pre_call_utils import ( move_guardrails_to_metadata, ) from litellm.litellm_core_utils.core_helpers import get_litellm_metadata_from_kwargs -from litellm.litellm_core_utils.litellm_logging import get_standard_logging_metadata -from litellm.proxy.spend_tracking.spend_tracking_utils import _get_spend_logs_metadata +from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload from litellm.litellm_core_utils.internal_call_metadata import MODEL_ACCESS_GROUP_METADATA_KEY from litellm.litellm_core_utils.redact_messages import _get_turn_off_message_logging_from_dynamic_params from litellm.litellm_core_utils.get_provider_specific_headers import ( @@ -6794,7 +6793,17 @@ async def test_add_litellm_data_to_request_stamps_used_client_oauth_token(path, return updated[metadata_variable_name] def spend_log_row_metadata(request_metadata: dict) -> dict: - return dict(_get_spend_logs_metadata(dict(get_standard_logging_metadata(metadata=request_metadata)))) + row = get_logging_payload( + kwargs={ + "model": "claude-sonnet-5", + "custom_llm_provider": "anthropic", + "litellm_params": {"metadata": request_metadata}, + }, + response_obj={}, + start_time=datetime.now(timezone.utc), + end_time=datetime.now(timezone.utc), + ) + return json.loads(row["metadata"]) seat_row = spend_log_row_metadata( await metadata_for({"Authorization": _OAUTH_TOKEN, "x-litellm-api-key": "Bearer sk-virtual-key"})