diff --git a/litellm/proxy/auth/auth_exception_handler.py b/litellm/proxy/auth/auth_exception_handler.py index a42187b3a44..153fd092f95 100644 --- a/litellm/proxy/auth/auth_exception_handler.py +++ b/litellm/proxy/auth/auth_exception_handler.py @@ -36,17 +36,35 @@ else: Span = Any -def _with_requester_ip_address(request_data: dict[str, object], requester_ip: str | None) -> dict[str, object]: - """Auth gate rejections are raised before `add_litellm_data_to_request` records the - caller IP, so their failure logs would otherwise carry no IP nor key/user identity.""" +def _with_auth_failure_metadata( + request_data: dict[str, object], + requester_ip: str | None, + spend_logs_metadata: Mapping[str, object] | None, +) -> dict[str, object]: + """Add metadata normally populated after auth without mutating the request body.""" + raw_spend_metadata: Final = request_data.get("metadata") + spend_metadata_base: Final[Mapping[str, object]] = ( # pyright: ignore[reportUnknownVariableType] # request JSON mappings have string keys + raw_spend_metadata if isinstance(raw_spend_metadata, Mapping) else EMPTY_MAPPING + ) + request_data_with_spend_metadata: Final = ( + { + **request_data, + "metadata": {**spend_metadata_base, "spend_logs_metadata": spend_logs_metadata}, + } + if spend_logs_metadata is not None + else request_data + ) # mutable-ok: logging hooks require plain dictionaries if not requester_ip: - return request_data - key: Final = "litellm_metadata" if "litellm_metadata" in request_data else "metadata" - metadata: Final = request_data.get(key) + return request_data_with_spend_metadata + key: Final = "litellm_metadata" if "litellm_metadata" in request_data_with_spend_metadata else "metadata" + metadata: Final = request_data_with_spend_metadata.get(key) base: Final[Mapping[str, object]] = metadata if isinstance(metadata, Mapping) else EMPTY_MAPPING if base.get("requester_ip_address"): - return request_data - return {**request_data, key: {**base, "requester_ip_address": requester_ip}} # mutable-ok: logging needs dicts + return request_data_with_spend_metadata + return { + **request_data_with_spend_metadata, + key: {**base, "requester_ip_address": requester_ip}, + } # mutable-ok: logging needs dicts class UserAPIKeyAuthExceptionHandler: @@ -156,8 +174,17 @@ class UserAPIKeyAuthExceptionHandler: ) # Allow callbacks to transform the error response + from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup + + spend_logs_metadata: Final = LiteLLMProxyRequestSetup.get_spend_logs_metadata_from_request_headers( + dict(request.headers) + ) transformed_exception: Final = await proxy_logging_obj.post_call_failure_hook( - request_data=_with_requester_ip_address(request_data, requester_ip), + request_data=_with_auth_failure_metadata( + request_data=request_data, + requester_ip=requester_ip, + spend_logs_metadata=spend_logs_metadata, + ), original_exception=e, user_api_key_dict=user_api_key_dict, error_type=ProxyErrorTypes.auth_error, diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 064b53e07b7..8e03424af3a 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -9,6 +9,7 @@ from types import MappingProxyType from typing import TYPE_CHECKING, Any, Final from fastapi import HTTPException, Request +from pydantic import TypeAdapter from pydantic import ValidationError as PydanticValidationError from starlette.datastructures import Headers @@ -78,6 +79,7 @@ _CODEX_CLIENT_PREFIX_RE: Final = re.compile(r"^codex[-_ /]", re.IGNORECASE) _SESSION_ID_VALUE_RE: Final = re.compile(r"^[a-zA-Z0-9_\-]{8,}$") _SHA256_HEX_RE: Final = re.compile(r"^[0-9a-f]{64}$") +_SPEND_LOGS_METADATA_ADAPTER: Final = TypeAdapter(dict[str, object]) # W3C Trace Context traceparent header: https://www.w3.org/TR/trace-context/ # e.g. "00-4bf92f3577b34da6a3ce929d0e0e4736-00f067aa0ba902b7-01" @@ -995,16 +997,19 @@ class LiteLLMProxyRequestSetup: return None @staticmethod - def _get_spend_logs_metadata_from_request_headers(headers: dict) -> dict | None: + def get_spend_logs_metadata_from_request_headers( + headers: Mapping[str, object], + ) -> dict[str, object] | None: # mutable-ok: request metadata is stored in a TypedDict """ Get the `spend_logs_metadata` from the request headers. """ - from litellm.litellm_core_utils.safe_json_loads import safe_json_loads - spend_logs_metadata_header: Final = headers.get("x-litellm-spend-logs-metadata", None) - if spend_logs_metadata_header is not None: - return safe_json_loads(spend_logs_metadata_header) - return None + if not isinstance(spend_logs_metadata_header, str): + return None + try: + return _SPEND_LOGS_METADATA_ADAPTER.validate_json(spend_logs_metadata_header) + except PydanticValidationError: + return None @staticmethod def _get_forwardable_headers( @@ -1236,7 +1241,7 @@ class LiteLLMProxyRequestSetup: from litellm.proxy._types import LitellmMetadataFromRequestHeaders metadata_from_headers: Final = LitellmMetadataFromRequestHeaders() - spend_logs_metadata: Final = LiteLLMProxyRequestSetup._get_spend_logs_metadata_from_request_headers(headers) + spend_logs_metadata: Final = LiteLLMProxyRequestSetup.get_spend_logs_metadata_from_request_headers(headers) if spend_logs_metadata is not None: metadata_from_headers["spend_logs_metadata"] = spend_logs_metadata diff --git a/tests/test_litellm/proxy/auth/test_auth_exception_handler.py b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py index 90b3b29d919..766d98e11d2 100644 --- a/tests/test_litellm/proxy/auth/test_auth_exception_handler.py +++ b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py @@ -528,6 +528,79 @@ def _http_request(client_host: str | None = "10.1.2.3", headers: dict[str, str] ) +@pytest.mark.asyncio +@pytest.mark.parametrize( + "header_value, expected_spend_logs_metadata", + [ + pytest.param( + json.dumps({"my_request_id": "req_rejected_2"}), + {"my_request_id": "req_rejected_2"}, + id="valid_header_overrides_body", + ), + pytest.param( + "{invalid-json", + {"source": "request-body"}, + id="invalid_header_keeps_body", + ), + pytest.param( + "null", + {"source": "request-body"}, + id="null_header_keeps_body", + ), + ], +) +async def test_budget_auth_failure_logs_spend_metadata_from_request_header( + header_value: str, + expected_spend_logs_metadata: dict[str, str], +) -> None: + request_data: dict[str, object] = { + "model": "gpt-4o", + "metadata": { + "existing": "keep-me", + "spend_logs_metadata": {"source": "request-body"}, + }, + } + + with ( + patch( # test-quality-ok: identity seeding is outside the failure-log payload contract + "litellm.proxy.auth.auth_exception_handler.seed_request_identity" + ), + patch( # test-quality-ok: capture the exact payload sent to the spend logger + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ) as mock_hook, + patch( # test-quality-ok: disable the unrelated database-outage fallback + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": False}, + ), + ): + with pytest.raises(ProxyException): + await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + BudgetExceededError(message="Budget exceeded", current_cost=100, max_budget=100), + _http_request( + headers={ + "x-litellm-spend-logs-metadata": header_value, + } + ), + request_data, + "/v1/chat/completions", + None, + "sk-over-budget", + ) + + logged_request_data = mock_hook.call_args.kwargs["request_data"] + assert logged_request_data["metadata"]["spend_logs_metadata"] == expected_spend_logs_metadata + assert logged_request_data["metadata"]["existing"] == "keep-me" + assert request_data == { + "model": "gpt-4o", + "metadata": { + "existing": "keep-me", + "spend_logs_metadata": {"source": "request-body"}, + }, + } + + @pytest.mark.asyncio @pytest.mark.parametrize( "auth_error, general_settings, request_kwargs, expected_ip",