fix(proxy): preserve spend metadata on auth failures

This commit is contained in:
Zhexun Hu 2026-08-27 16:26:48 +08:00
parent cd63c7e5a7
commit e7c42eccec
3 changed files with 121 additions and 16 deletions

View file

@ -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,

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
@ -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

View file

@ -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",