diff --git a/litellm/proxy/auth/auth_exception_handler.py b/litellm/proxy/auth/auth_exception_handler.py index 603e72463bc..233679126f8 100644 --- a/litellm/proxy/auth/auth_exception_handler.py +++ b/litellm/proxy/auth/auth_exception_handler.py @@ -2,12 +2,14 @@ Handles Authentication Errors """ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final from fastapi import HTTPException, Request, status import litellm from litellm._logging import verbose_proxy_logger +from litellm.constants import EMPTY_MAPPING from litellm.integrations.otel.runtime import seed_request_identity from litellm.proxy._types import ( LitellmUserRoles, @@ -33,12 +35,25 @@ 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.""" + 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) + 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 + + class UserAPIKeyAuthExceptionHandler: @staticmethod async def _handle_authentication_error( e: Exception, request: Request, - request_data: dict, + request_data: dict[str, object], route: str, parent_otel_span: Span | None, api_key: str, @@ -92,7 +107,7 @@ class UserAPIKeyAuthExceptionHandler: # raise the exception to the caller requester_ip: Final = _get_request_ip_address( request=request, - use_x_forwarded_for=general_settings.get("use_x_forwarded_for", False), + use_x_forwarded_for=general_settings.get("use_x_forwarded_for") is True, ) verbose_proxy_logger.exception( "litellm.proxy.proxy_server.user_api_key_auth(): Exception occured - %s\nRequester IP Address:%s", @@ -129,11 +144,14 @@ class UserAPIKeyAuthExceptionHandler: resolve_llm_provider_for_rate_limit, ) - _, e.llm_provider = resolve_llm_provider_for_rate_limit(request_data.get("model")) + budget_model: Final = request_data.get("model") + _, e.llm_provider = resolve_llm_provider_for_rate_limit( + budget_model if isinstance(budget_model, str) else None + ) # Allow callbacks to transform the error response transformed_exception: Final = await proxy_logging_obj.post_call_failure_hook( - request_data=request_data, + request_data=_with_requester_ip_address(request_data, requester_ip), original_exception=e, user_api_key_dict=user_api_key_dict, error_type=ProxyErrorTypes.auth_error, 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 27798ec0bff..b4725a81823 100644 --- a/tests/test_litellm/proxy/auth/test_auth_exception_handler.py +++ b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py @@ -31,6 +31,7 @@ sys.path.insert( ) # Adds the parent directory to the system path from litellm._logging import verbose_proxy_logger +from litellm.exceptions import BudgetExceededError from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler @@ -511,3 +512,196 @@ async def test_auth_failure_without_resolved_identity_still_logs(): assert logged.api_key != "sk-unknown" assert logged.api_key == UserAPIKeyAuth(api_key="sk-unknown").api_key assert logged.request_route == "/v1/chat/completions" + + +def _http_request(client_host: str | None = "10.1.2.3", headers: dict[str, str] | None = None) -> Request: + return Request( + { + "type": "http", + "http_version": "1.1", + "method": "POST", + "scheme": "http", + "path": "/v1/chat/completions", + "raw_path": b"/v1/chat/completions", + "query_string": b"", + "root_path": "", + "server": ("testserver", 80), + "client": (client_host, 51234) if client_host is not None else None, + "headers": [(k.lower().encode(), v.encode()) for k, v in (headers or {}).items()], + } + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "auth_error, general_settings, request_kwargs, expected_ip", + [ + pytest.param( + ProxyException( + message="Invalid API key", + type=ProxyErrorTypes.auth_error, + param=None, + code=status.HTTP_401_UNAUTHORIZED, + ), + {"allow_requests_on_db_unavailable": False}, + {}, + "10.1.2.3", + id="401_socket_peer", + ), + pytest.param( + ProxyException( + message="Invalid API key", + type=ProxyErrorTypes.auth_error, + param=None, + code=status.HTTP_401_UNAUTHORIZED, + ), + {"allow_requests_on_db_unavailable": False, "use_x_forwarded_for": True}, + {"headers": {"x-forwarded-for": "203.0.113.9"}}, + "203.0.113.9", + id="401_x_forwarded_for", + ), + pytest.param( + BudgetExceededError(message="Budget exceeded", current_cost=100, max_budget=100), + {"allow_requests_on_db_unavailable": False}, + {}, + "10.1.2.3", + id="429_budget_exceeded", + ), + ], +) +async def test_auth_failure_logs_requester_ip_address( + auth_error: Exception, + general_settings: dict[str, bool], + request_kwargs: dict[str, dict[str, str]], + expected_ip: str, +) -> None: + """401s and budget 429s are rejected before `add_litellm_data_to_request` stamps + the caller IP, so without this the failure logs (spend logs, prometheus client_ip) + had no IP, and a 401 rarely carries a key or user identity either.""" + with ( + patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ) as mock_hook, + patch("litellm.proxy.proxy_server.general_settings", general_settings), + ): + with pytest.raises(ProxyException): + await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + auth_error, + _http_request(**request_kwargs), + {"model": "gpt-4o"}, + "/v1/chat/completions", + None, + "sk-bad-key", + ) + + logged_request_data = mock_hook.call_args[1]["request_data"] + assert logged_request_data["metadata"]["requester_ip_address"] == expected_ip + + +@pytest.mark.asyncio +async def test_auth_failure_keeps_existing_requester_ip_address(): + """An IP already recorded upstream (e.g. a trusted-proxy resolved value) wins over + the socket peer.""" + with ( + patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ) as mock_hook, + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": False}, + ), + ): + with pytest.raises(ProxyException): + await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + ProxyException( + message="Invalid API key", + type=ProxyErrorTypes.auth_error, + param=None, + code=status.HTTP_401_UNAUTHORIZED, + ), + _http_request(), + {"metadata": {"requester_ip_address": "198.51.100.4"}}, + "/v1/chat/completions", + None, + "sk-bad-key", + ) + + logged_request_data = mock_hook.call_args[1]["request_data"] + assert logged_request_data["metadata"]["requester_ip_address"] == "198.51.100.4" + + +@pytest.mark.asyncio +async def test_auth_failure_ip_uses_litellm_metadata_when_present(): + """Routes that keep proxy metadata under `litellm_metadata` (e.g. /responses) must + get the IP there, since that is the dict the logging layer reads for them.""" + with ( + patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ) as mock_hook, + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": False}, + ), + ): + with pytest.raises(ProxyException): + await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + ProxyException( + message="Invalid API key", + type=ProxyErrorTypes.auth_error, + param=None, + code=status.HTTP_401_UNAUTHORIZED, + ), + _http_request(), + {"litellm_metadata": {}, "metadata": {"user_supplied": "keep-me"}}, + "/v1/responses", + None, + "sk-bad-key", + ) + + logged_request_data = mock_hook.call_args[1]["request_data"] + assert logged_request_data["litellm_metadata"]["requester_ip_address"] == "10.1.2.3" + assert logged_request_data["metadata"] == {"user_supplied": "keep-me"} + + +@pytest.mark.asyncio +async def test_auth_failure_ip_stamp_does_not_mutate_callers_request_data(): + """The handler must not rewrite the caller's dict; the IP is for the failure log only.""" + request_data = {"model": "gpt-4o"} + + with ( + patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ), + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": False}, + ), + ): + with pytest.raises(ProxyException): + await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + ProxyException( + message="Invalid API key", + type=ProxyErrorTypes.auth_error, + param=None, + code=status.HTTP_401_UNAUTHORIZED, + ), + _http_request(), + request_data, + "/v1/chat/completions", + None, + "sk-bad-key", + ) + + assert request_data == {"model": "gpt-4o"}