Merge pull request #14761 from uzaxirr/feat/sdk-additional-headers

feat: Add SDK support for additional headers
This commit is contained in:
Krish Dholakia 2025-09-21 21:09:51 -07:00 • committed by GitHub
commit 8628c265b9
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 803 additions and 78 deletions

View file

@ -0,0 +1,328 @@
# SDK Header Support
LiteLLM SDK provides comprehensive support for passing additional headers with API requests. This is essential for enterprise environments using API gateways, service meshes, and multi-tenant architectures.
## Overview
Headers can be passed to LiteLLM in three ways, with the following priority order:
1. **Request-specific headers** (highest priority)
2. **extra_headers parameter**
3. **Global litellm.headers** (lowest priority)
When the same header key is specified in multiple places, the higher priority value will be used.
## Usage Methods
### 1. Global Headers (litellm.headers)
Set headers that will be included in all API requests:
```python
import litellm
# Set global headers for all requests
litellm.headers = {
"X-API-Gateway-Key": "your-gateway-key",
"X-Company-ID": "acme-corp",
"X-Environment": "production"
}
# Now all completion calls will include these headers
response = litellm.completion(
model="claude-3-5-sonnet-latest",
messages=[{"role": "user", "content": "Hello"}]
)
```
### 2. Per-Request Headers (extra_headers)
Pass headers for specific requests using the `extra_headers` parameter:
```python
import litellm
response = litellm.completion(
model="gpt-4",
messages=[{"role": "user", "content": "Hello"}],
extra_headers={
"X-Request-ID": "req-12345",
"X-Tenant-ID": "tenant-abc",
"X-Custom-Auth": "bearer-token-xyz"
}
)
```
### 3. Request Headers (headers parameter)
Use the `headers` parameter for the highest priority header control:
```python
import litellm
response = litellm.completion(
model="claude-3-5-sonnet-latest",
messages=[{"role": "user", "content": "Hello"}],
headers={
"X-Priority-Header": "high-priority-value",
"Authorization": "Bearer custom-token"
}
)
```
### 4. Combining All Methods
You can combine all three methods. Headers will be merged with the priority order:
```python
import litellm
# Global headers (lowest priority)
litellm.headers = {
"X-Company-ID": "acme-corp",
"X-Shared-Header": "global-value"
}
response = litellm.completion(
model="gpt-4",
messages=[{"role": "user", "content": "Hello"}],
extra_headers={
"X-Request-ID": "req-12345",
"X-Shared-Header": "extra-value" # Overrides global
},
headers={
"X-Priority-Header": "important",
"X-Shared-Header": "request-value" # Overrides both global and extra
}
)
# Final headers sent to API:
# {
# "X-Company-ID": "acme-corp", # From global
# "X-Request-ID": "req-12345", # From extra_headers
# "X-Priority-Header": "important", # From headers
# "X-Shared-Header": "request-value" # From headers (highest priority)
# }
```
## Enterprise Use Cases
### API Gateway Integration (Apigee, Kong, AWS API Gateway)
```python
import litellm
# Set up headers for API gateway routing and authentication
litellm.headers = {
"X-API-Gateway-Key": "your-gateway-key",
"X-Route-Version": "v2"
}
# Per-tenant requests
response = litellm.completion(
model="claude-3-5-sonnet-latest",
messages=[{"role": "user", "content": "Analyze this data"}],
extra_headers={
"X-Tenant-ID": "tenant-123",
"X-Department": "engineering"
}
)
```
### Service Mesh (Istio, Linkerd)
```python
import litellm
response = litellm.completion(
model="gpt-4",
messages=[{"role": "user", "content": "Hello"}],
extra_headers={
"X-Trace-ID": "trace-abc-123",
"X-Service-Name": "ai-service",
"X-Version": "1.2.3"
}
)
```
### Multi-Tenant SaaS Applications
```python
import litellm
def make_ai_request(user_id, tenant_id, content):
return litellm.completion(
model="claude-3-5-sonnet-latest",
messages=[{"role": "user", "content": content}],
extra_headers={
"X-User-ID": user_id,
"X-Tenant-ID": tenant_id,
"X-Request-Time": str(int(time.time()))
}
)
# Usage
response = make_ai_request("user-456", "tenant-org-1", "Help me write code")
```
### Request Tracing and Debugging
```python
import litellm
import uuid
def traced_completion(model, messages, **kwargs):
trace_id = str(uuid.uuid4())
return litellm.completion(
model=model,
messages=messages,
extra_headers={
"X-Trace-ID": trace_id,
"X-Debug-Mode": "true",
"X-Source-Service": "my-app"
},
**kwargs
)
# Usage
response = traced_completion(
model="gpt-4",
messages=[{"role": "user", "content": "Debug this issue"}]
)
```
### Custom Authentication
```python
import litellm
def get_custom_auth_token():
# Your custom authentication logic
return "custom-auth-token"
response = litellm.completion(
model="claude-3-5-sonnet-latest",
messages=[{"role": "user", "content": "Hello"}],
headers={
"X-Custom-Auth": get_custom_auth_token(),
"X-Auth-Type": "custom"
}
)
```
## Provider Support
Headers are supported across all LiteLLM providers including:
- **OpenAI** (GPT models)
- **Anthropic** (Claude models)
- **Cohere**
- **Hugging Face**
- **Custom providers**
- **Azure OpenAI**
- **AWS Bedrock**
- **Google Vertex AI**
Each provider will receive your custom headers along with their required authentication and API-specific headers.
## Best Practices
### 1. Use Meaningful Header Names
```python
# Good
extra_headers = {
"X-Request-ID": "req-12345",
"X-Tenant-ID": "org-456"
}
# Avoid
extra_headers = {
"custom1": "value1",
"h2": "value2"
}
```
### 2. Include Tracing Information
```python
extra_headers = {
"X-Trace-ID": trace_id,
"X-Span-ID": span_id,
"X-Service-Name": "ai-service"
}
```
### 3. Handle Sensitive Information Carefully
```python
# Don't log sensitive headers
import os
if os.getenv("ENVIRONMENT") != "production":
extra_headers["X-Debug-User"] = user_id
```
### 4. Use Environment-Specific Headers
```python
import os
environment = os.getenv("ENVIRONMENT", "development")
litellm.headers = {
"X-Environment": environment,
"X-Service-Version": os.getenv("SERVICE_VERSION", "unknown")
}
```
## Troubleshooting
### Headers Not Being Passed
If your headers aren't reaching the API:
1. **Check Header Names**: Ensure header names don't conflict with provider-specific headers
2. **Verify Priority**: Remember that `headers` > `extra_headers` > `litellm.headers`
3. **Test with Logging**: Enable verbose logging to see what headers are being sent
```python
import litellm
# Enable debug logging
litellm.set_verbose = True
response = litellm.completion(
model="gpt-4",
messages=[{"role": "user", "content": "test"}],
extra_headers={"X-Debug": "test"}
)
```
### Gateway or Proxy Issues
If using API gateways or proxies:
1. **Check Gateway Requirements**: Verify required headers for your gateway
2. **Test Direct vs Gateway**: Compare direct API calls vs gateway calls
3. **Validate Header Format**: Some gateways have header format requirements
## Security Considerations
1. **Don't Log Sensitive Headers**: Avoid logging authentication tokens or personal data
2. **Use HTTPS**: Always use secure connections when passing sensitive headers
3. **Validate Header Values**: Sanitize user-provided header values
4. **Rotate Keys**: Regularly rotate any API keys passed in headers
```python
import litellm
import re
def safe_header_value(value):
# Remove potentially dangerous characters
return re.sub(r'[^\w\-.]', '', str(value))
response = litellm.completion(
model="gpt-4",
messages=[{"role": "user", "content": "Hello"}],
extra_headers={
"X-User-ID": safe_header_value(user_id)
}
)
```

