mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge pull request #14761 from uzaxirr/feat/sdk-additional-headers
feat: Add SDK support for additional headers
This commit is contained in:
commit
8628c265b9
4 changed files with 803 additions and 78 deletions
328
docs/my-website/docs/sdk/headers.md
Normal file
328
docs/my-website/docs/sdk/headers.md
Normal 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)
|
||||
}
|
||||
)
|
||||
```
|
||||
195
examples/sdk_headers_example.py
Normal file
195
examples/sdk_headers_example.py
Normal 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")
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue