mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-09 22:31:41 +00:00
fix(proxy): preserve spend metadata on auth failures
This commit is contained in:
parent
cd63c7e5a7
commit
e7c42eccec
3 changed files with 121 additions and 16 deletions
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue