mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #19881 from jayy-77/feat/user-agent-customization-issue-19017
feat: add User-Agent customization support
This commit is contained in:
commit
1a7fcfb713
3 changed files with 120 additions and 8 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue