diff --git a/.circleci/config.yml b/.circleci/config.yml
index c9de5adcc9f..7debc582915 100644
--- a/.circleci/config.yml
+++ b/.circleci/config.yml
@@ -1913,7 +1913,7 @@ jobs:
-e APORIA_API_BASE_1=$APORIA_API_BASE_1 \
-e AWS_ACCESS_KEY_ID=$AWS_ACCESS_KEY_ID \
-e AWS_SECRET_ACCESS_KEY=$AWS_SECRET_ACCESS_KEY \
- -e DEFAULT_NUM_WORKERS_LITELLM_PROXY=1
+ -e DEFAULT_NUM_WORKERS_LITELLM_PROXY=1 \
-e USE_DDTRACE=True \
-e DD_API_KEY=$DD_API_KEY \
-e DD_SITE=$DD_SITE \
diff --git a/cookbook/veo_video_generation.py b/cookbook/veo_video_generation.py
new file mode 100644
index 00000000000..64a7207feb1
--- /dev/null
+++ b/cookbook/veo_video_generation.py
@@ -0,0 +1,311 @@
+#!/usr/bin/env python3
+"""
+Complete example for Veo video generation through LiteLLM proxy.
+
+This script demonstrates how to:
+1. Generate videos using Google's Veo model
+2. Poll for completion status
+3. Download the generated video file
+
+Requirements:
+- LiteLLM proxy running with Google AI Studio pass-through configured
+- Google AI Studio API key with Veo access
+"""
+
+import json
+import os
+import time
+import requests
+from typing import Optional
+
+
+class VeoVideoGenerator:
+ """Complete Veo video generation client using LiteLLM proxy."""
+
+ def __init__(self, base_url: str = "http://localhost:4000/gemini/v1beta",
+ api_key: str = "sk-1234"):
+ """
+ Initialize the Veo video generator.
+
+ Args:
+ base_url: Base URL for the LiteLLM proxy with Gemini pass-through
+ api_key: API key for LiteLLM proxy authentication
+ """
+ self.base_url = base_url
+ self.api_key = api_key
+ self.headers = {
+ "x-goog-api-key": api_key,
+ "Content-Type": "application/json"
+ }
+
+ def generate_video(self, prompt: str) -> Optional[str]:
+ """
+ Initiate video generation with Veo.
+
+ Args:
+ prompt: Text description of the video to generate
+
+ Returns:
+ Operation name if successful, None otherwise
+ """
+ print(f"š¬ Generating video with prompt: '{prompt}'")
+
+ url = f"{self.base_url}/models/veo-3.0-generate-preview:predictLongRunning"
+ payload = {
+ "instances": [{
+ "prompt": prompt
+ }]
+ }
+
+ try:
+ response = requests.post(url, headers=self.headers, json=payload)
+ response.raise_for_status()
+
+ data = response.json()
+ operation_name = data.get("name")
+
+ if operation_name:
+ print(f"ā
Video generation started: {operation_name}")
+ return operation_name
+ else:
+ print("ā No operation name returned")
+ print(f"Response: {json.dumps(data, indent=2)}")
+ return None
+
+ except requests.RequestException as e:
+ print(f"ā Failed to start video generation: {e}")
+ if hasattr(e, 'response') and e.response is not None:
+ try:
+ error_data = e.response.json()
+ print(f"Error details: {json.dumps(error_data, indent=2)}")
+ except:
+ print(f"Error response: {e.response.text}")
+ return None
+
+ def wait_for_completion(self, operation_name: str, max_wait_time: int = 600) -> Optional[str]:
+ """
+ Poll operation status until video generation is complete.
+
+ Args:
+ operation_name: Name of the operation to monitor
+ max_wait_time: Maximum time to wait in seconds (default: 10 minutes)
+
+ Returns:
+ Video URI if successful, None otherwise
+ """
+ print("ā³ Waiting for video generation to complete...")
+
+ operation_url = f"{self.base_url}/{operation_name}"
+ start_time = time.time()
+ poll_interval = 10 # Start with 10 seconds
+
+ while time.time() - start_time < max_wait_time:
+ try:
+ print(f"š Polling status... ({int(time.time() - start_time)}s elapsed)")
+
+ response = requests.get(operation_url, headers=self.headers)
+ response.raise_for_status()
+
+ data = response.json()
+
+ # Check for errors
+ if "error" in data:
+ print("ā Error in video generation:")
+ print(json.dumps(data["error"], indent=2))
+ return None
+
+ # Check if operation is complete
+ is_done = data.get("done", False)
+
+ if is_done:
+ print("š Video generation complete!")
+
+ try:
+ # Extract video URI from nested response
+ video_uri = data["response"]["generateVideoResponse"]["generatedSamples"][0]["video"]["uri"]
+ print(f"š¹ Video URI: {video_uri}")
+ return video_uri
+ except KeyError as e:
+ print(f"ā Could not extract video URI: {e}")
+ print("Full response:")
+ print(json.dumps(data, indent=2))
+ return None
+
+ # Wait before next poll, with exponential backoff
+ time.sleep(poll_interval)
+ poll_interval = min(poll_interval * 1.2, 30) # Cap at 30 seconds
+
+ except requests.RequestException as e:
+ print(f"ā Error polling operation status: {e}")
+ time.sleep(poll_interval)
+
+ print(f"ā° Timeout after {max_wait_time} seconds")
+ return None
+
+ def download_video(self, video_uri: str, output_filename: str = "generated_video.mp4") -> bool:
+ """
+ Download the generated video file.
+
+ Args:
+ video_uri: URI of the video to download (from Google's response)
+ output_filename: Local filename to save the video
+
+ Returns:
+ True if download successful, False otherwise
+ """
+ print(f"ā¬ļø Downloading video...")
+ print(f"Original URI: {video_uri}")
+
+ # Convert Google URI to LiteLLM proxy URI
+ # Example: files/abc123 -> /gemini/v1beta/files/abc123:download?alt=media
+ if video_uri.startswith("files/"):
+ download_path = f"{video_uri}:download?alt=media"
+ else:
+ download_path = video_uri
+
+ litellm_download_url = f"{self.base_url}/{download_path}"
+ print(f"Download URL: {litellm_download_url}")
+
+ try:
+ # Download with streaming and redirect handling
+ response = requests.get(
+ litellm_download_url,
+ headers=self.headers,
+ stream=True,
+ allow_redirects=True # Handle redirects automatically
+ )
+ response.raise_for_status()
+
+ # Save video file
+ with open(output_filename, 'wb') as f:
+ downloaded_size = 0
+ for chunk in response.iter_content(chunk_size=8192):
+ if chunk:
+ f.write(chunk)
+ downloaded_size += len(chunk)
+
+ # Progress indicator for large files
+ if downloaded_size % (1024 * 1024) == 0: # Every MB
+ print(f"š¦ Downloaded {downloaded_size / (1024*1024):.1f} MB...")
+
+ # Verify file was created and has content
+ if os.path.exists(output_filename):
+ file_size = os.path.getsize(output_filename)
+ if file_size > 0:
+ print(f"ā
Video downloaded successfully!")
+ print(f"š Saved as: {output_filename}")
+ print(f"š File size: {file_size / (1024*1024):.2f} MB")
+ return True
+ else:
+ print("ā Downloaded file is empty")
+ os.remove(output_filename)
+ return False
+ else:
+ print("ā File was not created")
+ return False
+
+ except requests.RequestException as e:
+ print(f"ā Download failed: {e}")
+ if hasattr(e, 'response') and e.response is not None:
+ print(f"Status code: {e.response.status_code}")
+ print(f"Response headers: {dict(e.response.headers)}")
+ return False
+
+ def generate_and_download(self, prompt: str, output_filename: str = None) -> bool:
+ """
+ Complete workflow: generate video and download it.
+
+ Args:
+ prompt: Text description for video generation
+ output_filename: Output filename (auto-generated if None)
+
+ Returns:
+ True if successful, False otherwise
+ """
+ # Auto-generate filename if not provided
+ if output_filename is None:
+ timestamp = int(time.time())
+ safe_prompt = "".join(c for c in prompt[:30] if c.isalnum() or c in (' ', '-', '_')).rstrip()
+ output_filename = f"veo_video_{safe_prompt.replace(' ', '_')}_{timestamp}.mp4"
+
+ print("=" * 60)
+ print("š¬ VEO VIDEO GENERATION WORKFLOW")
+ print("=" * 60)
+
+ # Step 1: Generate video
+ operation_name = self.generate_video(prompt)
+ if not operation_name:
+ return False
+
+ # Step 2: Wait for completion
+ video_uri = self.wait_for_completion(operation_name)
+ if not video_uri:
+ return False
+
+ # Step 3: Download video
+ success = self.download_video(video_uri, output_filename)
+
+ if success:
+ print("=" * 60)
+ print("š SUCCESS! Video generation complete!")
+ print(f"š Video saved as: {output_filename}")
+ print("=" * 60)
+ else:
+ print("=" * 60)
+ print("ā FAILED! Video generation or download failed")
+ print("=" * 60)
+
+ return success
+
+
+def main():
+ """
+ Example usage of the VeoVideoGenerator.
+
+ Configure these environment variables:
+ - LITELLM_BASE_URL: Your LiteLLM proxy URL (default: http://localhost:4000/gemini/v1beta)
+ - LITELLM_API_KEY: Your LiteLLM API key (default: sk-1234)
+ """
+
+ # Configuration from environment or defaults
+ base_url = os.getenv("LITELLM_BASE_URL", "http://localhost:4000/gemini/v1beta")
+ api_key = os.getenv("LITELLM_API_KEY", "sk-1234")
+
+ print("š Starting Veo Video Generation Example")
+ print(f"š” Using LiteLLM proxy at: {base_url}")
+
+ # Initialize generator
+ generator = VeoVideoGenerator(base_url=base_url, api_key=api_key)
+
+ # Example prompts - try different ones!
+ example_prompts = [
+ "A cat playing with a ball of yarn in a sunny garden",
+ "Ocean waves crashing against rocky cliffs at sunset",
+ "A bustling city street with people walking and cars passing by",
+ "A peaceful forest with sunlight filtering through the trees"
+ ]
+
+ # Use first example or get from user
+ prompt = example_prompts[0]
+ print(f"š¬ Using prompt: '{prompt}'")
+
+ # Generate and download video
+ success = generator.generate_and_download(prompt)
+
+ if success:
+ print("\nā
Example completed successfully!")
+ print("š” Try modifying the prompt in the script for different videos!")
+ else:
+ print("\nā Example failed!")
+ print("š§ Check your LiteLLM proxy configuration and Google AI Studio API key")
+
+ # Troubleshooting tips
+ print("\nš Troubleshooting:")
+ print("1. Ensure LiteLLM proxy is running with Google AI Studio pass-through")
+ print("2. Verify your Google AI Studio API key has Veo access")
+ print("3. Check that your prompt meets Veo's content guidelines")
+ print("4. Review the LiteLLM proxy logs for detailed error information")
+
+
+if __name__ == "__main__":
+ main()
diff --git a/docs/my-website/docs/observability/callbacks.md b/docs/my-website/docs/observability/callbacks.md
index 69cb0d053ee..040d83697d3 100644
--- a/docs/my-website/docs/observability/callbacks.md
+++ b/docs/my-website/docs/observability/callbacks.md
@@ -4,9 +4,14 @@
liteLLM provides `input_callbacks`, `success_callbacks` and `failure_callbacks`, making it easy for you to send data to a particular provider depending on the status of your responses.
+:::tip
+**New to LiteLLM Callbacks?** Check out our comprehensive [Callback Management Guide](./callback_management.md) to understand when to use different callback hooks like `async_log_success_event` vs `async_post_call_success_hook`.
+:::
+
liteLLM supports:
- [Custom Callback Functions](https://docs.litellm.ai/docs/observability/custom_callback)
+- [Callback Management Guide](./callback_management.md) - **Comprehensive guide for choosing the right hooks**
- [Lunary](https://lunary.ai/docs)
- [Langfuse](https://langfuse.com/docs)
- [LangSmith](https://www.langchain.com/langsmith)
diff --git a/docs/my-website/docs/observability/custom_callback.md b/docs/my-website/docs/observability/custom_callback.md
index cc586b2e5d9..c206c23d0f4 100644
--- a/docs/my-website/docs/observability/custom_callback.md
+++ b/docs/my-website/docs/observability/custom_callback.md
@@ -4,7 +4,6 @@
**For PROXY** [Go Here](../proxy/logging.md#custom-callback-class-async)
:::
-
## Callback Class
You can create a custom callback class to precisely log events as they occur in litellm.
@@ -57,6 +56,17 @@ def async completion():
asyncio.run(completion())
```
+## Common Hooks
+
+- `async_log_success_event` - Log successful API calls
+- `async_log_failure_event` - Log failed API calls
+- `log_pre_api_call` - Log before API call
+- `log_post_api_call` - Log after API call
+
+**Proxy-only hooks** (only work with LiteLLM Proxy):
+- `async_post_call_success_hook` - Access user data + modify responses
+- `async_pre_call_hook` - Modify requests before sending
+
## Callback Functions
If you just want to log on a specific event (e.g. on input) - you can use callback functions.
@@ -174,260 +184,87 @@ async def test_chat_openai():
asyncio.run(test_chat_openai())
```
-:::info
+## What's Available in kwargs?
-We're actively trying to expand this to other event types. [Tell us if you need this!](https://github.com/BerriAI/litellm/issues/1007)
-:::
-
-## What's in kwargs?
-
-Notice we pass in a kwargs argument to custom callback.
-```python
-def custom_callback(
- kwargs, # kwargs to completion
- completion_response, # response from completion
- start_time, end_time # start/end time
-):
- # Your custom code here
- print("LITELLM: in custom callback function")
- print("kwargs", kwargs)
- print("completion_response", completion_response)
- print("start_time", start_time)
- print("end_time", end_time)
-```
-
-This is a dictionary containing all the model-call details (the params we receive, the values we send to the http endpoint, the response we receive, stacktrace in case of errors, etc.).
-
-This is all logged in the [model_call_details via our Logger](https://github.com/BerriAI/litellm/blob/fc757dc1b47d2eb9d0ea47d6ad224955b705059d/litellm/utils.py#L246).
-
-Here's exactly what you can expect in the kwargs dictionary:
-```shell
-### DEFAULT PARAMS ###
-"model": self.model,
-"messages": self.messages,
-"optional_params": self.optional_params, # model-specific params passed in
-"litellm_params": self.litellm_params, # litellm-specific params passed in (e.g. metadata passed to completion call)
-"start_time": self.start_time, # datetime object of when call was started
-
-### PRE-API CALL PARAMS ### (check via kwargs["log_event_type"]="pre_api_call")
-"input" = input # the exact prompt sent to the LLM API
-"api_key" = api_key # the api key used for that LLM API
-"additional_args" = additional_args # any additional details for that API call (e.g. contains optional params sent)
-
-### POST-API CALL PARAMS ### (check via kwargs["log_event_type"]="post_api_call")
-"original_response" = original_response # the original http response received (saved via response.text)
-
-### ON-SUCCESS PARAMS ### (check via kwargs["log_event_type"]="successful_api_call")
-"complete_streaming_response" = complete_streaming_response # the complete streamed response (only set if `completion(..stream=True)`)
-"end_time" = end_time # datetime object of when call was completed
-
-### ON-FAILURE PARAMS ### (check via kwargs["log_event_type"]="failed_api_call")
-"exception" = exception # the Exception raised
-"traceback_exception" = traceback_exception # the traceback generated via `traceback.format_exc()`
-"end_time" = end_time # datetime object of when call was completed
-```
-
-
-### Cache hits
-
-Cache hits are logged in success events as `kwarg["cache_hit"]`.
-
-Here's an example of accessing it:
-
- ```python
- import litellm
-from litellm.integrations.custom_logger import CustomLogger
-from litellm import completion, acompletion, Cache
-
-class MyCustomHandler(CustomLogger):
- async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
- print(f"On Success")
- print(f"Value of Cache hit: {kwargs['cache_hit']"})
-
-async def test_async_completion_azure_caching():
- customHandler_caching = MyCustomHandler()
- litellm.cache = Cache(type="redis", host=os.environ['REDIS_HOST'], port=os.environ['REDIS_PORT'], password=os.environ['REDIS_PASSWORD'])
- litellm.callbacks = [customHandler_caching]
- unique_time = time.time()
- response1 = await litellm.acompletion(model="azure/chatgpt-v-2",
- messages=[{
- "role": "user",
- "content": f"Hi š - i'm async azure {unique_time}"
- }],
- caching=True)
- await asyncio.sleep(1)
- print(f"customHandler_caching.states pre-cache hit: {customHandler_caching.states}")
- response2 = await litellm.acompletion(model="azure/chatgpt-v-2",
- messages=[{
- "role": "user",
- "content": f"Hi š - i'm async azure {unique_time}"
- }],
- caching=True)
- await asyncio.sleep(1) # success callbacks are done in parallel
- print(f"customHandler_caching.states post-cache hit: {customHandler_caching.states}")
- assert len(customHandler_caching.errors) == 0
- assert len(customHandler_caching.states) == 4 # pre, post, success, success
- ```
-
-### Get complete streaming response
-
-LiteLLM will pass you the complete streaming response in the final streaming chunk as part of the kwargs for your custom callback function.
+The kwargs dictionary contains all the details about your API call:
```python
-# litellm.set_verbose = False
- def custom_callback(
- kwargs, # kwargs to completion
- completion_response, # response from completion
- start_time, end_time # start/end time
- ):
- # print(f"streaming response: {completion_response}")
- if "complete_streaming_response" in kwargs:
- print(f"Complete Streaming Response: {kwargs['complete_streaming_response']}")
-
- # Assign the custom callback function
- litellm.success_callback = [custom_callback]
-
- response = completion(model="claude-instant-1", messages=messages, stream=True)
- for idx, chunk in enumerate(response):
- pass
-```
-
-
-### Log additional metadata
-
-LiteLLM accepts a metadata dictionary in the completion call. You can pass additional metadata into your completion call via `completion(..., metadata={"key": "value"})`.
-
-Since this is a [litellm-specific param](https://github.com/BerriAI/litellm/blob/b6a015404eed8a0fa701e98f4581604629300ee3/litellm/main.py#L235), it's accessible via kwargs["litellm_params"]
-
-```python
-from litellm import completion
-import os, litellm
-
-## set ENV variables
-os.environ["OPENAI_API_KEY"] = "your-api-key"
-
-messages = [{ "content": "Hello, how are you?","role": "user"}]
-
-def custom_callback(
- kwargs, # kwargs to completion
- completion_response, # response from completion
- start_time, end_time # start/end time
-):
- print(kwargs["litellm_params"]["metadata"])
+def custom_callback(kwargs, completion_response, start_time, end_time):
+ # Access common data
+ model = kwargs.get("model")
+ messages = kwargs.get("messages", [])
+ cost = kwargs.get("response_cost", 0)
+ cache_hit = kwargs.get("cache_hit", False)
-
-# Assign the custom callback function
-litellm.success_callback = [custom_callback]
-
-response = litellm.completion(model="gpt-3.5-turbo", messages=messages, metadata={"hello": "world"})
+ # Access metadata you passed in
+ metadata = kwargs.get("litellm_params", {}).get("metadata", {})
```
-## Examples
+**Key fields in kwargs:**
+- `model` - The model name
+- `messages` - Input messages
+- `response_cost` - Calculated cost
+- `cache_hit` - Whether response was cached
+- `litellm_params.metadata` - Your custom metadata
-### Custom Callback to track costs for Streaming + Non-Streaming
-By default, the response cost is accessible in the logging object via `kwargs["response_cost"]` on success (sync + async)
+## Practical Examples
+
+### Track API Costs
```python
+def track_cost_callback(kwargs, completion_response, start_time, end_time):
+ cost = kwargs["response_cost"] # litellm calculates this for you
+ print(f"Request cost: ${cost}")
-# Step 1. Write your custom callback function
-def track_cost_callback(
- kwargs, # kwargs to completion
- completion_response, # response from completion
- start_time, end_time # start/end time
-):
- try:
- response_cost = kwargs["response_cost"] # litellm calculates response cost for you
- print("regular response_cost", response_cost)
- except:
- pass
-
-# Step 2. Assign the custom callback function
litellm.success_callback = [track_cost_callback]
-# Step 3. Make litellm.completion call
-response = completion(
- model="gpt-3.5-turbo",
- messages=[
- {
- "role": "user",
- "content": "Hi š - i'm openai"
- }
- ]
-)
-
-print(response)
+response = completion(model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello"}])
```
-### Custom Callback to log transformed Input to LLMs
+### Log Inputs to LLMs
```python
-def get_transformed_inputs(
- kwargs,
-):
+def get_transformed_inputs(kwargs):
params_to_model = kwargs["additional_args"]["complete_input_dict"]
print("params to model", params_to_model)
litellm.input_callback = [get_transformed_inputs]
-def test_chat_openai():
- try:
- response = completion(model="claude-2",
- messages=[{
- "role": "user",
- "content": "Hi š - i'm openai"
- }])
-
- print(response)
-
- except Exception as e:
- print(e)
- pass
+response = completion(model="claude-2", messages=[{"role": "user", "content": "Hello"}])
```
-#### Output
-```shell
-params to model {'model': 'claude-2', 'prompt': "\n\nHuman: Hi š - i'm openai\n\nAssistant: ", 'max_tokens_to_sample': 256}
+### Send to External Service
+```python
+import requests
+
+def send_to_analytics(kwargs, completion_response, start_time, end_time):
+ data = {
+ "model": kwargs.get("model"),
+ "cost": kwargs.get("response_cost", 0),
+ "duration": (end_time - start_time).total_seconds()
+ }
+ requests.post("https://your-analytics.com/api", json=data)
+
+litellm.success_callback = [send_to_analytics]
```
-### Custom Callback to write to Mixpanel
+## Common Issues
+
+### Callback Not Called
+Make sure you:
+1. Register callbacks correctly: `litellm.callbacks = [MyHandler()]`
+2. Use the right hook names (check spelling)
+3. Don't use proxy-only hooks in library mode
+
+### Performance Issues
+- Use async hooks for I/O operations
+- Don't block in callback functions
+- Handle exceptions properly:
```python
-import mixpanel
-import litellm
-from litellm import completion
-
-def custom_callback(
- kwargs, # kwargs to completion
- completion_response, # response from completion
- start_time, end_time # start/end time
-):
- # Your custom code here
- mixpanel.track("LLM Response", {"llm_response": completion_response})
-
-
-# Assign the custom callback function
-litellm.success_callback = [custom_callback]
-
-response = completion(
- model="gpt-3.5-turbo",
- messages=[
- {
- "role": "user",
- "content": "Hi š - i'm openai"
- }
- ]
-)
-
-print(response)
-
+class SafeHandler(CustomLogger):
+ async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
+ try:
+ await external_service(response_obj)
+ except Exception as e:
+ print(f"Callback error: {e}") # Log but don't break the flow
```
-
-
-
-
-
-
-
-
-
-
-
diff --git a/docs/my-website/docs/pass_through/google_ai_studio.md b/docs/my-website/docs/pass_through/google_ai_studio.md
index c3671f58d36..3de7c54aa7a 100644
--- a/docs/my-website/docs/pass_through/google_ai_studio.md
+++ b/docs/my-website/docs/pass_through/google_ai_studio.md
@@ -230,6 +230,13 @@ curl -X POST "https://generativelanguage.googleapis.com/v1beta/models/gemini-1.5
```
+## **Example 4: Video Generation with Veo**
+
+Generate videos using Google's Veo model through LiteLLM pass-through routes.
+
+[**ā Complete Veo Video Generation Guide**](../proxy/veo_video_generation.md)
+
+
## Advanced
Pre-requisites
diff --git a/docs/my-website/docs/pass_through/intro.md b/docs/my-website/docs/pass_through/intro.md
index 3d6286afcc5..38218224f11 100644
--- a/docs/my-website/docs/pass_through/intro.md
+++ b/docs/my-website/docs/pass_through/intro.md
@@ -11,3 +11,43 @@ These endpoints are useful for 2 scenarios:
## How is your request handled?
The request is passed through to the provider's endpoint. The response is then passed back to the client. **No translation is done.**
+
+### Request Forwarding Process
+
+1. **Request Reception**: LiteLLM receives your request at `/provider/endpoint`
+2. **Authentication**: Your LiteLLM API key is validated and mapped to the provider's API key
+3. **Request Transformation**: Request is reformatted for the target provider's API
+4. **Forwarding**: Request is sent to the actual provider endpoint
+5. **Response Handling**: Provider response is returned directly to you
+
+### Authentication Flow
+
+```mermaid
+graph LR
+ A[Client Request] --> B[LiteLLM Proxy]
+ B --> C[Validate LiteLLM API Key]
+ C --> D[Map to Provider API Key]
+ D --> E[Forward to Provider]
+ E --> F[Return Response]
+```
+
+**Key Points:**
+- Use your **LiteLLM API key** in requests, not the provider's key
+- LiteLLM handles the provider authentication internally
+- Same authentication works across all passthrough endpoints
+
+### Error Handling
+
+**Provider Errors**: Forwarded directly to you with original error codes and messages
+
+**LiteLLM Errors**:
+- `401`: Invalid LiteLLM API key
+- `404`: Provider or endpoint not supported
+- `500`: Internal routing/forwarding errors
+
+### Benefits
+
+- **Unified Authentication**: One API key for all providers
+- **Centralized Logging**: All requests logged through LiteLLM
+- **Cost Tracking**: Usage tracked across all endpoints
+- **Access Control**: Same permissions apply to passthrough endpoints
diff --git a/docs/my-website/docs/proxy/call_hooks.md b/docs/my-website/docs/proxy/call_hooks.md
index b4e22027d19..aef33f8c708 100644
--- a/docs/my-website/docs/proxy/call_hooks.md
+++ b/docs/my-website/docs/proxy/call_hooks.md
@@ -6,6 +6,10 @@ import Image from '@theme/IdealImage';
- Reject data before making llm api calls / before returning the response
- Enforce 'user' param for all openai endpoint calls
+:::tip
+**Understanding Callback Hooks?** Check out our [Callback Management Guide](../observability/callback_management.md) to understand the differences between proxy-specific hooks like `async_pre_call_hook` and general logging hooks like `async_log_success_event`.
+:::
+
See a complete example with our [parallel request rate limiter](https://github.com/BerriAI/litellm/blob/main/litellm/proxy/hooks/parallel_request_limiter.py)
## Quick Start
diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md
index 7eb355d39f2..2d4b9655c5d 100644
--- a/docs/my-website/docs/proxy/config_settings.md
+++ b/docs/my-website/docs/proxy/config_settings.md
@@ -335,12 +335,15 @@ router_settings:
| ANTHROPIC_API_KEY | API key for Anthropic service
| ANTHROPIC_API_BASE | Base URL for Anthropic API. Default is https://api.anthropic.com
| AWS_ACCESS_KEY_ID | Access Key ID for AWS services
+| AWS_BATCH_ROLE_ARN | ARN of the AWS IAM role for batch operations
| AWS_DEFAULT_REGION | Default AWS region for service interactions when AWS_REGION is not set
| AWS_PROFILE_NAME | AWS CLI profile name to be used
| AWS_REGION | AWS region for service interactions (takes precedence over AWS_DEFAULT_REGION)
| AWS_REGION_NAME | Default AWS region for service interactions
| AWS_ROLE_ARN | ARN of the AWS IAM role to assume for authentication
| AWS_ROLE_NAME | Role name for AWS IAM usage
+| AWS_S3_BUCKET_NAME | Name of the AWS S3 bucket for file operations
+| AWS_S3_OUTPUT_BUCKET_NAME | Name of the AWS S3 output bucket for batch operations
| AWS_SECRET_ACCESS_KEY | Secret Access Key for AWS services
| AWS_SESSION_NAME | Name for AWS session
| AWS_WEB_IDENTITY_TOKEN | Web identity token for AWS
diff --git a/docs/my-website/docs/proxy/cost_tracking.md b/docs/my-website/docs/proxy/cost_tracking.md
index 19e3344f21b..7f0a13f763f 100644
--- a/docs/my-website/docs/proxy/cost_tracking.md
+++ b/docs/my-website/docs/proxy/cost_tracking.md
@@ -505,11 +505,11 @@ litellm_settings:
### Disable user-agent tracking
-You can disable user-agent tracking by setting `litellm_settings.disable_user_agent_tracking` to `true`.
+You can disable user-agent tracking by setting `litellm_settings.disable_add_user_agent_to_request_tags` to `true`.
```yaml
litellm_settings:
- disable_user_agent_tracking: true
+ disable_add_user_agent_to_request_tags: true
```
## ⨠(Enterprise) Generate Spend Reports
diff --git a/docs/my-website/docs/proxy/load_balancing.md b/docs/my-website/docs/proxy/load_balancing.md
index fd95b57c1ba..67f41d231db 100644
--- a/docs/my-website/docs/proxy/load_balancing.md
+++ b/docs/my-website/docs/proxy/load_balancing.md
@@ -13,6 +13,23 @@ For more details on routing strategies / params, see [Routing](../routing.md)
:::
+## How Load Balancing Works
+
+LiteLLM automatically distributes requests across multiple deployments of the same model using its built-in router. the proxy routes traffic to optimize performance and reliability.
+
+"simple-shuffle" routing strategy is used by default
+
+### Routing Strategies
+
+| Strategy | Description | When to Use |
+|----------|-------------|-------------|
+| **simple-shuffle** (recommended) | Randomly distributes requests | General purpose, good for even load distribution |
+| **least-busy** | Routes to deployment with fewest active requests | High concurrency scenarios |
+| **usage-based-routing** (bad for perf) | Routes to deployment with lowest current usage (RPM/TPM) | When you want to respect rate limits evenly |
+| **latency-based-routing** | Routes to fastest responding deployment | Latency-critical applications |
+| **cost-based-routing** | Routes to deployment with lowest cost | Cost-sensitive applications |
+
+
## Quick Start - Load Balancing
#### Step 1 - Set deployments on config
@@ -106,49 +123,13 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \
]
}'
```
-
-
-
-```python
-from langchain.chat_models import ChatOpenAI
-from langchain.prompts.chat import (
- ChatPromptTemplate,
- HumanMessagePromptTemplate,
- SystemMessagePromptTemplate,
-)
-from langchain.schema import HumanMessage, SystemMessage
-import os
-
-os.environ["OPENAI_API_KEY"] = "anything"
-
-chat = ChatOpenAI(
- openai_api_base="http://0.0.0.0:4000",
- model="gpt-3.5-turbo",
-)
-
-messages = [
- SystemMessage(
- content="You are a helpful assistant that im using to make a test request to."
- ),
- HumanMessage(
- content="test from litellm. tell me why it's amazing in 1 sentence"
- ),
-]
-response = chat(messages)
-
-print(response)
-```
-
-
-
-
### Test - Loadbalancing
In this request, the following will occur:
1. A rate limit exception will be raised
-2. LiteLLM proxy will retry the request on the model group (default is 3).
+2. LiteLLM proxy will retry the request on the model group (default retries are 3).
```bash
curl -X POST 'http://0.0.0.0:4000/chat/completions' \
@@ -256,4 +237,16 @@ model_group_alias: Optional[Dict[str, Union[str, RouterModelGroupAliasItem]]] =
class RouterModelGroupAliasItem(TypedDict):
model: str
hidden: bool # if 'True', don't return on `/v1/models`, `/v1/model/info`, `/v1/model_group/info`
-```
\ No newline at end of file
+```
+
+### When You'll See Load Balancing in Action
+
+**Immediate Effects:**
+
+- Different deployments serve subsequent requests (visible in logs)
+- Better response times during high traffic
+
+**Observable Benefits:**
+- **Higher throughput**: More requests handled simultaneously across deployments
+- **Improved reliability**: If one deployment fails, traffic automatically routes to healthy ones
+- **Better resource utilization**: Load spread evenly across all available deployments
diff --git a/docs/my-website/docs/proxy/prometheus.md b/docs/my-website/docs/proxy/prometheus.md
index dc7030949bd..8bbf737540d 100644
--- a/docs/my-website/docs/proxy/prometheus.md
+++ b/docs/my-website/docs/proxy/prometheus.md
@@ -63,7 +63,7 @@ Use this for for tracking per [user, key, team, etc.](virtual_keys)
| Metric Name | Description |
|----------------------|--------------------------------------|
-| `litellm_spend_metric` | Total Spend, per `"user", "key", "model", "team", "end-user"` |
+| `litellm_spend_metric` | Total Spend, per `"end_user", "hashed_api_key", "api_key_alias", "model", "team", "team_alias", "user"` |
| `litellm_total_tokens_metric` | input + output tokens per `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model"` |
| `litellm_input_tokens_metric` | input tokens per `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model"` |
| `litellm_output_tokens_metric` | output tokens per `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model"` |
@@ -73,9 +73,9 @@ Use this for for tracking per [user, key, team, etc.](virtual_keys)
| Metric Name | Description |
|----------------------|--------------------------------------|
-| `litellm_team_max_budget_metric` | Max Budget for Team Labels: `"team_id", "team_alias"`|
-| `litellm_remaining_team_budget_metric` | Remaining Budget for Team (A team created on LiteLLM) Labels: `"team_id", "team_alias"`|
-| `litellm_team_budget_remaining_hours_metric` | Hours before the team budget is reset Labels: `"team_id", "team_alias"`|
+| `litellm_team_max_budget_metric` | Max Budget for Team Labels: `"team", "team_alias"`|
+| `litellm_remaining_team_budget_metric` | Remaining Budget for Team (A team created on LiteLLM) Labels: `"team", "team_alias"`|
+| `litellm_team_budget_remaining_hours_metric` | Hours before the team budget is reset Labels: `"team", "team_alias"`|
### Virtual Key - Budget
@@ -119,8 +119,8 @@ Use this to track overall LiteLLM Proxy usage.
| Metric Name | Description |
|----------------------|--------------------------------------|
-| `litellm_proxy_failed_requests_metric` | Total number of failed responses from proxy - the client did not get a success response from litellm proxy. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "exception_status", "exception_class"` |
-| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code"` |
+| `litellm_proxy_failed_requests_metric` | Total number of failed responses from proxy - the client did not get a success response from litellm proxy. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "exception_status", "exception_class", "route"` |
+| `litellm_proxy_total_requests_metric` | Total number of requests made to the proxy server - track number of client side requests. Labels: `"end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "status_code", "user_email", "route"` |
## LLM Provider Metrics
@@ -155,7 +155,7 @@ Use this for LLM API Error monitoring and tracking remaining rate limits and tok
| Metric Name | Description |
|----------------------|--------------------------------------|
| `litellm_remaining_requests_metric` | Track `x-ratelimit-remaining-requests` returned from LLM API Deployment. Labels: `"model_group", "api_provider", "api_base", "litellm_model_name", "hashed_api_key", "api_key_alias"` |
-| `litellm_remaining_tokens` | Track `x-ratelimit-remaining-tokens` return from LLM API Deployment. Labels: `"model_group", "api_provider", "api_base", "litellm_model_name", "hashed_api_key", "api_key_alias"` |
+| `litellm_remaining_tokens_metric` | Track `x-ratelimit-remaining-tokens` return from LLM API Deployment. Labels: `"model_group", "api_provider", "api_base", "litellm_model_name", "hashed_api_key", "api_key_alias"` |
### Deployment State
| Metric Name | Description |
@@ -167,16 +167,22 @@ Use this for LLM API Error monitoring and tracking remaining rate limits and tok
| Metric Name | Description |
|----------------------|--------------------------------------|
-| `litellm_deployment_cooled_down` | Number of times a deployment has been cooled down by LiteLLM load balancing logic. Labels: `"litellm_model_name", "model_id", "api_base", "api_provider", "exception_status"` |
+| `litellm_deployment_cooled_down` | Number of times a deployment has been cooled down by LiteLLM load balancing logic. Labels: `"litellm_model_name", "model_id", "api_base", "api_provider"` |
| `litellm_deployment_successful_fallbacks` | Number of successful fallback requests from primary model -> fallback model. Labels: `"requested_model", "fallback_model", "hashed_api_key", "api_key_alias", "team", "team_alias", "exception_status", "exception_class"` |
| `litellm_deployment_failed_fallbacks` | Number of failed fallback requests from primary model -> fallback model. Labels: `"requested_model", "fallback_model", "hashed_api_key", "api_key_alias", "team", "team_alias", "exception_status", "exception_class"` |
+## Request Counting Metrics
+
+| Metric Name | Description |
+|----------------------|--------------------------------------|
+| `litellm_requests_metric` | Total number of requests tracked per endpoint. Labels: `"end_user", "hashed_api_key", "api_key_alias", "model", "team", "team_alias", "user", "user_email"` |
+
## Request Latency Metrics
| Metric Name | Description |
|----------------------|--------------------------------------|
| `litellm_request_total_latency_metric` | Total latency (seconds) for a request to LiteLLM Proxy Server - tracked for labels "end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model" |
-| `litellm_overhead_latency_metric` | Latency overhead (seconds) added by LiteLLM processing - tracked for labels "end_user", "hashed_api_key", "api_key_alias", "requested_model", "team", "team_alias", "user", "model" |
+| `litellm_overhead_latency_metric` | Latency overhead (seconds) added by LiteLLM processing - tracked for labels "model_group", "api_provider", "api_base", "litellm_model_name", "hashed_api_key", "api_key_alias" |
| `litellm_llm_api_latency_metric` | Latency (seconds) for just the LLM API call - tracked for labels "model", "hashed_api_key", "api_key_alias", "team", "team_alias", "requested_model", "end_user", "user" |
| `litellm_llm_api_time_to_first_token_metric` | Time to first token for LLM API call - tracked for labels `model`, `hashed_api_key`, `api_key_alias`, `team`, `team_alias` [Note: only emitted for streaming requests] |
@@ -486,7 +492,6 @@ Here is a screenshot of the metrics you can monitor with the LiteLLM Grafana Das
| Metric Name | Description |
|----------------------|--------------------------------------|
| `litellm_llm_api_failed_requests_metric` | **deprecated** use `litellm_proxy_failed_requests_metric` |
-| `litellm_requests_metric` | **deprecated** use `litellm_proxy_total_requests_metric` |
diff --git a/docs/my-website/docs/proxy/veo_video_generation.md b/docs/my-website/docs/proxy/veo_video_generation.md
new file mode 100644
index 00000000000..14c263bf847
--- /dev/null
+++ b/docs/my-website/docs/proxy/veo_video_generation.md
@@ -0,0 +1,163 @@
+import Tabs from '@theme/Tabs';
+import TabItem from '@theme/TabItem';
+
+# Veo Video Generation with Google AI Studio
+
+Generate videos using Google's Veo model through LiteLLM's pass-through endpoints.
+
+## Quick Start
+
+LiteLLM allows you to use Google AI Studio's Veo video generation API through pass-through routes with zero configuration.
+
+### 1. Add Google AI Studio API Key to your environment
+
+```bash
+export GEMINI_API_KEY="your_google_ai_studio_api_key"
+```
+
+### 2. Start LiteLLM Proxy
+
+```bash
+litellm
+
+# RUNNING on http://0.0.0.0:4000
+```
+
+### 3. Generate Video
+
+
+
+
+```python
+import requests
+import time
+import json
+
+# Configuration
+BASE_URL = "http://localhost:4000/gemini/v1beta"
+API_KEY = "anything" # Use "anything" as the key
+
+headers = {
+ "x-goog-api-key": API_KEY,
+ "Content-Type": "application/json"
+}
+
+# Step 1: Initiate video generation
+def generate_video(prompt):
+ url = f"{BASE_URL}/models/veo-3.0-generate-preview:predictLongRunning"
+ payload = {
+ "instances": [{
+ "prompt": prompt
+ }]
+ }
+
+ response = requests.post(url, headers=headers, json=payload)
+ response.raise_for_status()
+
+ data = response.json()
+ return data.get("name") # Operation name
+
+# Step 2: Poll for completion
+def wait_for_completion(operation_name):
+ operation_url = f"{BASE_URL}/{operation_name}"
+
+ while True:
+ response = requests.get(operation_url, headers=headers)
+ response.raise_for_status()
+
+ data = response.json()
+
+ if data.get("done", False):
+ # Extract video URI
+ video_uri = data["response"]["generateVideoResponse"]["generatedSamples"][0]["video"]["uri"]
+ return video_uri
+
+ time.sleep(10) # Wait 10 seconds before next poll
+
+# Step 3: Download video
+def download_video(video_uri, filename="generated_video.mp4"):
+ # Replace Google URL with LiteLLM proxy URL
+ litellm_url = video_uri.replace(
+ "https://generativelanguage.googleapis.com/v1beta",
+ BASE_URL
+ )
+
+ response = requests.get(litellm_url, headers=headers, stream=True)
+ response.raise_for_status()
+
+ with open(filename, 'wb') as f:
+ for chunk in response.iter_content(chunk_size=8192):
+ if chunk:
+ f.write(chunk)
+
+ return filename
+
+# Complete workflow
+prompt = "A cat playing with a ball of yarn in a sunny garden"
+
+print("Generating video...")
+operation_name = generate_video(prompt)
+
+print("Waiting for completion...")
+video_uri = wait_for_completion(operation_name)
+
+print("Downloading video...")
+filename = download_video(video_uri)
+
+print(f"Video saved as: {filename}")
+```
+
+
+
+
+
+```bash
+# Step 1: Initiate video generation
+curl -X POST "http://localhost:4000/gemini/v1beta/models/veo-3.0-generate-preview:predictLongRunning" \
+ -H "x-goog-api-key: anything" \
+ -H "Content-Type: application/json" \
+ -d '{
+ "instances": [{
+ "prompt": "A cat playing with a ball of yarn in a sunny garden"
+ }]
+ }'
+
+# Response will include operation name:
+# {"name": "operations/generate_12345"}
+
+# Step 2: Poll for completion
+curl -X GET "http://localhost:4000/gemini/v1beta/operations/generate_12345" \
+ -H "x-goog-api-key: anything"
+
+# Step 3: Download video (when done=true)
+curl -X GET "http://localhost:4000/gemini/v1beta/files/VIDEO_ID:download?alt=media" \
+ -H "x-goog-api-key: anything" \
+ --output generated_video.mp4
+```
+
+
+
+
+## Complete Example
+
+For a full working example with error handling and logging, see our [Veo Video Generation Cookbook](https://github.com/BerriAI/litellm/blob/main/cookbook/veo_video_generation.py).
+
+## How It Works
+
+1. **Video Generation Request**: Send a prompt to Veo's `predictLongRunning` endpoint
+2. **Operation Polling**: Monitor the long-running operation until completion
+3. **File Download**: Download the generated video through LiteLLM's pass-through with automatic redirect handling
+
+LiteLLM handles:
+- ā
Authentication with Google AI Studio
+- ā
Request routing and proxying
+- ā
Automatic redirect handling for file downloads
+
+## Configuration Options
+
+### Environment Variables
+
+```bash
+export GEMINI_API_KEY="your_google_ai_studio_api_key"
+```
+
diff --git a/docs/my-website/docs/reasoning_content.md b/docs/my-website/docs/reasoning_content.md
index f9cab01639d..5ddb5aefd47 100644
--- a/docs/my-website/docs/reasoning_content.md
+++ b/docs/my-website/docs/reasoning_content.md
@@ -20,6 +20,7 @@ Supported Providers:
- Vertex AI (`vertex_ai/`)
- Perplexity (`perplexity/`)
- Mistral AI (Magistral models) (`mistral/`)
+- Groq (`groq/`)
LiteLLM will standardize the `reasoning_content` in the response and `thinking_blocks` in the assistant message.
diff --git a/litellm/__init__.py b/litellm/__init__.py
index 6a184d70b5a..d411ff0ad45 100644
--- a/litellm/__init__.py
+++ b/litellm/__init__.py
@@ -450,6 +450,7 @@ vertex_vision_models: Set = set()
vertex_chat_models: Set = set()
vertex_code_chat_models: Set = set()
vertex_ai_image_models: Set = set()
+vertex_ai_video_models: Set = set()
vertex_text_models: Set = set()
vertex_code_text_models: Set = set()
vertex_embedding_models: Set = set()
@@ -605,6 +606,9 @@ def add_known_models():
elif value.get("litellm_provider") == "vertex_ai-image-models":
key = key.replace("vertex_ai/", "")
vertex_ai_image_models.add(key)
+ elif value.get("litellm_provider") == "vertex_ai-video-models":
+ key = key.replace("vertex_ai/", "")
+ vertex_ai_video_models.add(key)
elif value.get("litellm_provider") == "vertex_ai-openai_models":
key = key.replace("vertex_ai/", "")
vertex_openai_models.add(key)
diff --git a/litellm/batches/main.py b/litellm/batches/main.py
index 3ea0f95157f..0d250779da3 100644
--- a/litellm/batches/main.py
+++ b/litellm/batches/main.py
@@ -14,13 +14,15 @@ import asyncio
import contextvars
import os
from functools import partial
-from typing import Any, Coroutine, Dict, Literal, Optional, Union
+from typing import Any, Coroutine, Dict, Literal, Optional, Union, cast
import httpx
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.azure.batches.handler import AzureBatchesAPI
+from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
+from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
from litellm.llms.openai.openai import OpenAIBatchesAPI
from litellm.llms.vertex_ai.batches.handler import VertexAIBatchPrediction
from litellm.secret_managers.main import get_secret_str
@@ -31,13 +33,19 @@ from litellm.types.llms.openai import (
RetrieveBatchRequest,
)
from litellm.types.router import GenericLiteLLMParams
-from litellm.types.utils import LiteLLMBatch
-from litellm.utils import client, get_litellm_params, supports_httpx_timeout
+from litellm.types.utils import LiteLLMBatch, LlmProviders
+from litellm.utils import (
+ ProviderConfigManager,
+ client,
+ get_litellm_params,
+ supports_httpx_timeout,
+)
####### ENVIRONMENT VARIABLES ###################
openai_batches_instance = OpenAIBatchesAPI()
azure_batches_instance = AzureBatchesAPI()
vertex_ai_batches_instance = VertexAIBatchPrediction(gcs_bucket_name="")
+base_llm_http_handler = BaseLLMHTTPHandler()
#################################################
@@ -46,7 +54,7 @@ async def acreate_batch(
completion_window: Literal["24h"],
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions"],
input_file_id: str,
- custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
+ custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock"] = "openai",
metadata: Optional[Dict[str, str]] = None,
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
@@ -94,7 +102,7 @@ def create_batch(
completion_window: Literal["24h"],
endpoint: Literal["/v1/chat/completions", "/v1/embeddings", "/v1/completions"],
input_file_id: str,
- custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
+ custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock"] = "openai",
metadata: Optional[Dict[str, str]] = None,
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
@@ -111,8 +119,8 @@ def create_batch(
proxy_server_request = kwargs.get("proxy_server_request", None)
model_info = kwargs.get("model_info", None)
_is_async = kwargs.pop("acreate_batch", False) is True
- litellm_params = get_litellm_params(**kwargs)
- litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj", None)
+ litellm_params = dict(GenericLiteLLMParams(**kwargs))
+ litellm_logging_obj: LiteLLMLoggingObj = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj", None))
### TIMEOUT LOGIC ###
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
litellm_logging_obj.update_environment_variables(
@@ -142,6 +150,7 @@ def create_batch(
timeout = float(timeout) # type: ignore
elif timeout is None:
timeout = 600.0
+
_create_batch_request = CreateBatchRequest(
completion_window=completion_window,
@@ -151,6 +160,27 @@ def create_batch(
extra_headers=extra_headers,
extra_body=extra_body,
)
+ provider_config = ProviderConfigManager.get_provider_batches_config(
+ model="",
+ provider=LlmProviders(custom_llm_provider),
+ )
+ if provider_config is not None:
+ response = base_llm_http_handler.create_batch(
+ provider_config=provider_config,
+ litellm_params=litellm_params,
+ create_batch_data=_create_batch_request,
+ headers=extra_headers or {},
+ api_base=optional_params.api_base,
+ api_key=optional_params.api_key,
+ logging_obj=litellm_logging_obj,
+ _is_async=_is_async,
+ client=client
+ if client is not None
+ and isinstance(client, (HTTPHandler, AsyncHTTPHandler))
+ else None,
+ timeout=timeout,
+ )
+ return response
api_base: Optional[str] = None
if custom_llm_provider == "openai":
# for deepinfra/perplexity/anyscale/groq we check in get_llm_provider and pass in the api base from there
@@ -322,20 +352,21 @@ def retrieve_batch(
"""
try:
optional_params = GenericLiteLLMParams(**kwargs)
- litellm_logging_obj: LiteLLMLoggingObj = kwargs.get("litellm_logging_obj", None)
+ litellm_logging_obj: Optional[LiteLLMLoggingObj] = kwargs.get("litellm_logging_obj", None)
### TIMEOUT LOGIC ###
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
litellm_params = get_litellm_params(
custom_llm_provider=custom_llm_provider,
**kwargs,
)
- litellm_logging_obj.update_environment_variables(
- model=None,
- user=None,
- optional_params=optional_params.model_dump(),
- litellm_params=litellm_params,
- custom_llm_provider=custom_llm_provider,
- )
+ if litellm_logging_obj is not None:
+ litellm_logging_obj.update_environment_variables(
+ model=None,
+ user=None,
+ optional_params=optional_params.model_dump(),
+ litellm_params=litellm_params,
+ custom_llm_provider=custom_llm_provider,
+ )
if (
timeout is not None
diff --git a/litellm/files/main.py b/litellm/files/main.py
index 5d0dc05771a..299e52895bf 100644
--- a/litellm/files/main.py
+++ b/litellm/files/main.py
@@ -50,7 +50,7 @@ vertex_ai_files_instance = VertexAIFilesHandler()
async def acreate_file(
file: FileTypes,
purpose: Literal["assistants", "batch", "fine-tune"],
- custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
+ custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock"] = "openai",
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
**kwargs,
@@ -94,7 +94,7 @@ async def acreate_file(
def create_file(
file: FileTypes,
purpose: Literal["assistants", "batch", "fine-tune"],
- custom_llm_provider: Optional[Literal["openai", "azure", "vertex_ai"]] = None,
+ custom_llm_provider: Optional[Literal["openai", "azure", "vertex_ai", "bedrock"]] = None,
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
**kwargs,
@@ -109,7 +109,7 @@ def create_file(
try:
_is_async = kwargs.pop("acreate_file", False) is True
optional_params = GenericLiteLLMParams(**kwargs)
- litellm_params_dict = get_litellm_params(**kwargs)
+ litellm_params_dict = dict(**kwargs)
logging_obj = cast(
Optional[LiteLLMLoggingObj], kwargs.get("litellm_logging_obj")
)
diff --git a/litellm/litellm_core_utils/get_llm_provider_logic.py b/litellm/litellm_core_utils/get_llm_provider_logic.py
index 2049480e264..c784568cc9f 100644
--- a/litellm/litellm_core_utils/get_llm_provider_logic.py
+++ b/litellm/litellm_core_utils/get_llm_provider_logic.py
@@ -320,6 +320,7 @@ def get_llm_provider( # noqa: PLR0915
or model in litellm.vertex_embedding_models
or model in litellm.vertex_vision_models
or model in litellm.vertex_ai_image_models
+ or model in litellm.vertex_ai_video_models
):
custom_llm_provider = "vertex_ai"
## ai21
diff --git a/litellm/llms/base_llm/__init__.py b/litellm/llms/base_llm/__init__.py
index 187c985fd67..665e242969c 100644
--- a/litellm/llms/base_llm/__init__.py
+++ b/litellm/llms/base_llm/__init__.py
@@ -1,5 +1,6 @@
from .anthropic_messages.transformation import BaseAnthropicMessagesConfig
from .audio_transcription.transformation import BaseAudioTranscriptionConfig
+from .batches.transformation import BaseBatchesConfig
from .chat.transformation import BaseConfig
from .embedding.transformation import BaseEmbeddingConfig
from .image_edit.transformation import BaseImageEditConfig
@@ -12,4 +13,5 @@ __all__ = [
"BaseAnthropicMessagesConfig",
"BaseEmbeddingConfig",
"BaseImageEditConfig",
+ "BaseBatchesConfig",
]
diff --git a/litellm/llms/base_llm/batches/transformation.py b/litellm/llms/base_llm/batches/transformation.py
new file mode 100644
index 00000000000..1d3e54fae67
--- /dev/null
+++ b/litellm/llms/base_llm/batches/transformation.py
@@ -0,0 +1,176 @@
+import types
+from abc import ABC, abstractmethod
+from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union
+
+import httpx
+from httpx import Headers
+
+from litellm.types.llms.openai import (
+ AllMessageValues,
+ CreateBatchRequest,
+)
+from litellm.types.utils import LiteLLMBatch, LlmProviders
+
+if TYPE_CHECKING:
+ from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
+
+ from ..chat.transformation import BaseLLMException as _BaseLLMException
+
+ LiteLLMLoggingObj = _LiteLLMLoggingObj
+ BaseLLMException = _BaseLLMException
+else:
+ LiteLLMLoggingObj = Any
+ BaseLLMException = Any
+
+
+class BaseBatchesConfig(ABC):
+ """
+ Abstract base class for batch processing configurations across different LLM providers.
+
+ This class defines the interface that all provider-specific batch configurations
+ must implement to work with LiteLLM's unified batch processing system.
+ """
+
+ def __init__(self):
+ pass
+
+ @property
+ @abstractmethod
+ def custom_llm_provider(self) -> LlmProviders:
+ """Return the LLM provider type for this configuration."""
+ pass
+
+ @classmethod
+ def get_config(cls):
+ """Get configuration dictionary for this class."""
+ return {
+ k: v
+ for k, v in cls.__dict__.items()
+ if not k.startswith("__")
+ and not k.startswith("_abc")
+ and not isinstance(
+ v,
+ (
+ types.FunctionType,
+ types.BuiltinFunctionType,
+ classmethod,
+ staticmethod,
+ ),
+ )
+ and v is not None
+ }
+
+ @abstractmethod
+ def validate_environment(
+ self,
+ headers: dict,
+ model: str,
+ messages: List[AllMessageValues],
+ optional_params: dict,
+ litellm_params: dict,
+ api_key: Optional[str] = None,
+ api_base: Optional[str] = None,
+ ) -> dict:
+ """
+ Validate and prepare environment-specific headers and parameters.
+
+ Args:
+ headers: HTTP headers dictionary
+ model: Model name
+ messages: List of messages
+ optional_params: Optional parameters
+ litellm_params: LiteLLM parameters
+ api_key: API key
+ api_base: API base URL
+
+ Returns:
+ Updated headers dictionary
+ """
+ pass
+
+ @abstractmethod
+ def get_complete_batch_url(
+ self,
+ api_base: Optional[str],
+ api_key: Optional[str],
+ model: str,
+ optional_params: Dict,
+ litellm_params: Dict,
+ data: CreateBatchRequest,
+ ) -> str:
+ """
+ Get the complete URL for batch creation request.
+
+ Args:
+ api_base: Base API URL
+ api_key: API key
+ model: Model name
+ optional_params: Optional parameters
+ litellm_params: LiteLLM parameters
+ data: Batch creation request data
+
+ Returns:
+ Complete URL for the batch request
+ """
+ pass
+
+ @abstractmethod
+ def transform_create_batch_request(
+ self,
+ model: str,
+ create_batch_data: CreateBatchRequest,
+ optional_params: dict,
+ litellm_params: dict,
+ ) -> Union[bytes, str, Dict[str, Any]]:
+ """
+ Transform the batch creation request to provider-specific format.
+
+ Args:
+ model: Model name
+ create_batch_data: Batch creation request data
+ optional_params: Optional parameters
+ litellm_params: LiteLLM parameters
+
+ Returns:
+ Transformed request data
+ """
+ pass
+
+ @abstractmethod
+ def transform_create_batch_response(
+ self,
+ model: Optional[str],
+ raw_response: httpx.Response,
+ logging_obj: LiteLLMLoggingObj,
+ litellm_params: dict,
+ ) -> LiteLLMBatch:
+ """
+ Transform provider-specific batch response to LiteLLM format.
+
+ Args:
+ model: Model name
+ raw_response: Raw HTTP response
+ logging_obj: Logging object
+ litellm_params: LiteLLM parameters
+
+ Returns:
+ LiteLLM batch object
+ """
+ pass
+
+ @abstractmethod
+ def get_error_class(
+ self, error_message: str, status_code: int, headers: Union[Dict, Headers]
+ ) -> "BaseLLMException":
+ """
+ Get the appropriate error class for this provider.
+
+ Args:
+ error_message: Error message
+ status_code: HTTP status code
+ headers: Response headers
+
+ Returns:
+ Provider-specific exception class
+ """
+ pass
diff --git a/litellm/llms/base_llm/files/transformation.py b/litellm/llms/base_llm/files/transformation.py
index 5c37a8b7547..35b76479cdc 100644
--- a/litellm/llms/base_llm/files/transformation.py
+++ b/litellm/llms/base_llm/files/transformation.py
@@ -35,6 +35,16 @@ class BaseFilesConfig(BaseConfig):
def custom_llm_provider(self) -> LlmProviders:
pass
+ @property
+ def file_upload_http_method(self) -> str:
+ """
+ HTTP method to use for file uploads.
+ Override this in provider configs if they need different methods.
+ Default is POST (used by most providers like OpenAI, Anthropic).
+ S3-based providers like Bedrock should return "PUT".
+ """
+ return "POST"
+
@abstractmethod
def get_supported_openai_params(
self, model: str
diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py
new file mode 100644
index 00000000000..ce580ebc624
--- /dev/null
+++ b/litellm/llms/bedrock/batches/transformation.py
@@ -0,0 +1,254 @@
+import os
+import time
+from typing import Any, Dict, List, Literal, Optional, Union, cast
+
+from httpx import Headers, Response
+
+from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig
+from litellm.llms.base_llm.chat.transformation import BaseLLMException
+from litellm.types.llms.bedrock import (
+ BedrockBatchJobStatus,
+ BedrockCreateBatchRequest,
+ BedrockCreateBatchResponse,
+ BedrockInputDataConfig,
+ BedrockOutputDataConfig,
+ BedrockS3InputDataConfig,
+ BedrockS3OutputDataConfig,
+)
+from litellm.types.llms.openai import (
+ AllMessageValues,
+ CreateBatchRequest,
+)
+from litellm.types.utils import LiteLLMBatch, LlmProviders
+
+from ..base_aws_llm import BaseAWSLLM
+from ..common_utils import CommonBatchFilesUtils
+
+
+class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig):
+ """
+ Config for Bedrock Batches - handles batch job creation and management for Bedrock
+ """
+
+ def __init__(self):
+ super().__init__()
+ self.common_utils = CommonBatchFilesUtils()
+
+ @property
+ def custom_llm_provider(self) -> LlmProviders:
+ return LlmProviders.BEDROCK
+
+ def validate_environment(
+ self,
+ headers: dict,
+ model: str,
+ messages: List[AllMessageValues],
+ optional_params: dict,
+ litellm_params: dict,
+ api_key: Optional[str] = None,
+ api_base: Optional[str] = None,
+ ) -> dict:
+ """
+ Validate and prepare environment for Bedrock batch requests.
+ AWS credentials are handled by BaseAWSLLM.
+ """
+ # Add any Bedrock-specific headers if needed
+ return headers
+
+ def get_complete_batch_url(
+ self,
+ api_base: Optional[str],
+ api_key: Optional[str],
+ model: str,
+ optional_params: Dict,
+ litellm_params: Dict,
+ data: CreateBatchRequest,
+ ) -> str:
+ """
+ Get the complete URL for Bedrock batch creation.
+ Bedrock batch jobs are created via the model invocation job API.
+ """
+ aws_region_name = self._get_aws_region_name(optional_params, model)
+
+ # Bedrock model invocation job endpoint
+ # Format: https://bedrock.{region}.amazonaws.com/model-invocation-job
+ bedrock_endpoint = f"https://bedrock.{aws_region_name}.amazonaws.com/model-invocation-job"
+
+ return bedrock_endpoint
+
+
+
+
+
+
+
+ def transform_create_batch_request(
+ self,
+ model: str,
+ create_batch_data: CreateBatchRequest,
+ optional_params: dict,
+ litellm_params: dict,
+ ) -> Dict[str, Any]:
+ """
+ Transform the batch creation request to Bedrock format.
+
+ Bedrock batch inference requires:
+ - modelId: The Bedrock model ID
+ - jobName: Unique name for the batch job
+ - inputDataConfig: Configuration for input data (S3 location)
+ - outputDataConfig: Configuration for output data (S3 location)
+ - roleArn: IAM role ARN for the batch job
+ """
+ # Get required parameters
+ input_file_id = create_batch_data.get("input_file_id")
+ if not input_file_id:
+ raise ValueError("input_file_id is required for Bedrock batch creation")
+
+ # Extract S3 information from file ID using common utility
+ input_bucket, input_key = self.common_utils.parse_s3_uri(input_file_id)
+
+ # Get output S3 configuration
+ output_bucket = litellm_params.get("s3_output_bucket_name") or os.getenv("AWS_S3_OUTPUT_BUCKET_NAME")
+ if not output_bucket:
+ # Use same bucket as input if no output bucket specified
+ output_bucket = input_bucket
+
+ # Get IAM role ARN
+ role_arn = (
+ litellm_params.get("aws_batch_role_arn")
+ or optional_params.get("aws_batch_role_arn")
+ or os.getenv("AWS_BATCH_ROLE_ARN")
+ )
+ if not role_arn:
+ raise ValueError(
+ "AWS IAM role ARN is required for Bedrock batch jobs. "
+ "Set 'aws_batch_role_arn' in litellm_params or AWS_BATCH_ROLE_ARN env var"
+ )
+
+ # Get the actual Bedrock model ID using common utility
+ bedrock_model_id = self.common_utils.extract_model_from_s3_file_path(input_file_id, optional_params)
+
+ if not bedrock_model_id:
+ raise ValueError("Could not determine Bedrock model ID. Ensure the model is specified in the input file or passed as a parameter.")
+
+ # Generate job name with the correct model ID using common utility
+ job_name = self.common_utils.generate_unique_job_name(bedrock_model_id, prefix="litellm")
+ output_key = f"litellm-batch-outputs/{job_name}/"
+
+ # Build input data config
+ input_data_config: BedrockInputDataConfig = {
+ "s3InputDataConfig": BedrockS3InputDataConfig(
+ s3Uri=f"s3://{input_bucket}/{input_key}"
+ )
+ }
+
+ # Build output data config
+ output_data_config: BedrockOutputDataConfig = {
+ "s3OutputDataConfig": BedrockS3OutputDataConfig(
+ s3Uri=f"s3://{output_bucket}/{output_key}"
+ )
+ }
+
+ # Create Bedrock batch request with proper typing
+ bedrock_request: BedrockCreateBatchRequest = {
+ "modelId": bedrock_model_id,
+ "jobName": job_name,
+ "inputDataConfig": input_data_config,
+ "outputDataConfig": output_data_config,
+ "roleArn": role_arn
+ }
+
+ # Add optional parameters if provided
+ completion_window = create_batch_data.get("completion_window")
+ if completion_window:
+ # Map OpenAI completion window to Bedrock timeout
+ # OpenAI uses "24h", Bedrock expects timeout in hours
+ if completion_window == "24h":
+ bedrock_request["timeoutDurationInHours"] = 24
+
+ # For Bedrock, we need to return a pre-signed request with AWS auth headers
+ # Use common utility for AWS signing
+ endpoint_url = f"https://bedrock.{self._get_aws_region_name(optional_params, model)}.amazonaws.com/model-invocation-job"
+ signed_headers, signed_data = self.common_utils.sign_aws_request(
+ service_name="bedrock",
+ data=bedrock_request,
+ endpoint_url=endpoint_url,
+ optional_params=optional_params,
+ method="POST"
+ )
+
+ # Return a pre-signed request format that the HTTP handler can use
+ return {
+ "method": "POST",
+ "url": endpoint_url,
+ "headers": signed_headers,
+ "data": signed_data.decode('utf-8')
+ }
+
+ def transform_create_batch_response(
+ self,
+ model: Optional[str],
+ raw_response: Response,
+ logging_obj: Any,
+ litellm_params: dict,
+ ) -> LiteLLMBatch:
+ """
+ Transform Bedrock batch creation response to LiteLLM format.
+ """
+ try:
+ response_data: BedrockCreateBatchResponse = raw_response.json()
+ except Exception as e:
+ raise ValueError(f"Failed to parse Bedrock batch response: {e}")
+
+ # Extract information from typed Bedrock response
+ job_arn = response_data.get("jobArn", "")
+ status: BedrockBatchJobStatus = response_data.get("status", "Submitted")
+
+ # Map Bedrock status to OpenAI-compatible status
+ status_mapping: Dict[BedrockBatchJobStatus, str] = {
+ "Submitted": "validating",
+ "InProgress": "in_progress",
+ "Completed": "completed",
+ "Failed": "failed",
+ "Stopping": "cancelling",
+ "Stopped": "cancelled"
+ }
+
+ openai_status = cast(Literal["validating", "failed", "in_progress", "finalizing", "completed", "expired", "cancelling", "cancelled"], status_mapping.get(status, "validating"))
+
+ # Get original request data from litellm_params if available
+ original_request = litellm_params.get("original_batch_request", {})
+
+ # Create LiteLLM batch object
+ return LiteLLMBatch(
+ id=job_arn, # Use ARN as the batch ID
+ object="batch",
+ endpoint=original_request.get("endpoint", "/v1/chat/completions"),
+ errors=None,
+ input_file_id=original_request.get("input_file_id", ""),
+ completion_window=original_request.get("completion_window", "24h"),
+ status=openai_status,
+ output_file_id=None, # Will be populated when job completes
+ error_file_id=None,
+ created_at=int(time.time()),
+ in_progress_at=int(time.time()) if status == "InProgress" else None,
+ expires_at=None,
+ finalizing_at=None,
+ completed_at=None,
+ failed_at=None,
+ expired_at=None,
+ cancelling_at=None,
+ cancelled_at=None,
+ request_counts=None,
+ metadata=original_request.get("metadata", {}),
+ )
+
+ def get_error_class(
+ self, error_message: str, status_code: int, headers: Union[Dict, Headers]
+ ) -> BaseLLMException:
+ """
+ Get Bedrock-specific error class using common utility.
+ """
+ return self.common_utils.get_error_class(error_message, status_code, headers)
+
+
diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py
index c76fc0a80c3..831a6da93b3 100644
--- a/litellm/llms/bedrock/common_utils.py
+++ b/litellm/llms/bedrock/common_utils.py
@@ -6,6 +6,9 @@ import json
import os
from typing import TYPE_CHECKING, Dict, List, Literal, Optional, Union
+if TYPE_CHECKING:
+ from litellm.types.llms.bedrock import BedrockCreateBatchRequest
+
import httpx
import litellm
@@ -608,3 +611,218 @@ def get_anthropic_beta_from_headers(headers: dict) -> List[str]:
# Split comma-separated values and strip whitespace
return [beta.strip() for beta in anthropic_beta_header.split(",")]
+
+
+class CommonBatchFilesUtils:
+ """
+ Common utilities for Bedrock batch and file operations.
+ Provides shared functionality to reduce code duplication between batches and files.
+ """
+
+ def __init__(self):
+ # Import here to avoid circular imports
+ from .base_aws_llm import BaseAWSLLM
+ self._base_aws = BaseAWSLLM()
+
+ def get_bedrock_model_id_from_litellm_model(self, model: str) -> str:
+ """
+ Extract the actual Bedrock model ID from LiteLLM model name.
+
+ Args:
+ model: LiteLLM model name (e.g., "bedrock/anthropic.claude-3-sonnet-20240229-v1:0")
+
+ Returns:
+ Bedrock model ID (e.g., "anthropic.claude-3-sonnet-20240229-v1:0")
+ """
+ if model.startswith("bedrock/"):
+ return model[8:] # Remove "bedrock/" prefix
+ return model
+
+ def parse_s3_uri(self, s3_uri: str) -> tuple:
+ """
+ Parse S3 URI into bucket and key components.
+
+ Args:
+ s3_uri: S3 URI (e.g., "s3://bucket/key/path")
+
+ Returns:
+ Tuple of (bucket, key)
+
+ Raises:
+ ValueError: If URI format is invalid
+ """
+ if not s3_uri.startswith("s3://"):
+ raise ValueError(f"Invalid S3 URI format: {s3_uri}")
+
+ s3_parts = s3_uri[5:].split("/", 1) # Remove "s3://" and split on first "/"
+ if len(s3_parts) != 2:
+ raise ValueError(f"Invalid S3 URI format: {s3_uri}")
+
+ return s3_parts[0], s3_parts[1] # bucket, key
+
+ def extract_model_from_s3_file_path(self, s3_uri: str, optional_params: dict) -> str:
+ """
+ Extract model ID from S3 file path.
+
+ The Bedrock file transformation creates S3 objects with the model name embedded:
+ Format: s3://bucket/litellm-bedrock-files-{model}-{uuid}.jsonl
+ """
+ # Check if model is provided in optional_params first
+ if "model" in optional_params and optional_params["model"]:
+ return self.get_bedrock_model_id_from_litellm_model(optional_params["model"])
+
+ # Extract model from S3 URI path
+ # Expected format: s3://bucket/litellm-bedrock-files-{model}-{uuid}.jsonl
+ try:
+ bucket, object_key = self.parse_s3_uri(s3_uri)
+
+ # Extract model from object key if it follows our naming pattern
+ if object_key.startswith("litellm-bedrock-files-"):
+ # Remove prefix and suffix to get model part
+ model_part = object_key[22:] # Remove "litellm-bedrock-files-"
+ # Find the last dash before the UUID
+ parts = model_part.split("-")
+ if len(parts) > 1:
+ # Reconstruct model name (everything except the last UUID part and .jsonl)
+ model_name = "-".join(parts[:-1])
+ if model_name.endswith(".jsonl"):
+ model_name = model_name[:-6] # Remove .jsonl
+ return model_name
+ except Exception:
+ pass
+
+ # Fallback to default model
+ return "anthropic.claude-3-5-sonnet-20240620-v1:0"
+
+ def sign_aws_request(
+ self,
+ service_name: str,
+ data: Union[str, dict, "BedrockCreateBatchRequest"],
+ endpoint_url: str,
+ optional_params: dict,
+ method: str = "POST",
+ ) -> tuple:
+ """
+ Sign AWS request using Signature Version 4.
+
+ Args:
+ service_name: AWS service name ("bedrock" or "s3")
+ data: Request data (string or dict)
+ endpoint_url: Full endpoint URL
+ optional_params: Optional parameters containing AWS credentials
+ method: HTTP method (default: POST)
+
+ Returns:
+ Tuple of (signed_headers, signed_data)
+ """
+ try:
+ from botocore.auth import SigV4Auth
+ from botocore.awsrequest import AWSRequest
+ except ImportError:
+ raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
+
+ # Get AWS credentials using existing methods
+ aws_region_name = self._base_aws._get_aws_region_name(
+ optional_params=optional_params, model=""
+ )
+ credentials = self._base_aws.get_credentials(
+ aws_access_key_id=optional_params.get("aws_access_key_id"),
+ aws_secret_access_key=optional_params.get("aws_secret_access_key"),
+ aws_session_token=optional_params.get("aws_session_token"),
+ aws_region_name=aws_region_name,
+ aws_session_name=optional_params.get("aws_session_name"),
+ aws_profile_name=optional_params.get("aws_profile_name"),
+ aws_role_name=optional_params.get("aws_role_name"),
+ aws_web_identity_token=optional_params.get("aws_web_identity_token"),
+ aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
+ )
+
+ # Prepare the request data
+ if isinstance(data, dict):
+ import json
+ request_data = json.dumps(data)
+ else:
+ request_data = data
+
+ # Prepare headers
+ headers = {"Content-Type": "application/json"}
+
+ # Create AWS request and sign it
+ sigv4 = SigV4Auth(credentials, service_name, aws_region_name)
+ request = AWSRequest(
+ method=method.upper(), url=endpoint_url, data=request_data, headers=headers
+ )
+ sigv4.add_auth(request)
+ prepped = request.prepare()
+
+ return dict(prepped.headers), request_data.encode('utf-8') if isinstance(request_data, str) else request_data
+
+ def generate_unique_job_name(self, model: str, prefix: str = "litellm") -> str:
+ """
+ Generate a unique job name for AWS services.
+ AWS services often have length limits, so this creates a concise name.
+
+ Args:
+ model: Model name to include in the job name
+ prefix: Prefix for the job name
+
+ Returns:
+ Unique job name (⤠63 characters for Bedrock compatibility)
+ """
+ import fastuuid as uuid
+ unique_id = str(uuid.uuid4())[:8]
+ # Format: {prefix}-batch-{model}-{uuid}
+ # Example: litellm-batch-claude-266c398e
+ job_name = f"{prefix}-batch-{unique_id}"
+
+ return job_name
+
+ def get_s3_bucket_and_key_from_config(
+ self,
+ litellm_params: dict,
+ optional_params: dict,
+ bucket_env_var: str = "AWS_S3_BUCKET_NAME",
+ key_prefix: str = "litellm"
+ ) -> tuple:
+ """
+ Get S3 bucket and generate a unique key from configuration.
+
+ Args:
+ litellm_params: LiteLLM parameters
+ optional_params: Optional parameters
+ bucket_env_var: Environment variable name for bucket
+ key_prefix: Prefix for the S3 key
+
+ Returns:
+ Tuple of (bucket_name, object_key)
+ """
+ import time
+ import uuid
+
+ # Get bucket name
+ bucket_name = (
+ litellm_params.get("s3_bucket_name")
+ or optional_params.get("s3_bucket_name")
+ or os.getenv(bucket_env_var)
+ )
+ if not bucket_name:
+ raise ValueError(f"S3 bucket name is required. Set 's3_bucket_name' parameter or {bucket_env_var} env var")
+
+ # Generate unique object key
+ timestamp = int(time.time())
+ unique_id = str(uuid.uuid4())[:8]
+ object_key = f"{key_prefix}-{timestamp}-{unique_id}"
+
+ return bucket_name, object_key
+
+ def get_error_class(
+ self, error_message: str, status_code: int, headers: Union[Dict, httpx.Headers]
+ ) -> BaseLLMException:
+ """
+ Get Bedrock-specific error class.
+ """
+ return BedrockError(
+ status_code=status_code,
+ message=error_message,
+ headers=headers
+ )
diff --git a/litellm/llms/bedrock/files/transformation.py b/litellm/llms/bedrock/files/transformation.py
new file mode 100644
index 00000000000..83bbad7e1e8
--- /dev/null
+++ b/litellm/llms/bedrock/files/transformation.py
@@ -0,0 +1,607 @@
+import json
+import os
+import time
+import uuid
+from typing import Any, Dict, List, Optional, Tuple, Union
+
+from httpx import Headers, Response
+
+from litellm.litellm_core_utils.prompt_templates.common_utils import extract_file_data
+from litellm.llms.base_llm.chat.transformation import BaseLLMException
+from litellm.llms.base_llm.files.transformation import (
+ BaseFilesConfig,
+ LiteLLMLoggingObj,
+)
+from litellm.types.llms.openai import (
+ AllMessageValues,
+ CreateFileRequest,
+ FileTypes,
+ OpenAICreateFileRequestOptionalParams,
+ OpenAIFileObject,
+ PathLike,
+)
+from litellm.types.utils import ExtractedFileData, LlmProviders
+
+from ..base_aws_llm import BaseAWSLLM
+from ..common_utils import BedrockError
+
+
+class BedrockFilesConfig(BaseAWSLLM, BaseFilesConfig):
+ """
+ Config for Bedrock Files - handles S3 uploads for Bedrock batch processing
+ """
+
+ def __init__(self):
+ self.jsonl_transformation = BedrockJsonlFilesTransformation()
+ super().__init__()
+
+ @property
+ def custom_llm_provider(self) -> LlmProviders:
+ return LlmProviders.BEDROCK
+
+ @property
+ def file_upload_http_method(self) -> str:
+ """
+ Bedrock files are uploaded to S3, which requires PUT requests
+ """
+ return "PUT"
+
+ def validate_environment(
+ self,
+ headers: dict,
+ model: str,
+ messages: List[AllMessageValues],
+ optional_params: dict,
+ litellm_params: dict,
+ api_key: Optional[str] = None,
+ api_base: Optional[str] = None,
+ ) -> dict:
+ # No additional headers needed for S3 uploads - AWS credentials handled by BaseAWSLLM
+ return headers
+
+
+
+ def _get_content_from_openai_file(self, openai_file_content: FileTypes) -> str:
+ """
+ Helper to extract content from various OpenAI file types and return as string.
+
+ Handles:
+ - Direct content (str, bytes, IO[bytes])
+ - Tuple formats: (filename, content, [content_type], [headers])
+ - PathLike objects
+ """
+ content: Union[str, bytes] = b""
+ # Extract file content from tuple if necessary
+ if isinstance(openai_file_content, tuple):
+ # Take the second element which is always the file content
+ file_content = openai_file_content[1]
+ else:
+ file_content = openai_file_content
+
+ # Handle different file content types
+ if isinstance(file_content, str):
+ # String content can be used directly
+ content = file_content
+ elif isinstance(file_content, bytes):
+ # Bytes content can be decoded
+ content = file_content
+ elif isinstance(file_content, PathLike): # PathLike
+ with open(str(file_content), "rb") as f:
+ content = f.read()
+ elif hasattr(file_content, "read"): # IO[bytes]
+ # File-like objects need to be read
+ content = file_content.read()
+
+ # Ensure content is string
+ if isinstance(content, bytes):
+ content = content.decode("utf-8")
+
+ return content
+
+ def _get_s3_object_name_from_batch_jsonl(
+ self,
+ openai_jsonl_content: List[Dict[str, Any]],
+ ) -> str:
+ """
+ Gets a unique S3 object name for the Bedrock batch processing job
+
+ named as: litellm-bedrock-files/{model}/{uuid}
+ """
+ _model = openai_jsonl_content[0].get("body", {}).get("model", "")
+ # Remove bedrock/ prefix if present
+ if _model.startswith("bedrock/"):
+ _model = _model[8:]
+ object_name = f"litellm-bedrock-files-{_model}-{uuid.uuid4()}.jsonl"
+ return object_name
+
+ def get_object_name(
+ self, extracted_file_data: ExtractedFileData, purpose: str
+ ) -> str:
+ """
+ Get the object name for the request
+ """
+ extracted_file_data_content = extracted_file_data.get("content")
+
+ if extracted_file_data_content is None:
+ raise ValueError("file content is required")
+
+ if purpose == "batch":
+ ## 1. If jsonl, check if there's a model name
+ file_content = self._get_content_from_openai_file(
+ extracted_file_data_content
+ )
+
+ # Split into lines and parse each line as JSON
+ openai_jsonl_content = [
+ json.loads(line) for line in file_content.splitlines() if line.strip()
+ ]
+ if len(openai_jsonl_content) > 0:
+ return self._get_s3_object_name_from_batch_jsonl(openai_jsonl_content)
+
+ ## 2. If not jsonl, return the filename
+ filename = extracted_file_data.get("filename")
+ if filename:
+ return filename
+ ## 3. If no file name, return timestamp
+ return str(int(time.time()))
+
+ def get_complete_file_url(
+ self,
+ api_base: Optional[str],
+ api_key: Optional[str],
+ model: str,
+ optional_params: Dict,
+ litellm_params: Dict,
+ data: CreateFileRequest,
+ ) -> str:
+ """
+ Get the complete S3 URL for the file upload request
+ """
+ bucket_name = litellm_params.get("s3_bucket_name") or os.getenv("AWS_S3_BUCKET_NAME")
+ if not bucket_name:
+ raise ValueError("S3 bucket_name is required. Set 's3_bucket_name' in litellm_params or AWS_S3_BUCKET_NAME env var")
+
+ aws_region_name = self._get_aws_region_name(optional_params, model)
+
+ file_data = data.get("file")
+ purpose = data.get("purpose")
+ if file_data is None:
+ raise ValueError("file is required")
+ if purpose is None:
+ raise ValueError("purpose is required")
+ extracted_file_data = extract_file_data(file_data)
+ object_name = self.get_object_name(extracted_file_data, purpose)
+
+ # S3 endpoint URL format
+ s3_endpoint_url = optional_params.get("s3_endpoint_url") or f"https://s3.{aws_region_name}.amazonaws.com"
+
+ return f"{s3_endpoint_url}/{bucket_name}/{object_name}"
+
+ def get_supported_openai_params(
+ self, model: str
+ ) -> List[OpenAICreateFileRequestOptionalParams]:
+ return []
+
+ def map_openai_params(
+ self,
+ non_default_params: dict,
+ optional_params: dict,
+ model: str,
+ drop_params: bool,
+ ) -> dict:
+ return optional_params
+
+ def _get_bedrock_provider_from_model(self, model: str) -> Optional[str]:
+ """
+ Extract provider from Bedrock model name
+ """
+ if model.startswith("anthropic."):
+ return "anthropic"
+ elif model.startswith("cohere."):
+ return "cohere"
+ elif model.startswith("meta.") or model.startswith("llama"):
+ return "meta"
+ elif model.startswith("mistral."):
+ return "mistral"
+ elif model.startswith("ai21."):
+ return "ai21"
+ elif model.startswith("amazon."):
+ return "amazon"
+ else:
+ return None
+
+ def _map_openai_to_bedrock_params(
+ self,
+ openai_request_body: Dict[str, Any],
+ provider: Optional[str] = None,
+ ) -> Dict[str, Any]:
+ """
+ Transform OpenAI request body to Bedrock-compatible modelInput parameters using existing transformation logic
+ """
+ _model = openai_request_body.get("model", "")
+ messages = openai_request_body.get("messages", [])
+
+ # Use existing Anthropic transformation logic for Anthropic models
+ if provider == "anthropic":
+ from litellm.llms.bedrock.chat.invoke_transformations.anthropic_claude3_transformation import (
+ AmazonAnthropicClaudeConfig,
+ )
+
+ anthropic_config = AmazonAnthropicClaudeConfig()
+
+ # Extract optional params (everything except model and messages)
+ optional_params = {k: v for k, v in openai_request_body.items() if k not in ["model", "messages"]}
+
+ # Transform using existing Anthropic logic
+ bedrock_params = anthropic_config.transform_request(
+ model=_model,
+ messages=messages,
+ optional_params=optional_params,
+ litellm_params={},
+ headers={}
+ )
+
+ return bedrock_params
+ else:
+ # For other providers, use basic mapping
+ bedrock_params = {
+ "messages": messages,
+ **{k: v for k, v in openai_request_body.items() if k not in ["model", "messages"]}
+ }
+ return bedrock_params
+
+ def _transform_openai_jsonl_content_to_bedrock_jsonl_content(
+ self, openai_jsonl_content: List[Dict[str, Any]]
+ ) -> List[Dict[str, Any]]:
+ """
+ Transforms OpenAI JSONL content to Bedrock batch format
+
+ Bedrock batch format: { "recordId": "alphanumeric string", "modelInput": {JSON body} }
+ Example:
+ {
+ "recordId": "CALL0000001",
+ "modelInput": {
+ "anthropic_version": "bedrock-2023-05-31",
+ "max_tokens": 1024,
+ "messages": [
+ {
+ "role": "user",
+ "content": [{"type": "text", "text": "Hello"}]
+ }
+ ]
+ }
+ }
+ """
+
+ bedrock_jsonl_content = []
+ for idx, _openai_jsonl_content in enumerate(openai_jsonl_content):
+ # Extract the request body from OpenAI format
+ openai_body = _openai_jsonl_content.get("body", {})
+ model = openai_body.get("model", "")
+
+ # Determine provider from model name
+ provider = self._get_bedrock_provider_from_model(model)
+
+ # Transform to Bedrock modelInput format
+ model_input = self._map_openai_to_bedrock_params(
+ openai_request_body=openai_body,
+ provider=provider
+ )
+
+ # Create Bedrock batch record
+ record_id = _openai_jsonl_content.get("custom_id", f"CALL{str(idx).zfill(7)}")
+ bedrock_record = {
+ "recordId": record_id,
+ "modelInput": model_input
+ }
+
+ bedrock_jsonl_content.append(bedrock_record)
+ return bedrock_jsonl_content
+
+ def transform_create_file_request(
+ self,
+ model: str,
+ create_file_data: CreateFileRequest,
+ optional_params: dict,
+ litellm_params: dict,
+ ) -> Union[bytes, str, dict]:
+ """
+ Transform file request and return a pre-signed request for S3.
+ This keeps the HTTP handler clean by doing all the signing here.
+ """
+ file_data = create_file_data.get("file")
+ if file_data is None:
+ raise ValueError("file is required")
+ extracted_file_data = extract_file_data(file_data)
+ extracted_file_data_content = extracted_file_data.get("content")
+
+ # Get and transform the file content
+ if (
+ create_file_data.get("purpose") == "batch"
+ and extracted_file_data.get("content_type") == "application/jsonl"
+ and extracted_file_data_content is not None
+ ):
+ ## Transform JSONL content to Bedrock format
+ original_file_content = self._get_content_from_openai_file(
+ extracted_file_data_content
+ )
+ openai_jsonl_content = [
+ json.loads(line) for line in original_file_content.splitlines() if line.strip()
+ ]
+ bedrock_jsonl_content = (
+ self._transform_openai_jsonl_content_to_bedrock_jsonl_content(
+ openai_jsonl_content
+ )
+ )
+ file_content = "\n".join(json.dumps(item) for item in bedrock_jsonl_content)
+ elif isinstance(extracted_file_data_content, bytes):
+ file_content = extracted_file_data_content.decode('utf-8')
+ elif isinstance(extracted_file_data_content, str):
+ file_content = extracted_file_data_content
+ else:
+ raise ValueError("Unsupported file content type")
+
+ # Get the S3 URL for upload
+ api_base = self.get_complete_file_url(
+ api_base=None,
+ api_key=None,
+ model=model,
+ optional_params=optional_params,
+ litellm_params=litellm_params,
+ data=create_file_data,
+ )
+
+ # Sign the request and return a pre-signed request object
+ signed_headers, signed_body = self._sign_s3_request(
+ content=file_content,
+ api_base=api_base,
+ optional_params=optional_params,
+ )
+
+ # Return a dict that tells the HTTP handler exactly what to do
+ return {
+ "method": "PUT",
+ "url": api_base,
+ "headers": signed_headers,
+ "data": signed_body or file_content,
+ }
+
+ def _sign_s3_request(
+ self,
+ content: str,
+ api_base: str,
+ optional_params: dict,
+ ) -> Tuple[dict, str]:
+ """
+ Sign S3 PUT request using the same proven logic as S3Logger.
+ Reuses the exact pattern from litellm/integrations/s3_v2.py
+ """
+ try:
+ import hashlib
+
+ import requests
+ from botocore.auth import SigV4Auth
+ from botocore.awsrequest import AWSRequest
+ except ImportError:
+ raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.")
+
+ # Get AWS credentials using existing methods
+ aws_region_name = self._get_aws_region_name(
+ optional_params=optional_params, model=""
+ )
+ credentials = self.get_credentials(
+ aws_access_key_id=optional_params.get("aws_access_key_id"),
+ aws_secret_access_key=optional_params.get("aws_secret_access_key"),
+ aws_session_token=optional_params.get("aws_session_token"),
+ aws_region_name=aws_region_name,
+ aws_session_name=optional_params.get("aws_session_name"),
+ aws_profile_name=optional_params.get("aws_profile_name"),
+ aws_role_name=optional_params.get("aws_role_name"),
+ aws_web_identity_token=optional_params.get("aws_web_identity_token"),
+ aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
+ )
+
+ # Calculate SHA256 hash of the content (REQUIRED for S3)
+ content_hash = hashlib.sha256(content.encode("utf-8")).hexdigest()
+
+ # Prepare headers with required S3 headers (same as s3_v2.py)
+ request_headers = {
+ "Content-Type": "application/json", # JSONL files are JSON content
+ "x-amz-content-sha256": content_hash, # REQUIRED by S3
+ "Content-Language": "en",
+ "Cache-Control": "private, immutable, max-age=31536000, s-maxage=0",
+ }
+
+ # Use requests.Request to prepare the request (same pattern as s3_v2.py)
+ req = requests.Request("PUT", api_base, data=content, headers=request_headers)
+ prepped = req.prepare()
+
+ # Sign the request with S3 service
+ aws_request = AWSRequest(
+ method=prepped.method,
+ url=prepped.url,
+ data=prepped.body,
+ headers=prepped.headers,
+ )
+
+ # Get region name for non-LLM API calls (same as s3_v2.py)
+ signing_region = self.get_aws_region_name_for_non_llm_api_calls(
+ aws_region_name=aws_region_name
+ )
+
+ SigV4Auth(credentials, "s3", signing_region).add_auth(aws_request)
+
+ # Return signed headers and body
+ signed_body = aws_request.body
+ if isinstance(signed_body, bytes):
+ signed_body = signed_body.decode('utf-8')
+ elif signed_body is None:
+ signed_body = content # Fallback to original content
+
+ return dict(aws_request.headers), signed_body
+
+ def transform_create_file_response(
+ self,
+ model: Optional[str],
+ raw_response: Response,
+ logging_obj: LiteLLMLoggingObj,
+ litellm_params: dict,
+ ) -> OpenAIFileObject:
+ """
+ Transform S3 File upload response into OpenAI-style FileObject
+ """
+ # For S3 uploads, we typically get an ETag and other metadata
+ response_headers = raw_response.headers
+
+ # Extract S3 object information from the response
+ # S3 PUT object returns ETag and other metadata in headers
+ content_length = response_headers.get("Content-Length", "0")
+
+ # Extract bucket and key from the request URL or litellm_params
+ bucket_name = litellm_params.get("s3_bucket_name") or os.getenv("AWS_S3_BUCKET_NAME")
+
+ # Generate file ID in S3 format
+ object_key = getattr(logging_obj, 'object_key', None) or f"file-{int(time.time())}"
+ file_id = f"s3://{bucket_name}/{object_key}"
+
+ # Extract filename from object key
+ filename = object_key.split("/")[-1] if "/" in object_key else object_key
+
+ return OpenAIFileObject(
+ purpose="batch", # Default purpose for Bedrock files
+ id=file_id,
+ filename=filename,
+ created_at=int(time.time()), # Current timestamp
+ status="uploaded",
+ bytes=int(content_length) if content_length.isdigit() else 0,
+ object="file",
+ )
+
+ def get_error_class(
+ self, error_message: str, status_code: int, headers: Union[Dict, Headers]
+ ) -> BaseLLMException:
+ return BedrockError(
+ status_code=status_code, message=error_message, headers=headers
+ )
+
+
+class BedrockJsonlFilesTransformation:
+ """
+ Transforms OpenAI /v1/files/* requests to Bedrock S3 file uploads for batch processing
+ """
+
+ def transform_openai_file_content_to_bedrock_file_content(
+ self, openai_file_content: Optional[FileTypes] = None
+ ) -> Tuple[str, str]:
+ """
+ Transforms OpenAI FileContentRequest to Bedrock S3 file format
+ """
+
+ if openai_file_content is None:
+ raise ValueError("contents of file are None")
+ # Read the content of the file
+ file_content = self._get_content_from_openai_file(openai_file_content)
+
+ # Split into lines and parse each line as JSON
+ openai_jsonl_content = [
+ json.loads(line) for line in file_content.splitlines() if line.strip()
+ ]
+ bedrock_jsonl_content = (
+ self._transform_openai_jsonl_content_to_bedrock_jsonl_content(
+ openai_jsonl_content
+ )
+ )
+ bedrock_jsonl_string = "\n".join(
+ json.dumps(item) for item in bedrock_jsonl_content
+ )
+ object_name = self._get_s3_object_name(
+ openai_jsonl_content=openai_jsonl_content
+ )
+ return bedrock_jsonl_string, object_name
+
+ def _transform_openai_jsonl_content_to_bedrock_jsonl_content(
+ self, openai_jsonl_content: List[Dict[str, Any]]
+ ):
+ """
+ Delegate to the main BedrockFilesConfig transformation method
+ """
+ config = BedrockFilesConfig()
+ return config._transform_openai_jsonl_content_to_bedrock_jsonl_content(openai_jsonl_content)
+
+ def _get_s3_object_name(
+ self,
+ openai_jsonl_content: List[Dict[str, Any]],
+ ) -> str:
+ """
+ Gets a unique S3 object name for the Bedrock batch processing job
+
+ named as: litellm-bedrock-files-{model}-{uuid}
+ """
+ _model = openai_jsonl_content[0].get("body", {}).get("model", "")
+ # Remove bedrock/ prefix if present
+ if _model.startswith("bedrock/"):
+ _model = _model[8:]
+ object_name = f"litellm-bedrock-files-{_model}-{uuid.uuid4()}.jsonl"
+ return object_name
+
+
+
+ def _get_content_from_openai_file(self, openai_file_content: FileTypes) -> str:
+ """
+ Helper to extract content from various OpenAI file types and return as string.
+
+ Handles:
+ - Direct content (str, bytes, IO[bytes])
+ - Tuple formats: (filename, content, [content_type], [headers])
+ - PathLike objects
+ """
+ content: Union[str, bytes] = b""
+ # Extract file content from tuple if necessary
+ if isinstance(openai_file_content, tuple):
+ # Take the second element which is always the file content
+ file_content = openai_file_content[1]
+ else:
+ file_content = openai_file_content
+
+ # Handle different file content types
+ if isinstance(file_content, str):
+ # String content can be used directly
+ content = file_content
+ elif isinstance(file_content, bytes):
+ # Bytes content can be decoded
+ content = file_content
+ elif isinstance(file_content, PathLike): # PathLike
+ with open(str(file_content), "rb") as f:
+ content = f.read()
+ elif hasattr(file_content, "read"): # IO[bytes]
+ # File-like objects need to be read
+ content = file_content.read()
+
+ # Ensure content is string
+ if isinstance(content, bytes):
+ content = content.decode("utf-8")
+
+ return content
+
+ def transform_s3_bucket_response_to_openai_file_object(
+ self, create_file_data: CreateFileRequest, s3_upload_response: Dict[str, Any]
+ ) -> OpenAIFileObject:
+ """
+ Transforms S3 Bucket upload file response to OpenAI FileObject
+ """
+ # S3 response typically contains ETag, key, etc.
+ object_key = s3_upload_response.get("Key", "")
+ bucket_name = s3_upload_response.get("Bucket", "")
+
+ # Extract filename from object key
+ filename = object_key.split("/")[-1] if "/" in object_key else object_key
+
+ return OpenAIFileObject(
+ purpose=create_file_data.get("purpose", "batch"),
+ id=f"s3://{bucket_name}/{object_key}",
+ filename=filename,
+ created_at=int(time.time()), # Current timestamp
+ status="uploaded",
+ bytes=s3_upload_response.get("ContentLength", 0),
+ object="file",
+ )
diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py
index 4d8781fff2a..36b543086f5 100644
--- a/litellm/llms/custom_httpx/http_handler.py
+++ b/litellm/llms/custom_httpx/http_handler.py
@@ -212,6 +212,7 @@ class AsyncHTTPHandler:
verify=ssl_config,
cert=cert,
headers=headers,
+ follow_redirects=True,
)
async def close(self):
@@ -687,6 +688,7 @@ class HTTPHandler:
verify=ssl_config,
cert=cert,
headers=headers,
+ follow_redirects=True,
)
else:
self.client = client
diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py
index 2faea53901c..13133a56aad 100644
--- a/litellm/llms/custom_httpx/llm_http_handler.py
+++ b/litellm/llms/custom_httpx/llm_http_handler.py
@@ -28,6 +28,7 @@ from litellm.llms.base_llm.audio_transcription.transformation import (
BaseAudioTranscriptionConfig,
)
from litellm.llms.base_llm.base_model_iterator import MockResponseIterator
+from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig
from litellm.llms.base_llm.chat.transformation import BaseConfig
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
from litellm.llms.base_llm.files.transformation import BaseFilesConfig
@@ -58,6 +59,7 @@ from litellm.types.llms.anthropic_messages.anthropic_response import (
AnthropicMessagesResponse,
)
from litellm.types.llms.openai import (
+ CreateBatchRequest,
CreateFileRequest,
OpenAIFileObject,
ResponseInputParam,
@@ -66,7 +68,12 @@ from litellm.types.llms.openai import (
from litellm.types.rerank import OptionalRerankParams, RerankResponse
from litellm.types.responses.main import DeleteResponseResult
from litellm.types.router import GenericLiteLLMParams
-from litellm.types.utils import EmbeddingResponse, FileTypes, TranscriptionResponse
+from litellm.types.utils import (
+ EmbeddingResponse,
+ FileTypes,
+ LiteLLMBatch,
+ TranscriptionResponse,
+)
from litellm.types.vector_stores import (
VectorStoreCreateOptionalRequestParams,
VectorStoreCreateResponse,
@@ -2212,15 +2219,38 @@ class BaseLLMHTTPHandler:
else:
sync_httpx_client = client
- if isinstance(transformed_request, str) or isinstance(
- transformed_request, bytes
- ):
- upload_response = sync_httpx_client.post(
- url=api_base,
- headers=headers,
- data=transformed_request,
+ if isinstance(transformed_request, dict) and "method" in transformed_request:
+ # Handle pre-signed requests (e.g., from Bedrock S3 uploads)
+ upload_response = getattr(sync_httpx_client, transformed_request["method"].lower())(
+ url=transformed_request["url"],
+ headers=transformed_request["headers"],
+ data=transformed_request["data"],
timeout=timeout,
)
+ elif isinstance(transformed_request, str) or isinstance(
+ transformed_request, bytes
+ ):
+ # Handle traditional file uploads
+ # Ensure transformed_request is a string for httpx compatibility
+ if isinstance(transformed_request, bytes):
+ transformed_request = transformed_request.decode('utf-8')
+
+ # Use the HTTP method specified by the provider config
+ http_method = provider_config.file_upload_http_method.upper()
+ if http_method == "PUT":
+ upload_response = sync_httpx_client.put(
+ url=api_base,
+ headers=headers,
+ data=transformed_request,
+ timeout=timeout,
+ )
+ else: # Default to POST
+ upload_response = sync_httpx_client.post(
+ url=api_base,
+ headers=headers,
+ data=transformed_request,
+ timeout=timeout,
+ )
else:
try:
# Step 1: Initial request to get upload URL
@@ -2280,16 +2310,52 @@ class BaseLLMHTTPHandler:
)
else:
async_httpx_client = client
+
+ #########################################################
+ # Debug Logging
+ #########################################################
+ logging_obj.pre_call(
+ input="",
+ api_key="",
+ additional_args={
+ "complete_input_dict": transformed_request,
+ "api_base": api_base,
+ "headers": headers,
+ },
+ )
- if isinstance(transformed_request, str) or isinstance(
- transformed_request, bytes
- ):
- upload_response = await async_httpx_client.post(
- url=api_base,
- headers=headers,
- data=transformed_request,
+ if isinstance(transformed_request, dict) and "method" in transformed_request:
+ # Handle pre-signed requests (e.g., from Bedrock S3 uploads)
+ upload_response = await getattr(async_httpx_client, transformed_request["method"].lower())(
+ url=transformed_request["url"],
+ headers=transformed_request["headers"],
+ data=transformed_request["data"],
timeout=timeout,
)
+ elif isinstance(transformed_request, str) or isinstance(
+ transformed_request, bytes
+ ):
+ # Handle traditional file uploads
+ # Ensure transformed_request is a string for httpx compatibility
+ if isinstance(transformed_request, bytes):
+ transformed_request = transformed_request.decode('utf-8')
+
+ # Use the HTTP method specified by the provider config
+ http_method = provider_config.file_upload_http_method.upper()
+ if http_method == "PUT":
+ upload_response = await async_httpx_client.put(
+ url=api_base,
+ headers=headers,
+ data=transformed_request,
+ timeout=timeout,
+ )
+ else: # Default to POST
+ upload_response = await async_httpx_client.post(
+ url=api_base,
+ headers=headers,
+ data=transformed_request,
+ timeout=timeout,
+ )
else:
try:
# Step 1: Initial request to get upload URL
@@ -2330,6 +2396,188 @@ class BaseLLMHTTPHandler:
litellm_params=litellm_params,
)
+ def create_batch(
+ self,
+ create_batch_data: "CreateBatchRequest",
+ litellm_params: dict,
+ provider_config: "BaseBatchesConfig",
+ headers: dict,
+ api_base: Optional[str],
+ api_key: Optional[str],
+ logging_obj: "LiteLLMLoggingObj",
+ _is_async: bool = False,
+ client: Optional[Union["HTTPHandler", "AsyncHTTPHandler"]] = None,
+ timeout: Optional[Union[float, httpx.Timeout]] = None,
+ ) -> Union["LiteLLMBatch", Coroutine[Any, Any, "LiteLLMBatch"]]:
+ """
+ Creates a batch using provider-specific batch creation process
+ """
+ # get config from model, custom llm provider
+ headers = provider_config.validate_environment(
+ api_key=api_key,
+ headers=headers,
+ model="",
+ messages=[],
+ optional_params={},
+ litellm_params=litellm_params,
+ )
+
+ api_base = provider_config.get_complete_batch_url(
+ api_base=api_base,
+ api_key=api_key,
+ model="",
+ optional_params={},
+ litellm_params=litellm_params,
+ data=create_batch_data,
+ )
+ if api_base is None:
+ raise ValueError("api_base is required for create_batch")
+
+ # Get the transformed request data
+ transformed_request = provider_config.transform_create_batch_request(
+ model="",
+ create_batch_data=create_batch_data,
+ litellm_params=litellm_params,
+ optional_params={},
+ )
+
+ if _is_async:
+ return self.async_create_batch(
+ transformed_request=transformed_request,
+ litellm_params=litellm_params,
+ provider_config=provider_config,
+ headers=headers,
+ api_base=api_base,
+ logging_obj=logging_obj,
+ client=client,
+ timeout=timeout,
+ create_batch_data=create_batch_data,
+ )
+
+ if client is None or not isinstance(client, HTTPHandler):
+ sync_httpx_client = _get_httpx_client()
+ else:
+ sync_httpx_client = client
+
+ try:
+ if isinstance(transformed_request, dict) and "method" in transformed_request:
+ # Handle pre-signed requests (e.g., from Bedrock with AWS auth)
+ batch_response = getattr(sync_httpx_client, transformed_request["method"].lower())(
+ url=transformed_request["url"],
+ headers=transformed_request["headers"],
+ data=transformed_request["data"],
+ timeout=timeout,
+ )
+ elif isinstance(transformed_request, dict):
+ # For other providers that use JSON requests
+ batch_response = sync_httpx_client.post(
+ url=api_base,
+ headers={**headers, "Content-Type": "application/json"},
+ json=transformed_request,
+ timeout=timeout,
+ )
+ else:
+ # Handle other request types if needed
+ batch_response = sync_httpx_client.post(
+ url=api_base,
+ headers=headers,
+ data=transformed_request,
+ timeout=timeout,
+ )
+ except Exception as e:
+ verbose_logger.exception(f"Error creating batch: {e}")
+ raise self._handle_error(
+ e=e,
+ provider_config=provider_config,
+ )
+
+ # Store original request for response transformation
+ litellm_params_with_request = {**litellm_params, "original_batch_request": create_batch_data}
+
+ return provider_config.transform_create_batch_response(
+ model=None,
+ raw_response=batch_response,
+ logging_obj=logging_obj,
+ litellm_params=litellm_params_with_request,
+ )
+
+ async def async_create_batch(
+ self,
+ transformed_request: Union[bytes, str, dict],
+ litellm_params: dict,
+ provider_config: "BaseBatchesConfig",
+ headers: dict,
+ api_base: str,
+ logging_obj: "LiteLLMLoggingObj",
+ client: Optional[Union["HTTPHandler", "AsyncHTTPHandler"]] = None,
+ timeout: Optional[Union[float, httpx.Timeout]] = None,
+ create_batch_data: Optional["CreateBatchRequest"] = None,
+ ):
+ """
+ Async version of create_batch
+ """
+ if client is None or not isinstance(client, AsyncHTTPHandler):
+ async_httpx_client = get_async_httpx_client(
+ llm_provider=provider_config.custom_llm_provider
+ )
+ else:
+ async_httpx_client = client
+
+ #########################################################
+ # Debug Logging
+ #########################################################
+ logging_obj.pre_call(
+ input="",
+ api_key="",
+ additional_args={
+ "complete_input_dict": transformed_request,
+ "api_base": api_base,
+ "headers": headers,
+ },
+ )
+
+ try:
+ if isinstance(transformed_request, dict) and "method" in transformed_request:
+ # Handle pre-signed requests (e.g., from Bedrock with AWS auth)
+ batch_response = await getattr(async_httpx_client, transformed_request["method"].lower())(
+ url=transformed_request["url"],
+ headers=transformed_request["headers"],
+ data=transformed_request["data"],
+ timeout=timeout,
+ )
+ elif isinstance(transformed_request, dict):
+ # For other providers that use JSON requests
+ batch_response = await async_httpx_client.post(
+ url=api_base,
+ headers={**headers, "Content-Type": "application/json"},
+ json=transformed_request,
+ timeout=timeout,
+ )
+ else:
+ # Handle other request types if needed
+ batch_response = await async_httpx_client.post(
+ url=api_base,
+ headers=headers,
+ data=transformed_request,
+ timeout=timeout,
+ )
+ except Exception as e:
+ verbose_logger.exception(f"Error creating batch: {e}")
+ raise self._handle_error(
+ e=e,
+ provider_config=provider_config,
+ )
+
+ # Store original request for response transformation (for async version)
+ litellm_params_with_request = {**litellm_params, "original_batch_request": create_batch_data or {}}
+
+ return provider_config.transform_create_batch_response(
+ model=None,
+ raw_response=batch_response,
+ logging_obj=logging_obj,
+ litellm_params=litellm_params_with_request,
+ )
+
def list_files(self):
"""
Lists all files
@@ -2381,6 +2629,7 @@ class BaseLLMHTTPHandler:
BaseVectorStoreConfig,
BaseGoogleGenAIGenerateContentConfig,
BaseAnthropicMessagesConfig,
+ BaseBatchesConfig,
"BasePassthroughConfig",
],
):
diff --git a/litellm/llms/groq/chat/transformation.py b/litellm/llms/groq/chat/transformation.py
index 86fa323f9e3..165301efb5c 100644
--- a/litellm/llms/groq/chat/transformation.py
+++ b/litellm/llms/groq/chat/transformation.py
@@ -6,6 +6,8 @@ from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, cast,
import httpx
from pydantic import BaseModel
+import litellm
+from litellm._logging import verbose_logger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import (
@@ -55,6 +57,10 @@ class GroqChatConfig(OpenAILikeChatConfig):
if key != "self" and value is not None:
setattr(self.__class__, key, value)
+ @property
+ def custom_llm_provider(self) -> Optional[str]:
+ return "groq"
+
@classmethod
def get_config(cls):
return super().get_config()
@@ -65,6 +71,15 @@ class GroqChatConfig(OpenAILikeChatConfig):
base_params.remove("max_retries")
except ValueError:
pass
+
+ try:
+ if litellm.supports_reasoning(
+ model=model, custom_llm_provider=self.custom_llm_provider
+ ):
+ base_params.append("reasoning_effort")
+ except Exception as e:
+ verbose_logger.debug(f"Error checking if model supports reasoning: {e}")
+
return base_params
@overload
diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py
index ee30c1749a8..ea29970f0aa 100644
--- a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py
+++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py
@@ -1,5 +1,6 @@
# What is this?
## API Handler for calling Vertex AI Partner Models
+from enum import Enum
from typing import Callable, Optional, Union
import httpx # type: ignore
@@ -27,6 +28,16 @@ class VertexAIError(Exception):
self.message
) # Call the base class constructor with the parameters it needs
+class PartnerModelPrefixes(str, Enum):
+ META_PREFIX = "meta/"
+ DEEPSEEK_PREFIX = "deepseek-ai"
+ MISTRAL_PREFIX = "mistral"
+ CODERESTAL_PREFIX = "codestral"
+ JAMBA_PREFIX = "jamba"
+ CLAUDE_PREFIX = "claude"
+ QWEN_PREFIX = "qwen"
+ GPT_OSS_PREFIX = "openai/gpt-oss-"
+
class VertexAIPartnerModels(VertexBase):
def __init__(self) -> None:
@@ -42,14 +53,14 @@ class VertexAIPartnerModels(VertexBase):
bool: True if the model string is a Vertex AI Partner Model, False otherwise
"""
if (
- model.startswith("meta/")
- or model.startswith("deepseek-ai")
- or model.startswith("mistral")
- or model.startswith("codestral")
- or model.startswith("jamba")
- or model.startswith("claude")
- or model.startswith("qwen")
- or model.startswith("openai")
+ model.startswith(PartnerModelPrefixes.META_PREFIX)
+ or model.startswith(PartnerModelPrefixes.DEEPSEEK_PREFIX)
+ or model.startswith(PartnerModelPrefixes.MISTRAL_PREFIX)
+ or model.startswith(PartnerModelPrefixes.CODERESTAL_PREFIX)
+ or model.startswith(PartnerModelPrefixes.JAMBA_PREFIX)
+ or model.startswith(PartnerModelPrefixes.CLAUDE_PREFIX)
+ or model.startswith(PartnerModelPrefixes.QWEN_PREFIX)
+ or model.startswith(PartnerModelPrefixes.GPT_OSS_PREFIX)
):
return True
return False
@@ -58,9 +69,9 @@ class VertexAIPartnerModels(VertexBase):
def should_use_openai_handler(model: str):
OPENAI_LIKE_VERTEX_PROVIDERS = [
"llama",
- "deepseek-ai",
- "qwen",
- "openai",
+ PartnerModelPrefixes.DEEPSEEK_PREFIX,
+ PartnerModelPrefixes.QWEN_PREFIX,
+ PartnerModelPrefixes.GPT_OSS_PREFIX,
]
if any(provider in model for provider in OPENAI_LIKE_VERTEX_PROVIDERS):
return True
diff --git a/litellm/main.py b/litellm/main.py
index 9c2aa678369..d0377490942 100644
--- a/litellm/main.py
+++ b/litellm/main.py
@@ -356,7 +356,7 @@ async def acompletion(
logprobs: Optional[bool] = None,
top_logprobs: Optional[int] = None,
deployment_id=None,
- reasoning_effort: Optional[Literal["minimal", "low", "medium", "high"]] = None,
+ reasoning_effort: Optional[Literal["none", "minimal", "low", "medium", "high", "default"]] = None,
safety_identifier: Optional[str] = None,
# set api_base, api_version, api_key
base_url: Optional[str] = None,
@@ -897,7 +897,7 @@ def completion( # type: ignore # noqa: PLR0915
logit_bias: Optional[dict] = None,
user: Optional[str] = None,
# openai v1.0+ new params
- reasoning_effort: Optional[Literal["minimal", "low", "medium", "high"]] = None,
+ reasoning_effort: Optional[Literal["none", "minimal", "low", "medium", "high", "default"]] = None,
response_format: Optional[Union[dict, Type[BaseModel]]] = None,
seed: Optional[int] = None,
tools: Optional[List] = None,
diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json
index 7cedffd93f3..a7586124509 100644
--- a/litellm/model_prices_and_context_window_backup.json
+++ b/litellm/model_prices_and_context_window_backup.json
@@ -6157,21 +6157,7 @@
"supports_tool_choice": true,
"source": "https://inference-docs.cerebras.ai/support/pricing"
},
- "cerebras/openai/gpt-oss-20b": {
- "max_tokens": 32768,
- "max_input_tokens": 131072,
- "max_output_tokens": 32768,
- "input_cost_per_token": 7e-08,
- "output_cost_per_token": 3e-07,
- "litellm_provider": "cerebras",
- "mode": "chat",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_response_schema": true,
- "supports_reasoning": true,
- "supports_tool_choice": true,
- "source": "https://inference-docs.cerebras.ai/support/pricing"
- },
+
"cerebras/openai/gpt-oss-120b": {
"max_tokens": 32768,
"max_input_tokens": 131072,
@@ -9498,6 +9484,48 @@
"source": "https://aistudio.google.com",
"supports_tool_choice": true
},
+ "gemini/veo-3.0-generate-preview": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.75,
+ "litellm_provider": "gemini",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
+ "gemini/veo-3.0-fast-generate-preview": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.40,
+ "litellm_provider": "gemini",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
+ "gemini/veo-2.0-generate-001": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.35,
+ "litellm_provider": "gemini",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
"vertex_ai/claude-opus-4-1": {
"max_tokens": 4096,
"max_input_tokens": 200000,
@@ -10315,6 +10343,48 @@
"mode": "image_generation",
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing"
},
+ "vertex_ai/veo-3.0-generate-preview": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.75,
+ "litellm_provider": "vertex_ai-video-models",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
+ "vertex_ai/veo-3.0-fast-generate-preview": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.40,
+ "litellm_provider": "vertex_ai-video-models",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
+ "vertex_ai/veo-2.0-generate-001": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.35,
+ "litellm_provider": "vertex_ai-video-models",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
"text-embedding-004": {
"max_tokens": 2048,
"max_input_tokens": 2048,
diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py
index 7206d3dcb0e..0595c44d69d 100644
--- a/litellm/proxy/_types.py
+++ b/litellm/proxy/_types.py
@@ -388,7 +388,11 @@ class LiteLLMRoutes(enum.Enum):
]
# NOTE: ROUTES ONLY FOR MASTER KEY - only the Master Key should be able to Reset Spend
- master_key_only_routes = ["/global/spend/reset"]
+ master_key_only_routes = [
+ "/global/spend/reset",
+ "/memory-usage-in-mem-cache",
+ "/memory-usage-in-mem-cache-items",
+ ]
key_management_routes = [
KeyManagementRoutes.KEY_GENERATE,
diff --git a/litellm/proxy/common_utils/debug_utils.py b/litellm/proxy/common_utils/debug_utils.py
index 8096d782dec..16ab2cc8058 100644
--- a/litellm/proxy/common_utils/debug_utils.py
+++ b/litellm/proxy/common_utils/debug_utils.py
@@ -5,10 +5,12 @@ import os
import tracemalloc
from collections import Counter
-from fastapi import APIRouter
+from fastapi import APIRouter, Depends
from litellm import get_secret_str
from litellm._logging import verbose_proxy_logger
+from litellm.proxy._types import UserAPIKeyAuth
+from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
router = APIRouter()
@@ -84,7 +86,9 @@ if os.environ.get("LITELLM_PROFILE", "false").lower() == "true":
@router.get("/memory-usage-in-mem-cache", include_in_schema=False)
-async def memory_usage_in_mem_cache():
+async def memory_usage_in_mem_cache(
+ _: UserAPIKeyAuth = Depends(user_api_key_auth),
+):
# returns the size of all in-memory caches on the proxy server
"""
1. user_api_key_cache
@@ -121,7 +125,9 @@ async def memory_usage_in_mem_cache():
@router.get("/memory-usage-in-mem-cache-items", include_in_schema=False)
-async def memory_usage_in_mem_cache_items():
+async def memory_usage_in_mem_cache_items(
+ _: UserAPIKeyAuth = Depends(user_api_key_auth),
+):
# returns the size of all in-memory caches on the proxy server
"""
1. user_api_key_cache
diff --git a/litellm/proxy/hooks/parallel_request_limiter_v3.py b/litellm/proxy/hooks/parallel_request_limiter_v3.py
index b04d14bcc8b..b3840761d2a 100644
--- a/litellm/proxy/hooks/parallel_request_limiter_v3.py
+++ b/litellm/proxy/hooks/parallel_request_limiter_v3.py
@@ -68,6 +68,32 @@ end
return results
"""
+TOKEN_INCREMENT_SCRIPT = """
+local results = {}
+
+-- Process each key/increment_value/ttl triplet
+for i = 1, #KEYS do
+ local key = KEYS[i]
+ local increment_value = tonumber(ARGV[i * 2 - 1])
+ local ttl_seconds = tonumber(ARGV[i * 2])
+
+ -- Increment the value
+ local new_value = redis.call('INCRBYFLOAT', key, increment_value)
+
+ -- Handle TTL: only set expire if ttl_seconds > 0 and key has no current TTL
+ -- ttl_seconds can be 0 (no TTL) or positive (set TTL)
+ if ttl_seconds and ttl_seconds > 0 then
+ local current_ttl = redis.call('TTL', key)
+ if current_ttl == -1 then
+ redis.call('EXPIRE', key, ttl_seconds)
+ end
+ end
+
+ table.insert(results, new_value)
+end
+
+return results
+"""
class RateLimitDescriptorRateLimitObject(TypedDict, total=False):
requests_per_unit: Optional[int]
@@ -109,8 +135,14 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
BATCH_RATE_LIMITER_SCRIPT
)
)
+ self.token_increment_script = (
+ self.internal_usage_cache.dual_cache.redis_cache.async_register_script(
+ TOKEN_INCREMENT_SCRIPT
+ )
+ )
else:
self.batch_rate_limiter_script = None
+ self.token_increment_script = None
self.window_size = int(os.getenv("LITELLM_RATE_LIMIT_WINDOW_SIZE", 60))
@@ -567,6 +599,62 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
return pipeline_operations
+ async def async_increment_tokens_with_ttl_preservation(
+ self,
+ pipeline_operations: List["RedisPipelineIncrementOperation"],
+ parent_otel_span: Optional[Span] = None,
+ ) -> None:
+ """
+ Increment token counters using Lua script to preserve existing TTL.
+ This prevents TTL reset on every token increment.
+ """
+ if not pipeline_operations:
+ return
+
+ # Check if script is available
+ if self.token_increment_script is None:
+ verbose_proxy_logger.debug("TTL preservation script not available, using regular pipeline")
+ await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
+ increment_list=pipeline_operations,
+ litellm_parent_otel_span=parent_otel_span,
+ )
+ return
+
+ try:
+ # Use Lua script for all operations
+ keys = []
+ args = []
+
+ for op in pipeline_operations:
+ # Convert None TTL to 0 for Lua script
+ ttl_value = op["ttl"] if op["ttl"] is not None else 0
+
+ verbose_proxy_logger.debug(
+ f"Executing TTL-preserving increment for key={op['key']}, "
+ f"increment={op['increment_value']}, ttl={ttl_value}"
+ )
+ keys.append(op["key"])
+ args.extend([op["increment_value"], ttl_value])
+
+ await self.token_increment_script(
+ keys=keys,
+ args=args,
+ )
+
+ verbose_proxy_logger.debug(
+ f"Successfully executed TTL-preserving increment for {len(pipeline_operations)} keys"
+ )
+
+ except Exception as e:
+ verbose_proxy_logger.warning(
+ f"TTL preservation failed, falling back to regular pipeline: {str(e)}"
+ )
+ # Fallback to regular pipeline on error
+ await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
+ increment_list=pipeline_operations,
+ litellm_parent_otel_span=parent_otel_span,
+ )
+
def get_rate_limit_type(self) -> Literal["output", "input", "total"]:
from litellm.proxy.proxy_server import general_settings
@@ -713,9 +801,9 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
# Execute all increments in a single pipeline
if pipeline_operations:
- await self.internal_usage_cache.dual_cache.async_increment_cache_pipeline(
- increment_list=pipeline_operations,
- litellm_parent_otel_span=litellm_parent_otel_span,
+ await self.async_increment_tokens_with_ttl_preservation(
+ pipeline_operations=pipeline_operations,
+ parent_otel_span=litellm_parent_otel_span,
)
except Exception as e:
diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py
index 557aa3ceb9a..95f1eccffe4 100644
--- a/litellm/proxy/litellm_pre_call_utils.py
+++ b/litellm/proxy/litellm_pre_call_utils.py
@@ -677,11 +677,6 @@ async def add_litellm_data_to_request( # noqa: PLR0915
from litellm.proxy.proxy_server import llm_router, premium_user
from litellm.types.proxy.litellm_pre_call_utils import SecretFields
- safe_add_api_version_from_query_params(data, request)
- _metadata_variable_name = _get_metadata_variable_name(request)
- if data.get(_metadata_variable_name, None) is None:
- data[_metadata_variable_name] = {}
-
_headers = clean_headers(
request.headers,
@@ -692,6 +687,24 @@ async def add_litellm_data_to_request( # noqa: PLR0915
),
)
+ ##########################################################
+ # Init - Proxy Server Request
+ # we do this as soon as entering so we track the original request
+ ##########################################################
+ data["proxy_server_request"] = {
+ "url": str(request.url),
+ "method": request.method,
+ "headers": _headers,
+ "body": copy.copy(data), # use copy instead of deepcopy
+ }
+
+ safe_add_api_version_from_query_params(data, request)
+ _metadata_variable_name = _get_metadata_variable_name(request)
+ if data.get(_metadata_variable_name, None) is None:
+ data[_metadata_variable_name] = {}
+
+
+
data.update(
LiteLLMProxyRequestSetup.add_litellm_data_for_backend_llm_call(
headers=_headers,
@@ -721,13 +734,6 @@ async def add_litellm_data_to_request( # noqa: PLR0915
if "user" not in data:
data["user"] = user
- # Include original request and headers in the data
- data["proxy_server_request"] = {
- "url": str(request.url),
- "method": request.method,
- "headers": _headers,
- "body": copy.copy(data), # use copy instead of deepcopy
- }
data["secret_fields"] = SecretFields(raw_headers=dict(request.headers))
diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py
index ddce7481ce7..3868c9df694 100644
--- a/litellm/proxy/management_endpoints/key_management_endpoints.py
+++ b/litellm/proxy/management_endpoints/key_management_endpoints.py
@@ -1566,14 +1566,12 @@ async def generate_key_helper_fn( # noqa: PLR0915
if duration is None: # allow tokens that never expire
expires = None
else:
- duration_s = duration_in_seconds(duration=duration)
- expires = datetime.now(timezone.utc) + timedelta(seconds=duration_s)
+ expires = get_budget_reset_time(budget_duration=duration)
if key_budget_duration is None: # one-time budget
key_reset_at = None
else:
- duration_s = duration_in_seconds(duration=key_budget_duration)
- key_reset_at = datetime.now(timezone.utc) + timedelta(seconds=duration_s)
+ key_reset_at = get_budget_reset_time(budget_duration=key_budget_duration)
if budget_duration is None: # one-time budget
reset_at = None
diff --git a/litellm/proxy/management_endpoints/scim/scim_transformations.py b/litellm/proxy/management_endpoints/scim/scim_transformations.py
index bb07cdbd770..1bb59888405 100644
--- a/litellm/proxy/management_endpoints/scim/scim_transformations.py
+++ b/litellm/proxy/management_endpoints/scim/scim_transformations.py
@@ -121,15 +121,16 @@ class ScimTransformations:
if isinstance(team, dict):
team = LiteLLM_TeamTable(**team)
- # Get team members
+ # Get team members with proper display names
scim_members: List[SCIMMember] = []
for member in team.members_with_roles or []:
if isinstance(member, dict):
member = Member(**member)
+
scim_members.append(
SCIMMember(
value=ScimTransformations._get_scim_member_value(member),
- display=member.user_email,
+ display=ScimTransformations._get_scim_member_display(member),
)
)
@@ -151,6 +152,24 @@ class ScimTransformations:
@staticmethod
def _get_scim_member_value(member: Member) -> str:
- if member.user_email:
+ """
+ Get the SCIM member value. Use user_email if available, otherwise use user_id.
+ SCIM member value should be the unique identifier for the user.
+ """
+ if hasattr(member, "user_email") and member.user_email:
return member.user_email
+ elif hasattr(member, "user_id"):
+ return member.user_id or ScimTransformations.DEFAULT_SCIM_MEMBER_VALUE
+ return ScimTransformations.DEFAULT_SCIM_MEMBER_VALUE
+
+ @staticmethod
+ def _get_scim_member_display(member: Member) -> str:
+ """
+ Get the SCIM member display. Use user_email if available, otherwise use user_id.
+ SCIM member display should be the display name for the user.
+ """
+ if hasattr(member, "user_email") and member.user_email:
+ return member.user_email
+ elif hasattr(member, "user_id"):
+ return member.user_id or ScimTransformations.DEFAULT_SCIM_MEMBER_VALUE
return ScimTransformations.DEFAULT_SCIM_MEMBER_VALUE
diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py
index e1d94ceaff2..b9929c3b433 100644
--- a/litellm/proxy/management_endpoints/scim/scim_v2.py
+++ b/litellm/proxy/management_endpoints/scim/scim_v2.py
@@ -22,6 +22,7 @@ import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.proxy._types import (
+ LiteLLM_TeamTable,
LiteLLM_UserTable,
LitellmUserRoles,
Member,
@@ -237,6 +238,23 @@ async def _handle_team_membership_changes(user_id: str, existing_teams: List[str
)
+async def _get_team_member_user_ids_from_team(team: LiteLLM_TeamTable) -> List[str]:
+ """
+ Get the IDs of the members from a team.
+
+ Use one source of truth for the member IDs: team.members_with_roles
+
+ """
+ member_user_ids: List[str] = []
+ for member in team.members_with_roles or []:
+ if hasattr(member, "user_id") and member.user_id is not None:
+ member_user_ids.append(member.user_id)
+ elif isinstance(member, dict) and "user_id" in member:
+ user_id = member.get("user_id")
+ if user_id is not None:
+ member_user_ids.append(user_id)
+ return member_user_ids
+
# Dependency to set the correct SCIM Content-Type
async def set_scim_content_type(response: Response):
"""Sets the Content-Type header to application/scim+json"""
@@ -253,6 +271,12 @@ async def set_scim_content_type(response: Response):
)
async def get_service_provider_config(request: Request):
"""Return SCIM Service Provider Configuration."""
+ verbose_proxy_logger.debug(
+ "SCIM ServiceProviderConfig request: method=%s url=%s headers=%s",
+ request.method,
+ request.url,
+ dict(request.headers),
+ )
meta = {
"resourceType": "ServiceProviderConfig",
"location": str(request.url),
@@ -275,6 +299,12 @@ async def get_users(
"""
Get a list of users according to SCIM v2 protocol
"""
+ verbose_proxy_logger.debug(
+ "SCIM GET USERS request: startIndex=%s count=%s filter=%s",
+ startIndex,
+ count,
+ filter,
+ )
try:
prisma_client = await _get_prisma_client_or_raise_exception()
# Parse filter if provided (basic support)
@@ -334,6 +364,7 @@ async def get_user(
"""
Get a single user by ID according to SCIM v2 protocol
"""
+ verbose_proxy_logger.debug("SCIM GET USER request for user_id=%s", user_id)
try:
user = await _check_user_exists(user_id)
@@ -357,7 +388,9 @@ async def create_user(
Create a user according to SCIM v2 protocol
"""
try:
- verbose_proxy_logger.debug("SCIM CREATE USER request: %s", user)
+ verbose_proxy_logger.debug(
+ "SCIM CREATE USER request: %s", user.model_dump()
+ )
prisma_client = await _get_prisma_client_or_raise_exception()
# Extract data from SCIM user
@@ -435,7 +468,11 @@ async def update_user(
"""
Update a user according to SCIM v2 protocol (full replacement)
"""
- verbose_proxy_logger.debug("SCIM PUT USER request: %s", user)
+ verbose_proxy_logger.debug(
+ "SCIM PUT USER request for user_id=%s: %s",
+ user_id,
+ user.model_dump(),
+ )
try:
prisma_client = await _get_prisma_client_or_raise_exception()
@@ -497,6 +534,9 @@ async def delete_user(
"""
Delete a user according to SCIM v2 protocol
"""
+ verbose_proxy_logger.debug(
+ "SCIM DELETE USER request for user_id=%s", user_id
+ )
try:
prisma_client = await _get_prisma_client_or_raise_exception()
existing_user = await _check_user_exists(user_id)
@@ -691,7 +731,11 @@ async def patch_user(
"""
Patch a user according to SCIM v2 protocol
"""
- verbose_proxy_logger.debug("SCIM PATCH USER request: %s", patch_ops)
+ verbose_proxy_logger.debug(
+ "SCIM PATCH USER request for user_id=%s: %s",
+ user_id,
+ patch_ops.model_dump(),
+ )
try:
prisma_client = await _get_prisma_client_or_raise_exception()
@@ -744,6 +788,12 @@ async def get_groups(
"""
Get a list of groups according to SCIM v2 protocol
"""
+ verbose_proxy_logger.debug(
+ "SCIM GET GROUPS request: startIndex=%s count=%s filter=%s",
+ startIndex,
+ count,
+ filter,
+ )
try:
prisma_client = await _get_prisma_client_or_raise_exception()
# Parse filter if provided (basic support)
@@ -814,6 +864,9 @@ async def get_group(
"""
Get a single group by ID according to SCIM v2 protocol
"""
+ verbose_proxy_logger.debug(
+ "SCIM GET GROUP request for group_id=%s", group_id
+ )
try:
team = await _check_team_exists(group_id)
@@ -839,6 +892,10 @@ async def create_group(
"""
Create a group according to SCIM v2 protocol
"""
+ verbose_proxy_logger.debug(
+ "SCIM CREATE GROUP request: %s",
+ group.model_dump(),
+ )
try:
prisma_client = await _get_prisma_client_or_raise_exception()
@@ -892,78 +949,51 @@ async def update_group(
"""
Update a group according to SCIM v2 protocol
"""
+ verbose_proxy_logger.debug(
+ "SCIM PUT GROUP request for group_id=%s: %s",
+ group_id,
+ group.model_dump(),
+ )
try:
prisma_client = await _get_prisma_client_or_raise_exception()
existing_team = await _check_team_exists(group_id)
# Extract valid member IDs
member_ids = await _extract_group_member_ids(group)
+ verbose_proxy_logger.debug(f"SCIM PUT GROUP member_ids: {member_ids}")
- # Update team in database
+ # Prepare update data
existing_metadata = existing_team.metadata if existing_team.metadata else {}
updated_metadata = {**existing_metadata, "scim_data": group.model_dump()}
+ update_data = {
+ "team_alias": group.displayName,
+ "metadata": safe_dumps(updated_metadata),
+ }
+
+ # Update team in database
updated_team = await prisma_client.db.litellm_teamtable.update(
where={"team_id": group_id},
- data={
- "team_alias": group.displayName,
- "members": member_ids,
- "metadata": safe_dumps(updated_metadata),
- },
+ data=update_data,
)
- # Handle user-team relationships
- current_members = existing_team.members or []
-
- # Add new members to team
- for member_id in member_ids:
- if member_id not in current_members:
- user = await prisma_client.db.litellm_usertable.find_unique(
- where={"user_id": member_id}
- )
- if user:
- current_user_teams = user.teams or []
- if group_id not in current_user_teams:
- await prisma_client.db.litellm_usertable.update(
- where={"user_id": member_id},
- data={"teams": {"push": group_id}},
- )
-
- # Remove former members from team
- for member_id in current_members:
- if member_id not in member_ids:
- user = await prisma_client.db.litellm_usertable.find_unique(
- where={"user_id": member_id}
- )
- if user:
- current_user_teams = user.teams or []
- if group_id in current_user_teams:
- new_teams = [t for t in current_user_teams if t != group_id]
- await prisma_client.db.litellm_usertable.update(
- where={"user_id": member_id}, data={"teams": new_teams}
- )
-
- # Get updated members for response
- members = await _get_team_members_display(member_ids)
-
- team_created_at = (
- updated_team.created_at.isoformat() if updated_team.created_at else None
- )
- team_updated_at = (
- updated_team.updated_at.isoformat() if updated_team.updated_at else None
+ # Handle user-team relationship changes using the same approach as patch_group
+ current_members = set(await _get_team_member_user_ids_from_team(existing_team))
+ verbose_proxy_logger.debug(f"SCIM PUT GROUP current_members: {current_members}")
+ final_members = set(member_ids)
+ verbose_proxy_logger.debug(f"SCIM PUT GROUP final_members: {final_members}")
+
+ await _handle_group_membership_changes(
+ group_id=group_id,
+ current_members=current_members,
+ final_members=final_members,
)
- return SCIMGroup(
- schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
- id=group_id,
- displayName=updated_team.team_alias or group_id,
- members=members,
- meta={
- "resourceType": "Group",
- "created": team_created_at,
- "lastModified": team_updated_at,
- },
+ # Convert to SCIM format and return
+ scim_group = await ScimTransformations.transform_litellm_team_to_scim_group(
+ updated_team
)
+ return scim_group
except Exception as e:
raise handle_exception_on_proxy(e)
@@ -980,6 +1010,9 @@ async def delete_group(
"""
Delete a group according to SCIM v2 protocol
"""
+ verbose_proxy_logger.debug(
+ "SCIM DELETE GROUP request for group_id=%s", group_id
+ )
try:
prisma_client = await _get_prisma_client_or_raise_exception()
existing_team = await _check_team_exists(group_id)
@@ -1135,7 +1168,11 @@ async def patch_group(
"""
Patch a group according to SCIM v2 protocol
"""
- verbose_proxy_logger.debug("SCIM PATCH GROUP request: %s", patch_ops)
+ verbose_proxy_logger.debug(
+ "SCIM PATCH GROUP request for group_id=%s: %s",
+ group_id,
+ patch_ops.model_dump(),
+ )
try:
prisma_client = await _get_prisma_client_or_raise_exception()
@@ -1147,7 +1184,7 @@ async def patch_group(
)
# Track current members for comparison
- current_members = set(existing_team.members or [])
+ current_members = set(await _get_team_member_user_ids_from_team(existing_team))
# Apply updates to the database
updated_team = await _apply_group_patch_updates(
diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py
new file mode 100644
index 00000000000..dd772ffa502
--- /dev/null
+++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/openai_passthrough_logging_handler.py
@@ -0,0 +1,383 @@
+"""
+OpenAI Passthrough Logging Handler
+
+Handles cost tracking and logging for OpenAI passthrough endpoints, specifically /chat/completions.
+"""
+
+from datetime import datetime
+from typing import List, Optional, Union
+from urllib.parse import urlparse
+
+import httpx
+
+import litellm
+from litellm._logging import verbose_proxy_logger
+from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
+from litellm.litellm_core_utils.litellm_logging import (
+ get_standard_logging_object_payload,
+)
+from litellm.llms.openai.openai import OpenAIConfig
+from litellm.llms.openai.openai import OpenAIConfig as OpenAIConfigType
+from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
+from litellm.proxy.pass_through_endpoints.llm_provider_handlers.base_passthrough_logging_handler import (
+ BasePassthroughLoggingHandler,
+)
+from litellm.proxy.pass_through_endpoints.success_handler import (
+ PassThroughEndpointLogging,
+)
+from litellm.types.passthrough_endpoints.pass_through_endpoints import (
+ EndpointType,
+ PassthroughStandardLoggingPayload,
+)
+from litellm.types.utils import LlmProviders
+from litellm.utils import ModelResponse, TextCompletionResponse
+
+
+class OpenAIPassthroughLoggingHandler(BasePassthroughLoggingHandler):
+ """
+ OpenAI-specific passthrough logging handler that provides cost tracking for /chat/completions endpoints.
+ """
+
+ @property
+ def llm_provider_name(self) -> LlmProviders:
+ return LlmProviders.OPENAI
+
+ @staticmethod
+ def get_provider_config(model: str) -> OpenAIConfigType:
+ """Get OpenAI provider configuration for the given model."""
+ return OpenAIConfig()
+
+ @staticmethod
+ def is_openai_chat_completions_route(url_route: str) -> bool:
+ """Check if the URL route is an OpenAI chat completions endpoint."""
+ if not url_route:
+ return False
+ parsed_url = urlparse(url_route)
+ return bool(
+ parsed_url.hostname
+ and (
+ "api.openai.com" in parsed_url.hostname
+ or "openai.azure.com" in parsed_url.hostname
+ )
+ and "/v1/chat/completions" in parsed_url.path
+ )
+
+ @staticmethod
+ def _get_user_from_metadata(
+ passthrough_logging_payload: PassthroughStandardLoggingPayload,
+ ) -> Optional[str]:
+ """Extract user information from passthrough logging payload."""
+ request_body = passthrough_logging_payload.get("request_body")
+ if request_body:
+ return request_body.get("user")
+ return None
+
+ @staticmethod
+ def openai_passthrough_handler(
+ httpx_response: httpx.Response,
+ response_body: dict,
+ logging_obj: LiteLLMLoggingObj,
+ url_route: str,
+ result: str,
+ start_time: datetime,
+ end_time: datetime,
+ cache_hit: bool,
+ request_body: dict,
+ **kwargs,
+ ) -> PassThroughEndpointLoggingTypedDict:
+ """
+ Handle OpenAI passthrough logging with cost tracking for chat completions.
+ """
+ # Only handle chat completions endpoints
+ if not OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route(
+ url_route
+ ):
+ # For non-chat-completions endpoints, use the base handler without cost tracking
+ base_handler = OpenAIPassthroughLoggingHandler()
+ return base_handler.passthrough_chat_handler(
+ httpx_response=httpx_response,
+ response_body=response_body,
+ logging_obj=logging_obj,
+ url_route=url_route,
+ result=result,
+ start_time=start_time,
+ end_time=end_time,
+ cache_hit=cache_hit,
+ request_body=request_body,
+ **kwargs,
+ )
+
+ # Extract model from request or response
+ model = request_body.get("model", response_body.get("model", ""))
+ if not model:
+ verbose_proxy_logger.warning(
+ "No model found in request or response for OpenAI passthrough cost tracking"
+ )
+ base_handler = OpenAIPassthroughLoggingHandler()
+ return base_handler.passthrough_chat_handler(
+ httpx_response=httpx_response,
+ response_body=response_body,
+ logging_obj=logging_obj,
+ url_route=url_route,
+ result=result,
+ start_time=start_time,
+ end_time=end_time,
+ cache_hit=cache_hit,
+ request_body=request_body,
+ **kwargs,
+ )
+
+ try:
+ # Transform the response to LiteLLM format for cost calculation
+ provider_config = OpenAIPassthroughLoggingHandler.get_provider_config(
+ model=model
+ )
+ litellm_model_response: ModelResponse = provider_config.transform_response(
+ raw_response=httpx_response,
+ model_response=litellm.ModelResponse(),
+ model=model,
+ messages=request_body.get("messages", []),
+ logging_obj=logging_obj,
+ optional_params=request_body.get("optional_params", {}),
+ api_key="",
+ request_data=request_body,
+ encoding=litellm.encoding,
+ json_mode=request_body.get("response_format", {}).get("type")
+ == "json_object",
+ litellm_params={},
+ )
+
+ # Calculate cost using LiteLLM's cost calculator
+ response_cost = litellm.completion_cost(
+ completion_response=litellm_model_response,
+ model=model,
+ custom_llm_provider="openai",
+ )
+
+ # Update kwargs with cost information
+ kwargs["response_cost"] = response_cost
+ kwargs["model"] = model
+ kwargs["custom_llm_provider"] = "openai"
+
+ # Extract user information for tracking
+ passthrough_logging_payload: Optional[
+ PassthroughStandardLoggingPayload
+ ] = kwargs.get("passthrough_logging_payload")
+ if passthrough_logging_payload:
+ user = OpenAIPassthroughLoggingHandler._get_user_from_metadata(
+ passthrough_logging_payload=passthrough_logging_payload,
+ )
+ if user:
+ kwargs.setdefault("litellm_params", {})
+ kwargs["litellm_params"].update(
+ {"proxy_server_request": {"body": {"user": user}}}
+ )
+
+ # Create standard logging object
+ get_standard_logging_object_payload(
+ kwargs=kwargs,
+ init_response_obj=litellm_model_response,
+ start_time=start_time,
+ end_time=end_time,
+ logging_obj=logging_obj,
+ status="success",
+ )
+
+ # Update logging object with cost information
+ logging_obj.model_call_details["model"] = model
+ logging_obj.model_call_details["custom_llm_provider"] = "openai"
+ logging_obj.model_call_details["response_cost"] = response_cost
+
+ verbose_proxy_logger.debug(
+ f"OpenAI passthrough cost tracking - Model: {model}, Cost: ${response_cost:.6f}"
+ )
+
+ return {
+ "result": litellm_model_response,
+ "kwargs": kwargs,
+ }
+
+ except Exception as e:
+ verbose_proxy_logger.error(
+ f"Error in OpenAI passthrough cost tracking: {str(e)}"
+ )
+ # Fall back to base handler without cost tracking
+ base_handler = OpenAIPassthroughLoggingHandler()
+ return base_handler.passthrough_chat_handler(
+ httpx_response=httpx_response,
+ response_body=response_body,
+ logging_obj=logging_obj,
+ url_route=url_route,
+ result=result,
+ start_time=start_time,
+ end_time=end_time,
+ cache_hit=cache_hit,
+ request_body=request_body,
+ **kwargs,
+ )
+
+ def _build_complete_streaming_response(
+ self,
+ all_chunks: list,
+ litellm_logging_obj: LiteLLMLoggingObj,
+ model: str,
+ ) -> Optional[Union[ModelResponse, TextCompletionResponse]]:
+ """
+ Builds complete response from raw chunks for OpenAI streaming responses.
+
+ - Converts str chunks to generic chunks
+ - Converts generic chunks to litellm chunks (OpenAI format)
+ - Builds complete response from litellm chunks
+ """
+ try:
+ # OpenAI's response iterator to parse chunks
+ from litellm.llms.openai.openai import OpenAIChatCompletionResponseIterator
+
+ openai_iterator = OpenAIChatCompletionResponseIterator(
+ streaming_response=None,
+ sync_stream=False,
+ )
+
+ all_openai_chunks = []
+ for chunk_str in all_chunks:
+ try:
+ # Parse the string chunk using the base iterator's string parser
+ from litellm.llms.base_llm.base_model_iterator import (
+ BaseModelResponseIterator,
+ )
+
+ # Convert string chunk to dict
+ stripped_json_chunk = (
+ BaseModelResponseIterator._string_to_dict_parser(
+ str_line=chunk_str
+ )
+ )
+
+ if stripped_json_chunk:
+ # Parse the chunk using OpenAI's chunk parser
+ transformed_chunk = openai_iterator.chunk_parser(
+ chunk=stripped_json_chunk
+ )
+ if transformed_chunk is not None:
+ all_openai_chunks.append(transformed_chunk)
+
+ except (StopIteration, StopAsyncIteration, Exception) as e:
+ verbose_proxy_logger.debug(f"Error parsing streaming chunk: {e}")
+ continue
+
+ if not all_openai_chunks:
+ verbose_proxy_logger.warning(
+ "No valid chunks found in streaming response"
+ )
+ return None
+
+ # Build complete response from chunks
+ complete_streaming_response = litellm.stream_chunk_builder(
+ chunks=all_openai_chunks
+ )
+
+ return complete_streaming_response
+
+ except Exception as e:
+ verbose_proxy_logger.error(
+ f"Error building complete streaming response: {str(e)}"
+ )
+ return None
+
+ @staticmethod
+ def _handle_logging_openai_collected_chunks(
+ litellm_logging_obj: LiteLLMLoggingObj,
+ passthrough_success_handler_obj: PassThroughEndpointLogging,
+ url_route: str,
+ request_body: dict,
+ endpoint_type: EndpointType,
+ start_time: datetime,
+ all_chunks: List[str],
+ end_time: datetime,
+ ) -> PassThroughEndpointLoggingTypedDict:
+ """
+ Handle logging for collected OpenAI streaming chunks with cost tracking.
+ """
+ try:
+ # Extract model from request body
+ model = request_body.get("model", "gpt-4o")
+
+ # Build complete response from chunks using our streaming handler
+ handler = OpenAIPassthroughLoggingHandler()
+ complete_response = handler._build_complete_streaming_response(
+ all_chunks=all_chunks,
+ litellm_logging_obj=litellm_logging_obj,
+ model=model,
+ )
+
+ if complete_response is None:
+ verbose_proxy_logger.warning(
+ "Failed to build complete response from OpenAI streaming chunks"
+ )
+ return {
+ "result": None,
+ "kwargs": {},
+ }
+
+ # Calculate cost using LiteLLM's cost calculator
+ response_cost = litellm.completion_cost(
+ completion_response=complete_response,
+ model=model,
+ custom_llm_provider="openai",
+ )
+
+ # Prepare kwargs for logging
+ kwargs = {
+ "response_cost": response_cost,
+ "model": model,
+ "custom_llm_provider": "openai",
+ }
+
+ # Extract user information for tracking
+ passthrough_logging_payload: Optional[
+ PassthroughStandardLoggingPayload
+ ] = litellm_logging_obj.model_call_details.get(
+ "passthrough_logging_payload"
+ )
+ if passthrough_logging_payload:
+ user = OpenAIPassthroughLoggingHandler._get_user_from_metadata(
+ passthrough_logging_payload=passthrough_logging_payload,
+ )
+ if user:
+ kwargs.setdefault("litellm_params", {})
+ kwargs["litellm_params"].update(
+ {"proxy_server_request": {"body": {"user": user}}}
+ )
+
+ # Create standard logging object
+ get_standard_logging_object_payload(
+ kwargs=kwargs,
+ init_response_obj=complete_response,
+ start_time=start_time,
+ end_time=end_time,
+ logging_obj=litellm_logging_obj,
+ status="success",
+ )
+
+ # Update logging object with cost information
+ litellm_logging_obj.model_call_details["model"] = model
+ litellm_logging_obj.model_call_details["custom_llm_provider"] = "openai"
+ litellm_logging_obj.model_call_details["response_cost"] = response_cost
+
+ verbose_proxy_logger.debug(
+ f"OpenAI streaming passthrough cost tracking - Model: {model}, Cost: ${response_cost:.6f}"
+ )
+
+ return {
+ "result": complete_response,
+ "kwargs": kwargs,
+ }
+
+ except Exception as e:
+ verbose_proxy_logger.error(
+ f"Error in OpenAI streaming passthrough cost tracking: {str(e)}"
+ )
+ return {
+ "result": None,
+ "kwargs": {},
+ }
diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py
index adedcaf781d..fccc65b8bd4 100644
--- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py
+++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py
@@ -314,6 +314,12 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
return EndpointType.VERTEX_AI
elif parsed_url.hostname == "api.anthropic.com":
return EndpointType.ANTHROPIC
+ elif (
+ parsed_url.hostname == "api.openai.com"
+ or parsed_url.hostname == "openai.azure.com"
+ or (parsed_url.hostname and "openai.com" in parsed_url.hostname)
+ ):
+ return EndpointType.OPENAI
return EndpointType.GENERIC
@staticmethod
@@ -415,10 +421,10 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
for field_name, field_value in form_data.items():
if isinstance(field_value, (StarletteUploadFile, UploadFile)):
- files[field_name] = (
- await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
- upload_file=field_value
- )
+ files[
+ field_name
+ ] = await HttpPassThroughEndpointHelpers._build_request_files_from_upload_file(
+ upload_file=field_value
)
else:
form_data_dict[field_name] = field_value
@@ -497,9 +503,9 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
"passthrough_logging_payload": passthrough_logging_payload,
}
- logging_obj.model_call_details["passthrough_logging_payload"] = (
- passthrough_logging_payload
- )
+ logging_obj.model_call_details[
+ "passthrough_logging_payload"
+ ] = passthrough_logging_payload
return kwargs
@@ -531,10 +537,10 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
subpath = subpath[1:]
return base_target + subpath
-
+
@staticmethod
def _update_stream_param_based_on_request_body(
- parsed_body: dict,
+ parsed_body: dict,
stream: Optional[bool] = None,
) -> Optional[bool]:
"""
@@ -699,9 +705,11 @@ async def pass_through_request( # noqa: PLR0915
"headers": headers,
},
)
- stream = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
- parsed_body=_parsed_body,
- stream=stream,
+ stream = (
+ HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
+ parsed_body=_parsed_body,
+ stream=stream,
+ )
)
if stream:
diff --git a/litellm/proxy/pass_through_endpoints/streaming_handler.py b/litellm/proxy/pass_through_endpoints/streaming_handler.py
index 08b49bac383..2d5b0a686ce 100644
--- a/litellm/proxy/pass_through_endpoints/streaming_handler.py
+++ b/litellm/proxy/pass_through_endpoints/streaming_handler.py
@@ -14,6 +14,9 @@ from litellm.types.utils import StandardPassThroughResponseObject
from .llm_provider_handlers.anthropic_passthrough_logging_handler import (
AnthropicPassthroughLoggingHandler,
)
+from .llm_provider_handlers.openai_passthrough_logging_handler import (
+ OpenAIPassthroughLoggingHandler,
+)
from .llm_provider_handlers.vertex_passthrough_logging_handler import (
VertexPassthroughLoggingHandler,
)
@@ -78,6 +81,7 @@ class PassThroughStreamingHandler:
Supported endpoint types:
- Anthropic
- Vertex AI
+ - OpenAI
"""
all_chunks = PassThroughStreamingHandler._convert_raw_bytes_to_str_lines(
raw_bytes
@@ -119,6 +123,23 @@ class PassThroughStreamingHandler:
vertex_passthrough_logging_handler_result["result"]
)
kwargs = vertex_passthrough_logging_handler_result["kwargs"]
+ elif endpoint_type == EndpointType.OPENAI:
+ openai_passthrough_logging_handler_result = (
+ OpenAIPassthroughLoggingHandler._handle_logging_openai_collected_chunks(
+ litellm_logging_obj=litellm_logging_obj,
+ passthrough_success_handler_obj=passthrough_success_handler_obj,
+ url_route=url_route,
+ request_body=request_body,
+ endpoint_type=endpoint_type,
+ start_time=start_time,
+ all_chunks=all_chunks,
+ end_time=end_time,
+ )
+ )
+ standard_logging_response_object = (
+ openai_passthrough_logging_handler_result["result"]
+ )
+ kwargs = openai_passthrough_logging_handler_result["kwargs"]
if standard_logging_response_object is None:
standard_logging_response_object = StandardPassThroughResponseObject(
diff --git a/litellm/proxy/pass_through_endpoints/success_handler.py b/litellm/proxy/pass_through_endpoints/success_handler.py
index ce576d5ac71..5b54f94f8b5 100644
--- a/litellm/proxy/pass_through_endpoints/success_handler.py
+++ b/litellm/proxy/pass_through_endpoints/success_handler.py
@@ -162,9 +162,32 @@ class PassThroughEndpointLogging:
cohere_passthrough_logging_handler_result["result"]
)
kwargs = cohere_passthrough_logging_handler_result["kwargs"]
- return_dict["standard_logging_response_object"] = (
- standard_logging_response_object
- )
+ elif self.is_openai_route(url_route):
+ from .llm_provider_handlers.openai_passthrough_logging_handler import (
+ OpenAIPassthroughLoggingHandler,
+ )
+
+ openai_passthrough_logging_handler_result = (
+ OpenAIPassthroughLoggingHandler.openai_passthrough_handler(
+ httpx_response=httpx_response,
+ response_body=response_body or {},
+ logging_obj=logging_obj,
+ url_route=url_route,
+ result=result,
+ start_time=start_time,
+ end_time=end_time,
+ cache_hit=cache_hit,
+ request_body=request_body,
+ **kwargs,
+ )
+ )
+ standard_logging_response_object = (
+ openai_passthrough_logging_handler_result["result"]
+ )
+ kwargs = openai_passthrough_logging_handler_result["kwargs"]
+ return_dict[
+ "standard_logging_response_object"
+ ] = standard_logging_response_object
return_dict["kwargs"] = kwargs
return return_dict
@@ -185,9 +208,9 @@ class PassThroughEndpointLogging:
standard_logging_response_object: Optional[
PassThroughEndpointLoggingResultValues
] = None
- logging_obj.model_call_details["passthrough_logging_payload"] = (
- passthrough_logging_payload
- )
+ logging_obj.model_call_details[
+ "passthrough_logging_payload"
+ ] = passthrough_logging_payload
if self.is_assemblyai_route(url_route):
if (
AssemblyAIPassthroughLoggingHandler._should_log_request(
@@ -286,6 +309,16 @@ class PassThroughEndpointLogging:
return True
return False
+ def is_openai_route(self, url_route: str):
+ """Check if the URL route is an OpenAI API route."""
+ if not url_route:
+ return False
+ parsed_url = urlparse(url_route)
+ return parsed_url.hostname and (
+ "api.openai.com" in parsed_url.hostname
+ or "openai.azure.com" in parsed_url.hostname
+ )
+
def _set_cost_per_request(
self,
logging_obj: LiteLLMLoggingObj,
@@ -305,8 +338,8 @@ class PassThroughEndpointLogging:
kwargs["response_cost"] = passthrough_logging_payload.get(
"cost_per_request"
)
- logging_obj.model_call_details["response_cost"] = (
- passthrough_logging_payload.get("cost_per_request")
- )
+ logging_obj.model_call_details[
+ "response_cost"
+ ] = passthrough_logging_payload.get("cost_per_request")
return kwargs
diff --git a/litellm/responses/main.py b/litellm/responses/main.py
index 11c88c526ac..47ecbcf02c0 100644
--- a/litellm/responses/main.py
+++ b/litellm/responses/main.py
@@ -1029,4 +1029,4 @@ def list_input_items(
original_exception=e,
completion_kwargs=local_vars,
extra_kwargs=kwargs,
- )
+ )
\ No newline at end of file
diff --git a/litellm/responses/utils.py b/litellm/responses/utils.py
index b66fd0d547e..336f2b5a947 100644
--- a/litellm/responses/utils.py
+++ b/litellm/responses/utils.py
@@ -379,4 +379,4 @@ class ResponseAPILoggingUtils:
prompt_tokens=prompt_tokens,
completion_tokens=completion_tokens,
total_tokens=prompt_tokens + completion_tokens,
- )
+ )
\ No newline at end of file
diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py
index 068c9db035b..5853f9de2aa 100644
--- a/litellm/types/llms/bedrock.py
+++ b/litellm/types/llms/bedrock.py
@@ -573,3 +573,84 @@ class AmazonDeepSeekR1StreamingResponse(TypedDict):
generation_token_count: int
stop_reason: Optional[str]
prompt_token_count: int
+
+
+################ Bedrock Batch Types #################
+
+
+class BedrockS3InputDataConfig(TypedDict):
+ """S3 input data configuration for Bedrock batch jobs."""
+ s3Uri: str
+
+
+class BedrockInputDataConfig(TypedDict):
+ """Input data configuration for Bedrock batch jobs."""
+ s3InputDataConfig: BedrockS3InputDataConfig
+
+
+class BedrockS3OutputDataConfig(TypedDict):
+ """S3 output data configuration for Bedrock batch jobs."""
+ s3Uri: str
+
+
+class BedrockOutputDataConfig(TypedDict):
+ """Output data configuration for Bedrock batch jobs."""
+ s3OutputDataConfig: BedrockS3OutputDataConfig
+
+
+class BedrockCreateBatchRequest(TypedDict, total=False):
+ """
+ Request structure for creating a Bedrock batch inference job.
+
+ Reference: https://docs.aws.amazon.com/bedrock/latest/APIReference/API_CreateModelInvocationJob.html
+ """
+ jobName: str
+ roleArn: str
+ modelId: str
+ inputDataConfig: BedrockInputDataConfig
+ outputDataConfig: BedrockOutputDataConfig
+ timeoutDurationInHours: Optional[int]
+ clientRequestToken: Optional[str]
+ tags: Optional[List[dict]]
+
+
+BedrockBatchJobStatus = Literal[
+ "Submitted",
+ "InProgress",
+ "Completed",
+ "Failed",
+ "Stopping",
+ "Stopped"
+]
+
+
+class BedrockCreateBatchResponse(TypedDict):
+ """
+ Response structure from creating a Bedrock batch inference job.
+
+ Reference: https://docs.aws.amazon.com/bedrock/latest/APIReference/API_CreateModelInvocationJob.html
+ """
+ jobArn: str
+ jobName: str
+ status: BedrockBatchJobStatus
+
+
+class BedrockGetBatchResponse(TypedDict, total=False):
+ """
+ Response structure from getting a Bedrock batch inference job.
+
+ Reference: https://docs.aws.amazon.com/bedrock/latest/APIReference/API_GetModelInvocationJob.html
+ """
+ jobArn: str
+ jobName: str
+ modelId: str
+ roleArn: str
+ status: BedrockBatchJobStatus
+ message: Optional[str]
+ submitTime: Optional[str]
+ lastModifiedTime: Optional[str]
+ endTime: Optional[str]
+ inputDataConfig: BedrockInputDataConfig
+ outputDataConfig: BedrockOutputDataConfig
+ timeoutDurationInHours: Optional[int]
+ clientRequestToken: Optional[str]
diff --git a/litellm/types/passthrough_endpoints/pass_through_endpoints.py b/litellm/types/passthrough_endpoints/pass_through_endpoints.py
index 8a91dd14ebd..39facdd8e6f 100644
--- a/litellm/types/passthrough_endpoints/pass_through_endpoints.py
+++ b/litellm/types/passthrough_endpoints/pass_through_endpoints.py
@@ -5,6 +5,7 @@ from typing import Optional, TypedDict
class EndpointType(str, Enum):
VERTEX_AI = "vertex-ai"
ANTHROPIC = "anthropic"
+ OPENAI = "openai"
GENERIC = "generic"
diff --git a/litellm/types/router.py b/litellm/types/router.py
index 864fdbf79b8..8bed2bdccda 100644
--- a/litellm/types/router.py
+++ b/litellm/types/router.py
@@ -216,6 +216,10 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
auto_router_default_model: Optional[str] = None
auto_router_embedding_model: Optional[str] = None
+ # Batch/File API Params
+ s3_bucket_name: Optional[str] = None
+ gcs_bucket_name: Optional[str] = None
+
def __init__(
self,
custom_llm_provider: Optional[str] = None,
@@ -265,6 +269,9 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
auto_router_config: Optional[str] = None,
auto_router_default_model: Optional[str] = None,
auto_router_embedding_model: Optional[str] = None,
+ # Batch/File API Params
+ s3_bucket_name: Optional[str] = None,
+ gcs_bucket_name: Optional[str] = None,
**params,
):
args = locals()
diff --git a/litellm/utils.py b/litellm/utils.py
index 405d4cb98e4..aa582fa751e 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -234,6 +234,7 @@ from litellm.llms.base_llm.base_utils import (
BaseLLMModelInfo,
type_to_response_format_param,
)
+from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig
from litellm.llms.base_llm.chat.transformation import BaseConfig
from litellm.llms.base_llm.completion.transformation import BaseTextCompletionConfig
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
@@ -2395,7 +2396,7 @@ def _should_drop_param(k, additional_drop_params) -> bool:
def _get_non_default_params(
- passed_params: dict, default_params: dict, additional_drop_params: Optional[bool]
+ passed_params: dict, default_params: dict, additional_drop_params: Optional[list]
) -> dict:
non_default_params = {}
for k, v in passed_params.items():
@@ -2509,7 +2510,7 @@ def get_optional_params_image_gen(
user: Optional[str] = None,
input_fidelity: Optional[str] = None,
custom_llm_provider: Optional[str] = None,
- additional_drop_params: Optional[bool] = None,
+ additional_drop_params: Optional[list] = None,
provider_config: Optional[BaseImageGenerationConfig] = None,
drop_params: Optional[bool] = None,
**kwargs,
@@ -2628,9 +2629,20 @@ def get_optional_params_image_gen(
) # Default to square if size not recognized
optional_params["aspectRatio"] = aspect_ratio
- for k in passed_params.keys():
- if k not in default_params.keys():
- optional_params[k] = passed_params[k]
+ openai_params: list[str] = list(default_params.keys())
+ if provider_config is not None:
+ supported_params = provider_config.get_supported_openai_params(
+ model=model or ""
+ )
+ openai_params = list(supported_params)
+
+ optional_params = add_provider_specific_params_to_optional_params(
+ optional_params=optional_params,
+ passed_params=passed_params,
+ custom_llm_provider=custom_llm_provider or "",
+ openai_params=openai_params,
+ additional_drop_params=additional_drop_params,
+ )
return optional_params
@@ -7287,6 +7299,20 @@ class ProviderConfigManager:
from litellm.llms.vertex_ai.files.transformation import VertexAIFilesConfig
return VertexAIFilesConfig()
+ elif LlmProviders.BEDROCK == provider:
+ from litellm.llms.bedrock.files.transformation import BedrockFilesConfig
+
+ return BedrockFilesConfig()
+ return None
+
+ @staticmethod
+ def get_provider_batches_config(
+ model: str,
+ provider: LlmProviders,
+ ) -> Optional[BaseBatchesConfig]:
+ if LlmProviders.BEDROCK == provider:
+ from litellm.llms.bedrock.batches.transformation import BedrockBatchesConfig
+ return BedrockBatchesConfig()
return None
@staticmethod
diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json
index 7cedffd93f3..a7586124509 100644
--- a/model_prices_and_context_window.json
+++ b/model_prices_and_context_window.json
@@ -6157,21 +6157,7 @@
"supports_tool_choice": true,
"source": "https://inference-docs.cerebras.ai/support/pricing"
},
- "cerebras/openai/gpt-oss-20b": {
- "max_tokens": 32768,
- "max_input_tokens": 131072,
- "max_output_tokens": 32768,
- "input_cost_per_token": 7e-08,
- "output_cost_per_token": 3e-07,
- "litellm_provider": "cerebras",
- "mode": "chat",
- "supports_function_calling": true,
- "supports_parallel_function_calling": true,
- "supports_response_schema": true,
- "supports_reasoning": true,
- "supports_tool_choice": true,
- "source": "https://inference-docs.cerebras.ai/support/pricing"
- },
+
"cerebras/openai/gpt-oss-120b": {
"max_tokens": 32768,
"max_input_tokens": 131072,
@@ -9498,6 +9484,48 @@
"source": "https://aistudio.google.com",
"supports_tool_choice": true
},
+ "gemini/veo-3.0-generate-preview": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.75,
+ "litellm_provider": "gemini",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
+ "gemini/veo-3.0-fast-generate-preview": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.40,
+ "litellm_provider": "gemini",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
+ "gemini/veo-2.0-generate-001": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.35,
+ "litellm_provider": "gemini",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
"vertex_ai/claude-opus-4-1": {
"max_tokens": 4096,
"max_input_tokens": 200000,
@@ -10315,6 +10343,48 @@
"mode": "image_generation",
"source": "https://cloud.google.com/vertex-ai/generative-ai/pricing"
},
+ "vertex_ai/veo-3.0-generate-preview": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.75,
+ "litellm_provider": "vertex_ai-video-models",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
+ "vertex_ai/veo-3.0-fast-generate-preview": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.40,
+ "litellm_provider": "vertex_ai-video-models",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
+ "vertex_ai/veo-2.0-generate-001": {
+ "max_tokens": 1024,
+ "max_input_tokens": 1024,
+ "output_cost_per_second": 0.35,
+ "litellm_provider": "vertex_ai-video-models",
+ "mode": "video_generation",
+ "supported_modalities": [
+ "text"
+ ],
+ "supported_output_modalities": [
+ "video"
+ ],
+ "source": "https://ai.google.dev/gemini-api/docs/video"
+ },
"text-embedding-004": {
"max_tokens": 2048,
"max_input_tokens": 2048,
diff --git a/pyproject.toml b/pyproject.toml
index 2b9650a1527..9f5d876cf2c 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm"
-version = "1.76.1"
+version = "1.76.3"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI"]
license = "MIT"
@@ -156,7 +156,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
-version = "1.76.1"
+version = "1.76.3"
version_files = [
"pyproject.toml:^version"
]
diff --git a/tests/batches_tests/bedrock_batch_completions.jsonl b/tests/batches_tests/bedrock_batch_completions.jsonl
new file mode 100644
index 00000000000..3037b1031ea
--- /dev/null
+++ b/tests/batches_tests/bedrock_batch_completions.jsonl
@@ -0,0 +1,3 @@
+{"custom_id": "request-1", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0", "messages": [{"role": "system", "content": "You are a helpful assistant."},{"role": "user", "content": "Hello world!"}],"max_tokens": 10}}
+{"custom_id": "request-2", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0", "messages": [{"role": "system", "content": "You are an unhelpful assistant."},{"role": "user", "content": "Hello world!"}],"max_tokens": 10}}
+
diff --git a/tests/batches_tests/test_bedrock_files_and_batches.py b/tests/batches_tests/test_bedrock_files_and_batches.py
new file mode 100644
index 00000000000..edf8c3c7745
--- /dev/null
+++ b/tests/batches_tests/test_bedrock_files_and_batches.py
@@ -0,0 +1,70 @@
+
+# What is this?
+## Unit Tests for OpenAI Batches API
+import asyncio
+import json
+import os
+import sys
+import traceback
+import tempfile
+from dotenv import load_dotenv
+
+load_dotenv()
+sys.path.insert(
+ 0, os.path.abspath("../..")
+) # Adds the parent directory to the system-path
+
+
+import pytest
+from typing import Optional
+import litellm
+
+
+@pytest.mark.asyncio()
+async def test_async_create_file():
+ """
+ 1. Create File for Batch completion
+ 2. Create Batch Request
+ 3. Retrieve the specific batch
+ """
+ litellm._turn_on_debug()
+ print("Testing async create batch")
+
+ file_name = "bedrock_batch_completions.jsonl"
+ _current_dir = os.path.dirname(os.path.abspath(__file__))
+ file_path = os.path.join(_current_dir, file_name)
+ file_obj = await litellm.acreate_file(
+ file=open(file_path, "rb"),
+ purpose="batch",
+ custom_llm_provider="bedrock",
+ s3_bucket_name="litellm-proxy",
+ )
+
+@pytest.mark.asyncio()
+async def test_async_file_and_batch():
+ """
+ Test file retrieval
+ """
+ litellm._turn_on_debug()
+ file_name = "bedrock_batch_completions.jsonl"
+ _current_dir = os.path.dirname(os.path.abspath(__file__))
+ file_path = os.path.join(_current_dir, file_name)
+ file_obj = await litellm.acreate_file(
+ file=open(file_path, "rb"),
+ purpose="batch",
+ custom_llm_provider="bedrock",
+ s3_bucket_name="litellm-proxy",
+ )
+ print("CREATED FILE RESPONSE=", file_obj)
+
+ # create batch
+ create_batch_response = await litellm.acreate_batch(
+ completion_window="24h",
+ endpoint="/v1/chat/completions",
+ input_file_id=file_obj.id,
+ metadata={"key1": "value1", "key2": "value2"},
+ custom_llm_provider="bedrock",
+ aws_batch_role_arn="arn:aws:iam::888602223428:role/service-role/AmazonBedrockExecutionRoleForAgents_BB9HNW6V4CV"
+ )
+ print("CREATED BATCH RESPONSE=", create_batch_response)
+
diff --git a/tests/llm_responses_api_testing/base_responses_api.py b/tests/llm_responses_api_testing/base_responses_api.py
index f8fbe53e03d..f254a197aaf 100644
--- a/tests/llm_responses_api_testing/base_responses_api.py
+++ b/tests/llm_responses_api_testing/base_responses_api.py
@@ -528,7 +528,6 @@ class BaseResponsesAPITest(ABC):
# Validate final response structure
validate_responses_api_response(final_response, final_chunk=True)
assert final_response.output is not None
- assert len(final_response.output) > 0
@pytest.mark.asyncio
async def test_text_format_to_text_conversion(self):
diff --git a/tests/llm_translation/test_groq.py b/tests/llm_translation/test_groq.py
index dd1ae0aed6d..9a9dae9ee89 100644
--- a/tests/llm_translation/test_groq.py
+++ b/tests/llm_translation/test_groq.py
@@ -1,6 +1,16 @@
-from base_llm_unit_tests import BaseLLMChatTest
+import os
+import sys
+import pytest
+
+# sys.path.insert(
+# 0, os.path.abspath("../..")
+# ) # Adds the parent directory to the system path
+
+from base_llm_unit_tests import BaseLLMChatTest
+from litellm.llms.groq.chat.transformation import GroqChatConfig
+
class TestGroq(BaseLLMChatTest):
def get_base_completion_call_args(self) -> dict:
return {
@@ -10,3 +20,9 @@ class TestGroq(BaseLLMChatTest):
def test_tool_call_no_arguments(self, tool_call_no_arguments):
"""Test that tool calls with no arguments is translated correctly. Relevant issue: https://github.com/BerriAI/litellm/issues/6833"""
pass
+
+ @pytest.mark.parametrize("model", ["groq/qwen/qwen3-32b", "groq/openai/gpt-oss-20b", "groq/openai/gpt-oss-120b"])
+ def test_reasoning_effort_in_supported_params(self, model):
+ """Test that reasoning_effort is in the list of supported parameters for Groq"""
+ supported_params = GroqChatConfig().get_supported_openai_params(model=model)
+ assert "reasoning_effort" in supported_params
diff --git a/tests/local_testing/test_amazing_vertex_completion.py b/tests/local_testing/test_amazing_vertex_completion.py
index 9b6fa868677..a27fe738c7f 100644
--- a/tests/local_testing/test_amazing_vertex_completion.py
+++ b/tests/local_testing/test_amazing_vertex_completion.py
@@ -839,7 +839,6 @@ from test_completion import response_format_tests
"model,region",
[
("vertex_ai/mistral-large-2411", "us-central1"),
- ("vertex_ai/mistral-nemo@2407", "us-central1"),
("vertex_ai/qwen/qwen3-coder-480b-a35b-instruct-maas", "us-south1"),
("vertex_ai/openai/gpt-oss-20b-maas", "us-central1"),
],
diff --git a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py
index 694a49159c0..da4218a9547 100644
--- a/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py
+++ b/tests/test_litellm/proxy/hooks/test_parallel_request_limiter_v3.py
@@ -934,3 +934,204 @@ async def test_team_member_rate_limits_v3():
assert team_member_descriptor["value"] == f"{_team_id}:{_user_id}", "Team member value should combine team_id and user_id"
assert team_member_descriptor["rate_limit"]["requests_per_unit"] == 10, "Team member RPM limit should be set"
assert team_member_descriptor["rate_limit"]["tokens_per_unit"] == 1000, "Team member TPM limit should be set"
+
+
+@pytest.mark.asyncio
+async def test_async_increment_tokens_with_ttl_preservation():
+ """
+ Test TTL preservation functionality for token increment operations.
+
+ This test verifies that:
+ 1. Keys are created with proper TTL on first increment
+ 2. TTL is preserved on subsequent increments (not reset)
+ 3. Both TTL and non-TTL operations work correctly in the same call
+
+ Environment variables required:
+ - REDIS_HOST: Redis server hostname
+ - REDIS_PORT: Redis server port
+ - REDIS_PASSWORD: Redis password (optional)
+
+ Test scenario:
+ 1. First call: Create keys with TTL=60s and TTL=None
+ 2. Wait 2 seconds
+ 3. Second call: Increment same keys
+ 4. Verify TTL decreased but wasn't reset to 60s
+ """
+ import os
+ import time
+ from litellm.caching.redis_cache import RedisCache
+ from litellm.types.caching import RedisPipelineIncrementOperation
+
+ # Skip test if Redis environment variables are not set
+ redis_host = os.getenv("REDIS_HOST")
+ redis_port = os.getenv("REDIS_PORT")
+ redis_password = os.getenv("REDIS_PASSWORD")
+
+ if not redis_host or not redis_port:
+ pytest.skip("Redis environment variables (REDIS_HOST, REDIS_PORT) not set")
+
+ # Setup Redis cache
+ redis_cache = RedisCache(
+ host=redis_host,
+ port=int(redis_port),
+ password=redis_password,
+ )
+
+ local_cache = DualCache(redis_cache=redis_cache)
+ parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
+ internal_usage_cache=InternalUsageCache(local_cache)
+ )
+
+ # Verify Redis connection is working
+ try:
+ await redis_cache.ping()
+ except Exception as e:
+ pytest.skip(f"Redis connection failed: {str(e)}")
+
+ # Test keys
+ test_key_with_ttl = "test_ttl_preservation:with_ttl"
+ test_key_without_ttl = "test_ttl_preservation:without_ttl"
+
+ try:
+ # Clean up any existing test keys
+ try:
+ await redis_cache.async_delete_cache(test_key_with_ttl)
+ await redis_cache.async_delete_cache(test_key_without_ttl)
+ except Exception:
+ # Keys might not exist, ignore cleanup errors
+ pass
+
+ # First increment: Create operations with mixed TTL scenarios
+ pipeline_operations_first = [
+ RedisPipelineIncrementOperation(
+ key=test_key_with_ttl,
+ increment_value=10.0,
+ ttl=60
+ ),
+ RedisPipelineIncrementOperation(
+ key=test_key_without_ttl,
+ increment_value=5.0,
+ ttl=None # No TTL
+ )
+ ]
+
+ # Execute first increment
+ await parallel_request_handler.async_increment_tokens_with_ttl_preservation(
+ pipeline_operations=pipeline_operations_first
+ )
+
+ # Verify keys exist and check initial TTL
+ ttl_after_first = await redis_cache.async_get_ttl(test_key_with_ttl)
+ value_after_first_with_ttl = await redis_cache.async_get_cache(test_key_with_ttl)
+ value_after_first_without_ttl = await redis_cache.async_get_cache(test_key_without_ttl)
+
+ assert value_after_first_with_ttl == 10.0, "First increment should set value to 10.0"
+ assert value_after_first_without_ttl == 5.0, "First increment should set value to 5.0"
+ assert ttl_after_first is not None and ttl_after_first > 0, "Key with TTL should have positive TTL after first increment"
+ assert ttl_after_first <= 60, "TTL should not exceed the set value"
+
+ # Check TTL for key without TTL (should be None, meaning no expiry)
+ ttl_no_ttl_key = await redis_cache.async_get_ttl(test_key_without_ttl)
+ assert ttl_no_ttl_key is None, "Key without TTL should have no expiry (None from async_get_ttl)"
+
+ # Wait a moment to ensure TTL decreases
+ await asyncio.sleep(2)
+
+ # Second increment: Same operations to test TTL preservation
+ pipeline_operations_second = [
+ RedisPipelineIncrementOperation(
+ key=test_key_with_ttl,
+ increment_value=15.0,
+ ttl=60 # Same TTL value
+ ),
+ RedisPipelineIncrementOperation(
+ key=test_key_without_ttl,
+ increment_value=7.0,
+ ttl=None # No TTL
+ )
+ ]
+
+ # Execute second increment
+ await parallel_request_handler.async_increment_tokens_with_ttl_preservation(
+ pipeline_operations=pipeline_operations_second
+ )
+
+ # Verify TTL preservation and value updates
+ ttl_after_second = await redis_cache.async_get_ttl(test_key_with_ttl)
+ value_after_second_with_ttl = await redis_cache.async_get_cache(test_key_with_ttl)
+ value_after_second_without_ttl = await redis_cache.async_get_cache(test_key_without_ttl)
+
+ assert value_after_second_with_ttl == 25.0, "Second increment should update value to 25.0"
+ assert value_after_second_without_ttl == 12.0, "Second increment should update value to 12.0"
+
+ # Critical test: TTL should be preserved (not reset to 60)
+ assert ttl_after_second is not None, "TTL should still exist"
+ assert ttl_after_second < ttl_after_first, "TTL should have decreased (not been reset)"
+ assert ttl_after_second > 0, "TTL should still be positive"
+
+ # TTL should not be close to the original 60 seconds (proving it wasn't reset)
+ assert ttl_after_second < 59, "TTL should be significantly less than original, proving preservation"
+
+ # Key without TTL should still have no expiry
+ ttl_no_ttl_key_after_second = await redis_cache.async_get_ttl(test_key_without_ttl)
+ assert ttl_no_ttl_key_after_second is None, "Key without TTL should still have no expiry"
+
+ finally:
+ # Clean up test keys
+ try:
+ await redis_cache.async_delete_cache(test_key_with_ttl)
+ await redis_cache.async_delete_cache(test_key_without_ttl)
+ except Exception:
+ # Ignore cleanup errors
+ pass
+
+ # Properly close Redis connections to prevent warnings
+ try:
+ await redis_cache.disconnect()
+ except Exception:
+ # Ignore disconnect errors
+ pass
+
+
+@pytest.mark.asyncio
+async def test_async_increment_tokens_fallback_behavior():
+ """
+ Test fallback behavior when Lua script is not available.
+ """
+ from litellm.types.caching import RedisPipelineIncrementOperation
+
+ local_cache = DualCache()
+ parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
+ internal_usage_cache=InternalUsageCache(local_cache)
+ )
+
+ # Mock the token_increment_script to None to simulate unavailable script
+ parallel_request_handler.token_increment_script = None
+
+ # Mock the fallback method
+ fallback_called = False
+ original_method = parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline
+
+ async def mock_fallback(*args, **kwargs):
+ nonlocal fallback_called
+ fallback_called = True
+ return await original_method(*args, **kwargs)
+
+ parallel_request_handler.internal_usage_cache.dual_cache.async_increment_cache_pipeline = mock_fallback
+
+ # Test operations
+ pipeline_operations = [
+ RedisPipelineIncrementOperation(
+ key="test_fallback_key",
+ increment_value=10.0,
+ ttl=60
+ )
+ ]
+
+ # Execute increment
+ await parallel_request_handler.async_increment_tokens_with_ttl_preservation(
+ pipeline_operations=pipeline_operations
+ )
+
+ # Verify fallback was called
+ assert fallback_called, "Fallback method should be called when Lua script is not available"
diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py
index f2b374657e6..be24444afaf 100644
--- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py
+++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_transformations.py
@@ -224,10 +224,10 @@ class TestScimTransformations:
result = ScimTransformations._get_scim_member_value(member_with_email)
assert result == member_with_email.user_email
- # Member without email
+ # Member without email should fall back to user_id
member_without_email = Member(user_id="user-456", user_email=None, role="user")
result = ScimTransformations._get_scim_member_value(member_without_email)
- assert result == ScimTransformations.DEFAULT_SCIM_MEMBER_VALUE
+ assert result == member_without_email.user_id
class TestSCIMPatchOperations:
diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py
index 208c8774675..959275787c8 100644
--- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py
@@ -10,10 +10,13 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import (
create_user,
get_service_provider_config,
patch_user,
+ update_group,
update_user,
)
from litellm.types.proxy.management_endpoints.scim_v2 import (
SCIMFeature,
+ SCIMGroup,
+ SCIMMember,
SCIMPatchOp,
SCIMPatchOperation,
SCIMServiceProviderConfig,
@@ -678,4 +681,233 @@ async def test_update_group_metadata_serialization_issue(mocker):
parsed_metadata = json.loads(metadata)
assert "existing_key" in parsed_metadata
assert "scim_data" in parsed_metadata
- assert parsed_metadata["existing_key"] == "existing_value"
\ No newline at end of file
+
+
+@pytest.mark.asyncio
+async def test_team_membership_management(mocker):
+ """
+ Test that team membership changes work correctly:
+ - Adding members to team
+ - Removing members from team
+ - members_with_roles is used as source of truth
+ """
+ from litellm.proxy._types import Member
+ from litellm.proxy.management_endpoints.scim.scim_v2 import (
+ _get_team_member_user_ids_from_team,
+ _handle_group_membership_changes,
+ patch_team_membership,
+ )
+
+ # Mock team with members_with_roles as source of truth
+ mock_team = mocker.MagicMock()
+ mock_team.members_with_roles = [
+ Member(user_id="user1", role="user"),
+ Member(user_id="user2", role="user")
+ ]
+ mock_team.members = ["user1", "user2", "user3"] # This should be ignored
+
+ # Test that members_with_roles is source of truth
+ member_ids = await _get_team_member_user_ids_from_team(mock_team)
+ assert set(member_ids) == {"user1", "user2"}
+ assert "user3" not in member_ids # Should not be included even though in members
+
+ # Mock patch_team_membership function
+ mock_patch_team_membership = mocker.patch(
+ "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership",
+ AsyncMock()
+ )
+
+ # Test adding and removing members
+ group_id = "test-group-id"
+ current_members = {"user1", "user2"}
+ final_members = {"user2", "user3", "user4"} # Remove user1, add user3 and user4
+
+ await _handle_group_membership_changes(
+ group_id=group_id,
+ current_members=current_members,
+ final_members=final_members
+ )
+
+ # Verify patch_team_membership was called correctly
+ assert mock_patch_team_membership.call_count == 3
+
+ # Check calls for adding members
+ add_calls = [call for call in mock_patch_team_membership.call_args_list
+ if call[1]["teams_ids_to_add_user_to"] == [group_id]]
+ assert len(add_calls) == 2 # user3 and user4
+
+ add_user_ids = {call[1]["user_id"] for call in add_calls}
+ assert add_user_ids == {"user3", "user4"}
+
+ # Check calls for removing members
+ remove_calls = [call for call in mock_patch_team_membership.call_args_list
+ if call[1]["teams_ids_to_remove_user_from"] == [group_id]]
+ assert len(remove_calls) == 1 # user1
+
+ remove_user_ids = {call[1]["user_id"] for call in remove_calls}
+ assert remove_user_ids == {"user1"}
+
+ # Verify all calls have correct structure
+ for call in mock_patch_team_membership.call_args_list:
+ assert "user_id" in call[1]
+ assert "teams_ids_to_add_user_to" in call[1]
+ assert "teams_ids_to_remove_user_from" in call[1]
+ # Each call should either add OR remove, not both
+ add_teams = call[1]["teams_ids_to_add_user_to"]
+ remove_teams = call[1]["teams_ids_to_remove_user_from"]
+ assert (len(add_teams) > 0) != (len(remove_teams) > 0) # XOR - one should be empty
+
+
+@pytest.mark.asyncio
+async def test_update_group_e2e(mocker):
+ """
+ End-to-end test for update_group endpoint:
+ - Updates group metadata (displayName)
+ - Handles complete member replacement (add/remove members)
+ - Verifies members_with_roles is updated as source of truth
+ - Tests the full flow from SCIM request to database updates
+ """
+ from litellm.proxy._types import LiteLLM_TeamTable, Member
+ from litellm.proxy.management_endpoints.scim.scim_transformations import (
+ ScimTransformations,
+ )
+ from litellm.proxy.utils import safe_dumps
+
+ # Setup test data
+ group_id = "test-team-123"
+
+ # Mock existing team in database
+ existing_team = LiteLLM_TeamTable(
+ team_id=group_id,
+ team_alias="Old Team Name",
+ members=["user1", "user2"], # This should be ignored
+ members_with_roles=[
+ Member(user_id="user1", role="user"),
+ Member(user_id="user2", role="user")
+ ],
+ metadata={"existing_key": "existing_value"}
+ )
+
+ # Mock updated SCIM group request
+ scim_group_update = SCIMGroup(
+ schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
+ id=group_id,
+ displayName="Updated Team Name",
+ members=[
+ SCIMMember(value="user2", display="User Two"), # Keep user2
+ SCIMMember(value="user3", display="User Three"), # Add user3
+ SCIMMember(value="user4", display="User Four") # Add user4
+ ]
+ )
+
+ # Mock prisma client
+ mock_prisma_client = mocker.MagicMock()
+ mock_prisma_client.db = mocker.MagicMock()
+ mock_prisma_client.db.litellm_teamtable = mocker.MagicMock()
+ mock_prisma_client.db.litellm_usertable = mocker.MagicMock()
+
+ # Mock database operations
+ mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=existing_team)
+
+ # Mock the updated team that gets returned from database
+ updated_team = LiteLLM_TeamTable(
+ team_id=group_id,
+ team_alias="Updated Team Name",
+ members=["user2", "user3", "user4"],
+ members_with_roles=[
+ Member(user_id="user2", role="user"),
+ Member(user_id="user3", role="user"),
+ Member(user_id="user4", role="user")
+ ],
+ metadata={
+ "existing_key": "existing_value",
+ "scim_data": scim_group_update.model_dump()
+ }
+ )
+ mock_prisma_client.db.litellm_teamtable.update = AsyncMock(return_value=updated_team)
+
+ # Mock user validation (all users exist)
+ mock_user = mocker.MagicMock()
+ mock_user.user_id = "test-user"
+ mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=mock_user)
+
+ # Mock dependencies
+ mocker.patch(
+ "litellm.proxy.management_endpoints.scim.scim_v2._get_prisma_client_or_raise_exception",
+ AsyncMock(return_value=mock_prisma_client)
+ )
+
+ # Mock patch_team_membership to track membership changes
+ mock_patch_team_membership = mocker.patch(
+ "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership",
+ AsyncMock()
+ )
+
+ # Mock SCIM transformation
+ expected_scim_response = SCIMGroup(
+ schemas=["urn:ietf:params:scim:schemas:core:2.0:Group"],
+ id=group_id,
+ displayName="Updated Team Name",
+ members=[
+ SCIMMember(value="user2", display="user2"),
+ SCIMMember(value="user3", display="user3"),
+ SCIMMember(value="user4", display="user4")
+ ]
+ )
+ mocker.patch.object(
+ ScimTransformations,
+ "transform_litellm_team_to_scim_group",
+ AsyncMock(return_value=expected_scim_response)
+ )
+
+ # Execute the update_group function
+ result = await update_group(group_id=group_id, group=scim_group_update)
+
+ # Verify database update was called with correct data
+ mock_prisma_client.db.litellm_teamtable.update.assert_called_once()
+ update_call_args = mock_prisma_client.db.litellm_teamtable.update.call_args
+
+ # Check the update parameters
+ assert update_call_args[1]["where"]["team_id"] == group_id
+ update_data = update_call_args[1]["data"]
+ assert update_data["team_alias"] == "Updated Team Name"
+
+ # Verify metadata includes both existing data and SCIM data
+ metadata_str = update_data["metadata"]
+ import json
+ metadata = json.loads(metadata_str)
+ assert metadata["existing_key"] == "existing_value"
+ assert "scim_data" in metadata
+ assert metadata["scim_data"]["displayName"] == "Updated Team Name"
+
+ # Verify team membership changes were handled correctly
+ assert mock_patch_team_membership.call_count == 3 # Remove user1, add user3, add user4
+
+ # Check membership changes
+ call_args_list = mock_patch_team_membership.call_args_list
+
+ # Find remove operation (user1)
+ remove_calls = [call for call in call_args_list
+ if call[1]["teams_ids_to_remove_user_from"] == [group_id]]
+ assert len(remove_calls) == 1
+ assert remove_calls[0][1]["user_id"] == "user1"
+ assert remove_calls[0][1]["teams_ids_to_add_user_to"] == []
+
+ # Find add operations (user3, user4)
+ add_calls = [call for call in call_args_list
+ if call[1]["teams_ids_to_add_user_to"] == [group_id]]
+ assert len(add_calls) == 2
+ add_user_ids = {call[1]["user_id"] for call in add_calls}
+ assert add_user_ids == {"user3", "user4"}
+
+ # Verify all add calls have empty remove lists
+ for call in add_calls:
+ assert call[1]["teams_ids_to_remove_user_from"] == []
+
+ # Verify the response
+ assert result.id == group_id
+ assert result.displayName == "Updated Team Name"
+ assert len(result.members) == 3
+
+ # Verify SCIM transformation was called with updated team
+ ScimTransformations.transform_litellm_team_to_scim_group.assert_called_once_with(updated_team)
\ No newline at end of file
diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py
index e0102f8cd7a..893e5767ecd 100644
--- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py
+++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py
@@ -118,9 +118,9 @@ async def test_key_token_handling(monkeypatch):
@pytest.mark.asyncio
-async def test_budget_reset_at_first_of_month(monkeypatch):
+async def test_budget_reset_and_expires_at_first_of_month(monkeypatch):
"""
- Test that when budget_duration is "1mo", budget_reset_at is set to first of next month
+ Test that when budget_duration, duration, and key_budget_duration are "1mo", budget_reset_at and expires are set to first of next month
"""
mock_prisma_client = AsyncMock()
mock_insert_data = AsyncMock(
@@ -152,10 +152,12 @@ async def test_budget_reset_at_first_of_month(monkeypatch):
# Use monkeypatch to set the prisma_client
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client)
- # Test key generation with budget_duration="1mo"
+ # Test key generation with budget_duration="1mo", duration="1mo", key_budget_duration="1mo"
response = await generate_key_helper_fn(
request_type="user",
budget_duration="1mo",
+ duration="1mo",
+ key_budget_duration="1mo",
user_id="test_user",
)
@@ -171,17 +173,17 @@ async def test_budget_reset_at_first_of_month(monkeypatch):
expected_month = now.month + 1
expected_year = now.year
- # Parse the response date
- response_date = response["budget_reset_at"]
-
- # Verify budget_reset_at is set to first of next month
- assert (
- response_date.year == expected_year
- ), f"Expected year {expected_year}, got {response_date.year}"
- assert (
- response_date.month == expected_month
- ), f"Expected month {expected_month}, got {response_date.month}"
- assert response_date.day == 1, f"Expected day 1, got {response_date.day}"
+ # Verify budget_reset_at, expires is set to first of next month
+ for key in ["budget_reset_at", "expires"]:
+ response_date = response.get(key)
+ assert response_date is not None, f"{key} not found in response"
+ assert (
+ response_date.year == expected_year
+ ), f"Expected year {expected_year}, got {response_date.year} for {key}"
+ assert (
+ response_date.month == expected_month
+ ), f"Expected month {expected_month}, got {response_date.month} for {key}"
+ assert response_date.day == 1, f"Expected day 1, got {response_date.day} for {key}"
@pytest.mark.asyncio
diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py
new file mode 100644
index 00000000000..6d5e80910ba
--- /dev/null
+++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_openai_passthrough_logging_handler.py
@@ -0,0 +1,451 @@
+import json
+import os
+import sys
+from datetime import datetime
+from typing import Any, Dict, List
+from unittest.mock import MagicMock, patch, AsyncMock
+
+import pytest
+import httpx
+
+sys.path.insert(
+ 0, os.path.abspath("../../..")
+) # Adds the parent directory to the system path
+
+from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
+from litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler import (
+ OpenAIPassthroughLoggingHandler,
+)
+from litellm.proxy.pass_through_endpoints.success_handler import (
+ PassThroughEndpointLogging,
+)
+from litellm.types.passthrough_endpoints.pass_through_endpoints import (
+ PassthroughStandardLoggingPayload,
+)
+
+
+class TestOpenAIPassthroughLoggingHandler:
+ """Test the OpenAI passthrough logging handler for cost tracking."""
+
+ def setup_method(self):
+ """Set up test fixtures"""
+ self.start_time = datetime.now()
+ self.end_time = datetime.now()
+ self.handler = OpenAIPassthroughLoggingHandler()
+
+ # Mock OpenAI chat completions response
+ self.mock_openai_response = {
+ "id": "chatcmpl-123",
+ "object": "chat.completion",
+ "created": 1677652288,
+ "model": "gpt-4o-2024-08-06",
+ "choices": [
+ {
+ "index": 0,
+ "message": {
+ "role": "assistant",
+ "content": "Hello! How can I help you today?"
+ },
+ "finish_reason": "stop"
+ }
+ ],
+ "usage": {
+ "prompt_tokens": 20,
+ "completion_tokens": 15,
+ "total_tokens": 35
+ }
+ }
+
+ def _create_mock_logging_obj(self) -> LiteLLMLoggingObj:
+ """Create a mock logging object"""
+ mock_logging_obj = MagicMock()
+ mock_logging_obj.model_call_details = {}
+ return mock_logging_obj
+
+ def _create_mock_httpx_response(self, response_data: dict = None) -> httpx.Response:
+ """Create a mock httpx response"""
+ if response_data is None:
+ response_data = self.mock_openai_response
+
+ mock_response = MagicMock(spec=httpx.Response)
+ mock_response.status_code = 200
+ mock_response.text = json.dumps(response_data)
+ mock_response.json.return_value = response_data
+ mock_response.headers = {"content-type": "application/json"}
+ return mock_response
+
+ def _create_passthrough_logging_payload(self, user: str = "test_user") -> PassthroughStandardLoggingPayload:
+ """Create a mock passthrough logging payload"""
+ return PassthroughStandardLoggingPayload(
+ url="https://api.openai.com/v1/chat/completions",
+ request_body={"model": "gpt-4o", "messages": [{"role": "user", "content": "Hello"}]},
+ request_method="POST",
+ )
+
+ def test_llm_provider_name(self):
+ """Test that the handler returns the correct provider name"""
+ assert self.handler.llm_provider_name == "openai"
+
+ def test_get_provider_config(self):
+ """Test that the handler returns an OpenAI config"""
+ config = OpenAIPassthroughLoggingHandler.get_provider_config(model="gpt-4o")
+ assert config is not None
+ # Verify it's an OpenAI config by checking if it has the expected methods
+ assert hasattr(config, 'transform_response')
+
+ def test_is_openai_chat_completions_route(self):
+ """Test OpenAI chat completions route detection"""
+ # Positive cases
+ assert OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route("https://api.openai.com/v1/chat/completions") == True
+ assert OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route("https://openai.azure.com/v1/chat/completions") == True
+
+ # Negative cases
+ assert OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route("https://api.openai.com/v1/models") == False
+ assert OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route("http://localhost:4000/openai/v1/chat/completions") == False
+ assert OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route("https://api.anthropic.com/v1/messages") == False
+ assert OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route("") == False
+
+ @patch('litellm.completion_cost')
+ @patch('litellm.litellm_core_utils.litellm_logging.get_standard_logging_object_payload')
+ def test_openai_passthrough_handler_success(self, mock_get_standard_logging, mock_completion_cost):
+ """Test successful cost tracking for OpenAI chat completions"""
+ # Arrange
+ mock_completion_cost.return_value = 0.000045
+ mock_get_standard_logging.return_value = {"test": "logging_payload"}
+
+ mock_httpx_response = self._create_mock_httpx_response()
+ mock_logging_obj = self._create_mock_logging_obj()
+ passthrough_payload = self._create_passthrough_logging_payload()
+
+ kwargs = {
+ "passthrough_logging_payload": passthrough_payload,
+ "model": "gpt-4o",
+ }
+
+ # Act
+ result = OpenAIPassthroughLoggingHandler.openai_passthrough_handler(
+ httpx_response=mock_httpx_response,
+ response_body=self.mock_openai_response,
+ logging_obj=mock_logging_obj,
+ url_route="https://api.openai.com/v1/chat/completions",
+ result="",
+ start_time=self.start_time,
+ end_time=self.end_time,
+ cache_hit=False,
+ request_body={"model": "gpt-4o", "messages": [{"role": "user", "content": "Hello"}]},
+ **kwargs
+ )
+
+ # Assert
+ assert result is not None
+ assert "result" in result
+ assert "kwargs" in result
+ assert result["kwargs"]["response_cost"] == 0.000045
+ assert result["kwargs"]["model"] == "gpt-4o"
+ assert result["kwargs"]["custom_llm_provider"] == "openai"
+
+ # Verify cost calculation was called
+ mock_completion_cost.assert_called_once()
+
+ # Verify logging object was updated
+ assert mock_logging_obj.model_call_details["response_cost"] == 0.000045
+ assert mock_logging_obj.model_call_details["model"] == "gpt-4o"
+ assert mock_logging_obj.model_call_details["custom_llm_provider"] == "openai"
+
+ @patch('litellm.completion_cost')
+ def test_openai_passthrough_handler_non_chat_completions(self, mock_completion_cost):
+ """Test that non-chat-completions routes fall back to base handler"""
+ # Arrange
+ mock_httpx_response = self._create_mock_httpx_response()
+ mock_logging_obj = self._create_mock_logging_obj()
+ passthrough_payload = self._create_passthrough_logging_payload()
+
+ kwargs = {
+ "passthrough_logging_payload": passthrough_payload,
+ "model": "gpt-4o",
+ }
+
+ # Act - Use a non-chat-completions route
+ result = OpenAIPassthroughLoggingHandler.openai_passthrough_handler(
+ httpx_response=mock_httpx_response,
+ response_body={"id": "file-123", "object": "file"},
+ logging_obj=mock_logging_obj,
+ url_route="https://api.openai.com/v1/files",
+ result="",
+ start_time=self.start_time,
+ end_time=self.end_time,
+ cache_hit=False,
+ request_body={"purpose": "fine-tune"},
+ **kwargs
+ )
+
+ # Assert - Should fall back to base handler for non-chat-completions
+ assert result is not None
+ assert "result" in result
+ assert "kwargs" in result
+ # Cost calculation may be called by the base handler fallback
+ # The important thing is that our specific OpenAI handler logic didn't run
+
+ @patch('litellm.completion_cost')
+ @patch('litellm.litellm_core_utils.litellm_logging.get_standard_logging_object_payload')
+ def test_openai_passthrough_handler_with_user_tracking(self, mock_get_standard_logging, mock_completion_cost):
+ """Test cost tracking with user information"""
+ # Arrange
+ mock_completion_cost.return_value = 0.000123
+ mock_get_standard_logging.return_value = {"test": "logging_payload"}
+
+ mock_httpx_response = self._create_mock_httpx_response()
+ mock_logging_obj = self._create_mock_logging_obj()
+
+ # Create payload with user information
+ passthrough_payload = PassthroughStandardLoggingPayload(
+ url="https://api.openai.com/v1/chat/completions",
+ request_body={
+ "model": "gpt-4o",
+ "messages": [{"role": "user", "content": "Hello"}],
+ "user": "test_user_123"
+ },
+ request_method="POST",
+ )
+
+ kwargs = {
+ "passthrough_logging_payload": passthrough_payload,
+ "model": "gpt-4o",
+ }
+
+ # Act
+ result = OpenAIPassthroughLoggingHandler.openai_passthrough_handler(
+ httpx_response=mock_httpx_response,
+ response_body=self.mock_openai_response,
+ logging_obj=mock_logging_obj,
+ url_route="https://api.openai.com/v1/chat/completions",
+ result="",
+ start_time=self.start_time,
+ end_time=self.end_time,
+ cache_hit=False,
+ request_body={"model": "gpt-4o", "messages": [{"role": "user", "content": "Hello"}], "user": "test_user_123"},
+ **kwargs
+ )
+
+ # Assert
+ assert result is not None
+ assert "result" in result
+ assert "kwargs" in result
+ assert result["kwargs"]["response_cost"] == 0.000123
+
+ # Verify user information is included in litellm_params
+ assert "litellm_params" in result["kwargs"]
+ assert "proxy_server_request" in result["kwargs"]["litellm_params"]
+ assert "body" in result["kwargs"]["litellm_params"]["proxy_server_request"]
+ assert result["kwargs"]["litellm_params"]["proxy_server_request"]["body"]["user"] == "test_user_123"
+
+ @patch('litellm.completion_cost')
+ def test_openai_passthrough_handler_cost_calculation_error(self, mock_completion_cost):
+ """Test error handling in cost calculation"""
+ # Arrange
+ mock_completion_cost.side_effect = Exception("Cost calculation failed")
+
+ mock_httpx_response = self._create_mock_httpx_response()
+ mock_logging_obj = self._create_mock_logging_obj()
+ passthrough_payload = self._create_passthrough_logging_payload()
+
+ kwargs = {
+ "passthrough_logging_payload": passthrough_payload,
+ "model": "gpt-4o",
+ }
+
+ # Act
+ result = OpenAIPassthroughLoggingHandler.openai_passthrough_handler(
+ httpx_response=mock_httpx_response,
+ response_body=self.mock_openai_response,
+ logging_obj=mock_logging_obj,
+ url_route="https://api.openai.com/v1/chat/completions",
+ result="",
+ start_time=self.start_time,
+ end_time=self.end_time,
+ cache_hit=False,
+ request_body={"model": "gpt-4o", "messages": [{"role": "user", "content": "Hello"}]},
+ **kwargs
+ )
+
+ # Assert - Should fall back to base handler when cost calculation fails
+ assert result is not None
+ assert "result" in result
+ assert "kwargs" in result
+
+ def test_build_complete_streaming_response(self):
+ """Test the streaming response builder (placeholder implementation)"""
+ # This is a placeholder method that returns None for now
+ result = self.handler._build_complete_streaming_response(
+ all_chunks=["chunk1", "chunk2"],
+ litellm_logging_obj=self._create_mock_logging_obj(),
+ model="gpt-4o",
+ )
+
+ assert result is None # Placeholder implementation
+
+ @patch('litellm.completion_cost')
+ @patch('litellm.litellm_core_utils.litellm_logging.get_standard_logging_object_payload')
+ def test_different_models_cost_tracking(self, mock_get_standard_logging, mock_completion_cost):
+ """Test cost tracking for different OpenAI models"""
+ # Arrange
+ mock_get_standard_logging.return_value = {"test": "logging_payload"}
+
+ test_cases = [
+ ("gpt-4o", 0.000045),
+ ("gpt-4o-mini", 0.000015),
+ ("gpt-3.5-turbo", 0.000002),
+ ]
+
+ for model, expected_cost in test_cases:
+ mock_completion_cost.return_value = expected_cost
+
+ mock_httpx_response = self._create_mock_httpx_response()
+ mock_httpx_response.json.return_value = {
+ **self.mock_openai_response,
+ "model": model
+ }
+
+ mock_logging_obj = self._create_mock_logging_obj()
+ passthrough_payload = self._create_passthrough_logging_payload()
+
+ kwargs = {
+ "passthrough_logging_payload": passthrough_payload,
+ "model": model,
+ }
+
+ # Act
+ result = OpenAIPassthroughLoggingHandler.openai_passthrough_handler(
+ httpx_response=mock_httpx_response,
+ response_body={**self.mock_openai_response, "model": model},
+ logging_obj=mock_logging_obj,
+ url_route="https://api.openai.com/v1/chat/completions",
+ result="",
+ start_time=self.start_time,
+ end_time=self.end_time,
+ cache_hit=False,
+ request_body={"model": model, "messages": [{"role": "user", "content": "Hello"}]},
+ **kwargs
+ )
+
+ # Assert
+ assert result is not None
+ assert "result" in result
+ assert "kwargs" in result
+ assert result["kwargs"]["response_cost"] == expected_cost
+ assert result["kwargs"]["model"] == model
+ assert result["kwargs"]["custom_llm_provider"] == "openai"
+
+ def test_static_methods(self):
+ """Test that static methods work correctly"""
+ # Test static method calls
+ assert OpenAIPassthroughLoggingHandler.is_openai_chat_completions_route("https://api.openai.com/v1/chat/completions") == True
+ assert OpenAIPassthroughLoggingHandler.get_provider_config("gpt-4o") is not None
+
+
+class TestOpenAIPassthroughIntegration:
+ """Integration tests for OpenAI passthrough cost tracking"""
+
+ def setup_method(self):
+ """Set up test fixtures"""
+ self.handler = PassThroughEndpointLogging()
+
+ def test_is_openai_route_detection(self):
+ """Test OpenAI route detection in the main success handler"""
+ # Positive cases
+ assert self.handler.is_openai_route("https://api.openai.com/v1/chat/completions") == True
+ assert self.handler.is_openai_route("https://openai.azure.com/v1/chat/completions") == True
+ assert self.handler.is_openai_route("https://api.openai.com/v1/models") == True
+
+ # Negative cases
+ assert self.handler.is_openai_route("http://localhost:4000/openai/v1/chat/completions") == False
+ assert self.handler.is_openai_route("https://api.anthropic.com/v1/messages") == False
+ assert self.handler.is_openai_route("https://api.assemblyai.com/v2/transcript") == False
+ assert self.handler.is_openai_route("") == False
+
+ @patch('litellm.proxy.pass_through_endpoints.llm_provider_handlers.openai_passthrough_logging_handler.OpenAIPassthroughLoggingHandler.openai_passthrough_handler')
+ @pytest.mark.asyncio
+ async def test_success_handler_calls_openai_handler(self, mock_openai_handler):
+ """Test that the success handler calls our OpenAI handler for OpenAI routes"""
+ # Arrange
+ mock_openai_handler.return_value = {
+ "result": {"id": "chatcmpl-123"},
+ "kwargs": {
+ "response_cost": 0.000045,
+ "model": "gpt-4o",
+ "custom_llm_provider": "openai"
+ }
+ }
+
+ mock_httpx_response = MagicMock(spec=httpx.Response)
+ mock_httpx_response.text = '{"id": "chatcmpl-123", "choices": [{"message": {"content": "Hello"}}]}'
+
+ mock_logging_obj = AsyncMock()
+ mock_logging_obj.model_call_details = {}
+ mock_logging_obj.async_success_handler = AsyncMock()
+
+ passthrough_payload = PassthroughStandardLoggingPayload(
+ url="https://api.openai.com/v1/chat/completions",
+ request_body={"model": "gpt-4o", "messages": [{"role": "user", "content": "Hello"}]},
+ request_method="POST",
+ )
+
+ # Act
+ result = await self.handler.pass_through_async_success_handler(
+ httpx_response=mock_httpx_response,
+ response_body={"id": "chatcmpl-123", "choices": [{"message": {"content": "Hello"}}]},
+ logging_obj=mock_logging_obj,
+ url_route="https://api.openai.com/v1/chat/completions",
+ result="",
+ start_time=datetime.now(),
+ end_time=datetime.now(),
+ cache_hit=False,
+ request_body={"model": "gpt-4o", "messages": [{"role": "user", "content": "Hello"}]},
+ passthrough_logging_payload=passthrough_payload,
+ )
+
+ # Assert
+ mock_openai_handler.assert_called_once()
+ # The success handler returns None on success, which is expected
+ assert result is None
+
+ @pytest.mark.asyncio
+ async def test_success_handler_falls_back_for_non_openai_routes(self):
+ """Test that non-OpenAI routes don't call our handler"""
+ # Arrange
+ mock_httpx_response = MagicMock(spec=httpx.Response)
+ mock_httpx_response.text = '{"status": "success"}'
+ mock_httpx_response.headers = {"content-type": "application/json"}
+
+ mock_logging_obj = MagicMock()
+ mock_logging_obj.model_call_details = {}
+
+ passthrough_payload = PassthroughStandardLoggingPayload(
+ url="https://api.anthropic.com/v1/messages",
+ request_body={"model": "claude-3-sonnet", "messages": [{"role": "user", "content": "Hello"}]},
+ request_method="POST",
+ )
+
+ # Mock the _handle_logging method to capture calls
+ self.handler._handle_logging = AsyncMock()
+
+ # Act
+ result = await self.handler.pass_through_async_success_handler(
+ httpx_response=mock_httpx_response,
+ response_body={"status": "success"},
+ logging_obj=mock_logging_obj,
+ url_route="https://api.anthropic.com/v1/messages",
+ result="",
+ start_time=datetime.now(),
+ end_time=datetime.now(),
+ cache_hit=False,
+ request_body={"model": "claude-3-sonnet", "messages": [{"role": "user", "content": "Hello"}]},
+ passthrough_logging_payload=passthrough_payload,
+ )
+
+ # Assert - Should call the base handler, not our OpenAI handler
+ self.handler._handle_logging.assert_called_once()
+
+
+if __name__ == "__main__":
+ pytest.main([__file__])
diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py
index b8381201b0a..e296cb25f80 100644
--- a/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py
+++ b/tests/test_litellm/proxy/pass_through_endpoints/test_pass_through_endpoints.py
@@ -1245,3 +1245,58 @@ async def test_delete_pass_through_endpoint_empty_list():
# Verify the exception
assert exc_info.value.status_code == 400
assert "no pass-through endpoints setup" in str(exc_info.value.detail).lower()
+
+
+
+@pytest.mark.asyncio
+async def test_pass_through_with_httpbin_redirect():
+ """
+ Integration test using httpbin.org redirect endpoint to test real redirect handling.
+ This tests the actual redirect handling capability end-to-end using the full pass_through_request function.
+ """
+ from unittest.mock import MagicMock
+
+ from fastapi import Request
+ from starlette.datastructures import Headers, QueryParams
+
+ from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
+ pass_through_request,
+ )
+
+ # Create mock request
+ mock_request = MagicMock(spec=Request)
+ mock_request.method = "GET"
+ mock_request.headers = Headers({})
+ mock_request.query_params = QueryParams("")
+
+ # Mock the body method to return empty bytes for GET request
+ async def mock_body():
+ return b""
+ mock_request.body = mock_body
+
+ # Mock user API key dict
+ mock_user_api_key_dict = MagicMock()
+
+ try:
+ # Test with httpbin.org redirect endpoint
+ # This will redirect to httpbin.org/get
+ response = await pass_through_request(
+ request=mock_request,
+ target="https://httpbin.org/redirect/1",
+ custom_headers={},
+ user_api_key_dict=mock_user_api_key_dict
+ )
+
+ # Should get the final response (200) from /get endpoint, not the redirect (302)
+ assert response.status_code == 200
+
+ # The response should be from the /get endpoint
+ response_content = response.body.decode('utf-8')
+
+ # httpbin.org/get returns JSON with info about the request
+ assert '"url": "https://httpbin.org/get"' in response_content
+ print("GOT A Response from HTTPBIN=", response_content)
+ except Exception as e:
+ # If httpbin.org is not accessible, skip the test
+ import pytest
+ pytest.skip(f"Could not reach httpbin.org for integration test: {e}")
diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py
index 9e487939284..6f63b866220 100644
--- a/tests/test_litellm/test_utils.py
+++ b/tests/test_litellm/test_utils.py
@@ -549,6 +549,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"completion",
"embedding",
"image_generation",
+ "video_generation",
"moderation",
"rerank",
"responses",
@@ -636,7 +637,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"type": "array",
"items": {
"type": "string",
- "enum": ["text", "image", "audio", "code"],
+ "enum": ["text", "image", "audio", "code", "video"],
},
},
"supports_native_streaming": {"type": "boolean"},