diff --git a/litellm/litellm_core_utils/get_provider_specific_headers.py b/litellm/litellm_core_utils/get_provider_specific_headers.py index 0d5f77b71a0..2618aee9afa 100644 --- a/litellm/litellm_core_utils/get_provider_specific_headers.py +++ b/litellm/litellm_core_utils/get_provider_specific_headers.py @@ -1,15 +1,7 @@ from collections.abc import Sequence from typing import Final -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 +from litellm.types.utils import ProviderSpecificHeader class ProviderSpecificHeaderUtils: diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 319bec87a06..e2b3938bb4f 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -79,7 +79,6 @@ 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, @@ -5745,6 +5744,9 @@ class StandardLoggingPayloadSetup: - If the input metadata is None or not a dictionary, an empty StandardLoggingMetadata object is returned. - If 'user_api_key' is present in metadata and is a valid SHA256 hash, it's stored as 'user_api_key_hash'. """ + from litellm.llms.anthropic.common_utils import ( # noqa: PLC0415 # that module imports this one transitively + resolve_used_client_oauth_token, + ) prompt_management_metadata: StandardLoggingPromptManagementMetadata | None = None if litellm_params is not None: diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index bf2d588dd3a..cc7b50e6b4f 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -39,6 +39,7 @@ from litellm.types.llms.anthropic import ( ) from litellm.types.llms.openai import AllMessageValues from litellm.types.proxy.model_listing import ModelInfoResponse +from litellm.types.utils import LlmProviders _MessageT = TypeVar("_MessageT") @@ -225,6 +226,15 @@ def is_anthropic_oauth_key(value: str | None) -> bool: return value.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX) +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 + + def _merge_beta_headers(existing: str | None, new_beta: str) -> str: """Merge a new beta value into an existing comma-separated anthropic-beta header.""" if not existing: diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 09deac1a5d1..b6aebcf0ad8 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -13,6 +13,7 @@ from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, budget_reservation_from_metadata, get_litellm_metadata_from_kwargs, + get_metadata_variable_name_from_kwargs, ) from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup from litellm.litellm_core_utils.llm_cost_calc.guardrail_cost import guardrail_information_cost @@ -83,10 +84,12 @@ _CAPTURED_IDENTITY_CALL_TYPES: Final[frozenset[str]] = frozenset( str(CallTypes.aretrieve_batch), ) ) -_FAILURE_ROW_KEYS_LIFTED_FROM_LITELLM_METADATA: Final[tuple[str, ...]] = ( - "standard_logging_guardrail_information", - "used_client_oauth_token", -) + + +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)) + stamped: Final = proxy_metadata.get("used_client_oauth_token") if isinstance(proxy_metadata, dict) else None + return stamped if isinstance(stamped, bool) else None def _proxy_spend_writer() -> DBSpendUpdateWriter: @@ -195,17 +198,19 @@ class _ProxyDBLogger(CustomLogger): metadata=_metadata, original_exception=original_exception ) + _metadata["used_client_oauth_token"] = _proxy_stamped_used_client_oauth_token(request_data) + existing_metadata: Final[dict] = request_data.get("metadata", None) or {} existing_metadata.update(_metadata) litellm_metadata_bucket: Final = request_data.get("litellm_metadata") - existing_metadata.update( - (key, litellm_metadata_bucket[key]) - for key in _FAILURE_ROW_KEYS_LIFTED_FROM_LITELLM_METADATA - if isinstance(litellm_metadata_bucket, dict) - and key not in existing_metadata - and litellm_metadata_bucket.get(key) is not None - ) + if ( + isinstance(litellm_metadata_bucket, dict) + and "standard_logging_guardrail_information" not in existing_metadata + ): + guardrail_info: Final = litellm_metadata_bucket.get("standard_logging_guardrail_information") + if guardrail_info is not None: + existing_metadata["standard_logging_guardrail_information"] = guardrail_info if "litellm_params" not in request_data: request_data["litellm_params"] = {} diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 96dac16bd55..6d3ecaa2ca2 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -34,7 +34,6 @@ 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, @@ -46,6 +45,7 @@ from litellm.litellm_core_utils.url_utils import ( is_url_destination_allowed_by_host, provider_url_destination_candidates, ) +from litellm.llms.anthropic.common_utils import ANTHROPIC_OAUTH_FORWARD_PROVIDERS from litellm.proxy._types import ( AddTeamCallback, CommonProxyErrors, diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 39b86a2bf42..6b84d10f9f9 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -2429,13 +2429,15 @@ async def ui_view_spend_logs( default=None, description="Filter logs by cache state: 'hit' or 'miss'. Miss includes legacy rows with a null/unknown cache state", ), - used_client_oauth_token: bool | None = fastapi.Query( - default=None, - description=( - "Filter logs by the credential the upstream call used: true for a client-forwarded Anthropic OAuth token, " - "false for the deployment's configured key. Rows written before this flag existed match neither" + used_client_oauth_token: Annotated[ + bool | None, + fastapi.Query( + description=( + "Filter logs by the credential the upstream call used: true for a client-forwarded Anthropic OAuth " + "token, false for the deployment's configured key. Rows written before this flag existed match neither" + ), ), - ), + ] = None, span_type: str | None = fastapi.Query( default=None, description="Filter logs by span type: llm, agent, mcp, or batch", diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 3129d867824..81a00d43ee4 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -35,7 +35,6 @@ 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, @@ -44,6 +43,7 @@ from litellm.litellm_core_utils.litellm_logging import ( ) from litellm.litellm_core_utils.ptu_pricing import azure_spillover from litellm.litellm_core_utils.safe_json_dumps import safe_dumps, strip_null_bytes +from litellm.llms.anthropic.common_utils import resolve_used_client_oauth_token from litellm.proxy._types import SpendLogsMetadata, SpendLogsPayload, SpendLogsRouterMetadata from litellm.proxy.route_llm_request import ProxyModelNotFoundError from litellm.proxy.spend_tracking.spend_log_error_logger import spend_log_error 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 cf305115bf0..06fdcba7b70 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 @@ -201,6 +201,50 @@ async def test_async_post_call_failure_hook_carries_used_client_oauth_token_from assert json.loads(payload["metadata"])["used_client_oauth_token"] is expected +@pytest.mark.asyncio +@pytest.mark.parametrize( + "metadata_buckets, expected", + [ + ({"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"used_client_oauth_token": False}}, False), + ({"metadata": {"used_client_oauth_token": True}, "litellm_metadata": {"user_id": "caller"}}, None), + ({"metadata": {"used_client_oauth_token": "yes"}}, None), + ], +) +async def test_async_post_call_failure_hook_never_lets_caller_metadata_set_used_client_oauth_token( + metadata_buckets: dict, expected: bool | None +): + """ + On /v1/messages and /v1/responses the request's own metadata field belongs to the caller, so a + used_client_oauth_token they put there must never outrank the proxy's stamp or stand in for a missing one + """ + 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"), + ) + + 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