This commit is contained in:
RealJasonHu 2026-08-31 12:54:00 -07:00 • committed by GitHub
commit f5e7b86773
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
3 changed files with 138 additions and 16 deletions

View file

@ -19,6 +19,9 @@ from litellm.proxy._types import (
UserAPIKeyAuth,
)
from litellm.proxy.auth.auth_utils import _get_request_ip_address
from litellm.proxy.common_utils.http_parsing_utils import (
_safe_get_request_headers, # pyright: ignore[reportPrivateUsage] # canonical non-throwing request header reader
)
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
from litellm.types.services import ServiceTypes
@ -36,17 +39,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 +177,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(
_safe_get_request_headers(request)
)
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,

View file

@ -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
@ -81,6 +82,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"
@ -1019,16 +1021,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(
@ -1260,7 +1265,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

View file

@ -1,5 +1,6 @@
import asyncio
import json
from copy import deepcopy
from unittest.mock import AsyncMock, MagicMock, patch
import httpx
@ -528,6 +529,92 @@ def _http_request(client_host: str | None = "10.1.2.3", headers: dict[str, str]
)
@pytest.mark.asyncio
@pytest.mark.parametrize(
"header_value, has_body_metadata, expected_spend_logs_metadata",
[
pytest.param(
json.dumps({"my_request_id": "req_rejected_2"}),
True,
{"my_request_id": "req_rejected_2"},
id="valid_header_overrides_body",
),
pytest.param(
json.dumps({"my_request_id": "req_rejected_2"}),
False,
{"my_request_id": "req_rejected_2"},
id="valid_header_without_body_metadata",
),
pytest.param(
"{invalid-json",
True,
{"source": "request-body"},
id="invalid_header_keeps_body",
),
pytest.param(
"null",
True,
{"source": "request-body"},
id="null_header_keeps_body",
),
pytest.param(
None,
True,
{"source": "request-body"},
id="missing_headers_scope_keeps_body",
),
],
)
async def test_budget_auth_failure_logs_spend_metadata_from_request_header(
header_value: str | None,
has_body_metadata: bool,
expected_spend_logs_metadata: dict[str, str],
) -> None:
request_data: dict[str, object] = {"model": "gpt-4o"}
if has_body_metadata:
request_data["metadata"] = {
"existing": "keep-me",
"spend_logs_metadata": {"source": "request-body"},
}
original_request_data = deepcopy(request_data)
request = _http_request(
headers={"x-litellm-spend-logs-metadata": header_value} if header_value is not None else None
)
if header_value is None:
request.scope.pop("headers")
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),
request,
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"]["requester_ip_address"] == "10.1.2.3"
if has_body_metadata:
assert logged_request_data["metadata"]["existing"] == "keep-me"
assert request_data == original_request_data
@pytest.mark.asyncio
@pytest.mark.parametrize(
"auth_error, general_settings, request_kwargs, expected_ip",