View file

@ -0,0 +1,195 @@
#!/usr/bin/env python3
"""
Example demonstrating LiteLLM SDK header support for enterprise environments.
This example shows how to use additional headers with API gateways, service meshes,
and multi-tenant architectures.
"""
import litellm
import os
from typing import Dict, Any
def example_global_headers():
"""Example: Set global headers for all requests"""
print("=== Global Headers Example ===")
# Set global headers that will be included in all API requests
litellm.headers = {
"X-API-Gateway-Key": "your-gateway-key-here",
"X-Company-ID": "acme-corp",
"X-Environment": "production"
}
print("Global headers set:", litellm.headers)
# These headers will now be included in all completion calls
# (Note: This example doesn't actually make API calls)
print("Global headers will be included in all subsequent completion() calls")
def example_per_request_headers():
"""Example: Using extra_headers for specific requests"""
print("\n=== Per-Request Headers Example ===")
headers_to_send = {
"X-Request-ID": "req-12345",
"X-Tenant-ID": "tenant-abc",
"X-Custom-Auth": "bearer-token-xyz"
}
print("Per-request headers:", headers_to_send)
# Example of how you would use extra_headers in a real call
# response = litellm.completion(
# model="claude-3-5-sonnet-latest",
# messages=[{"role": "user", "content": "Hello"}],
# extra_headers=headers_to_send
# )
def example_header_priority():
"""Example: Demonstrating header priority and merging"""
print("\n=== Header Priority Example ===")
# Set global headers
litellm.headers = {
"X-Company-ID": "acme-corp",
"X-Shared-Header": "global-value"
}
# Headers that would be sent in a request
extra_headers = {
"X-Request-ID": "req-12345",
"X-Shared-Header": "extra-value" # Overrides global
}
request_headers = {
"X-Priority-Header": "important",
"X-Shared-Header": "request-value" # Overrides both global and extra
}
print("Global headers:", litellm.headers)
print("Extra headers:", extra_headers)
print("Request headers:", request_headers)
print("\nFinal headers would be:")
print(" X-Company-ID: acme-corp (from global)")
print(" X-Request-ID: req-12345 (from extra)")
print(" X-Priority-Header: important (from request)")
print(" X-Shared-Header: request-value (request wins - highest priority)")
def example_enterprise_api_gateway():
"""Example: Enterprise API Gateway scenario"""
print("\n=== Enterprise API Gateway Example ===")
# Simulate enterprise environment with Apigee or similar
gateway_config = {
"X-API-Gateway-Key": os.getenv("API_GATEWAY_KEY", "demo-key"),
"X-Route-Version": "v2",
"X-Rate-Limit-Group": "premium"
}
# Set gateway headers globally
litellm.headers = gateway_config
print("Gateway headers configured:", gateway_config)
# Function to make tenant-specific requests
def make_tenant_request(tenant_id: str, user_id: str, content: str) -> Dict[str, Any]:
"""Make an AI request with tenant-specific headers"""
tenant_headers = {
"X-Tenant-ID": tenant_id,
"X-User-ID": user_id,
"X-Request-Time": "2024-01-01T00:00:00Z",
"X-Service-Name": "ai-assistant"
}
print(f"Making request for tenant {tenant_id}, user {user_id}")
print("Tenant-specific headers:", tenant_headers)
# In a real scenario, this would make the actual API call:
# return litellm.completion(
# model="claude-3-5-sonnet-latest",
# messages=[{"role": "user", "content": content}],
# extra_headers=tenant_headers
# )
# For demo purposes, return mock data
return {"mock": "response", "headers_used": {**gateway_config, **tenant_headers}}
# Example usage
result = make_tenant_request("tenant-123", "user-456", "Analyze this data")
print("Response:", result)
def example_service_mesh():
"""Example: Service mesh integration (Istio, Linkerd)"""
print("\n=== Service Mesh Example ===")
service_mesh_headers = {
"X-Trace-ID": "trace-abc-123",
"X-Span-ID": "span-def-456",
"X-Service-Name": "ai-service",
"X-Version": "1.2.3",
"X-Cluster": "prod-us-west-2"
}
print("Service mesh headers:", service_mesh_headers)
# Example of using these headers for distributed tracing
# response = litellm.completion(
# model="gpt-4",
# messages=[{"role": "user", "content": "Hello"}],
# extra_headers=service_mesh_headers
# )
def example_debugging_and_monitoring():
"""Example: Request debugging and monitoring"""
print("\n=== Debugging and Monitoring Example ===")
import uuid
import time
# Generate unique identifiers for request tracking
trace_id = str(uuid.uuid4())
request_id = f"req-{int(time.time())}"
debug_headers = {
"X-Trace-ID": trace_id,
"X-Request-ID": request_id,
"X-Debug-Mode": "true",
"X-Source-Service": "customer-support-bot",
"X-Request-Priority": "high"
}
print("Debug headers:", debug_headers)
print(f"Trace ID: {trace_id}")
print(f"Request ID: {request_id}")
# These headers help with:
# 1. Distributed tracing across services
# 2. Request correlation in logs
# 3. Debug mode enablement
# 4. Priority-based routing
if __name__ == "__main__":
print("LiteLLM SDK Header Support Examples")
print("=" * 50)
example_global_headers()
example_per_request_headers()
example_header_priority()
example_enterprise_api_gateway()
example_service_mesh()
example_debugging_and_monitoring()
print("\n" + "=" * 50)
print("All examples completed!")
print("\nTo use in your application:")
print("1. Set litellm.headers for global headers")
print("2. Use extra_headers parameter for request-specific headers")
print("3. Use headers parameter for highest priority headers")
print("4. Headers are merged with priority: headers > extra_headers > litellm.headers")

