diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 4f86877a6c0..ac9dd5998e2 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -50,9 +50,21 @@ try: except Exception: version = "0.0.0" -headers = { - "User-Agent": f"litellm/{version}", -} +def get_default_headers() -> dict: + """ + Get default headers for HTTP requests. + + - Default: `User-Agent: litellm/{version}` + - Override: set `LITELLM_USER_AGENT` to fully override the header value. + """ + user_agent = os.environ.get("LITELLM_USER_AGENT") + if user_agent is not None: + return {"User-Agent": user_agent} + + return {"User-Agent": f"litellm/{version}"} + +# Initialize headers (User-Agent) +headers = get_default_headers() # https://www.python-httpx.org/advanced/timeouts _DEFAULT_TIMEOUT = httpx.Timeout(timeout=5.0, connect=5.0) @@ -371,13 +383,16 @@ class AsyncHTTPHandler: shared_session=shared_session, ) + # Get default headers (User-Agent, overridable via LITELLM_USER_AGENT) + default_headers = get_default_headers() + return httpx.AsyncClient( transport=transport, event_hooks=event_hooks, timeout=timeout, verify=ssl_config, cert=cert, - headers=headers, + headers=default_headers, follow_redirects=True, ) @@ -899,6 +914,9 @@ class HTTPHandler: # /path/to/client.pem cert = os.getenv("SSL_CERTIFICATE", litellm.ssl_certificate) + # Get default headers (User-Agent, overridable via LITELLM_USER_AGENT) + default_headers = get_default_headers() if not disable_default_headers else None + if client is None: transport = self._create_sync_transport() @@ -908,7 +926,7 @@ class HTTPHandler: timeout=timeout, verify=ssl_config, cert=cert, - headers=headers if not disable_default_headers else None, + headers=default_headers, follow_redirects=True, ) else: diff --git a/litellm/llms/custom_httpx/httpx_handler.py b/litellm/llms/custom_httpx/httpx_handler.py index 6f684ba01c2..491cd97f7db 100644 --- a/litellm/llms/custom_httpx/httpx_handler.py +++ b/litellm/llms/custom_httpx/httpx_handler.py @@ -1,3 +1,4 @@ +import os from typing import Optional, Union import httpx @@ -7,13 +8,22 @@ try: except Exception: version = "0.0.0" -headers = { - "User-Agent": f"litellm/{version}", -} +def get_default_headers() -> dict: + """ + Get default headers for HTTP requests. + - Default: `User-Agent: litellm/{version}` + - Override: set `LITELLM_USER_AGENT` to fully override the header value. + """ + user_agent = os.environ.get("LITELLM_USER_AGENT") + if user_agent is not None: + return {"User-Agent": user_agent} + + return {"User-Agent": f"litellm/{version}"} class HTTPHandler: def __init__(self, concurrent_limit=1000): + headers = get_default_headers() # Create a client with a connection pool self.client = httpx.AsyncClient( limits=httpx.Limits( diff --git a/tests/test_litellm/llms/custom_httpx/test_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_http_handler.py index 0b154474d48..65f08ef5021 100644 --- a/tests/test_litellm/llms/custom_httpx/test_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_http_handler.py @@ -471,3 +471,87 @@ def test_ssl_ecdh_curve(env_curve, litellm_curve, expected_curve, should_call, m assert isinstance(ssl_context, ssl.SSLContext) finally: litellm.ssl_ecdh_curve = original_value + + +def test_default_user_agent_is_litellm_version(monkeypatch): + from litellm._version import version + from litellm.llms.custom_httpx.http_handler import get_default_headers + + monkeypatch.delenv("LITELLM_USER_AGENT", raising=False) + + assert get_default_headers()["User-Agent"] == f"litellm/{version}" + + +def test_user_agent_can_be_overridden_via_env_var(monkeypatch): + from litellm.llms.custom_httpx.http_handler import get_default_headers + + monkeypatch.setenv("LITELLM_USER_AGENT", "Claude Code") + + assert get_default_headers()["User-Agent"] == "Claude Code" + + +def test_user_agent_env_var_can_be_empty_string(monkeypatch): + from litellm.llms.custom_httpx.http_handler import get_default_headers + + monkeypatch.setenv("LITELLM_USER_AGENT", "") + + assert get_default_headers()["User-Agent"] == "" + + +def test_user_agent_override_is_not_appended_to_default(monkeypatch): + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + monkeypatch.delenv("LITELLM_USER_AGENT", raising=False) + + handler = HTTPHandler() + try: + req = handler.client.build_request( + "GET", + "https://example.com", + headers={"user-agent": "Claude Code"}, + ) + + assert req.headers.get_list("User-Agent") == ["Claude Code"] + finally: + handler.close() + + +def test_sync_http_handler_uses_env_user_agent(monkeypatch): + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + monkeypatch.setenv("LITELLM_USER_AGENT", "Claude Code") + + handler = HTTPHandler() + try: + req = handler.client.build_request("GET", "https://example.com") + assert req.headers.get("User-Agent") == "Claude Code" + finally: + handler.close() + + +@pytest.mark.asyncio +async def test_async_http_handler_uses_env_user_agent(monkeypatch): + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + monkeypatch.setenv("LITELLM_USER_AGENT", "Claude Code") + + handler = AsyncHTTPHandler() + try: + req = handler.client.build_request("GET", "https://example.com") + assert req.headers.get("User-Agent") == "Claude Code" + finally: + await handler.close() + + +@pytest.mark.asyncio +async def test_httpx_handler_uses_env_user_agent(monkeypatch): + from litellm.llms.custom_httpx.httpx_handler import HTTPHandler + + monkeypatch.setenv("LITELLM_USER_AGENT", "Claude Code") + + handler = HTTPHandler() + try: + req = handler.client.build_request("GET", "https://example.com") + assert req.headers.get("User-Agent") == "Claude Code" + finally: + await handler.close()