diff --git a/litellm/identity/extractors/client.py b/litellm/identity/extractors/client.py index 53da72dc12a..72aadb8191b 100644 --- a/litellm/identity/extractors/client.py +++ b/litellm/identity/extractors/client.py @@ -6,7 +6,7 @@ The trust logic is delegated to ``IPAddressUtils.is_request_from_trusted_proxy`` so we stay in sync with the rest of the proxy. """ -from typing import Any, Dict, List, Optional +from typing import Any, Dict, List, Mapping, Optional from litellm.identity.context import ClientInfo @@ -29,32 +29,16 @@ def extract_client_info( request: Any, general_settings: Optional[Dict[str, Any]] = None, ) -> ClientInfo: - headers = getattr(request, "headers", {}) or {} + headers: Mapping[str, str] = getattr(request, "headers", {}) or {} forwarded_chain: List[str] = [] - ip: Optional[str] = None - # Headers in FastAPI are case-insensitive; normalize defensively for dicts. - def _get_header(name: str) -> Optional[str]: - try: - value = headers.get(name) - except AttributeError: - return None - if value is not None: - return value - try: - for key, val in headers.items(): - if isinstance(key, str) and key.lower() == name: - return val - except Exception: - return None - return None - - xff = _get_header("x-forwarded-for") + xff = headers.get("x-forwarded-for") if xff: forwarded_chain = _split_forwarded_chain(xff) from litellm.proxy.auth.ip_address_utils import IPAddressUtils + ip: Optional[str] if forwarded_chain and IPAddressUtils.is_request_from_trusted_proxy( request=request, general_settings=general_settings ): @@ -62,7 +46,7 @@ def extract_client_info( else: ip = _direct_client_host(request) - user_agent = _get_header("user-agent") + user_agent = headers.get("user-agent") return ClientInfo( ip=ip, diff --git a/tests/test_litellm/identity/extractors/test_client.py b/tests/test_litellm/identity/extractors/test_client.py index e0d9c2f0cbe..4cf55077062 100644 --- a/tests/test_litellm/identity/extractors/test_client.py +++ b/tests/test_litellm/identity/extractors/test_client.py @@ -45,3 +45,15 @@ def test_user_agent_passthrough(): req = _fake_request({"user-agent": "curl/8"}, client_host="127.0.0.1") info = extract_client_info(req, general_settings={}) assert info.user_agent == "curl/8" + + +def test_extract_client_info_uses_starlette_headers_case_insensitively(): + from starlette.datastructures import Headers + + req = _fake_request( + Headers({"X-Forwarded-For": "1.2.3.4", "User-Agent": "curl/8"}), + client_host="127.0.0.1", + ) + info = extract_client_info(req, general_settings={}) + assert info.forwarded_chain == ["1.2.3.4"] + assert info.user_agent == "curl/8"