View file

@ -1004,7 +1004,15 @@ def completion( # type: ignore # noqa: PLR0915
provider_specific_header = cast(
Optional[ProviderSpecificHeader], kwargs.get("provider_specific_header", None)
)
headers = kwargs.get("headers", None) or extra_headers
# Properly merge headers with priority: request headers > extra_headers > global litellm.headers
headers = {}
if litellm.headers is not None and isinstance(litellm.headers, dict):
headers.update(litellm.headers)
if extra_headers is not None and isinstance(extra_headers, dict):
headers.update(extra_headers)
request_headers = kwargs.get("headers", None)
if request_headers is not None and isinstance(request_headers, dict):
headers.update(request_headers)
ensure_alternating_roles: Optional[bool] = kwargs.get(
"ensure_alternating_roles", None
@ -1015,10 +1023,6 @@ def completion( # type: ignore # noqa: PLR0915
assistant_continue_message: Optional[ChatCompletionAssistantMessage] = kwargs.get(
"assistant_continue_message", None
)
if headers is None:
headers = {}
if extra_headers is not None:
headers.update(extra_headers)
num_retries = kwargs.get(
"num_retries", None
) ## alt. param for 'max_retries'. Use this to pass retries w/ instructor.
@ -1075,7 +1079,6 @@ def completion( # type: ignore # noqa: PLR0915
prompt_id=prompt_id, non_default_params=non_default_params
)
):
(
model,
messages,
@ -1428,8 +1431,7 @@ def completion( # type: ignore # noqa: PLR0915
"azure_ad_token_provider", None
)
headers = headers or litellm.headers
# Use the consolidated headers that were already merged at the top of the function
if extra_headers is not None:
optional_params["extra_headers"] = extra_headers
if max_retries is not None:
@ -1694,8 +1696,7 @@ def completion( # type: ignore # noqa: PLR0915
or get_secret("OPENAI_API_KEY")
)
headers = headers or litellm.headers
# Use the consolidated headers that were already merged at the top of the function
if extra_headers is not None:
optional_params["extra_headers"] = extra_headers
@ -2032,7 +2033,6 @@ def completion( # type: ignore # noqa: PLR0915
try:
if use_base_llm_http_handler:
response = base_llm_http_handler.completion(
model=model,
messages=messages,
@ -2411,12 +2411,8 @@ def completion( # type: ignore # noqa: PLR0915
or "https://api.cohere.ai/v1/chat"
)
headers = headers or litellm.headers or {}
if headers is None:
headers = {}
if extra_headers is not None:
headers.update(extra_headers)
# Use the consolidated headers that were already merged at the top of the function
# No need for additional merging here as it's already done
response = base_llm_http_handler.completion(
model=model,
@ -2512,15 +2508,10 @@ def completion( # type: ignore # noqa: PLR0915
)
elif custom_llm_provider == "compactifai":
api_key = (
api_key
or get_secret_str("COMPACTIFAI_API_KEY")
or litellm.api_key
api_key or get_secret_str("COMPACTIFAI_API_KEY") or litellm.api_key
)
api_base = (
api_base
or "https://api.compactif.ai/v1"
)
api_base = api_base or "https://api.compactif.ai/v1"
## COMPLETION CALL
response = base_llm_http_handler.completion(
@ -3106,9 +3097,9 @@ def completion( # type: ignore # noqa: PLR0915
"aws_region_name" not in optional_params
or optional_params["aws_region_name"] is None
):
optional_params["aws_region_name"] = (
aws_bedrock_client.meta.region_name
)
optional_params[
"aws_region_name"
] = aws_bedrock_client.meta.region_name
bedrock_route = BedrockModelInfo.get_bedrock_route(model)
if bedrock_route == "converse":
@ -3450,7 +3441,6 @@ def completion( # type: ignore # noqa: PLR0915
)
raise e
elif custom_llm_provider == "gradient_ai":
api_base = litellm.api_base or api_base
response = base_llm_http_handler.completion(
model=model,
@ -3810,7 +3800,7 @@ def embedding(
*,
aembedding: Literal[True],
**kwargs,
) -> Coroutine[Any, Any, EmbeddingResponse]:
) -> Coroutine[Any, Any, EmbeddingResponse]:
...
@ -3836,7 +3826,7 @@ def embedding(
*,
aembedding: Literal[False] = False,
**kwargs,
) -> EmbeddingResponse:
) -> EmbeddingResponse:
...
# fmt: on
@ -4147,10 +4137,8 @@ def embedding( # noqa: PLR0915
or litellm.api_key
)
if extra_headers is not None and isinstance(extra_headers, dict):
headers = extra_headers
else:
headers = {}
# Use the consolidated headers that were already merged at the top of the function
# No need for additional merging here as it's already done
response = base_llm_http_handler.embedding(
model=model,
@ -5089,9 +5077,9 @@ def adapter_completion(
new_kwargs = translation_obj.translate_completion_input_params(kwargs=kwargs)
response: Union[ModelResponse, CustomStreamWrapper] = completion(**new_kwargs) # type: ignore
translated_response: Optional[Union[BaseModel, AdapterCompletionStreamWrapper]] = (
None
)
translated_response: Optional[
Union[BaseModel, AdapterCompletionStreamWrapper]
] = None
if isinstance(response, ModelResponse):
translated_response = translation_obj.translate_completion_output_params(
response=response
@ -6079,9 +6067,9 @@ def stream_chunk_builder( # noqa: PLR0915
]
if len(content_chunks) > 0:
response["choices"][0]["message"]["content"] = (
processor.get_combined_content(content_chunks)
)
response["choices"][0]["message"][
"content"
] = processor.get_combined_content(content_chunks)
thinking_blocks = [
chunk
@ -6092,9 +6080,9 @@ def stream_chunk_builder( # noqa: PLR0915
]
if len(thinking_blocks) > 0:
response["choices"][0]["message"]["thinking_blocks"] = (
processor.get_combined_thinking_content(thinking_blocks)
)
response["choices"][0]["message"][
"thinking_blocks"
] = processor.get_combined_thinking_content(thinking_blocks)
reasoning_chunks = [
chunk
@ -6105,9 +6093,9 @@ def stream_chunk_builder( # noqa: PLR0915
]
if len(reasoning_chunks) > 0:
response["choices"][0]["message"]["reasoning_content"] = (
processor.get_combined_reasoning_content(reasoning_chunks)
)
response["choices"][0]["message"][
"reasoning_content"
] = processor.get_combined_reasoning_content(reasoning_chunks)
audio_chunks = [
chunk

View file

@ -5,7 +5,6 @@ import sys
import httpx
import pytest
import respx
from fastapi.testclient import TestClient
sys.path.insert(
0, os.path.abspath("../..")
@ -176,7 +175,7 @@ async def test_url_with_format_param(model, sync_mode, monkeypatch):
else:
response = await acompletion(**args, client=client)
print(response)
except Exception as e:
except Exception:
pass
mock_client.assert_called()
@ -269,7 +268,6 @@ def test_bedrock_latency_optimized_inference():
def test_custom_provider_with_extra_headers():
from litellm.llms.custom_httpx.http_handler import HTTPHandler
with patch.object(
litellm.llms.custom_httpx.http_handler.HTTPHandler, "post"
@ -286,7 +284,6 @@ def test_custom_provider_with_extra_headers():
def test_custom_provider_with_extra_body():
from litellm.llms.custom_httpx.http_handler import HTTPHandler
with patch.object(
litellm.llms.custom_httpx.http_handler.HTTPHandler, "post"
@ -1132,52 +1129,57 @@ def test_anthropic_disable_url_suffix_env_var():
# Test with environment variable disabled (default behavior)
with patch.dict(os.environ, {"ANTHROPIC_API_BASE": "https://api.example.com"}):
actual_api_base = None
with patch("litellm.main.anthropic_chat_completions") as mock_anthropic:
def capture_completion(**kwargs):
nonlocal actual_api_base
actual_api_base = kwargs.get("api_base")
mock_response = MagicMock()
mock_response.choices = [MagicMock()]
return mock_response
mock_anthropic.completion = capture_completion
# This should append /v1/messages
completion(
model="anthropic/claude-3-sonnet",
messages=[{"role": "user", "content": "test"}],
api_key="test-key"
api_key="test-key",
)
# Verify the api_base has /v1/messages appended
assert actual_api_base.endswith("/v1/messages")
assert actual_api_base == "https://api.example.com/v1/messages"
# Test with environment variable enabled
with patch.dict(os.environ, {
"ANTHROPIC_API_BASE": "https://api.example.com/custom/path",
"LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX": "true"
}):
with patch.dict(
os.environ,
{
"ANTHROPIC_API_BASE": "https://api.example.com/custom/path",
"LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX": "true",
},
):
actual_api_base = None
with patch("litellm.main.anthropic_chat_completions") as mock_anthropic:
def capture_completion(**kwargs):
nonlocal actual_api_base
actual_api_base = kwargs.get("api_base")
mock_response = MagicMock()
mock_response.choices = [MagicMock()]
return mock_response
mock_anthropic.completion = capture_completion
# This should NOT append /v1/messages
completion(
model="anthropic/claude-3-sonnet",
messages=[{"role": "user", "content": "test"}],
api_key="test-key"
api_key="test-key",
)
# Verify the api_base does not have /v1/messages appended
assert actual_api_base == "https://api.example.com/custom/path"
assert not actual_api_base.endswith("/v1/messages")
@ -1192,48 +1194,260 @@ def test_anthropic_text_disable_url_suffix_env_var():
# Test with environment variable disabled (default behavior)
with patch.dict(os.environ, {"ANTHROPIC_API_BASE": "https://api.example.com"}):
actual_api_base = None
with patch("litellm.main.base_llm_http_handler") as mock_handler:
def capture_completion(**kwargs):
nonlocal actual_api_base
actual_api_base = kwargs.get("api_base")
return MagicMock()
mock_handler.completion = capture_completion
# This should append /v1/complete
completion(
model="anthropic_text/claude-instant-1",
messages=[{"role": "user", "content": "test"}],
api_key="test-key"
api_key="test-key",
)
# Verify the api_base has /v1/complete appended
assert actual_api_base.endswith("/v1/complete")
assert actual_api_base == "https://api.example.com/v1/complete"
# Test with environment variable enabled
with patch.dict(os.environ, {
"ANTHROPIC_API_BASE": "https://api.example.com/custom/complete",
"LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX": "true"
}):
with patch.dict(
os.environ,
{
"ANTHROPIC_API_BASE": "https://api.example.com/custom/complete",
"LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX": "true",
},
):
actual_api_base = None
with patch("litellm.main.base_llm_http_handler") as mock_handler:
def capture_completion(**kwargs):
nonlocal actual_api_base
actual_api_base = kwargs.get("api_base")
return MagicMock()
mock_handler.completion = capture_completion
# This should NOT append /v1/complete
completion(
model="anthropic_text/claude-instant-1",
messages=[{"role": "user", "content": "test"}],
api_key="test-key"
api_key="test-key",
)
# Verify the api_base does not have /v1/complete appended
assert actual_api_base == "https://api.example.com/custom/complete"
assert not actual_api_base.endswith("/v1/complete")
# Test header handling functionality
def test_header_priority_and_merging():
"""Test that headers are properly merged with correct priority: request headers > extra_headers > global litellm.headers"""
import litellm
from unittest.mock import patch, MagicMock
# Store original headers to restore later
original_headers = litellm.headers
try:
# Set global headers
litellm.headers = {
"X-Global-Header": "global-value",
"X-Shared-Header": "global",
}
captured_headers = {}
with patch(
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post"
) as mock_post:
def capture_headers(*args, **kwargs):
captured_headers.update(kwargs.get("headers", {}))
mock_response = MagicMock()
mock_response.json.return_value = {
"choices": [{"message": {"content": "test"}}],
"usage": {"total_tokens": 10},
}
mock_response.status_code = 200
return mock_response
mock_post.side_effect = capture_headers
# Test header merging
try:
litellm.completion(
model="custom/test-model",
messages=[{"role": "user", "content": "test"}],
api_base="https://example.com/api",
extra_headers={
"X-Extra-Header": "extra-value",
"X-Shared-Header": "extra",
},
headers={
"X-Request-Header": "request-value",
"X-Shared-Header": "request",
},
)
except Exception:
# Expected since we're mocking
pass
# Verify header priority: request > extra > global
assert "X-Global-Header" in captured_headers
assert "X-Extra-Header" in captured_headers
assert "X-Request-Header" in captured_headers
assert captured_headers["X-Global-Header"] == "global-value"
assert captured_headers["X-Extra-Header"] == "extra-value"
assert captured_headers["X-Request-Header"] == "request-value"
# Request headers should override others
assert captured_headers["X-Shared-Header"] == "request"
finally:
# Restore original headers
litellm.headers = original_headers
def test_anthropic_header_passing():
"""Test that custom headers are properly passed to Anthropic API calls"""
from unittest.mock import patch, MagicMock
captured_headers = {}
with patch("litellm.llms.anthropic.chat.handler.HTTPHandler.post") as mock_post:
def capture_headers(*args, **kwargs):
captured_headers.update(kwargs.get("headers", {}))
mock_response = MagicMock()
mock_response.json.return_value = {
"content": [{"text": "test response"}],
"usage": {"input_tokens": 5, "output_tokens": 5},
}
mock_response.status_code = 200
return mock_response
mock_post.side_effect = capture_headers
try:
litellm.completion(
model="claude-3-5-sonnet-latest",
messages=[{"role": "user", "content": "Hello"}],
extra_headers={
"X-API-Gateway-Key": "gateway-123",
"X-Tenant-ID": "tenant-456",
},
)
except Exception:
# Expected since we're mocking
pass
# Verify custom headers are included along with anthropic headers
assert "X-API-Gateway-Key" in captured_headers
assert "X-Tenant-ID" in captured_headers
assert captured_headers["X-API-Gateway-Key"] == "gateway-123"
assert captured_headers["X-Tenant-ID"] == "tenant-456"
# Verify anthropic-specific headers are also present
assert "x-api-key" in captured_headers
assert "anthropic-version" in captured_headers
def test_openai_header_passing():
"""Test that custom headers are properly passed to OpenAI API calls"""
from unittest.mock import patch, MagicMock
captured_extra_headers = {}
with patch(
"litellm.llms.openai.chat.gpt_transformation.OpenAIGPTConfig.transform_request"
) as mock_transform:
with patch("litellm.completion_cost") as mock_cost:
mock_cost.return_value = 0.0
mock_transform.return_value = {"model": "gpt-4", "messages": []}
with patch("openai.OpenAI") as mock_openai:
mock_client = MagicMock()
mock_openai.return_value = mock_client
mock_response = MagicMock()
mock_response.choices = [MagicMock()]
mock_response.choices[0].message.content = "test response"
mock_response.usage.total_tokens = 10
mock_client.chat.completions.create.return_value = mock_response
try:
litellm.completion(
model="gpt-4",
messages=[{"role": "user", "content": "Hello"}],
extra_headers={
"X-Custom-Auth": "bearer-token",
"X-Request-ID": "req-789",
},
)
except Exception:
# Expected since we're mocking
pass
# Verify extra_headers were passed to the OpenAI client
call_args = mock_client.chat.completions.create.call_args
if call_args:
kwargs = call_args.kwargs
assert "extra_headers" in kwargs
extra_headers = kwargs["extra_headers"]
assert "X-Custom-Auth" in extra_headers
assert "X-Request-ID" in extra_headers
def test_global_headers_functionality():
"""Test that global litellm.headers work correctly"""
import litellm
from unittest.mock import patch, MagicMock
# Store original headers to restore later
original_headers = litellm.headers
try:
# Set global headers
litellm.headers = {"X-Company-ID": "acme-corp", "X-Environment": "production"}
captured_headers = {}
with patch(
"litellm.llms.custom_httpx.http_handler.HTTPHandler.post"
) as mock_post:
def capture_headers(*args, **kwargs):
captured_headers.update(kwargs.get("headers", {}))
mock_response = MagicMock()
mock_response.json.return_value = {
"choices": [{"message": {"content": "test"}}],
"usage": {"total_tokens": 10},
}
mock_response.status_code = 200
return mock_response
mock_post.side_effect = capture_headers
try:
litellm.completion(
model="custom/test-model",
messages=[{"role": "user", "content": "test"}],
api_base="https://example.com/api",
)
except Exception:
# Expected since we're mocking
pass
# Verify global headers are included
assert "X-Company-ID" in captured_headers
assert "X-Environment" in captured_headers
assert captured_headers["X-Company-ID"] == "acme-corp"
assert captured_headers["X-Environment"] == "production"
finally:
# Restore original headers
litellm.headers = original_headers