From 3986b073c9939b7780cde47c89783c953ea0b930 Mon Sep 17 00:00:00 2001 From: uzaxirr Date: Sun, 21 Sep 2025 14:58:14 +0530 Subject: [PATCH 1/2] feat: Add SDK support for additional headers --- docs/my-website/docs/sdk/headers.md | 328 ++++++++++++++++++++++++++++ examples/sdk_headers_example.py | 195 +++++++++++++++++ litellm/main.py | 28 ++- tests/test_litellm/test_main.py | 183 ++++++++++++++++ 4 files changed, 719 insertions(+), 15 deletions(-) create mode 100644 docs/my-website/docs/sdk/headers.md create mode 100644 examples/sdk_headers_example.py diff --git a/docs/my-website/docs/sdk/headers.md b/docs/my-website/docs/sdk/headers.md new file mode 100644 index 00000000000..a1fdaadc6f5 --- /dev/null +++ b/docs/my-website/docs/sdk/headers.md @@ -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) + } +) +``` \ No newline at end of file diff --git a/examples/sdk_headers_example.py b/examples/sdk_headers_example.py new file mode 100644 index 00000000000..4ce31df4a70 --- /dev/null +++ b/examples/sdk_headers_example.py @@ -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") \ No newline at end of file diff --git a/litellm/main.py b/litellm/main.py index 6100ab5de22..abdc03d3dd3 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -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. @@ -1428,8 +1432,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 +1697,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 @@ -2411,12 +2413,8 @@ def completion( # type: ignore # noqa: PLR0915 or "https://api.cohere.ai/v1/generate" ) - 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, diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index 954597dda25..8962f383e5b 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -1237,3 +1237,186 @@ def test_anthropic_text_disable_url_suffix_env_var(): # 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 as e: + # 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 as e: + # 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 as e: + # 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 as e: + # 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 From 50d717cde07732646576b12d0def78335e50e073 Mon Sep 17 00:00:00 2001 From: uzaxirr Date: Mon, 22 Sep 2025 04:46:08 +0530 Subject: [PATCH 2/2] Apply Black formatting and fix Ruff issues - Format code with Black to meet style requirements - Fix auto-fixable Ruff linting issues - Maintain header implementation functionality --- litellm/main.py | 46 +++++------ tests/test_litellm/test_main.py | 139 +++++++++++++++++++------------- 2 files changed, 104 insertions(+), 81 deletions(-) diff --git a/litellm/main.py b/litellm/main.py index abdc03d3dd3..549c108d7b4 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1079,7 +1079,6 @@ def completion( # type: ignore # noqa: PLR0915 prompt_id=prompt_id, non_default_params=non_default_params ) ): - ( model, messages, @@ -2034,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, @@ -2550,15 +2548,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( @@ -3144,9 +3137,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": @@ -3488,7 +3481,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, @@ -3848,7 +3840,7 @@ def embedding( *, aembedding: Literal[True], **kwargs, -) -> Coroutine[Any, Any, EmbeddingResponse]: +) -> Coroutine[Any, Any, EmbeddingResponse]: ... @@ -3874,7 +3866,7 @@ def embedding( *, aembedding: Literal[False] = False, **kwargs, -) -> EmbeddingResponse: +) -> EmbeddingResponse: ... # fmt: on @@ -5127,9 +5119,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 @@ -6117,9 +6109,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 @@ -6130,9 +6122,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 @@ -6143,9 +6135,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 diff --git a/tests/test_litellm/test_main.py b/tests/test_litellm/test_main.py index 8962f383e5b..631592f4c1f 100644 --- a/tests/test_litellm/test_main.py +++ b/tests/test_litellm/test_main.py @@ -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,53 @@ 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") @@ -1250,15 +1257,24 @@ def test_header_priority_and_merging(): try: # Set global headers - litellm.headers = {"X-Global-Header": "global-value", "X-Shared-Header": "global"} + 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: + 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.json.return_value = { + "choices": [{"message": {"content": "test"}}], + "usage": {"total_tokens": 10}, + } mock_response.status_code = 200 return mock_response @@ -1270,10 +1286,16 @@ def test_header_priority_and_merging(): 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"} + extra_headers={ + "X-Extra-Header": "extra-value", + "X-Shared-Header": "extra", + }, + headers={ + "X-Request-Header": "request-value", + "X-Shared-Header": "request", + }, ) - except Exception as e: + except Exception: # Expected since we're mocking pass @@ -1299,10 +1321,14 @@ def test_anthropic_header_passing(): 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.json.return_value = { + "content": [{"text": "test response"}], + "usage": {"input_tokens": 5, "output_tokens": 5}, + } mock_response.status_code = 200 return mock_response @@ -1314,10 +1340,10 @@ def test_anthropic_header_passing(): messages=[{"role": "user", "content": "Hello"}], extra_headers={ "X-API-Gateway-Key": "gateway-123", - "X-Tenant-ID": "tenant-456" - } + "X-Tenant-ID": "tenant-456", + }, ) - except Exception as e: + except Exception: # Expected since we're mocking pass @@ -1337,7 +1363,9 @@ def test_openai_header_passing(): captured_extra_headers = {} - with patch("litellm.llms.openai.chat.gpt_transformation.OpenAIGPTConfig.transform_request") as mock_transform: + 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": []} @@ -1357,10 +1385,10 @@ def test_openai_header_passing(): messages=[{"role": "user", "content": "Hello"}], extra_headers={ "X-Custom-Auth": "bearer-token", - "X-Request-ID": "req-789" - } + "X-Request-ID": "req-789", + }, ) - except Exception as e: + except Exception: # Expected since we're mocking pass @@ -1384,18 +1412,21 @@ def test_global_headers_functionality(): try: # Set global headers - litellm.headers = { - "X-Company-ID": "acme-corp", - "X-Environment": "production" - } + 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: + 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.json.return_value = { + "choices": [{"message": {"content": "test"}}], + "usage": {"total_tokens": 10}, + } mock_response.status_code = 200 return mock_response @@ -1405,9 +1436,9 @@ def test_global_headers_functionality(): litellm.completion( model="custom/test-model", messages=[{"role": "user", "content": "test"}], - api_base="https://example.com/api" + api_base="https://example.com/api", ) - except Exception as e: + except Exception: # Expected since we're mocking pass