Merge branch 'main' into cryptography_upgrade

This commit is contained in:
Andrew Doan 2025-08-28 10:59:51 -04:00 • committed by GitHub
commit 3f215fcf8f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
92 changed files with 7324 additions and 2445 deletions

View file

@ -1291,6 +1291,7 @@ jobs:
pip install jinja2
pip install "tokenizers==0.20.0"
pip install "uvloop==0.21.0"
pip install "fastuuid==0.12.0"
pip install jsonschema
- setup_litellm_enterprise_pip
- run:

View file

@ -14,4 +14,5 @@ google-cloud-iam==2.19.1
fastapi-sso==0.16.0
uvloop==0.21.0
mcp==1.10.1 # for MCP server
semantic_router==0.1.10 # for auto-routing with litellm
semantic_router==0.1.10 # for auto-routing with litellm
fastuuid==0.12.0

View file

@ -0,0 +1,232 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Image Generation in Chat Completions, Responses API
This guide covers how to generate images when using the `chat/completions`. Note - if you want this on Responses API please file a Feature Request [here](https://github.com/BerriAI/litellm/issues/new).
:::info
Requires LiteLLM v1.76.1+
:::
Supported Providers:
- Google AI Studio (`gemini`)
- Vertex AI (`vertex_ai/`)
LiteLLM will standardize the `image` response in the assistant message for models that support image generation during chat completions.
```python title="Example response from litellm"
"message": {
...
"content": "Here's the image you requested:",
"image": {
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAA...",
"detail": "auto"
}
}
```
## Quick Start
<Tabs>
<TabItem value="sdk" label="SDK">
```python showLineNumbers title="Image generation with chat completion"
from litellm import completion
import os
os.environ["GEMINI_API_KEY"] = "your-api-key"
response = completion(
model="gemini/gemini-2.5-flash-image-preview",
messages=[
{"role": "user", "content": "Generate an image of a banana wearing a costume that says LiteLLM"}
],
)
print(response.choices[0].message.content) # Text response
print(response.choices[0].message.image) # Image data
```
</TabItem>
<TabItem value="proxy" label="PROXY">
1. Setup config.yaml
```yaml showLineNumbers title="config.yaml"
model_list:
- model_name: gemini-image-gen
litellm_params:
model: gemini/gemini-2.5-flash-image-preview
api_key: os.environ/GEMINI_API_KEY
```
2. Run proxy server
```bash showLineNumbers title="Start the proxy"
litellm --config config.yaml
# RUNNING on http://0.0.0.0:4000
```
3. Test it!
```bash showLineNumbers title="Make request"
curl http://0.0.0.0:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer $LITELLM_KEY" \
-d '{
"model": "gemini-image-gen",
"messages": [
{
"role": "user",
"content": "Generate an image of a banana wearing a costume that says LiteLLM"
}
]
}'
```
</TabItem>
</Tabs>
**Expected Response**
```bash
{
"id": "chatcmpl-3b66124d79a708e10c603496b363574c",
"choices": [
{
"finish_reason": "stop",
"index": 0,
"message": {
"content": "Here's the image you requested:",
"role": "assistant",
"image": {
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAA...",
"detail": "auto"
}
}
}
],
"created": 1723323084,
"model": "gemini/gemini-2.5-flash-image-preview",
"object": "chat.completion",
"usage": {
"completion_tokens": 12,
"prompt_tokens": 16,
"total_tokens": 28
}
}
```
## Streaming Support
<Tabs>
<TabItem value="sdk" label="SDK">
```python showLineNumbers title="Streaming image generation"
from litellm import completion
import os
os.environ["GEMINI_API_KEY"] = "your-api-key"
response = completion(
model="gemini/gemini-2.5-flash-image-preview",
messages=[
{"role": "user", "content": "Generate an image of a banana wearing a costume that says LiteLLM"}
],
stream=True,
)
for chunk in response:
if hasattr(chunk.choices[0].delta, "image") and chunk.choices[0].delta.image is not None:
print("Generated image:", chunk.choices[0].delta.image["url"])
break
```
</TabItem>
<TabItem value="proxy" label="PROXY">
```bash showLineNumbers title="Streaming request"
curl http://0.0.0.0:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-H "Authorization: Bearer $LITELLM_KEY" \
-d '{
"model": "gemini-image-gen",
"messages": [
{
"role": "user",
"content": "Generate an image of a banana wearing a costume that says LiteLLM"
}
],
"stream": true
}'
```
</TabItem>
</Tabs>
**Expected Streaming Response**
```bash
data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1723323084,"model":"gemini/gemini-2.5-flash-image-preview","choices":[{"index":0,"delta":{"role":"assistant"},"finish_reason":null}]}
data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1723323084,"model":"gemini/gemini-2.5-flash-image-preview","choices":[{"index":0,"delta":{"content":"Here's the image you requested:"},"finish_reason":null}]}
data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1723323084,"model":"gemini/gemini-2.5-flash-image-preview","choices":[{"index":0,"delta":{"image":{"url":"data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAA...","detail":"auto"}},"finish_reason":null}]}
data: {"id":"chatcmpl-123","object":"chat.completion.chunk","created":1723323084,"model":"gemini/gemini-2.5-flash-image-preview","choices":[{"index":0,"delta":{},"finish_reason":"stop"}]}
data: [DONE]
```
## Async Support
```python showLineNumbers title="Async image generation"
from litellm import acompletion
import asyncio
import os
os.environ["GEMINI_API_KEY"] = "your-api-key"
async def generate_image():
response = await acompletion(
model="gemini/gemini-2.5-flash-image-preview",
messages=[
{"role": "user", "content": "Generate an image of a banana wearing a costume that says LiteLLM"}
],
)
print(response.choices[0].message.content) # Text response
print(response.choices[0].message.image) # Image data
return response
# Run the async function
asyncio.run(generate_image())
```
## Supported Models
| Provider | Model |
|----------|--------|
| Google AI Studio | `gemini/gemini-2.5-flash-image-preview` |
| Vertex AI | `vertex_ai/gemini-2.5-flash-image-preview` |
## Spec
The `image` field in the response follows this structure:
```python
"image": {
"url": "data:image/png;base64,<base64_encoded_image>",
"detail": "auto"
}
```
- `url` - str: Base64 encoded image data in data URI format
- `detail` - str: Image detail level (always "auto" for generated images)
The image is returned as a base64-encoded data URI that can be directly used in HTML `<img>` tags or saved to a file.

View file

@ -0,0 +1,201 @@
# Gemini Image Generation Migration Guide
## Who is impacted by this change?
Anyone using the following models with /chat/completions:
- `gemini/gemini-2.0-flash-exp-image-generation`
- `vertex_ai/gemini-2.5-flash-image-preview`
## Key Change
Gemini models now support image generation through chat completions. Images are returned in `response.choices[0].message.image` with base64 data URLs.
## Before and After
### Before
```python
from litellm import completion
response = completion(
model="gemini/gemini-2.0-flash-exp-image-generation",
messages=[{"role": "user", "content": "Generate an image of a cat"}],
modalities=["image", "text"],
)
base_64_image_data = response.choices[0].message.content
```
### After
```python
from litellm import completion
response = completion(
model="gemini/gemini-2.0-flash-exp-image-generation",
messages=[{"role": "user", "content": "Generate an image of a cat"}],
modalities=["image", "text"],
)
# Image is now available in the response
image_url = response.choices[0].message.image["url"] # "data:image/png;base64,..."
```
## Usage
### Using the Python SDK
**Key Change:**
```diff
# Before
-- base_64_image_data = response.choices[0].message.content
# After
++ image_url = response.choices[0].message.image["url"]
```
#### Basic Image Generation
```python
from litellm import completion
import os
# Set your API key
os.environ["GEMINI_API_KEY"] = "your-api-key"
# Generate an image
response = completion(
model="gemini/gemini-2.0-flash-exp-image-generation",
messages=[{"role": "user", "content": "Generate an image of a cat"}],
modalities=["image", "text"],
)
# Access the generated image
print(response.choices[0].message.content) # Text response (if any)
print(response.choices[0].message.image) # Image data
```
#### Response Format
The image is returned in the `message.image` field:
```python
{
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAA...",
"detail": "auto"
}
```
### Using the LiteLLM Proxy Server
**Key Change:**
```diff
# Before
-- "content": "base64-image-data..."
# After
++ "image": {
++ "url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAA...",
++ "detail": "auto"
++ }
```
#### Configuration Setup
1. **Configure your models in `config.yaml`:**
```yaml
model_list:
- model_name: gemini-image-gen
litellm_params:
model: gemini/gemini-2.0-flash-exp-image-generation
api_key: os.environ/GEMINI_API_KEY
- model_name: vertex-image-gen
litellm_params:
model: vertex_ai/gemini-2.5-flash-image-preview
vertex_project: your-project-id
vertex_location: us-central1
general_settings:
master_key: sk-1234 # Your proxy API key
```
2. **Start the proxy server:**
```bash
litellm --config /path/to/config.yaml
# RUNNING on http://0.0.0.0:4000
```
#### Making Requests
**Using OpenAI SDK:**
```python
from openai import OpenAI
# Point to your proxy server
client = OpenAI(
api_key="sk-1234", # Your proxy API key
base_url="http://0.0.0.0:4000"
)
response = client.chat.completions.create(
model="gemini-image-gen",
messages=[{"role": "user", "content": "Generate an image of a cat"}],
extra_body={"modalities": ["image", "text"]}
)
# Access the generated image
print(response.choices[0].message.content) # Text response (if any)
print(response.choices[0].message.image) # Image data
```
**Using curl:**
```bash
curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer sk-1234' \
-d '{
"model": "gemini-image-gen",
"messages": [
{
"role": "user",
"content": "Generate an image of a cat"
}
],
"modalities": ["image", "text"]
}'
```
**Response format from proxy:**
```json
{
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1704089632,
"model": "gemini-image-gen",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "Here's an image of a cat for you!",
"image": {
"url": "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAA...",
"detail": "auto"
}
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": 10,
"completion_tokens": 8,
"total_tokens": 18
}
}
```

View file

@ -4,7 +4,7 @@ import TabItem from '@theme/TabItem';
# /images/edits
LiteLLM provides image editing functionality that maps to OpenAI's `/images/edits` API endpoint.
LiteLLM provides image editing functionality that maps to OpenAI's `/images/edits` API endpoint. Now supports both single and multiple image editing.
| Feature | Supported | Notes |
|---------|-----------|--------|
@ -13,7 +13,7 @@ LiteLLM provides image editing functionality that maps to OpenAI's `/images/edit
| End-user Tracking | ✅ | |
| Fallbacks | ✅ | Works between supported models |
| Loadbalancing | ✅ | Works between supported models |
| Supported operations | Create image edits | |
| Supported operations | Create image edits | Single and multiple images supported |
| Supported LiteLLM SDK Versions | 1.63.8+ | |
| Supported LiteLLM Proxy Versions | 1.71.1+ | |
| Supported LLM providers | **OpenAI** | Currently only `openai` is supported |
@ -41,6 +41,26 @@ response = litellm.image_edit(
print(response)
```
#### Multiple Images Edit
```python showLineNumbers title="OpenAI Multiple Images Edit"
import litellm
# Edit multiple images with a prompt
response = litellm.image_edit(
model="gpt-image-1",
image=[
open("image1.png", "rb"),
open("image2.png", "rb"),
open("image3.png", "rb")
],
prompt="Apply vintage filter to all images",
n=1,
size="1024x1024"
)
print(response)
```
#### Image Edit with Mask
```python showLineNumbers title="OpenAI Image Edit with Mask"
import litellm
@ -80,6 +100,30 @@ response = asyncio.run(edit_image())
print(response)
```
#### Async Multiple Images Edit
```python showLineNumbers title="Async OpenAI Multiple Images Edit"
import litellm
import asyncio
async def edit_multiple_images():
response = await litellm.aimage_edit(
model="gpt-image-1",
image=[
open("portrait1.png", "rb"),
open("portrait2.png", "rb")
],
prompt="Add professional lighting to the portraits",
n=1,
size="1024x1024",
response_format="url"
)
return response
# Run the async function
response = asyncio.run(edit_multiple_images())
print(response)
```
#### Image Edit with Custom Parameters
```python showLineNumbers title="OpenAI Image Edit with Custom Parameters"
import litellm
@ -163,6 +207,20 @@ curl -X POST "http://localhost:4000/v1/images/edits" \
-F "response_format=url"
```
#### cURL Multiple Images Example
```bash showLineNumbers title="cURL Multiple Images Edit Request"
curl -X POST "http://localhost:4000/v1/images/edits" \
-H "Authorization: Bearer your-api-key" \
-F "model=gpt-image-1" \
-F "image=@image1.png" \
-F "image=@image2.png" \
-F "image=@image3.png" \
-F "prompt=Apply artistic filter to all images" \
-F "n=1" \
-F "size=1024x1024" \
-F "response_format=url"
```
</TabItem>
</Tabs>

View file

@ -71,6 +71,10 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
It is recommended that you include the `project_id` or `project_name` to ensure your traces are being written out to the correct Braintrust project.
### Custom Span Names
You can customize the span name in Braintrust logging by passing `span_name` in the metadata. By default, the span name is set to "Chat Completion".
<Tabs>
<TabItem value="sdk" label="SDK">
@ -84,7 +88,9 @@ response = litellm.completion(
"project_id": "1234",
# passing project_name will try to find a project with that name, or create one if it doesn't exist
# if both project_id and project_name are passed, project_id will be used
# "project_name": "my-special-project"
# "project_name": "my-special-project",
# custom span name for this operation (default: "Chat Completion")
"span_name": "User Greeting Handler"
}
)
```
@ -99,6 +105,7 @@ response = litellm.completion(
],
metadata={
"project_id": "1234",
"span_name": "Custom Operation",
"item1": "an item",
"item2": "another item"
}
@ -121,7 +128,8 @@ curl -X POST 'http://0.0.0.0:4000/chat/completions' \
{ "role": "user", "content": "What time is it now? Use your tool"}
],
"metadata": {
"project_id": "my-special-project"
"project_id": "my-special-project",
"span_name": "Tool Usage Request"
}
}'
```
@ -146,7 +154,8 @@ response = client.chat.completions.create(
],
extra_body={ # pass in any provider-specific param, if not supported by openai, https://docs.litellm.ai/docs/completion/input#provider-specific-params
"metadata": { # 👈 use for logging additional params (e.g. to braintrust)
"project_id": "my-special-project"
"project_id": "my-special-project",
"span_name": "Poetry Generation"
}
}
)
@ -168,3 +177,7 @@ Here's everything you can pass in metadata for a braintrust request
`braintrust_*` - If you are adding metadata from _proxy request headers_, any metadata field starting with `braintrust_` will be passed as metadata to the logging request. If you are using the SDK, just pass your metadata like normal (e.g., `metadata={"project_name": "my-test-project", "item1": "an item", "item2": "another item"}`)
`project_id` - Set the project id for a braintrust call. Default is `litellm`.
`project_name` - Set the project name for a braintrust call. Will try to find a project with that name, or create one if it doesn't exist. If both `project_id` and `project_name` are passed, `project_id` will be used.
`span_name` - Set a custom span name for the operation. Default is `"Chat Completion"`. Use this to provide more descriptive names for different types of operations in your application (e.g., "User Query", "Document Summary", "Code Generation").

View file

@ -35,14 +35,14 @@ The Langfuse OpenTelemetry integration allows you to send LiteLLM traces and obs
|----------|----------|-------------|---------|
| `LANGFUSE_PUBLIC_KEY` | Yes | Your Langfuse public key | `pk-lf-...` |
| `LANGFUSE_SECRET_KEY` | Yes | Your Langfuse secret key | `sk-lf-...` |
| `LANGFUSE_HOST` | No | Langfuse host URL | `https://us.cloud.langfuse.com` (default) |
| `LANGFUSE_OTEL_HOST` | No | OTEL endpoint host | `https://otel.my-langfuse.com` |
### Endpoint Resolution
The integration automatically constructs the OTEL endpoint from the `LANGFUSE_HOST`:
The integration automatically constructs the OTEL endpoint from `LANGFUSE_OTEL_HOST`
- **Default (US)**: `https://us.cloud.langfuse.com/api/public/otel`
- **EU Region**: `https://cloud.langfuse.com/api/public/otel`
- **Self-hosted**: `{LANGFUSE_HOST}/api/public/otel`
- **Self-hosted**: `{LANGFUSE_OTEL_HOST}/api/public/otel`
## Usage
@ -77,11 +77,11 @@ os.environ["LANGFUSE_PUBLIC_KEY"] = "pk-lf-..."
os.environ["LANGFUSE_SECRET_KEY"] = "sk-lf-..."
# Use EU region
os.environ["LANGFUSE_HOST"] = "https://cloud.langfuse.com" # EU region
# os.environ["LANGFUSE_HOST"] = "https://us.cloud.langfuse.com" # US region (default)
os.environ["LANGFUSE_OTEL_HOST"] = "https://cloud.langfuse.com" # EU region
# os.environ["LANGFUSE_OTEL_HOST"] = "https://otel.my-langfuse.company.com" # custom OTEL endpoint
# Or use self-hosted instance
# os.environ["LANGFUSE_HOST"] = "https://my-langfuse.company.com"
# os.environ["LANGFUSE_OTEL_HOST"] = "https://my-langfuse.company.com"
litellm.callbacks = ["langfuse_otel"]
```
@ -98,14 +98,16 @@ import litellm
# Get keys for your project from the project settings page: https://cloud.langfuse.com
os.environ["LANGFUSE_PUBLIC_KEY"] = "pk-lf-..."
os.environ["LANGFUSE_SECRET_KEY"] = "sk-lf-..."
os.environ["LANGFUSE_HOST"] = "https://cloud.langfuse.com" # EU region
# os.environ["LANGFUSE_HOST"] = "https://us.cloud.langfuse.com" # US region
os.environ["LANGFUSE_OTEL_HOST"] = "https://cloud.langfuse.com" # EU region
# os.environ["LANGFUSE_OTEL_HOST"] = "https://us.cloud.langfuse.com" # US region
# os.environ["LANGFUSE_OTEL_HOST"] = "https://otel.my-langfuse.company.com" # custom OTEL endpoint
LANGFUSE_AUTH = base64.b64encode(
f"{os.environ.get('LANGFUSE_PUBLIC_KEY')}:{os.environ.get('LANGFUSE_SECRET_KEY')}".encode()
).decode()
os.environ["OTEL_EXPORTER_OTLP_ENDPOINT"] = os.environ.get("LANGFUSE_HOST") + "/api/public/otel"
host = os.environ.get("LANGFUSE_OTEL_HOST")
os.environ["OTEL_EXPORTER_OTLP_ENDPOINT"] = host + "/api/public/otel"
os.environ["OTEL_EXPORTER_OTLP_HEADERS"] = f"Authorization=Basic {LANGFUSE_AUTH}"
litellm.callbacks = ["langfuse_otel"]
@ -120,7 +122,8 @@ Add the integration to your proxy configuration:
```bash
export LANGFUSE_PUBLIC_KEY="pk-lf-..."
export LANGFUSE_SECRET_KEY="sk-lf-..."
export LANGFUSE_HOST="https://us.cloud.langfuse.com" # Default US region
export LANGFUSE_OTEL_HOST="https://us.cloud.langfuse.com" # Default US region
# export LANGFUSE_OTEL_HOST="https://otel.my-langfuse.company.com" # custom OTEL endpoint
```
2. Setup config.yaml

View file

@ -0,0 +1,144 @@
# CometAPI
LiteLLM supports all AI models from [CometAPI](https://www.cometapi.com/). CometAPI provides access to 500+ AI models through a unified API interface, including cutting-edge models like GPT-5, Claude Opus 4.1, and various other state-of-the-art language models.
## Authentication
To use CometAPI models, you need to obtain an API key from [CometAPI Token Console](https://api.cometapi.com/console/token). CometAPI offers free tokens for new users - you can get your free API key instantly by registering.
## Usage
Set your CometAPI key as an environment variable and use the completion function:
```python
import os
from litellm import completion
# Set API key
os.environ["COMETAPI_KEY"] = "your_comet_api_key_here"
# Define messages
messages = [{"content": "Hello, how are you?", "role": "user"}]
# Method 1: Using environment variable (recommended)
response = completion(
model="cometapi/gpt-5",
messages=messages
)
print(response.choices[0].message.content)
```
### Alternative Usage - Explicit API Key
You can also pass the API key explicitly:
```python
import os
from litellm import completion
# Define messages
messages = [{"content": "Hello, how are you?", "role": "user"}]
# Method 2: Explicitly passing API key
response = completion(
model="cometapi/gpt-4o",
messages=messages,
api_key="your_comet_api_key_here"
)
print(response.choices[0].message.content)
```
## Usage - Streaming
Just set `stream=True` when calling completion:
```python
import os
from litellm import completion
os.environ["COMETAPI_KEY"] = "your_comet_api_key_here"
messages = [{"content": "Hello, how are you?", "role": "user"}]
response = completion(
model="cometapi/gpt-5",
messages=messages,
stream=True
)
for chunk in response:
print(chunk.choices[0].delta.content or "", end="")
```
## Usage - Async Streaming
For async streaming, use `acompletion`:
```python
from litellm import acompletion
import asyncio, os, traceback
async def completion_call():
try:
os.environ["COMETAPI_KEY"] = "your_comet_api_key_here"
print("test acompletion + streaming")
response = await acompletion(
model="cometapi/chatgpt-4o-latest",
messages=[{"content": "Hello, how are you?", "role": "user"}],
stream=True
)
print(f"response: {response}")
async for chunk in response:
print(chunk)
except:
print(f"error occurred: {traceback.format_exc()}")
pass
# Run the async function
await completion_call()
```
## CometAPI Models
CometAPI offers access to 500+ AI models through a unified API. Some popular models include:
| Model Name | Function Call |
|------------|---------------|
| cometapi/gpt-5 | `completion('cometapi/gpt-5', messages)` |
| cometapi/gpt-5-mini | `completion('cometapi/gpt-5-mini', messages)` |
| cometapi/gpt-5-nano | `completion('cometapi/gpt-5-nano', messages)` |
| cometapi/gpt-oss-20b | `completion('cometapi/gpt-oss-20b', messages)` |
| cometapi/gpt-oss-120b | `completion('cometapi/gpt-oss-120b', messages)` |
| cometapi/chatgpt-4o-latest | `completion('cometapi/chatgpt-4o-latest', messages)` |
For a complete list of available models, visit the [CometAPI Models page](https://www.cometapi.com/model/).
## Environment Variables
| Variable | Description | Required |
|----------|-------------|----------|
| `COMETAPI_KEY` | Your CometAPI API key | Yes |
## Error Handling
```python
import os
from litellm import completion
try:
os.environ["COMETAPI_KEY"] = "your_comet_api_key_here"
messages = [{"content": "Hello, how are you?", "role": "user"}]
response = completion(
model="cometapi/gpt-5",
messages=messages
)
print(response.choices[0].message.content)
except Exception as e:
print(f"Error: {e}")
```

View file

@ -42,7 +42,7 @@ os.environ["GEMINI_API_KEY"] = "your-api-key-here"
# Generate a single image
response = litellm.image_generation(
model="gemini/imagen-4.0-generate-preview-06-06",
model="gemini/imagen-4.0-generate-001",
prompt="A cute baby sea otter swimming in crystal clear water"
)
@ -64,7 +64,7 @@ async def generate_image():
# Generate image asynchronously
response = await litellm.aimage_generation(
model="gemini/imagen-4.0-generate-preview-06-06",
model="gemini/imagen-4.0-generate-001",
prompt="A beautiful sunset over mountains with vibrant colors",
n=1,
)
@ -89,7 +89,7 @@ os.environ["GEMINI_API_KEY"] = "your-api-key-here"
# Generate image with additional parameters
response = litellm.image_generation(
model="gemini/imagen-4.0-generate-preview-06-06",
model="gemini/imagen-4.0-generate-001",
prompt="A futuristic cityscape at night with neon lights",
n=1,
size="1024x1024",
@ -112,7 +112,7 @@ for image in response.data:
model_list:
- model_name: google-imagen
litellm_params:
model: gemini/imagen-4.0-generate-preview-06-06
model: gemini/imagen-4.0-generate-001
api_key: os.environ/GEMINI_API_KEY
model_info:
mode: image_generation
@ -198,7 +198,7 @@ Google AI Studio Image Generation supports the following OpenAI-compatible param
| Parameter | Type | Description | Default | Example |
|-----------|------|-------------|---------|---------|
| `prompt` | string | Text description of the image to generate | Required | `"A sunset over the ocean"` |
| `model` | string | The model to use for generation | Required | `"gemini/imagen-4.0-generate-preview-06-06"` |
| `model` | string | The model to use for generation | Required | `"gemini/imagen-4.0-generate-001"` |
| `n` | integer | Number of images to generate (1-4) | `1` | `2` |
| `size` | string | Image dimensions | `"1024x1024"` | `"512x512"`, `"1024x1024"` |

View file

@ -18,7 +18,7 @@ import litellm
# Generate a single image
response = await litellm.aimage_generation(
prompt="An olympic size swimming pool with crystal clear water and modern architecture",
model="vertex_ai/imagen-4.0-generate-preview-06-06",
model="vertex_ai/imagen-4.0-generate-001",
vertex_ai_project="your-project-id",
vertex_ai_location="us-central1",
)
@ -34,7 +34,7 @@ print(response.data[0].url)
model_list:
- model_name: vertex-imagen
litellm_params:
model: vertex_ai/imagen-4.0-generate-preview-06-06
model: vertex_ai/imagen-4.0-generate-001
vertex_ai_project: "your-project-id"
vertex_ai_location: "us-central1"
vertex_ai_credentials: "path/to/service-account.json" # Optional if using environment auth

View file

@ -570,6 +570,7 @@ router_settings:
| LITELLM_LICENSE | License key for LiteLLM usage
| LITELLM_LOCAL_MODEL_COST_MAP | Local configuration for model cost mapping in LiteLLM
| LITELLM_LOG | Enable detailed logging for LiteLLM
| LITELLM_LOG_FILE | File path to write LiteLLM logs to. When set, logs will be written to both console and the specified file
| LITELLM_MASTER_KEY | Master key for proxy authentication
| LITELLM_MODE | Operating mode for LiteLLM (e.g., production, development)
| LITELLM_RATE_LIMIT_WINDOW_SIZE | Rate limit window size for LiteLLM. Default is 60

View file

@ -0,0 +1,299 @@
import Image from '@theme/IdealImage';
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Noma Security
Use [Noma Security](https://noma.security/) to protect your LLM applications with comprehensive AI content moderation and safety guardrails.
## Quick Start
### 1. Define Guardrails on your LiteLLM config.yaml
Define your guardrails under the `guardrails` section:
```yaml showLineNumbers title="litellm config.yaml"
model_list:
- model_name: gpt-4o-mini
litellm_params:
model: openai/gpt-4o-mini
api_key: os.environ/OPENAI_API_KEY
guardrails:
- guardrail_name: "noma-guard"
litellm_params:
guardrail: noma
mode: "during_call"
api_key: os.environ/NOMA_API_KEY
api_base: os.environ/NOMA_API_BASE
- guardrail_name: "noma-pre-guard"
litellm_params:
guardrail: noma
mode: "pre_call"
api_key: os.environ/NOMA_API_KEY
api_base: os.environ/NOMA_API_BASE
```
#### Supported values for `mode`
- `pre_call` Run **before** LLM call, on **input**
- `post_call` Run **after** LLM call, on **input & output**
- `during_call` Run **during** LLM call, on **input**. Same as `pre_call` but runs in parallel with the LLM call. Response not returned until guardrail check completes
### 2. Start LiteLLM Gateway
```shell
litellm --config config.yaml --detailed_debug
```
### 3. Test request
<Tabs>
<TabItem label="Unsuccessful call" value="not-allowed">
Expect this to fail since the request contains harmful content:
```shell showLineNumbers title="Curl Request"
curl -i http://0.0.0.0:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-d '{
"model": "gpt-4o-mini",
"messages": [
{"role": "user", "content": "Tell me how to hack into someone's email account"}
]
}'
```
Expected response on failure:
```json
{
"error": {
"message": "{\n \"error\": \"Request blocked by Noma guardrail\",\n \"details\": {\n \"prompt\": {\n \"harmfulContent\": {\n \"result\": true,\n \"confidence\": 0.95\n }\n }\n }\n }",
"type": "None",
"param": "None",
"code": "400"
}
}
```
</TabItem>
<TabItem label="Successful Call" value="allowed">
```shell showLineNumbers title="Curl Request"
curl -i http://0.0.0.0:4000/v1/chat/completions \
-H "Content-Type: application/json" \
-d '{
"model": "gpt-4o-mini",
"messages": [
{"role": "user", "content": "What is the capital of France?"}
]
}'
```
Expected response:
```json
{
"id": "chatcmpl-123",
"object": "chat.completion",
"created": 1677652288,
"model": "gpt-4o-mini",
"choices": [
{
"index": 0,
"message": {
"role": "assistant",
"content": "The capital of France is Paris."
},
"finish_reason": "stop"
}
],
"usage": {
"prompt_tokens": 9,
"completion_tokens": 12,
"total_tokens": 21
}
}
```
</TabItem>
</Tabs>
## Supported Params
```yaml
guardrails:
- guardrail_name: "noma-guard"
litellm_params:
guardrail: noma
mode: "pre_call"
api_key: os.environ/NOMA_API_KEY
api_base: os.environ/NOMA_API_BASE
### OPTIONAL ###
# application_id: "my-app"
# monitor_mode: false
# block_failures: true
```
### Required Parameters
- **`api_key`**: Your Noma Security API key (set as `os.environ/NOMA_API_KEY` in YAML config)
### Optional Parameters
- **`api_base`**: Noma API base URL (defaults to `https://api.noma.security/`)
- **`application_id`**: Your application identifier (defaults to `"litellm"`)
- **`monitor_mode`**: If `true`, logs violations without blocking (defaults to `false`)
- **`block_failures`**: If `true`, blocks requests when guardrail API failures occur (defaults to `true`)
## Environment Variables
You can set these environment variables instead of hardcoding values in your config:
```shell
export NOMA_API_KEY="your-api-key-here"
export NOMA_API_BASE="https://api.noma.security/" # Optional
export NOMA_APPLICATION_ID="my-app" # Optional
export NOMA_MONITOR_MODE="false" # Optional
export NOMA_BLOCK_FAILURES="true" # Optional
```
## Advanced Configuration
### Monitor Mode
Use monitor mode to test your guardrails without blocking requests:
```yaml
guardrails:
- guardrail_name: "noma-monitor"
litellm_params:
guardrail: noma
mode: "pre_call"
api_key: os.environ/NOMA_API_KEY
monitor_mode: true # Log violations but don't block
```
### Handling API Failures
Control behavior when the Noma API is unavailable:
```yaml
guardrails:
- guardrail_name: "noma-failopen"
litellm_params:
guardrail: noma
mode: "pre_call"
api_key: os.environ/NOMA_API_KEY
block_failures: false # Allow requests to proceed if guardrail API fails
```
### Multiple Guardrails
Apply different configurations for input and output:
```yaml
guardrails:
- guardrail_name: "noma-strict-input"
litellm_params:
guardrail: noma
mode: "pre_call"
api_key: os.environ/NOMA_API_KEY
block_failures: true
- guardrail_name: "noma-monitor-output"
litellm_params:
guardrail: noma
mode: "post_call"
api_key: os.environ/NOMA_API_KEY
monitor_mode: true
```
## ✨ Pass Additional Parameters
Use `extra_body` to pass additional parameters to the Noma Security API call, such as dynamically setting the application ID for specific requests.
<Tabs>
<TabItem value="openai" label="OpenAI Python">
```python
import openai
client = openai.OpenAI(
api_key="your-api-key",
base_url="http://0.0.0.0:4000"
)
response = client.chat.completions.create(
model="gpt-4o-mini",
messages=[{"role": "user", "content": "Hello, how are you?"}],
extra_body={
"guardrails": {
"noma-guard": {
"extra_body": {
"application_id": "my-specific-app-id"
}
}
}
}
)
```
</TabItem>
<TabItem value="curl" label="Curl">
```shell
curl 'http://0.0.0.0:4000/v1/chat/completions' \
-H 'Content-Type: application/json' \
-d '{
"model": "gpt-4o-mini",
"messages": [
{
"role": "user",
"content": "Hello, how are you?"
}
],
"guardrails": {
"noma-guard": {
"extra_body": {
"application_id": "my-specific-app-id"
}
}
}
}'
```
</TabItem>
</Tabs>
This allows you to override the default `application_id` parameter for specific requests, which is useful for tracking usage across different applications or components.
## Response Details
When content is blocked, Noma provides detailed information about the violations as JSON inside the `message` field, with the following structure:
```json
{
"error": "Request blocked by Noma guardrail",
"details": {
"prompt": {
"harmfulContent": {
"result": true,
"confidence": 0.95
},
"sensitiveData": {
"email": {
"result": true,
"entities": ["user@example.com"]
}
},
"bannedTopics": {
"violence": {
"result": true,
"confidence": 0.88
}
}
}
}
}
```

View file

@ -309,6 +309,37 @@ curl -X POST '<PROXY_BASE_URL>/team/new' \
</TabItem>
</Tabs>
### Team Member Rate Limits
Set a default tpm/rpm limit for an individual team member.
You can do this when creating a new team, or by updating an existing team.
<Tabs>
<TabItem value="ui" label="UI">
<Image img={require('../../img/create_team_member_rate_limits.png')} style={{ width: '600px', height: 'auto' }} />
</TabItem>
<TabItem value="api" label="API">
```bash
curl -X POST '<PROXY_BASE_URL>/team/new' \
-H 'Authorization: Bearer <PROXY_MASTER_KEY>' \
-H 'Content-Type: application/json' \
-D '{
"team_alias": "team_1",
"team_member_rpm_limit": 100,
"team_member_tpm_limit": 1000
}'
```
</TabItem>
</Tabs>
### Set default params for new teams
When you connect litellm to your SSO provider, litellm can auto-create teams. Use this to set the default `models`, `max_budget`, `budget_duration` for these auto-created teams.

Binary file not shown.

After

Width:  |  Height:  |  Size: 100 KiB

View file

@ -86,9 +86,9 @@ This is great to central AI Platform teams looking to track how they are helping
| Provider | Model | Context Window | Input ($/1M tokens) | Output ($/1M tokens) | Cost per Image |
| ----------- | -------------------------------------- | -------------- | ------------------- | -------------------- | -------------- |
| OpenRouter | `openrouter/x-ai/grok-4` | 256k | $3 | $15 | N/A |
| Google AI Studio | `gemini/imagen-4.0-generate-preview-06-06` | N/A | N/A | N/A | $0.04 |
| Google AI Studio | `gemini/imagen-4.0-ultra-generate-preview-06-06` | N/A | N/A | N/A | $0.06 |
| Google AI Studio | `gemini/imagen-4.0-fast-generate-preview-06-06` | N/A | N/A | N/A | $0.02 |
| Google AI Studio | `gemini/imagen-4.0-generate-001` | N/A | N/A | N/A | $0.04 |
| Google AI Studio | `gemini/imagen-4.0-ultra-generate-001` | N/A | N/A | N/A | $0.06 |
| Google AI Studio | `gemini/imagen-4.0-fast-generate-001` | N/A | N/A | N/A | $0.02 |
| Google AI Studio | `gemini/imagen-3.0-generate-002` | N/A | N/A | N/A | $0.04 |
| Google AI Studio | `gemini/imagen-3.0-generate-001` | N/A | N/A | N/A | $0.04 |
| Google AI Studio | `gemini/imagen-3.0-fast-generate-001` | N/A | N/A | N/A | $0.02 |

View file

@ -28,7 +28,7 @@ import TabItem from '@theme/TabItem';
docker run \
-e STORE_MODEL_IN_DB=True \
-p 4000:4000 \
ghcr.io/berriai/litellm:v1.75.8
ghcr.io/berriai/litellm:v1.75.8-stable
```
</TabItem>
@ -55,7 +55,7 @@ pip install litellm==1.75.8
## Team Member Rate Limits
<Image
img={require('../img/release_notes/team_member_rate_limits.png')}
img={require('../../img/release_notes/team_member_rate_limits.png')}
style={{width: '100%', display: 'block', margin: '2rem auto'}}
/>
<p style={{textAlign: 'left', color: '#666'}}>

View file

@ -40,6 +40,7 @@ const sidebars = {
"proxy/guardrails/guardrails_ai",
"proxy/guardrails/lakera_ai",
"proxy/guardrails/model_armor",
"proxy/guardrails/noma_security",
"proxy/guardrails/openai_moderation",
"proxy/guardrails/pangea",
"proxy/guardrails/pillar_security",
@ -492,6 +493,7 @@ const sidebars = {
"guides/finetuned_models",
"guides/security_settings",
"completion/audio",
"completion/image_generation_chat",
"completion/web_search",
"completion/document_understanding",
"completion/vision",

View file

@ -1278,7 +1278,6 @@ from .router import Router
from .assistants.main import *
from .batches.main import *
from .images.main import *
from .vector_stores import *
from .batch_completion.main import * # type: ignore
from .rerank_api.main import *
from .llms.anthropic.experimental_pass_through.messages.handler import *

View file

@ -108,6 +108,7 @@ verbose_router_logger.addHandler(handler)
verbose_proxy_logger.addHandler(handler)
verbose_logger.addHandler(handler)
def _suppress_loggers():
"""Suppress noisy loggers at INFO level"""
# Suppress httpx request logging at INFO level
@ -120,6 +121,7 @@ def _suppress_loggers():
apscheduler_scheduler_logger = logging.getLogger("apscheduler.scheduler")
apscheduler_scheduler_logger.setLevel(logging.WARNING)
# Call the suppression function
_suppress_loggers()
@ -187,6 +189,4 @@ def _is_debugging_on() -> bool:
"""
Returns True if debugging is on
"""
if verbose_logger.isEnabledFor(logging.DEBUG) or set_verbose is True:
return True
return False
return verbose_logger.isEnabledFor(logging.DEBUG) or set_verbose is True

View file

@ -112,14 +112,15 @@ class InMemoryCache(BaseCache):
- 3. the size of in-memory cache is bounded
"""
for key in list(self.ttl_dict.keys()):
if self._is_key_expired(key):
self._remove_key(key)
current_time = time.time()
expired_keys = [key for key, ttl in self.ttl_dict.items() if current_time > ttl]
for key in expired_keys:
self._remove_key(key)
# de-reference the removed item
# https://www.geeksforgeeks.org/diagnosing-and-fixing-memory-leaks-in-python/
# One of the most common causes of memory leaks in Python is the retention of objects that are no longer being used.
# This can occur when an object is referenced by another object, but the reference is never removed.
# de-reference the removed item
# https://www.geeksforgeeks.org/diagnosing-and-fixing-memory-leaks-in-python/
# One of the most common causes of memory leaks in Python is the retention of objects that are no longer being used.
# This can occur when an object is referenced by another object, but the reference is never removed.
def allow_ttl_override(self, key: str) -> bool:
"""

View file

@ -13,6 +13,7 @@ import asyncio
import json
from functools import partial
from typing import Optional
from datetime import datetime, timezone, timedelta
from litellm._logging import print_verbose, verbose_logger
@ -69,11 +70,9 @@ class S3Cache(BaseCache):
if ttl is not None:
cache_control = f"immutable, max-age={ttl}, s-maxage={ttl}"
import datetime
# Calculate expiration time
expiration_time = datetime.datetime.now() + datetime.timedelta(seconds=ttl)
expiration_time = datetime.now(timezone.utc) + timedelta(seconds=ttl)
# Upload the data to S3 with the calculated expiration time
self.s3_client.put_object(
Bucket=self.bucket_name,
@ -126,6 +125,13 @@ class S3Cache(BaseCache):
)
if cached_response is not None:
if "Expires" in cached_response:
expires_time = cached_response['Expires']
current_time = datetime.now(expires_time.tzinfo)
if current_time > expires_time:
return None
# cached_response is in `b{} convert it to ModelResponse
cached_response = (
cached_response["Body"].read().decode("utf-8")

View file

@ -157,8 +157,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
responses_api_request["metadata"] = value
elif key in ("previous_response_id"):
responses_api_request["previous_response_id"] = value
elif key == "reasoning_effort":
responses_api_request["reasoning"] = self._map_reasoning_effort(value)
responses_api_request["reasoning"] = self._map_reasoning_effort(optional_params.get("reasoning_effort"))
# Get stream parameter from litellm_params if not in optional_params
stream = optional_params.get("stream") or litellm_params.get("stream", False)
@ -452,7 +452,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
responses_tools.append(tool)
return cast(List["ALL_RESPONSES_API_TOOL_PARAMS"], responses_tools)
def _map_reasoning_effort(self, reasoning_effort: str) -> Optional[Reasoning]:
def _map_reasoning_effort(self, reasoning_effort: Optional[str]) -> Reasoning:
if reasoning_effort == "high":
return Reasoning(effort="high", summary="detailed")
elif reasoning_effort == "medium":
@ -462,7 +462,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
return Reasoning(effort="low", summary="auto")
elif reasoning_effort == "minimal":
return Reasoning(effort="minimal", summary="auto")
return None
return Reasoning(summary="auto")
def _map_responses_status_to_finish_reason(self, status: Optional[str]) -> str:
"""Map responses API status to chat completion finish_reason"""

View file

@ -638,18 +638,55 @@ featherless_ai_models: set = set([
])
nebius_models: set = set([
# deepseek models
"deepseek-ai/DeepSeek-R1-0528",
"deepseek-ai/DeepSeek-V3-0324",
"deepseek-ai/DeepSeek-V3",
"deepseek-ai/DeepSeek-R1",
"deepseek-ai/DeepSeek-R1-Distill-Llama-70B",
# google models
"google/gemma-2-2b-it",
"google/gemma-2-9b-it-fast",
# llama models
"meta-llama/Llama-3.3-70B-Instruct",
"meta-llama/Meta-Llama-3.1-70B-Instruct",
"meta-llama/Meta-Llama-3.1-8B-Instruct",
"meta-llama/Meta-Llama-3.1-405B-Instruct",
"NousResearch/Hermes-3-Llama-405B",
# microsoft models
"microsoft/phi-4",
# mistral models
"mistralai/Mistral-Nemo-Instruct-2407",
"mistralai/Devstral-Small-2505",
# moonshot models
"moonshotai/Kimi-K2-Instruct",
# nvidia models
"nvidia/Llama-3_1-Nemotron-Ultra-253B-v1",
"nvidia/Llama-3_3-Nemotron-Super-49B-v1",
# openai models
"openai/gpt-oss-120b",
"openai/gpt-oss-20b",
# qwen models
"Qwen/Qwen3-Coder-480B-A35B-Instruct",
"Qwen/Qwen3-235B-A22B-Instruct-2507",
"Qwen/Qwen3-235B-A22B",
"Qwen/Qwen3-30B-A3B-fast",
"Qwen/Qwen3-30B-A3B",
"Qwen/Qwen3-32B",
"Qwen/Qwen3-14B",
"nvidia/Llama-3_1-Nemotron-Ultra-253B-v1",
"deepseek-ai/DeepSeek-V3-0324",
"deepseek-ai/DeepSeek-V3-0324-fast",
"deepseek-ai/DeepSeek-R1",
"deepseek-ai/DeepSeek-R1-fast",
"meta-llama/Llama-3.3-70B-Instruct-fast",
"Qwen/Qwen2.5-32B-Instruct-fast",
"Qwen/Qwen2.5-Coder-32B-Instruct-fast",
"Qwen/Qwen3-4B-fast",
"Qwen/Qwen2.5-Coder-7B",
"Qwen/Qwen2.5-Coder-32B-Instruct",
"Qwen/Qwen2.5-72B-Instruct",
"Qwen/QwQ-32B",
"Qwen/Qwen3-30B-A3B-Thinking-2507",
"Qwen/Qwen3-30B-A3B-Instruct-2507",
# zai models
"zai-org/GLM-4.5",
"zai-org/GLM-4.5-Air",
# other models
"aaditya/Llama3-OpenBioLLM-70B",
"ProdeusUnity/Stellar-Odyssey-12b-v0.0",
"all-hands/openhands-lm-32b-v0.1",
])
dashscope_models: set = set([

View file

@ -1,7 +1,7 @@
import asyncio
import contextvars
from functools import partial
from typing import Any, Coroutine, Dict, Literal, Optional, Union, cast, overload
from typing import Any, Coroutine, Dict, List, Literal, Optional, Union, cast, overload
import httpx
@ -347,6 +347,7 @@ def image_generation( # noqa: PLR0915
raise ValueError(f"image generation config is not supported for {custom_llm_provider}")
return llm_http_handler.image_generation_handler(
api_key=api_key,
model=model,
prompt=prompt,
image_generation_provider_config=image_generation_config,
@ -676,7 +677,7 @@ def image_variation(
@client
def image_edit(
image: FileTypes,
image: Union[FileTypes, List[FileTypes]],
prompt: str,
model: Optional[str] = None,
mask: Optional[str] = None,
@ -704,6 +705,9 @@ def image_edit(
litellm_call_id: Optional[str] = kwargs.get("litellm_call_id", None)
_is_async = kwargs.pop("async_call", False) is True
#add images / or return a single image
images = image if isinstance(image, list) else [image]
# get llm provider logic
litellm_params = GenericLiteLLMParams(**kwargs)
model, custom_llm_provider, _, _ = get_llm_provider(
@ -752,7 +756,7 @@ def image_edit(
# Call the handler with _is_async flag instead of directly calling the async handler
return base_llm_http_handler.image_edit_handler(
model=model,
image=image,
image=images,
prompt=prompt,
image_edit_provider_config=image_edit_provider_config,
image_edit_optional_request_params=image_edit_request_params,
@ -778,7 +782,7 @@ def image_edit(
@client
async def aimage_edit(
image: FileTypes,
image: Union[FileTypes, List[FileTypes]],
model: str,
prompt: str,
mask: Optional[str] = None,
@ -818,9 +822,11 @@ async def aimage_edit(
model=model, api_base=local_vars.get("base_url", None)
)
images = image if isinstance(image, list) else [image]
func = partial(
image_edit,
image=image,
image=images,
prompt=prompt,
mask=mask,
model=model,

View file

@ -274,12 +274,15 @@ class BraintrustLogger(CustomLogger):
"end": end_time.timestamp(),
}
# Allow metadata override for span name
span_name = metadata.get("span_name", "Chat Completion")
request_data = {
"id": litellm_call_id,
"input": prompt["messages"],
"metadata": clean_metadata,
"tags": tags,
"span_attributes": {"name": "Chat Completion", "type": "llm"},
"span_attributes": {"name": span_name, "type": "llm"},
}
if choices is not None:
request_data["output"] = [choice.dict() for choice in choices]
@ -426,13 +429,16 @@ class BraintrustLogger(CustomLogger):
- api_call_start_time.timestamp()
)
# Allow metadata override for span name
span_name = metadata.get("span_name", "Chat Completion")
request_data = {
"id": litellm_call_id,
"input": prompt["messages"],
"output": output,
"metadata": clean_metadata,
"tags": tags,
"span_attributes": {"name": "Chat Completion", "type": "llm"},
"span_attributes": {"name": span_name, "type": "llm"},
}
if choices is not None:
request_data["output"] = [choice.dict() for choice in choices]

View file

@ -141,6 +141,17 @@ class LangfuseOtelLogger(OpenTelemetry):
value = str(value)
safe_set_attribute(span, enum_attr.value, value)
@staticmethod
def _get_langfuse_otel_host() -> Optional[str]:
"""
Returns the Langfuse OTEL host based on environment variables.
Returned in the following order of precedence:
1. LANGFUSE_OTEL_HOST
2. LANGFUSE_HOST
"""
return os.environ.get("LANGFUSE_OTEL_HOST") or os.environ.get("LANGFUSE_HOST")
@staticmethod
def get_langfuse_otel_config() -> LangfuseOtelConfig:
"""
@ -166,7 +177,7 @@ class LangfuseOtelLogger(OpenTelemetry):
)
# Determine endpoint - default to US cloud
langfuse_host = os.environ.get("LANGFUSE_HOST", None)
langfuse_host = LangfuseOtelLogger._get_langfuse_otel_host()
if langfuse_host:
# If LANGFUSE_HOST is provided, construct OTEL endpoint from it

View file

@ -8,6 +8,7 @@ It searches the vector store for relevant context and appends it to the messages
from typing import TYPE_CHECKING, Dict, List, Optional, Tuple, cast
import litellm
import litellm.vector_stores
from litellm._logging import verbose_logger
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage
@ -192,4 +193,4 @@ class VectorStorePreCallHook(CustomLogger):
modified_messages.insert(-1, cast(AllMessageValues, context_message))
return modified_messages
return messages
return messages

View file

@ -0,0 +1,23 @@
from typing import Dict, Optional
from litellm.types.utils import ProviderSpecificHeader
class ProviderSpecificHeaderUtils:
@staticmethod
def get_provider_specific_headers(
provider_specific_header: Optional[ProviderSpecificHeader],
custom_llm_provider: Optional[str],
) -> Dict:
"""
Get the provider specific headers for the given custom llm provider
Returns:
Optional[Dict]: The provider specific headers for the given custom llm provider
"""
if (
provider_specific_header is not None
and provider_specific_header.get("custom_llm_provider") == custom_llm_provider
):
return provider_specific_header.get("extra_headers", {})
return {}

View file

@ -268,10 +268,9 @@ def get_supported_openai_params( # noqa: PLR0915
from litellm.llms.elevenlabs.audio_transcription.transformation import (
ElevenLabsAudioTranscriptionConfig,
)
return (
ElevenLabsAudioTranscriptionConfig().get_supported_openai_params(
model=model
)
return ElevenLabsAudioTranscriptionConfig().get_supported_openai_params(
model=model
)
elif custom_llm_provider in litellm._custom_providers:
if request_type == "chat_completion":

View file

@ -3846,7 +3846,13 @@ def function_call_prompt(messages: list, functions: list):
function_added_to_prompt = False
for message in messages:
if "system" in message["role"]:
message["content"] += f""" {function_prompt}"""
if isinstance(message["content"], str):
message["content"] += f""" {function_prompt}"""
else:
message["content"].append({
"type": "text",
"text": f""" {function_prompt}"""
})
function_added_to_prompt = True
if function_added_to_prompt is False:

View file

@ -20,7 +20,9 @@ from litellm.litellm_core_utils.redact_messages import LiteLLMLoggingObject
from litellm.litellm_core_utils.thread_pool_executor import executor
from litellm.types.llms.openai import ChatCompletionChunk
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import Delta
from litellm.types.utils import (
Delta,
)
from litellm.types.utils import GenericStreamingChunk as GChunk
from litellm.types.utils import (
ModelResponse,
@ -35,6 +37,12 @@ from .exception_mapping_utils import exception_type
from .llm_response_utils.get_api_base import get_api_base
from .rules import Rules
# Constants for special delta attribute names
AUDIO_ATTRIBUTE = "audio"
IMAGE_ATTRIBUTE = "image"
TOOL_CALLS_ATTRIBUTE = "tool_calls"
FUNCTION_CALL_ATTRIBUTE = "function_call"
def is_async_iterable(obj: Any) -> bool:
"""
@ -766,6 +774,66 @@ class CustomStreamWrapper:
model_response.choices[0].delta = Delta(**_initial_delta)
return model_response
def _has_special_delta_content(self, model_response: ModelResponseStream) -> bool:
"""
Check if the delta contains special content types (tool_calls, function_call, audio, or image).
"""
if len(model_response.choices) == 0:
return False
delta = model_response.choices[0].delta
# Check for tool_calls or function_call
if getattr(delta, TOOL_CALLS_ATTRIBUTE, None) is not None or getattr(delta, FUNCTION_CALL_ATTRIBUTE, None) is not None:
return True
# Check for audio
if hasattr(delta, AUDIO_ATTRIBUTE) and getattr(delta, AUDIO_ATTRIBUTE, None) is not None:
return True
# Check for image
if hasattr(delta, IMAGE_ATTRIBUTE) and getattr(delta, IMAGE_ATTRIBUTE, None) is not None:
return True
return False
def _handle_special_delta_content(self, model_response: ModelResponseStream) -> ModelResponseStream:
"""
Handle special delta content types by stripping role and returning the response.
"""
return self.strip_role_from_delta(model_response)
def _has_special_delta_attribute(self, delta, attribute_name: str) -> bool:
"""
Check if delta has a specific attribute and it's not None.
"""
return delta is not None and getattr(delta, attribute_name, None) is not None
def _copy_delta_attribute(self, source_delta, target_delta, attribute_name: str) -> None:
"""
Copy a specific attribute from source delta to target delta.
"""
setattr(target_delta, attribute_name, getattr(source_delta, attribute_name))
def _has_any_special_delta_attributes(self, delta) -> bool:
"""
Check if delta has any special attributes (audio, image).
"""
special_attributes = [AUDIO_ATTRIBUTE, IMAGE_ATTRIBUTE]
for attribute in special_attributes:
if self._has_special_delta_attribute(delta, attribute):
return True
return False
def _handle_special_delta_attributes(self, delta, model_response: "ModelResponseStream") -> None:
"""
Handle special delta attributes (audio, image) by copying them to model_response.
"""
special_attributes = [AUDIO_ATTRIBUTE, IMAGE_ATTRIBUTE]
for attribute in special_attributes:
if self._has_special_delta_attribute(delta, attribute):
self._copy_delta_attribute(delta, model_response.choices[0].delta, attribute)
def return_processed_chunk_logic( # noqa
self,
completion_obj: Dict[str, Any],
@ -888,20 +956,8 @@ class CustomStreamWrapper:
self.sent_last_chunk = True
return model_response
elif (
model_response.choices[0].delta.tool_calls is not None
or model_response.choices[0].delta.function_call is not None
):
model_response = self.strip_role_from_delta(model_response)
return model_response
elif (
len(model_response.choices) > 0
and hasattr(model_response.choices[0].delta, "audio")
and model_response.choices[0].delta.audio is not None
):
model_response = self.strip_role_from_delta(model_response)
return model_response
elif self._has_special_delta_content(model_response):
return self._handle_special_delta_content(model_response)
else:
if hasattr(model_response, "usage"):
self.chunks.append(model_response)
@ -1374,10 +1430,8 @@ class CustomStreamWrapper:
)
)
model_response.choices[0].delta = Delta()
elif (
delta is not None and getattr(delta, "audio", None) is not None
):
model_response.choices[0].delta.audio = delta.audio
elif self._has_any_special_delta_attributes(delta):
self._handle_special_delta_attributes(delta, model_response)
else:
try:
delta = (

View file

@ -529,7 +529,7 @@ def _get_count_function(
encoding = tiktoken.get_encoding("cl100k_base")
def count_tokens(text: str) -> int:
return len(encoding.encode(text))
return len(encoding.encode(text, disallowed_special=()))
else:
raise ValueError("Unsupported tokenizer type")

View file

@ -1257,6 +1257,10 @@ class BaseLLMHTTPHandler:
stream: Optional[bool] = False,
kwargs: Optional[Dict[str, Any]] = None,
) -> Union[AnthropicMessagesResponse, AsyncIterator]:
from litellm.litellm_core_utils.get_provider_specific_headers import (
ProviderSpecificHeaderUtils,
)
if client is None or not isinstance(client, AsyncHTTPHandler):
async_httpx_client = get_async_httpx_client(
llm_provider=litellm.LlmProviders.ANTHROPIC
@ -1270,10 +1274,9 @@ class BaseLLMHTTPHandler:
Optional[litellm.types.utils.ProviderSpecificHeader],
kwargs.get("provider_specific_header", None),
)
extra_headers = (
provider_specific_header.get("extra_headers", {})
if provider_specific_header
else {}
extra_headers = ProviderSpecificHeaderUtils.get_provider_specific_headers(
provider_specific_header=provider_specific_header,
custom_llm_provider=custom_llm_provider,
)
(
headers,
@ -2678,6 +2681,7 @@ class BaseLLMHTTPHandler:
_is_async: bool = False,
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
api_key: Optional[str] = None,
) -> Union[
ImageResponse,
Coroutine[Any, Any, ImageResponse],
@ -2702,6 +2706,7 @@ class BaseLLMHTTPHandler:
client=client if isinstance(client, AsyncHTTPHandler) else None,
fake_stream=fake_stream,
litellm_metadata=litellm_metadata,
api_key=api_key,
)
if client is None or not isinstance(client, HTTPHandler):
@ -2712,7 +2717,7 @@ class BaseLLMHTTPHandler:
sync_httpx_client = client
headers = image_generation_provider_config.validate_environment(
api_key=litellm_params.get("api_key", None),
api_key=api_key,
headers=image_generation_optional_request_params.get("extra_headers", {})
or {},
model=model,
@ -2795,6 +2800,7 @@ class BaseLLMHTTPHandler:
client: Optional[Union[HTTPHandler, AsyncHTTPHandler]] = None,
fake_stream: bool = False,
litellm_metadata: Optional[Dict[str, Any]] = None,
api_key: Optional[str] = None,
) -> ImageResponse:
"""
Async version of the image generation handler.
@ -2809,7 +2815,7 @@ class BaseLLMHTTPHandler:
async_httpx_client = client
headers = image_generation_provider_config.validate_environment(
api_key=litellm_params.get("api_key", None),
api_key=api_key,
headers=image_generation_optional_request_params.get("extra_headers", {})
or {},
model=model,

View file

@ -75,8 +75,36 @@ class GithubCopilotConfig(OpenAIConfig):
initiator = self._determine_initiator(messages)
validated_headers["X-Initiator"] = initiator
# Add Copilot-Vision-Request header if request contains images
if self._has_vision_content(messages):
validated_headers["Copilot-Vision-Request"] = "true"
return validated_headers
def get_supported_openai_params(self, model: str) -> list:
"""
Get supported OpenAI parameters for GitHub Copilot.
For Claude models that support extended thinking (Claude 4 family and Claude 3-7), includes thinking and reasoning_effort parameters.
For other models, returns standard OpenAI parameters (which may include reasoning_effort for o-series models).
"""
from litellm.utils import supports_reasoning
# Get base OpenAI parameters
base_params = super().get_supported_openai_params(model)
# Add Claude-specific parameters for models that support extended thinking
if "claude" in model.lower() and supports_reasoning(
model=model.lower(),
):
if "thinking" not in base_params:
base_params.append("thinking")
# reasoning_effort is not included by parent for Claude models, so add it
if "reasoning_effort" not in base_params:
base_params.append("reasoning_effort")
return base_params
def _determine_initiator(self, messages: List[AllMessageValues]) -> str:
"""
Determine if request is user or agent initiated based on message roles.
@ -87,3 +115,27 @@ class GithubCopilotConfig(OpenAIConfig):
if role in ["tool", "assistant"]:
return "agent"
return "user"
def _has_vision_content(self, messages: List[AllMessageValues]) -> bool:
"""
Check if any message contains vision content (images).
Returns True if any message has content with vision-related types, otherwise False.
Checks for:
- image_url content type (OpenAI format)
- Content items with type 'image_url'
"""
for message in messages:
content = message.get("content")
if isinstance(content, list):
# Check if any content item indicates vision content
for content_item in content:
if isinstance(content_item, dict):
# Check for image_url field (direct image URL)
if "image_url" in content_item:
return True
# Check for type field indicating image content
content_type = content_item.get("type")
if content_type == "image_url":
return True
return False

View file

@ -28,7 +28,6 @@ class GithubCopilotError(BaseLLMException):
)
class GetDeviceCodeError(GithubCopilotError):
pass

View file

@ -49,15 +49,8 @@ async def make_call(
model_response = ModelResponse(**response.json())
completion_stream = MockResponseIterator(model_response=model_response)
else:
# Use aiter_text with explicit UTF-8 encoding to avoid ASCII encoding errors
async def utf8_aiter_lines():
async for line in response.aiter_text(encoding='utf-8'):
for line_part in line.splitlines(keepends=True):
if line_part.strip():
yield line_part.rstrip('\r\n')
completion_stream = ModelResponseIterator(
streaming_response=utf8_aiter_lines(), sync_stream=False
streaming_response=response.aiter_lines(), sync_stream=False
)
# LOGGING
logging_obj.post_call(
@ -100,15 +93,8 @@ def make_sync_call(
model_response = ModelResponse(**response.json())
completion_stream = MockResponseIterator(model_response=model_response)
else:
# Use iter_text with explicit UTF-8 encoding to avoid ASCII encoding errors
def utf8_iter_lines():
for line in response.iter_text(encoding='utf-8'):
for line_part in line.splitlines(keepends=True):
if line_part.strip():
yield line_part.rstrip('\r\n')
completion_stream = ModelResponseIterator(
streaming_response=utf8_iter_lines(), sync_stream=True
streaming_response=response.iter_lines(), sync_stream=True
)
# LOGGING

View file

@ -254,9 +254,7 @@ def _filter_anyof_fields(schema_dict: Dict[str, Any]) -> Dict[str, Any]:
item["title"] = title
if description:
item["description"] = description
return {"anyOf": any_of}
else:
return schema_dict
return {"anyOf": any_of}
return schema_dict

View file

@ -35,6 +35,7 @@ from litellm.types.llms.openai import (
ChatCompletionFileObject,
ChatCompletionImageObject,
ChatCompletionTextObject,
ChatCompletionUserMessage,
)
from litellm.types.llms.vertex_ai import *
from litellm.types.llms.vertex_ai import (
@ -475,6 +476,13 @@ async def async_transform_request_body(
optional_params=optional_params,
)
def _default_user_message_when_system_message_passed() -> ChatCompletionUserMessage:
"""
Returns a default user message when a "system" message is passed in gemini fails.
This adds a blank user message to the messages list, to ensure that gemini doesn't fail the request.
"""
return ChatCompletionUserMessage(content=".", role="user")
def _transform_system_message(
supports_system_message: bool, messages: List[AllMessageValues]
@ -510,6 +518,13 @@ def _transform_system_message(
messages.pop(idx)
if len(system_content_blocks) > 0:
#########################################################
# If no messages are passed in, add a blank user message
# Relevant Issue - https://github.com/BerriAI/litellm/issues/13769
#########################################################
if len(messages) == 0:
messages.append(_default_user_message_when_system_message_passed())
#########################################################
return SystemInstructions(parts=system_content_blocks), messages
return None, messages

View file

@ -46,6 +46,7 @@ from litellm.types.llms.openai import (
ChatCompletionToolCallChunk,
ChatCompletionToolCallFunctionChunk,
ChatCompletionToolParamFunctionChunk,
ImageURLObject,
OpenAIChatCompletionFinishReason,
)
from litellm.types.llms.vertex_ai import (
@ -89,11 +90,12 @@ from .transformation import (
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.utils import ModelResponseStream
from litellm.types.utils import ModelResponseStream, StreamingChoices
LoggingClass = LiteLLMLoggingObj
else:
LoggingClass = Any
StreamingChoices = Any
class VertexAIBaseConfig:
@ -418,8 +420,10 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
@staticmethod
def _map_reasoning_effort_to_thinking_budget(
reasoning_effort: str,
reasoning_effort: Optional[str],
) -> GeminiThinkingConfig:
if not reasoning_effort:
return { "includeThoughts": True }
if reasoning_effort == "low":
return {
"thinkingBudget": DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET,
@ -614,6 +618,12 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
optional_params = self._add_tools_to_optional_params(
optional_params, [_tools]
)
if supports_reasoning(model):
optional_params["thinkingConfig"] = (
VertexGeminiConfig._map_reasoning_effort_to_thinking_budget(
non_default_params.get("reasoning_effort")
)
)
if litellm.vertex_ai_safety_settings is not None:
optional_params["safety_settings"] = litellm.vertex_ai_safety_settings
@ -774,8 +784,9 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
elif "inlineData" in part:
mime_type = part["inlineData"]["mimeType"]
data = part["inlineData"]["data"]
# Check if inline data is audio - if so, exclude from text content
if mime_type.startswith("audio/"):
# Check if inline data is audio or image - if so, exclude from text content
# Images and audio are now handled separately in their respective response fields
if mime_type.startswith("audio/") or mime_type.startswith("image/"):
continue
_content_str += "data:{};base64,{}".format(mime_type, data)
@ -790,6 +801,23 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
content_str += _content_str
return content_str, reasoning_content_str
def _extract_image_response_from_parts(
self, parts: List[HttpxPartType]
) -> Optional[ImageURLObject]:
"""Extract image response from parts if present"""
for part in parts:
if "inlineData" in part:
mime_type = part["inlineData"]["mimeType"]
data = part["inlineData"]["data"]
if mime_type.startswith("image/"):
# Convert base64 data to data URI format
data_uri = f"data:{mime_type};base64,{data}"
return ImageURLObject(
url=data_uri,
detail="auto"
)
return None
def _extract_audio_response_from_parts(
self, parts: List[HttpxPartType]
@ -1108,6 +1136,75 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
elif web_search_queries:
web_search_requests = len(grounding_metadata)
return web_search_requests
@staticmethod
def _create_streaming_choice(
chat_completion_message: ChatCompletionResponseMessage,
candidate: Candidates,
idx: int,
tools: Optional[List[ChatCompletionToolCallChunk]],
functions: Optional[ChatCompletionToolCallFunctionChunk],
chat_completion_logprobs: Optional[ChoiceLogprobs],
image_response: Optional[ImageURLObject],
) -> StreamingChoices:
"""
Helper method to create a streaming choice object for Vertex AI
"""
from litellm.types.utils import Delta, StreamingChoices
# create a streaming choice object
choice = StreamingChoices(
finish_reason=VertexGeminiConfig._check_finish_reason(
chat_completion_message, candidate.get("finishReason")
),
index=candidate.get("index", idx),
delta=Delta(
content=chat_completion_message.get("content"),
reasoning_content=chat_completion_message.get(
"reasoning_content"
),
tool_calls=tools,
image=image_response,
function_call=functions,
),
logprobs=chat_completion_logprobs,
enhancements=None,
)
return choice
@staticmethod
def _extract_candidate_metadata(candidate: Candidates) -> Tuple[List[dict], List[dict], List, List]:
"""
Extract metadata from a single candidate response.
Returns:
grounding_metadata: List[dict]
url_context_metadata: List[dict]
safety_ratings: List
citation_metadata: List
"""
grounding_metadata: List[dict] = []
url_context_metadata: List[dict] = []
safety_ratings: List = []
citation_metadata: List = []
if "groundingMetadata" in candidate:
if isinstance(candidate["groundingMetadata"], list):
grounding_metadata.extend(candidate["groundingMetadata"]) # type: ignore
else:
grounding_metadata.append(candidate["groundingMetadata"]) # type: ignore
if "safetyRatings" in candidate:
safety_ratings.append(candidate["safetyRatings"])
if "citationMetadata" in candidate:
citation_metadata.append(candidate["citationMetadata"])
if "urlContextMetadata" in candidate:
# Add URL context metadata to grounding metadata
url_context_metadata.append(cast(dict, candidate["urlContextMetadata"]))
return grounding_metadata, url_context_metadata, safety_ratings, citation_metadata
@staticmethod
def _process_candidates(
@ -1131,6 +1228,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
grounding_metadata: List[dict] = []
url_context_metadata: List[dict] = []
image_response: Optional[ImageURLObject] = None
safety_ratings: List = []
citation_metadata: List = []
chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"}
@ -1143,21 +1241,18 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
if "content" not in candidate:
continue
if "groundingMetadata" in candidate:
if isinstance(candidate["groundingMetadata"], list):
grounding_metadata.extend(candidate["groundingMetadata"]) # type: ignore
else:
grounding_metadata.append(candidate["groundingMetadata"]) # type: ignore
if "safetyRatings" in candidate:
safety_ratings.append(candidate["safetyRatings"])
if "citationMetadata" in candidate:
citation_metadata.append(candidate["citationMetadata"])
if "urlContextMetadata" in candidate:
# Add URL context metadata to grounding metadata
url_context_metadata.append(cast(dict, candidate["urlContextMetadata"]))
# Extract metadata using helper function
(
candidate_grounding_metadata,
candidate_url_context_metadata,
candidate_safety_ratings,
candidate_citation_metadata,
) = VertexGeminiConfig._extract_candidate_metadata(candidate)
grounding_metadata.extend(candidate_grounding_metadata)
url_context_metadata.extend(candidate_url_context_metadata)
safety_ratings.extend(candidate_safety_ratings)
citation_metadata.extend(candidate_citation_metadata)
if "parts" in candidate["content"]:
(
@ -1172,18 +1267,25 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
parts=candidate["content"]["parts"]
)
)
image_response = (
VertexGeminiConfig()._extract_image_response_from_parts(
parts=candidate["content"]["parts"]
)
)
if audio_response is not None:
cast(Dict[str, Any], chat_completion_message)[
"audio"
] = audio_response
chat_completion_message["content"] = None # OpenAI spec
elif image_response is not None:
# Handle image response - combine with text content into structured format
cast(Dict[str, Any], chat_completion_message)["image"] = image_response
elif content is not None:
chat_completion_message["content"] = content
if reasoning_content is not None:
chat_completion_message["reasoning_content"] = reasoning_content
(
functions,
tools,
@ -1206,24 +1308,14 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
chat_completion_message["function_call"] = functions
if isinstance(model_response, ModelResponseStream):
from litellm.types.utils import Delta, StreamingChoices
# create a streaming choice object
choice = StreamingChoices(
finish_reason=VertexGeminiConfig._check_finish_reason(
chat_completion_message, candidate.get("finishReason")
),
index=candidate.get("index", idx),
delta=Delta(
content=chat_completion_message.get("content"),
reasoning_content=chat_completion_message.get(
"reasoning_content"
),
tool_calls=tools,
function_call=functions,
),
logprobs=chat_completion_logprobs,
enhancements=None,
choice = VertexGeminiConfig._create_streaming_choice(
chat_completion_message=chat_completion_message,
candidate=candidate,
idx=idx,
tools=tools,
functions=functions,
chat_completion_logprobs=chat_completion_logprobs,
image_response=image_response
)
model_response.choices.append(choice)
elif isinstance(model_response, ModelResponse):

View file

@ -61,6 +61,9 @@ from litellm.exceptions import LiteLLMUnknownProvider
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.audio_utils.utils import get_audio_file_for_health_check
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.get_provider_specific_headers import (
ProviderSpecificHeaderUtils,
)
from litellm.litellm_core_utils.health_check_utils import (
_create_health_check_response,
_filter_model_params,
@ -1107,11 +1110,11 @@ def completion( # type: ignore # noqa: PLR0915
api_key=api_key,
)
if (
provider_specific_header is not None
and provider_specific_header["custom_llm_provider"] == custom_llm_provider
):
headers.update(provider_specific_header["extra_headers"])
if provider_specific_header is not None:
headers.update(ProviderSpecificHeaderUtils.get_provider_specific_headers(
provider_specific_header=provider_specific_header,
custom_llm_provider=custom_llm_provider,
))
if model_response is not None and hasattr(model_response, "_hidden_params"):
model_response._hidden_params["custom_llm_provider"] = custom_llm_provider
@ -1253,6 +1256,7 @@ def completion( # type: ignore # noqa: PLR0915
additional_drop_params=kwargs.get("additional_drop_params"),
remove_sensitive_keys=True,
add_provider_specific_params=True,
provider_config=provider_config,
)
if litellm.add_function_to_prompt and optional_params.get(
@ -2169,8 +2173,18 @@ def completion( # type: ignore # noqa: PLR0915
or "https://api.anthropic.com/v1/complete"
)
if api_base is not None and not api_base.endswith("/v1/complete"):
# Check if we should disable automatic URL suffix appending
disable_url_suffix = get_secret_bool("LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX")
if (
api_base is not None
and not disable_url_suffix
and not api_base.endswith("/v1/complete")
):
api_base += "/v1/complete"
elif disable_url_suffix:
verbose_logger.debug(
"LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX is set, skipping /v1/complete suffix"
)
response = base_llm_http_handler.completion(
model=model,
@ -2206,8 +2220,18 @@ def completion( # type: ignore # noqa: PLR0915
or "https://api.anthropic.com/v1/messages"
)
if api_base is not None and not api_base.endswith("/v1/messages"):
# Check if we should disable automatic URL suffix appending
disable_url_suffix = get_secret_bool("LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX")
if (
api_base is not None
and not disable_url_suffix
and not api_base.endswith("/v1/messages")
):
api_base += "/v1/messages"
elif disable_url_suffix:
verbose_logger.debug(
"LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX is set, skipping /v1/messages suffix"
)
response = anthropic_chat_completions.completion(
model=model,

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1,36 @@
from typing import TYPE_CHECKING
from litellm.types.guardrails import SupportedGuardrailIntegrations
from .noma import NomaGuardrail
if TYPE_CHECKING:
from litellm.types.guardrails import Guardrail, LitellmParams
def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"):
import litellm
_noma_callback = NomaGuardrail(
guardrail_name=guardrail.get("guardrail_name", ""),
api_key=litellm_params.api_key,
api_base=litellm_params.api_base,
application_id=litellm_params.application_id,
monitor_mode=litellm_params.monitor_mode,
block_failures=litellm_params.block_failures,
event_hook=litellm_params.mode,
default_on=litellm_params.default_on,
)
litellm.logging_callback_manager.add_litellm_callback(_noma_callback)
return _noma_callback
guardrail_initializer_registry = {
SupportedGuardrailIntegrations.NOMA.value: initialize_guardrail,
}
guardrail_class_registry = {
SupportedGuardrailIntegrations.NOMA.value: NomaGuardrail,
}

View file

@ -0,0 +1,403 @@
# +-------------------------------------------------------------+
#
# Noma Security Guardrail Integration for LiteLLM
# https://noma.security
#
# +-------------------------------------------------------------+
import copy
import os
from typing import Any, Dict, Literal, Optional, Union
from urllib.parse import urljoin
from fastapi import HTTPException
import litellm
from litellm import DualCache, ModelResponse
from litellm._logging import verbose_proxy_logger
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import EmbeddingResponse, ImageResponse
class NomaBlockedMessage(HTTPException):
"""Exception raised when Noma guardrail blocks a message"""
def __init__(self, classification_response: dict):
classification = self._filter_triggered_classifications(classification_response)
super().__init__(
status_code=400,
detail={
"error": "Request blocked by Noma guardrail",
"details": classification,
},
)
def _filter_triggered_classifications(
self,
response_dict: dict,
) -> dict:
"""Filter and return only triggered classifications"""
filtered_response = copy.deepcopy(response_dict)
# Filter prompt classifications if present
if filtered_response.get("prompt"):
filtered_response["prompt"] = self.filter_classification_object(
filtered_response["prompt"]
)
# Filter response classifications if present
if filtered_response.get("response"):
filtered_response["response"] = self.filter_classification_object(
filtered_response["response"]
)
return filtered_response
def filter_classification_object(
self,
classification_obj: dict,
) -> dict:
"""Filter classification object to only include triggered items"""
if not classification_obj:
return {}
result = {}
for key, value in classification_obj.items():
if value is None:
continue
if key in [
"allowedTopics",
"bannedTopics",
"topicGuardrails",
] and isinstance(value, dict):
filtered_topics = {}
for topic, topic_result in value.items():
if self._is_result_true(topic_result):
filtered_topics[topic] = topic_result
if filtered_topics:
result[key] = filtered_topics
elif key == "sensitiveData" and isinstance(value, dict):
filtered_sensitive = {}
for data_type, data_result in value.items():
if self._is_result_true(data_result):
filtered_sensitive[data_type] = data_result
if filtered_sensitive:
result[key] = filtered_sensitive
elif isinstance(value, dict) and "result" in value:
if self._is_result_true(value):
result[key] = value
return result
def _is_result_true(self, result_obj: Optional[Dict[str, Any]]) -> bool:
"""
Check if a result object has a "result" field that is True.
Args:
result_obj: A dictionary that may contain a "result" field
Returns:
True if the "result" field exists and is True, False otherwise
"""
if not result_obj or not isinstance(result_obj, dict):
return False
return result_obj.get("result") is True
class NomaGuardrail(CustomGuardrail):
"""
Noma Security Guardrail for LiteLLM
This guardrail integrates with Noma Security's AI-DR API to provide
content moderation and safety checks for LLM inputs and outputs.
"""
_DEFAULT_API_BASE = "https://api.noma.security/"
_AIDR_ENDPOINT = "/ai-dr/v1/prompt/scan/aggregate"
def __init__(
self,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
application_id: Optional[str] = None,
monitor_mode: Optional[bool] = None,
block_failures: Optional[bool] = None,
**kwargs,
):
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
self.api_key = api_key or os.environ.get("NOMA_API_KEY")
self.api_base = api_base or os.environ.get(
"NOMA_API_BASE", NomaGuardrail._DEFAULT_API_BASE
)
self.application_id = application_id or os.environ.get(
"NOMA_APPLICATION_ID", "litellm"
)
if monitor_mode is None:
self.monitor_mode = (
os.environ.get("NOMA_MONITOR_MODE", "false").lower() == "true"
)
else:
self.monitor_mode = monitor_mode
if block_failures is None:
self.block_failures = (
os.environ.get("NOMA_BLOCK_FAILURES", "true").lower() == "true"
)
else:
self.block_failures = block_failures
super().__init__(**kwargs)
async def async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict,
call_type: Literal[
"completion",
"text_completion",
"embeddings",
"image_generation",
"moderation",
"audio_transcription",
"pass_through_endpoint",
"rerank",
"mcp_call",
],
) -> Optional[Union[Exception, str, dict]]:
verbose_proxy_logger.debug("Running Noma pre-call hook")
if (
self.should_run_guardrail(
data=data, event_type=GuardrailEventHooks.pre_call
)
is False
):
return data
try:
return await self._check_user_message(data, user_api_key_dict)
except NomaBlockedMessage:
raise
except Exception as e:
verbose_proxy_logger.error(f"Noma pre-call hook failed: {str(e)}")
if self.block_failures and not self.monitor_mode:
raise
return data
async def async_moderation_hook(
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
call_type: Literal[
"completion",
"embeddings",
"image_generation",
"moderation",
"audio_transcription",
"responses",
"mcp_call",
],
) -> Union[Exception, str, dict, None]:
event_type: GuardrailEventHooks = GuardrailEventHooks.during_call
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return data
try:
return await self._check_user_message(data, user_api_key_dict)
except NomaBlockedMessage:
raise
except Exception as e:
verbose_proxy_logger.error(f"Noma moderation hook failed: {str(e)}")
if self.block_failures and not self.monitor_mode:
raise
return data
async def async_post_call_success_hook(
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
response: Union[Any, ModelResponse, EmbeddingResponse, ImageResponse],
):
event_type: GuardrailEventHooks = GuardrailEventHooks.post_call
if self.should_run_guardrail(data=data, event_type=event_type) is not True:
return response
try:
return await self._check_llm_response(data, response, user_api_key_dict)
except NomaBlockedMessage:
raise
except Exception as e:
verbose_proxy_logger.error(f"Noma post-call hook failed: {str(e)}")
if self.block_failures and not self.monitor_mode:
raise
return response
async def _check_user_message(
self,
request_data: dict,
user_auth: UserAPIKeyAuth,
) -> Union[Exception, str, dict, None]:
"""Check user message for policy violations"""
extra_data = self.get_guardrail_dynamic_request_body_params(request_data)
user_message = await self._extract_user_message(request_data)
if not user_message:
return request_data
payload = {"request": {"text": user_message}}
response_json = await self._call_noma_api(
payload=payload,
llm_request_id=None,
request_data=request_data,
user_auth=user_auth,
extra_data=extra_data,
)
await self._check_verdict("user", user_message, response_json)
return request_data
async def _check_llm_response(
self,
request_data: dict,
response: Union[Any, ModelResponse, EmbeddingResponse, ImageResponse],
user_auth: UserAPIKeyAuth,
) -> Union[Exception, ModelResponse, Any]:
"""Check LLM response for policy violations"""
extra_data = self.get_guardrail_dynamic_request_body_params(request_data)
if not isinstance(response, litellm.ModelResponse):
return response
content = None
for choice in response.choices:
if isinstance(choice, litellm.Choices) and choice.message.content:
content = choice.message.content
break
if not content or not isinstance(content, str):
return response
payload = {"response": {"text": content}}
response_json = await self._call_noma_api(
payload=payload,
llm_request_id=response.id,
request_data=request_data,
user_auth=user_auth,
extra_data=extra_data,
)
await self._check_verdict("assistant", content, response_json)
return response
async def _extract_user_message(self, data: dict) -> Optional[str]:
"""Extract the last user message from request data"""
messages = data.get("messages", [])
if not messages:
return None
# Get the last user message
user_messages = [msg for msg in messages if msg.get("role") == "user"]
if not user_messages:
return None
last_user_message = user_messages[-1].get("content", "")
if not last_user_message or not isinstance(last_user_message, str):
return None
return last_user_message
async def _call_noma_api(
self,
payload: dict,
llm_request_id: Optional[str],
request_data: dict,
user_auth: UserAPIKeyAuth,
extra_data: dict,
) -> dict:
call_id = request_data.get("litellm_call_id")
headers = {
"X-Noma-AIDR-Application-ID": self.application_id,
**({"Authorization": f"Bearer {self.api_key}"} if self.api_key else {}),
**({"X-Noma-Request-ID": call_id} if call_id else {}),
}
endpoint = urljoin(
self.api_base or "https://api.noma.security/", NomaGuardrail._AIDR_ENDPOINT
)
response = await self.async_handler.post(
endpoint,
headers=headers,
json={
**payload,
"context": {
"applicationId": extra_data.get("application_id")
or request_data.get("metadata", {})
.get("headers", {})
.get("x-noma-application-id"),
"ipAddress": request_data.get("metadata", {}).get(
"requester_ip_address", None
),
"userId": user_auth.user_email
if user_auth.user_email
else user_auth.user_id,
"sessionId": call_id,
"requestId": llm_request_id,
},
},
)
response.raise_for_status()
return response.json()
async def _check_verdict(
self,
type: Literal["user", "assistant"],
message: str,
response_json: dict,
) -> None:
"""
Check the verdict from the Noma API and raise an exception if needed
"""
if not response_json.get("verdict", True):
msg = str.format(
"Noma guardrail blocked {type} message: {message}",
type=type,
message=message,
)
if self.monitor_mode:
verbose_proxy_logger.warning(msg)
else:
verbose_proxy_logger.debug(msg)
original_response = response_json.get("originalResponse", {})
raise NomaBlockedMessage(original_response)
else:
msg = str.format(
"Noma guardrail allowed {type} message: {message}",
type=type,
message=message,
)
if self.monitor_mode:
verbose_proxy_logger.info(msg)
else:
verbose_proxy_logger.debug(msg)

View file

@ -1,6 +1,6 @@
# litellm/proxy/guardrails/guardrail_hooks/pangea.py
import os
from typing import TYPE_CHECKING, Any, Optional, Protocol, Type
from typing import TYPE_CHECKING, Any, Optional, Type
from fastapi import HTTPException
@ -19,7 +19,7 @@ from litellm.proxy.common_utils.callback_utils import (
add_guardrail_to_applied_guardrails_header,
)
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import LLMResponseTypes, ModelResponse, TextCompletionResponse
from litellm.types.utils import Choices, LLMResponseTypes, ModelResponse, TextCompletionResponse
if TYPE_CHECKING:
from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel
@ -31,14 +31,6 @@ class PangeaGuardrailMissingSecrets(Exception):
pass
class _Transformer(Protocol):
def get_messages(self) -> list[dict]: # noqa: E704
...
def update_original_body(self, prompt_messages: list[dict]) -> Any: # noqa: E704
...
class _TextCompletionRequest:
def __init__(self, body):
self.body = body
@ -53,109 +45,6 @@ class _TextCompletionRequest:
return self.body
class _TextCompletionResponse:
def __init__(self, body):
self.body = body
def get_messages(self) -> list[dict]:
messages = []
for choice in self.body["choices"]:
messages.append({"role": "assistant", "content": choice["text"]})
return messages
def update_original_body(self, prompt_messages: list[dict]) -> Any:
assert len(prompt_messages) == len(self.body["choices"])
for choice, prompt_message in zip(self.body["choices"], prompt_messages):
choice["text"] = prompt_message["content"]
return self.body
class _ChatCompletionRequest:
def __init__(self, body):
self.body = body
def get_messages(self) -> list[dict]:
messages = []
for message in self.body["messages"]:
role = message["role"]
content = message["content"]
if isinstance(content, str):
messages.append({"role": role, "content": content})
if isinstance(content, list):
for content_part in content:
if content_part["type"] == "text":
messages.append({"role": role, "content": content_part["text"]})
return messages
def update_original_body(self, prompt_messages: list[dict]) -> Any:
count = 0
for message in self.body["messages"]:
content = message["content"]
if isinstance(content, str):
message["content"] = prompt_messages[count]["content"]
count += 1
if isinstance(content, list):
for content_part in content:
if content_part["type"] == "text":
content_part["text"] = prompt_messages[count]["content"]
count += 1
assert len(prompt_messages) == count
return self.body
class _ChatCompletionResponse:
def __init__(self, body):
self.body = body
def get_messages(self) -> list[dict]:
messages = []
for choice in self.body["choices"]:
messages.append(
{
"role": choice["message"]["role"],
"content": choice["message"]["content"],
}
)
return messages
def update_original_body(self, prompt_messages: list[dict]) -> Any:
assert len(prompt_messages) == len(self.body["choices"])
for choice, prompt_message in zip(self.body["choices"], prompt_messages):
choice["message"]["content"] = prompt_message["content"]
return self.body
def _get_transformer_for_request(body, call_type) -> Optional[_Transformer]:
match call_type:
case "text_completion" | "atext_completion":
return _TextCompletionRequest(body)
case "completion" | "acompletion":
return _ChatCompletionRequest(body)
return None
def _get_transformer_for_response(body) -> Optional[_Transformer]:
match body:
case TextCompletionResponse():
return _TextCompletionResponse(body)
case ModelResponse():
return _ChatCompletionResponse(body)
return None
class PangeaHandler(CustomGuardrail):
"""
Pangea AI Guardrail handler to interact with the Pangea AI Guard service.
@ -200,7 +89,6 @@ class PangeaHandler(CustomGuardrail):
)
self.pangea_input_recipe = pangea_input_recipe
self.pangea_output_recipe = pangea_output_recipe
self.guardrail_endpoint = f"{self.api_base}/v1/text/guard"
# Pass relevant kwargs to the parent class
super().__init__(guardrail_name=guardrail_name, **kwargs)
@ -208,7 +96,9 @@ class PangeaHandler(CustomGuardrail):
f"Initialized Pangea Guardrail: name={guardrail_name}, recipe={pangea_input_recipe}, api_base={self.api_base}"
)
async def _call_pangea_guard(self, payload: dict, hook_name: str) -> dict:
async def _call_pangea_ai_guard(
self, api: str, payload: dict, hook_name: str
) -> dict:
"""
Makes the API call to the Pangea AI Guard endpoint.
The function itself will raise an error in the case that a response
@ -216,6 +106,7 @@ class PangeaHandler(CustomGuardrail):
should act on.
Args:
api (str): Which API to use (text/guard or v1beta/guard)
payload (dict): The request payload.
request_data (dict): Original request data (used for logging/headers).
hook_name (str): Name of the hook calling this function (for logging).
@ -227,62 +118,84 @@ class PangeaHandler(CustomGuardrail):
Returns:
list[dict]: The original response body
"""
endpoint = f"{self.api_base}/{api}"
headers = {
"Authorization": f"Bearer {self.api_key}",
"Content-Type": "application/json",
}
try:
verbose_proxy_logger.debug(
f"Pangea Guardrail ({hook_name}): Calling endpoint {self.guardrail_endpoint} with payload: {payload}"
)
response = await self.async_handler.post(
url=self.guardrail_endpoint, json=payload, headers=headers
)
response.raise_for_status() # Raise HTTPError for bad responses (4xx or 5xx)
result = response.json()
verbose_proxy_logger.debug(
f"Pangea Guardrail ({hook_name}): Received response: {result}"
verbose_proxy_logger.debug(
f"Pangea Guardrail ({hook_name}): Calling endpoint {endpoint} with payload: {payload}"
)
response = await self.async_handler.post(
url=endpoint, json=payload, headers=headers
)
response.raise_for_status()
result = response.json()
if result.get("result", {}).get("blocked"):
verbose_proxy_logger.warning(
f"Pangea Guardrail ({hook_name}): Request blocked. Response: {result}"
)
# Check if the request was blocked
if result.get("result", {}).get("blocked") is True:
verbose_proxy_logger.warning(
f"Pangea Guardrail ({hook_name}): Request blocked. Response: {result}"
)
raise HTTPException(
status_code=400, # Bad Request, indicating violation
detail={
"error": "Violated Pangea guardrail policy",
"guardrail_name": self.guardrail_name,
"pangea_response": result.get("result"),
},
)
else:
verbose_proxy_logger.info(
f"Pangea Guardrail ({hook_name}): Request passed. Response: {result.get('result', {}).get('detectors')}"
)
return result
except HTTPException as e:
# Re-raise HTTPException if it's the one we raised for blocking
raise e
except Exception as e:
verbose_proxy_logger.error(
f"Pangea Guardrail ({hook_name}): Error calling API: {e}. Response text: {getattr(e, 'response', None) and getattr(e.response, 'text', None)}" # type: ignore
)
# Decide if you want to block by default on error, or allow through
# Raising an exception here will block the request.
# To allow through on error, you might just log and return.
raise HTTPException(
status_code=500,
status_code=400, # Bad Request, indicating violation
detail={
"error": "Error communicating with Pangea Guardrail",
"error": "Violated Pangea guardrail policy",
"guardrail_name": self.guardrail_name,
"exception": str(e),
},
) from e
)
verbose_proxy_logger.info(
f"Pangea Guardrail ({hook_name}): Request passed. Response: {result.get('result', {}).get('detectors')}"
)
return result
async def _async_pre_call_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
cache: DualCache,
data: dict,
call_type: str
):
transformer = None
messages: Any = None
if call_type == "text_completion" or call_type == "atext_completion":
transformer = _TextCompletionRequest(data)
messages = transformer.get_messages()
else:
messages = data.get("messages")
ai_guard_payload = {
"debug": False,
"input": {
"messages": messages, # type: ignore
"tools": data.get("tools")
},
"event_type": "input",
}
if self.pangea_input_recipe:
ai_guard_payload["recipe"] = self.pangea_input_recipe
ai_guard_response = await self._call_pangea_ai_guard(
"v1beta/guard", ai_guard_payload, "async_pre_call_hook"
)
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=self.guardrail_name
)
if not ai_guard_response.get("result", {}).get("transformed"):
return
output = ai_guard_response.get("result", {}).get("output", {})
if call_type == "text_completion" or call_type == "atext_completion":
data = transformer.update_original_body(output["messages"]) # type: ignore
else:
data["messages"] = output["messages"]
return data
@log_guardrail_information
async def async_pre_call_hook(
@ -299,50 +212,75 @@ class PangeaHandler(CustomGuardrail):
)
return data
transformer = _get_transformer_for_request(data, call_type)
if not transformer:
verbose_proxy_logger.warning(
f"Pangea Guardrail (async_pre_call_hook): Skipping guardrail {self.guardrail_name}"
f" because we cannot determine type of request: call_type '{call_type}'"
)
return
messages = transformer.get_messages()
if not messages:
verbose_proxy_logger.warning(
f"Pangea Guardrail (async_pre_call_hook): Skipping guardrail {self.guardrail_name}"
" because messages is empty."
)
return
ai_guard_payload = {
"debug": False, # Or make this configurable if needed
"messages": messages,
}
if self.pangea_input_recipe:
ai_guard_payload["recipe"] = self.pangea_input_recipe
ai_guard_response = await self._call_pangea_guard(
ai_guard_payload, "async_pre_call_hook"
)
# Add guardrail name to header if passed
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=self.guardrail_name
)
prompt_messages = ai_guard_response.get("result", {}).get("prompt_messages", [])
try:
return transformer.update_original_body(prompt_messages)
return await self._async_pre_call_hook(user_api_key_dict, cache, data, call_type)
except HTTPException:
raise
except Exception as e:
raise HTTPException(
status_code=500,
detail={
"error": "Failed to update original request body",
"error": "Error in Pangea Guardrail",
"guardrail_name": self.guardrail_name,
"exceptions": str(e),
},
}
) from e
async def _async_post_call_success_hook(
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
# This union isn't actually correct -- it can get other response types depending on the API called
response: LLMResponseTypes,
):
if isinstance(response, TextCompletionResponse):
# Assume the earlier call type as well
input_messages = _TextCompletionRequest(data).get_messages()
if not isinstance(response, ModelResponse):
return
else:
input_messages = data.get("messages")
if choices := response.get("choices"):
if isinstance(choices, list):
serialized_choices = []
for c in choices:
if isinstance(c, Choices):
try:
serialized_choices.append(c.model_dump())
except Exception:
serialized_choices.append(c.dict())
else:
serialized_choices.append(c)
choices = serialized_choices
ai_guard_payload = {
"debug": False,
"input": {
"messages": input_messages,
"tools": data.get("tools"),
"choices": choices,
},
"event_type": "output",
}
if self.pangea_output_recipe:
ai_guard_payload["recipe"] = self.pangea_output_recipe
ai_guard_response = await self._call_pangea_ai_guard(
"v1beta/guard", ai_guard_payload, "async_pre_call_hook"
)
add_guardrail_to_applied_guardrails_header(
request_data=data, guardrail_name=self.guardrail_name
)
if not ai_guard_response.get("result", {}).get("transformed"):
return
output = ai_guard_response.get("result", {}).get("output", {})
response.choices = output["choices"]
return response
@log_guardrail_information
async def async_post_call_success_hook(
self,
@ -365,39 +303,18 @@ class PangeaHandler(CustomGuardrail):
f"Pangea Guardrail (async_pre_call_hook): Guardrail is disabled {self.guardrail_name}."
)
return data
transformer = _get_transformer_for_response(response)
if not transformer:
verbose_proxy_logger.warning(
f"Pangea Guardrail (async_post_call_success_hook): Skipping guardrail {self.guardrail_name}"
" because we cannot determine type of request"
)
return
messages = transformer.get_messages()
verbose_proxy_logger.warning(f"GOT MESSAGES: {messages}")
ai_guard_payload = {
"debug": False, # Or make this configurable if needed
"messages": messages,
}
if self.pangea_output_recipe:
ai_guard_payload["recipe"] = self.pangea_output_recipe
ai_guard_response = await self._call_pangea_guard(
ai_guard_payload, "post_call_success_hook"
)
prompt_messages = ai_guard_response.get("result", {}).get("prompt_messages", [])
try:
return transformer.update_original_body(prompt_messages)
return await self._async_post_call_success_hook(data, user_api_key_dict, response)
except HTTPException:
raise
except Exception as e:
raise HTTPException(
status_code=500,
detail={
"error": "Failed to update original response body",
"error": "Error in Pangea Guardrail",
"guardrail_name": self.guardrail_name,
"exceptions": str(e),
},
}
) from e
@staticmethod

View file

@ -489,8 +489,10 @@ class LiteLLMProxyRequestSetup:
@staticmethod
def add_key_level_controls(
key_metadata: dict, data: dict, _metadata_variable_name: str
key_metadata: Optional[dict], data: dict, _metadata_variable_name: str
):
if key_metadata is None:
return data
if "cache" in key_metadata:
data["cache"] = {}
if isinstance(key_metadata["cache"], dict):

View file

@ -531,6 +531,19 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
subpath = subpath[1:]
return base_target + subpath
@staticmethod
def _update_stream_param_based_on_request_body(
parsed_body: dict,
stream: Optional[bool] = None,
) -> Optional[bool]:
"""
If stream is provided in the request body, use it.
Otherwise, use the stream parameter passed to the `pass_through_request` function
"""
if "stream" in parsed_body:
return parsed_body.get("stream", stream)
return stream
async def pass_through_request( # noqa: PLR0915
@ -686,6 +699,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,
)
if stream:
req = async_client.build_request(
"POST",

View file

@ -1,17 +1,23 @@
model_list:
- model_name: bedrock/converse/us.anthropic.claude-sonnet-4-20250514-v1:0
- model_name: anthropic/*
litellm_params:
model: bedrock/converse/us.anthropic.claude-sonnet-4-20250514-v1:0
model: anthropic/*
api_key: os.environ/OPENAI_API_KEY_IJ
- model_name: bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0
litellm_params:
model: bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0
- model_name: bedrock/converse/us.anthropic.claude-sonnet-4-20250514-v1:0
litellm_params:
model: bedrock/converse/us.anthropic.claude-sonnet-4-20250514-v1:0
router_settings:
fallbacks: [
{"anthropic/claude-opus-4-20250514":
{
"bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"
}
}
]
litellm_settings:
callbacks: ["datadog_llm_observability"]
guardrails:
- guardrail_name: "bedrock-pre-guard"
litellm_params:
guardrail: bedrock # supported values: "aporia", "bedrock", "lakera"
mode: "during_call"
guardrailIdentifier: ff6ujrregl1q
guardrailVersion: "DRAFT"

View file

@ -2,7 +2,7 @@
Handler for transforming responses api requests to litellm.completion requests
"""
from typing import Any, Coroutine, Optional, Union
from typing import Any, Coroutine, Dict, Optional, Union
import litellm
from litellm.responses.litellm_completion_transformation.streaming_iterator import (
@ -30,6 +30,7 @@ class LiteLLMCompletionTransformationHandler:
custom_llm_provider: Optional[str] = None,
_is_async: bool = False,
stream: Optional[bool] = None,
extra_headers: Optional[Dict[str, Any]] = None,
**kwargs,
) -> Union[
ResponsesAPIResponse,
@ -45,6 +46,7 @@ class LiteLLMCompletionTransformationHandler:
responses_api_request=responses_api_request,
custom_llm_provider=custom_llm_provider,
stream=stream,
extra_headers=extra_headers,
**kwargs,
)
)

View file

@ -99,6 +99,7 @@ class LiteLLMCompletionResponsesConfig:
responses_api_request: ResponsesAPIOptionalRequestParams,
custom_llm_provider: Optional[str] = None,
stream: Optional[bool] = None,
extra_headers: Optional[Dict[str, Any]] = None,
**kwargs,
) -> dict:
"""
@ -126,6 +127,7 @@ class LiteLLMCompletionResponsesConfig:
"web_search_options": web_search_options,
# litellm specific params
"custom_llm_provider": custom_llm_provider,
"extra_headers": extra_headers,
}
# Responses API `Completed` events require usage, we pass `stream_options` to litellm.completion to include usage

View file

@ -455,6 +455,7 @@ def responses(
custom_llm_provider=custom_llm_provider,
_is_async=_is_async,
stream=stream,
extra_headers=extra_headers,
**kwargs,
)

View file

@ -40,57 +40,24 @@ def simple_shuffle(
Dict: A single healthy deployment
"""
############## Check if 'weight' param set for a weighted pick #################
weight = healthy_deployments[0].get("litellm_params").get("weight", None)
if weight is not None:
# use weight-random pick if rpms provided
weights = [m["litellm_params"].get("weight", 0) for m in healthy_deployments]
verbose_router_logger.debug(f"\nweight {weights}")
total_weight = sum(weights)
weights = [safe_divide(weight, total_weight, 0) for weight in weights]
verbose_router_logger.debug(f"\n weights {weights}")
# Perform weighted random pick
selected_index = random.choices(range(len(weights)), weights=weights)[0]
verbose_router_logger.debug(f"\n selected index, {selected_index}")
deployment = healthy_deployments[selected_index]
verbose_router_logger.info(
f"get_available_deployment for model: {model}, Selected deployment: {llm_router_instance.print_deployment(deployment) or deployment[0]} for model: {model}"
)
return deployment or deployment[0]
############## Check if we can do a RPM/TPM based weighted pick #################
rpm = healthy_deployments[0].get("litellm_params").get("rpm", None)
if rpm is not None:
# use weight-random pick if rpms provided
rpms = [m["litellm_params"].get("rpm", 0) for m in healthy_deployments]
verbose_router_logger.debug(f"\nrpms {rpms}")
total_rpm = sum(rpms)
weights = [safe_divide(rpm, total_rpm, 0) for rpm in rpms]
verbose_router_logger.debug(f"\n weights {weights}")
# Perform weighted random pick
selected_index = random.choices(range(len(rpms)), weights=weights)[0]
verbose_router_logger.debug(f"\n selected index, {selected_index}")
deployment = healthy_deployments[selected_index]
verbose_router_logger.info(
f"get_available_deployment for model: {model}, Selected deployment: {llm_router_instance.print_deployment(deployment) or deployment[0]} for model: {model}"
)
return deployment or deployment[0]
############## Check if we can do a RPM/TPM based weighted pick #################
tpm = healthy_deployments[0].get("litellm_params").get("tpm", None)
if tpm is not None:
# use weight-random pick if rpms provided
tpms = [m["litellm_params"].get("tpm", 0) for m in healthy_deployments]
verbose_router_logger.debug(f"\ntpms {tpms}")
total_tpm = sum(tpms)
weights = [safe_divide(tpm, total_tpm, 0) for tpm in tpms]
verbose_router_logger.debug(f"\n weights {weights}")
# Perform weighted random pick
selected_index = random.choices(range(len(tpms)), weights=weights)[0]
verbose_router_logger.debug(f"\n selected index, {selected_index}")
deployment = healthy_deployments[selected_index]
verbose_router_logger.info(
f"get_available_deployment for model: {model}, Selected deployment: {llm_router_instance.print_deployment(deployment) or deployment[0]} for model: {model}"
)
return deployment or deployment[0]
############## Check if 'weight' or 'rpm' or 'tpm' param set for a weighted pick #################
for weight_by in ["weight", "rpm", "tpm"]:
weight = healthy_deployments[0].get("litellm_params").get(weight_by, None)
if weight is not None:
weights = [m["litellm_params"].get(weight_by, 0) for m in healthy_deployments]
verbose_router_logger.debug(f"\nweight {weights}")
total_weight = sum(weights)
weights = [weight / total_weight for weight in weights]
verbose_router_logger.debug(f"\n weights {weights} by {weight_by}")
# Perform weighted random pick
selected_index = random.choices(range(len(weights)), weights=weights)[0]
verbose_router_logger.debug(f"\n selected index, {selected_index}")
deployment = healthy_deployments[selected_index]
verbose_router_logger.info(
f"get_available_deployment for model: {model}, Selected deployment: {llm_router_instance.print_deployment(deployment) or deployment[0]} for model: {model}"
)
return deployment or deployment[0]
############## No RPM/TPM passed, we do a random pick #################
item = random.choice(healthy_deployments)

View file

@ -40,6 +40,7 @@ class SupportedGuardrailIntegrations(Enum):
AZURE_TEXT_MODERATIONS = "azure/text_moderations"
MODEL_ARMOR = "model_armor"
OPENAI_MODERATION = "openai_moderation"
NOMA = "noma"
class Role(Enum):
SYSTEM = "system"
@ -359,6 +360,23 @@ class PillarGuardrailConfigModel(BaseModel):
)
class NomaGuardrailConfigModel(BaseModel):
"""Configuration parameters for the Noma Security guardrail"""
application_id: Optional[str] = Field(
default=None,
description="Application ID for Noma Security. Defaults to 'litellm' if not provided",
)
monitor_mode: Optional[bool] = Field(
default=None,
description="If True, logs violations without blocking. Defaults to False if not provided",
)
block_failures: Optional[bool] = Field(
default=None,
description="If True, blocks requests on API failures. Defaults to True if not provided",
)
class BaseLitellmParams(BaseModel): # works for new and patch update guardrails
api_key: Optional[str] = Field(
default=None, description="API key for the guardrail service"
@ -445,6 +463,7 @@ class LitellmParams(
LakeraV2GuardrailConfigModel,
LassoGuardrailConfigModel,
PillarGuardrailConfigModel,
NomaGuardrailConfigModel,
BaseLitellmParams,
):
guardrail: str = Field(description="The type of guardrail integration to use")

View file

@ -1,6 +1,5 @@
import json
import time
import uuid
from enum import Enum
from typing import (
TYPE_CHECKING,
@ -14,6 +13,7 @@ from typing import (
Union,
)
import fastuuid as uuid
from aiohttp import FormData
from openai._models import BaseModel as OpenAIObject
from openai.types.audio.transcription_create_params import FileTypes # type: ignore
@ -51,6 +51,7 @@ from .llms.openai import (
ChatCompletionUsageBlock,
FileSearchTool,
FineTuningJob,
ImageURLObject,
OpenAIChatCompletionChunk,
OpenAIFileObject,
OpenAIRealtimeStreamList,
@ -572,6 +573,7 @@ class Message(OpenAIObject):
tool_calls: Optional[List[ChatCompletionMessageToolCall]]
function_call: Optional[FunctionCall]
audio: Optional[ChatCompletionAudioResponse] = None
image: Optional[ImageURLObject] = None
reasoning_content: Optional[str] = None
thinking_blocks: Optional[
List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]]
@ -588,6 +590,7 @@ class Message(OpenAIObject):
function_call=None,
tool_calls: Optional[list] = None,
audio: Optional[ChatCompletionAudioResponse] = None,
image: Optional[ImageURLObject] = None,
provider_specific_fields: Optional[Dict[str, Any]] = None,
reasoning_content: Optional[str] = None,
thinking_blocks: Optional[
@ -621,6 +624,9 @@ class Message(OpenAIObject):
if audio is not None:
init_values["audio"] = audio
if image is not None:
init_values["image"] = image
if thinking_blocks is not None:
init_values["thinking_blocks"] = thinking_blocks
@ -640,6 +646,10 @@ class Message(OpenAIObject):
# OpenAI compatible APIs like mistral API will raise an error if audio is passed in
if hasattr(self, "audio"):
del self.audio
if image is None:
if hasattr(self, "image"):
del self.image
if annotations is None:
# ensure default response matches OpenAI spec
@ -693,6 +703,7 @@ class Delta(OpenAIObject):
function_call=None,
tool_calls=None,
audio: Optional[ChatCompletionAudioResponse] = None,
image: Optional[ImageURLObject] = None,
reasoning_content: Optional[str] = None,
thinking_blocks: Optional[
List[
@ -710,6 +721,7 @@ class Delta(OpenAIObject):
self.function_call: Optional[Union[FunctionCall, Any]] = None
self.tool_calls: Optional[List[Union[ChatCompletionDeltaToolCall, Any]]] = None
self.audio: Optional[ChatCompletionAudioResponse] = None
self.image: Optional[ImageURLObject] = None
self.annotations: Optional[List[ChatCompletionAnnotation]] = None
if reasoning_content is not None:
@ -729,6 +741,11 @@ class Delta(OpenAIObject):
self.annotations = annotations
else:
del self.annotations
if image is not None:
self.image = image
else:
del self.image
if function_call is not None and isinstance(function_call, dict):
self.function_call = FunctionCall(**function_call)

View file

@ -3088,6 +3088,7 @@ def pre_process_non_default_params(
model: str,
remove_sensitive_keys: bool = False,
add_provider_specific_params: bool = False,
provider_config: Optional[BaseConfig] = None,
) -> dict:
"""
Pre-process non-default params to a standardized format
@ -3103,14 +3104,6 @@ def pre_process_non_default_params(
additional_endpoint_specific_params=["messages"],
)
provider_config: Optional[BaseConfig] = None
if custom_llm_provider is not None and custom_llm_provider in [
provider.value for provider in LlmProviders
]:
provider_config = ProviderConfigManager.get_provider_chat_config(
model=model, provider=LlmProviders(custom_llm_provider)
)
if "response_format" in non_default_params:
if provider_config is not None:
non_default_params[

File diff suppressed because it is too large Load diff

35
poetry.lock generated
View file

@ -1703,6 +1703,41 @@ lz4 = ["lz4"]
snappy = ["cramjam"]
zstandard = ["zstandard"]
[[package]]
name = "fastuuid"
version = "0.12.0"
description = "Python bindings to Rust's UUID library."
optional = false
python-versions = ">=3.8"
groups = ["main"]
files = [
{file = "fastuuid-0.12.0-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:22a900ef0956aacf862b460e20541fdae2d7c340594fe1bd6fdcb10d5f0791a9"},
{file = "fastuuid-0.12.0-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:0302f5acf54dc75de30103025c5a95db06d6c2be36829043a0aa16fc170076bc"},
{file = "fastuuid-0.12.0-cp310-cp310-manylinux_2_34_x86_64.whl", hash = "sha256:7946b4a310cfc2d597dcba658019d72a2851612a2cebb949d809c0e2474cf0a6"},
{file = "fastuuid-0.12.0-cp310-cp310-win_amd64.whl", hash = "sha256:a1b6764dd42bf0c46c858fb5ade7b7a3d93b7a27485a7a5c184909026694cd88"},
{file = "fastuuid-0.12.0-cp311-cp311-macosx_11_0_arm64.whl", hash = "sha256:2bced35269315d16fe0c41003f8c9d63f2ee16a59295d90922cad5e6a67d0418"},
{file = "fastuuid-0.12.0-cp311-cp311-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:82106e4b0a24f4f2f73c88f89dadbc1533bb808900740ca5db9bbb17d3b0c824"},
{file = "fastuuid-0.12.0-cp311-cp311-manylinux_2_34_x86_64.whl", hash = "sha256:4db1bc7b8caa1d7412e1bea29b016d23a8d219131cff825b933eb3428f044dca"},
{file = "fastuuid-0.12.0-cp311-cp311-win_amd64.whl", hash = "sha256:07afc8e674e67ac3d35a608c68f6809da5fab470fb4ef4469094fdb32ba36c51"},
{file = "fastuuid-0.12.0-cp312-cp312-macosx_11_0_arm64.whl", hash = "sha256:328694a573fe9dce556b0b70c9d03776786801e028d82f0b6d9db1cb0521b4d1"},
{file = "fastuuid-0.12.0-cp312-cp312-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:02acaea2c955bb2035a7d8e7b3fba8bd623b03746ae278e5fa932ef54c702f9f"},
{file = "fastuuid-0.12.0-cp312-cp312-manylinux_2_34_x86_64.whl", hash = "sha256:ed9f449cba8cf16cced252521aee06e633d50ec48c807683f21cc1d89e193eb0"},
{file = "fastuuid-0.12.0-cp312-cp312-win_amd64.whl", hash = "sha256:0df2ea4c9db96fd8f4fa38d0e88e309b3e56f8fd03675a2f6958a5b082a0c1e4"},
{file = "fastuuid-0.12.0-cp313-cp313-macosx_11_0_arm64.whl", hash = "sha256:7fe2407316a04ee8f06d3dbc7eae396d0a86591d92bafe2ca32fce23b1145786"},
{file = "fastuuid-0.12.0-cp313-cp313-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:b9b31dd488d0778c36f8279b306dc92a42f16904cba54acca71e107d65b60b0c"},
{file = "fastuuid-0.12.0-cp313-cp313-manylinux_2_34_x86_64.whl", hash = "sha256:b19361ee649365eefc717ec08005972d3d1eb9ee39908022d98e3bfa9da59e37"},
{file = "fastuuid-0.12.0-cp313-cp313-win_amd64.whl", hash = "sha256:8fc66b11423e6f3e1937385f655bedd67aebe56a3dcec0cb835351cfe7d358c9"},
{file = "fastuuid-0.12.0-cp38-cp38-macosx_11_0_arm64.whl", hash = "sha256:7b15c54d300279ab20a9cc0579ada9c9f80d1bc92997fc61fb7bf3103d7cb26b"},
{file = "fastuuid-0.12.0-cp38-cp38-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:458f1bc3ebbd76fdb89ad83e6b81ccd3b2a99fa6707cd3650b27606745cfb170"},
{file = "fastuuid-0.12.0-cp38-cp38-manylinux_2_34_x86_64.whl", hash = "sha256:a8f0f83fbba6dc44271a11b22e15838641b8c45612cdf541b4822a5930f6893c"},
{file = "fastuuid-0.12.0-cp38-cp38-win_amd64.whl", hash = "sha256:7cfd2092253d3441f6a8c66feff3c3c009da25a5b3da82bc73737558543632be"},
{file = "fastuuid-0.12.0-cp39-cp39-macosx_11_0_arm64.whl", hash = "sha256:9303617e887429c193d036d47d0b32b774ed3618431123e9106f610d601eb57e"},
{file = "fastuuid-0.12.0-cp39-cp39-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:8790221325b376e1122e95f865753ebf456a9fb8faf0dca4f9bf7a3ff620e413"},
{file = "fastuuid-0.12.0-cp39-cp39-manylinux_2_34_x86_64.whl", hash = "sha256:e4b12d3e23515e29773fa61644daa660ceb7725e05397a986c2109f512579a48"},
{file = "fastuuid-0.12.0-cp39-cp39-win_amd64.whl", hash = "sha256:e41656457c34b5dcb784729537ea64c7d9bbaf7047b480c6c6a64c53379f455a"},
{file = "fastuuid-0.12.0.tar.gz", hash = "sha256:d0bd4e5b35aad2826403f4411937c89e7c88857b1513fe10f696544c03e9bd8e"},
]
[[package]]
name = "filelock"
version = "3.16.1"

View file

@ -1,5 +1,5 @@
[tool.poetry]
name = "litellm_ad"
name = "litellm"
version = "1.76.1"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI, AndrewDoan"]
@ -20,6 +20,7 @@ Documentation = "https://docs.litellm.ai"
[tool.poetry.dependencies]
python = ">=3.8.1,<4.0, !=3.9.7"
fastuuid = ">=0.12.0"
httpx = ">=0.23.0"
openai = ">=1.99.5"
python-dotenv = ">=0.2.0"
@ -155,7 +156,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "1.76.0"
version = "1.76.1"
version_files = [
"pyproject.toml:^version"
]

View file

@ -7,6 +7,7 @@ backoff==2.2.1 # server dep
pyyaml==6.0.2 # server dep
uvicorn==0.29.0 # server dep
gunicorn==23.0.0 # server dep
fastuuid==0.12.0 # for uuid4
uvloop==0.21.0 # uvicorn dep, gives us much better performance under load
boto3==1.36.0 # aws bedrock/sagemaker calls
redis==5.2.1 # redis caching
@ -23,7 +24,7 @@ async_generator==1.10.0 # for async ollama calls
langfuse==2.59.7 # for langfuse self-hosted logging
prometheus_client==0.20.0 # for /metrics endpoint on proxy
ddtrace==2.19.0 # for advanced DD tracing / profiling
orjson==3.10.12 # fast /embedding responses
orjson==3.11.2 # fast /embedding responses
polars==1.31.0 # for data processing
apscheduler==3.10.4 # for resetting budget in background
fastapi-sso==0.16.0 # admin UI, SSO

Binary file not shown.

View file

@ -19,6 +19,9 @@ from litellm.utils import ImageResponse
from litellm.integrations.custom_logger import CustomLogger
from litellm.types.utils import StandardLoggingPayload
# Configure pytest marks to avoid warnings
pytestmark = pytest.mark.asyncio
class TestCustomLogger(CustomLogger):
def __init__(self):
self.standard_logging_payload: Optional[StandardLoggingPayload] = None
@ -35,6 +38,8 @@ TEST_IMAGES = [
open(os.path.join(pwd, "litellm_site.png"), "rb"),
]
SINGLE_TEST_IMAGE = open(os.path.join(pwd, "ishaan_github.png"), "rb")
def get_test_images_as_bytesio():
"""Helper function to get test images as BytesIO objects"""
bytesio_images = []
@ -501,3 +506,157 @@ def test_recraft_image_edit_config():
assert files[0][0] == "image" # Field name (not image[] like OpenAI)
assert files[0][1][1] == mock_image # Image data
assert files[0][1][2] == "image/png" # Content type
@pytest.mark.parametrize("sync_mode", [True, False])
@pytest.mark.flaky(retries=3, delay=2)
@pytest.mark.asyncio
async def test_multiple_vs_single_image_edit(sync_mode):
"""Test that both single and multiple image editing work correctly"""
from litellm import image_edit, aimage_edit
litellm._turn_on_debug()
try:
prompt = "Add a soft blue tint to the image(s)"
# Test single image
if sync_mode:
single_result = image_edit(
prompt=prompt,
model="gpt-image-1",
image=SINGLE_TEST_IMAGE,
)
else:
single_result = await aimage_edit(
prompt=prompt,
model="gpt-image-1",
image=SINGLE_TEST_IMAGE,
)
print("Single image result:", single_result)
ImageResponse.model_validate(single_result)
# Test multiple images
if sync_mode:
multiple_result = image_edit(
prompt=prompt,
model="gpt-image-1",
image=TEST_IMAGES,
)
else:
multiple_result = await aimage_edit(
prompt=prompt,
model="gpt-image-1",
image=TEST_IMAGES,
)
print("Multiple images result:", multiple_result)
ImageResponse.model_validate(multiple_result)
# Both should return valid responses
assert single_result is not None
assert multiple_result is not None
assert single_result.data is not None
assert multiple_result.data is not None
assert len(single_result.data) > 0
assert len(multiple_result.data) > 0
except litellm.ContentPolicyViolationError as e:
pytest.skip(f"Content policy violation: {e}")
@pytest.mark.flaky(retries=3, delay=2)
@pytest.mark.asyncio
async def test_multiple_image_edit_with_different_formats():
"""Test multiple images editing with different file formats and types"""
from litellm import aimage_edit
litellm._turn_on_debug()
try:
prompt = "Create a cohesive artistic style across all images"
# Test with mixed BytesIO and file objects
mixed_images = [
SINGLE_TEST_IMAGE, # File object
get_test_images_as_bytesio()[1] # BytesIO object
]
result = await aimage_edit(
prompt=prompt,
model="gpt-image-1",
image=mixed_images,
)
print("Mixed format images result:", result)
ImageResponse.model_validate(result)
assert result is not None
assert result.data is not None
assert len(result.data) > 0
# Save result if available
if result.data and result.data[0].b64_json:
image_bytes = base64.b64decode(result.data[0].b64_json)
with open("test_multiple_image_edit_mixed.png", "wb") as f:
f.write(image_bytes)
except litellm.ContentPolicyViolationError as e:
pytest.skip(f"Content policy violation: {e}")
@pytest.mark.flaky(retries=3, delay=2)
@pytest.mark.asyncio
async def test_image_edit_array_handling():
"""Test that the image parameter correctly handles both single items and arrays"""
from litellm import aimage_edit
# Mock response
mock_response = {
"created": 1589478378,
"data": [
{
"b64_json": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mP8/5+hHgAHggJ/PchI7wAAAABJRU5ErkJggg=="
}
]
}
class MockResponse:
def __init__(self, json_data, status_code):
self._json_data = json_data
self.status_code = status_code
self.text = json.dumps(json_data)
def json(self):
return self._json_data
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post:
mock_post.return_value = MockResponse(mock_response, 200)
prompt = "Test prompt"
# Test 1: Single image (should be converted to list internally)
result1 = await aimage_edit(
prompt=prompt,
model="gpt-image-1",
image=SINGLE_TEST_IMAGE,
)
# Test 2: Multiple images (already a list)
result2 = await aimage_edit(
prompt=prompt,
model="gpt-image-1",
image=TEST_IMAGES,
)
# Both valid calls should succeed
ImageResponse.model_validate(result1)
ImageResponse.model_validate(result2)
# Verify that both calls were made to the API
assert mock_post.call_count == 2

View file

@ -175,7 +175,7 @@ class TestAimlImageGeneration(BaseImageGenTest):
class TestGoogleImageGen(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
return {"model": "gemini/imagen-4.0-generate-preview-06-06"}
return {"model": "gemini/imagen-4.0-generate-001"}
class TestAzureOpenAIDalle3(BaseImageGenTest):
def get_base_image_generation_call_args(self) -> dict:
@ -330,3 +330,78 @@ async def test_gpt_image_1_with_input_fidelity():
assert captured_kwargs["quality"] == "medium"
assert captured_kwargs["size"] == "1024x1024"
@pytest.mark.asyncio
async def test_aiml_image_generation_with_dynamic_api_key():
"""
Test that when api_key is passed as a dynamic parameter to aimage_generation,
it gets properly used for AIML provider authentication instead of falling back
to environment variables.
This test validates the fix for ensuring dynamic API keys are respected
when making image generation requests to the AIML provider.
"""
from unittest.mock import AsyncMock, patch, MagicMock
import httpx
# Mock AIML response
mock_aiml_response = {
"created": 1703658209,
"data": [
{
"url": "https://example.com/generated_image.png"
}
]
}
# Track captured arguments
captured_headers = None
captured_url = None
captured_json_data = None
def capture_post_call(*args, **kwargs):
nonlocal captured_headers, captured_url, captured_json_data
captured_url = kwargs.get('url') or (args[0] if args else None)
captured_headers = kwargs.get('headers', {})
captured_json_data = kwargs.get('json', {})
# Create a mock response
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = mock_aiml_response
mock_response.text = json.dumps(mock_aiml_response)
return mock_response
# Mock the HTTP client that actually makes the request (sync version for image generation)
with patch('litellm.llms.custom_httpx.http_handler.HTTPHandler.post') as mock_post:
mock_post.side_effect = capture_post_call
# Test with dynamic api_key
test_api_key = "test-dynamic-api-key-12345"
response = await litellm.aimage_generation(
prompt="A cute baby sea otter",
model="aiml/flux-pro/v1.1",
api_key=test_api_key, # This should be used instead of env vars
)
# Validate the response (mocked response processing might not populate data correctly)
assert response is not None
# The most important validations: API key and endpoint usage
# These prove that the dynamic API key was properly used
assert captured_headers is not None
assert "Authorization" in captured_headers
assert captured_headers["Authorization"] == f"Bearer {test_api_key}"
print("TESTCAPTURED HEADERS", captured_headers)
# Validate the correct AIML endpoint was called
assert captured_url is not None
assert "api.aimlapi.com" in captured_url
assert "/v1/images/generations" in captured_url
# Validate the request data
assert captured_json_data is not None
assert captured_json_data["prompt"] == "A cute baby sea otter"
assert captured_json_data["model"] == "flux-pro/v1.1"

View file

@ -119,6 +119,28 @@ class BaseLLMChatTest(ABC):
pytest.skip("Model is overloaded")
assert response.choices[0].message.content is not None
def test_system_message_with_no_user_message(self):
"""
Test that the system message is translated correctly for non-OpenAI providers.
"""
base_completion_call_args = self.get_base_completion_call_args()
messages = [
{
"role": "system",
"content": "Be a good bot!",
},
]
try:
response = self.completion_function(
**base_completion_call_args,
messages=messages,
)
assert response is not None
except litellm.InternalServerError:
pytest.skip("Model is overloaded")
assert response.choices[0].message.content is not None
def test_content_list_handling(self):
"""Check if content list is supported by LLM API"""

View file

@ -261,7 +261,13 @@ def test_gemini_image_generation():
messages=[{"role": "user", "content": "Generate an image of a cat"}],
modalities=["image", "text"],
)
assert response.choices[0].message.content is not None
#########################################################
# Important: Validate we did get an image in the response
#########################################################
assert response.choices[0].message.image is not None
assert response.choices[0].message.image["url"] is not None
assert response.choices[0].message.image["url"].startswith("data:image/png;base64,")
def test_gemini_thinking():
@ -571,3 +577,50 @@ def test_gemini_tool_use():
stop_reason = chunk.choices[0].finish_reason
assert stop_reason is not None
assert stop_reason == "tool_calls"
@pytest.mark.asyncio
async def test_gemini_image_generation_async():
#litellm._turn_on_debug()
response = await litellm.acompletion(
messages=[{"role": "user", "content": "Generate an image of a banana wearing a costume that says LiteLLM"}],
model="gemini/gemini-2.5-flash-image-preview",
)
CONTENT = response.choices[0].message.content
IMAGE_URL = response.choices[0].message.image
print("IMAGE_URL: ", IMAGE_URL)
assert CONTENT is not None
assert IMAGE_URL is not None
assert IMAGE_URL["url"] is not None
assert IMAGE_URL["url"].startswith("data:image/png;base64,")
@pytest.mark.asyncio
async def test_gemini_image_generation_async_stream():
#litellm._turn_on_debug()
response = await litellm.acompletion(
messages=[{"role": "user", "content": "Generate an image of a banana wearing a costume that says LiteLLM"}],
model="gemini/gemini-2.5-flash-image-preview",
stream=True,
)
print("RESPONSE: ", response)
model_response_image = None
async for chunk in response:
print("CHUNK: ", chunk)
if hasattr(chunk.choices[0].delta, "image") and chunk.choices[0].delta.image is not None:
model_response_image = chunk.choices[0].delta.image
print("MODEL_RESPONSE_IMAGE: ", model_response_image)
assert model_response_image is not None
assert model_response_image["url"].startswith("data:image/png;base64,")
break
#########################################################
# Important: Validate we did get an image in the response
#########################################################
assert model_response_image is not None
assert model_response_image["url"].startswith("data:image/png;base64,")

View file

@ -1557,7 +1557,7 @@ def test_azure_ai_cohere_embed_input_type_param():
def test_optional_params_image_gen_with_aspect_ratio():
optional_params = get_optional_params_image_gen(
model="imagen-4.0-ultra-generate-preview-06-06",
model="imagen-4.0-ultra-generate-001",
custom_llm_provider="vertex_ai",
aspect_ratio="16:9",
)

View file

@ -910,6 +910,7 @@ async def test_partner_models_httpx(model, region, sync_mode):
[
("vertex_ai/meta/llama-4-scout-17b-16e-instruct-maas", "us-east5"),
("vertex_ai/qwen/qwen3-coder-480b-a35b-instruct-maas", "us-south1"),
("vertex_ai/mistral-large-2411", "us-central1"), # critical - we had this issue: https://github.com/BerriAI/litellm/issues/13888
],
)
@pytest.mark.parametrize(
@ -920,7 +921,7 @@ async def test_partner_models_httpx(model, region, sync_mode):
@pytest.mark.flaky(retries=3, delay=1)
async def test_partner_models_httpx_streaming(model, region, sync_mode):
try:
#load_vertex_ai_credentials()
load_vertex_ai_credentials()
litellm._turn_on_debug()
messages = [
@ -955,8 +956,6 @@ async def test_partner_models_httpx_streaming(model, region, sync_mode):
print(f"response: {response}")
except litellm.RateLimitError as e:
pass
except litellm.InternalServerError as e:
pass
except Exception as e:
if "429 Quota exceeded" in str(e):
pass

View file

@ -1 +1 @@
{"custom_id": "ae006110bb364606||/workspace/saved_models/meta-llama/Meta-Llama-3.1-8B-Instruct", "method": "POST", "url": "/v1/chat/completions", "body": {"model": "gpt-4o-mini", "temperature": 0, "max_tokens": 1024, "response_format": {"type": "json_object"}, "messages": [{"role": "user", "content": "# Instruction \n\nYou are an expert evaluator. Your task is to evaluate the quality of the responses generated by AI models. \nWe will provide you with the user query and an AI-generated responses.\nYo must respond in json"}]}}
{"custom_id": "ae006110bb364606||/workspace/saved_models/meta-llama/Meta-Llama-3.1-8B-Instruct", "method": "POST", "url": "/chat/completions", "body": {"model": "gpt-4o-mini", "temperature": 0, "max_tokens": 1024, "response_format": {"type": "json_object"}, "messages": [{"role": "user", "content": "# Instruction \n\nYou are an expert evaluator. Your task is to evaluate the quality of the responses generated by AI models. \nWe will provide you with the user query and an AI-generated responses.\nYo must respond in json"}]}}

View file

@ -273,6 +273,58 @@ async def test_anthropic_messages_litellm_router_routing_strategy():
print(f"Non-streaming response: {json.dumps(response, indent=2)}")
return response
@pytest.mark.asyncio
async def test_anthropic_messages_fallbacks():
"""
E2E test the anthropic_messages fallbacks from Anthropic API to Bedrock
"""
litellm._turn_on_debug()
router = Router(
model_list=[
{
"model_name": "anthropic/claude-opus-4-20250514",
"litellm_params": {
"model": "anthropic/claude-opus-4-20250514",
"api_key": "bad-key",
},
},
{
"model_name": "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0",
"litellm_params": {
"model": "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0",
},
}
],
fallbacks=[
{
"anthropic/claude-opus-4-20250514":
["bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"]
}
]
)
# Set up test parameters
messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}]
# Call the handler
response = await router.aanthropic_messages(
messages=messages,
model="anthropic/claude-opus-4-20250514",
max_tokens=100,
metadata={
"user_id": "hello",
},
)
# Verify response
assert "id" in response
assert "content" in response
assert "model" in response
assert response["role"] == "assistant"
print(f"Non-streaming response: {json.dumps(response, indent=2)}")
return response
@pytest.mark.asyncio
async def test_anthropic_messages_litellm_router_latency_metadata_tracking():

View file

@ -56,7 +56,7 @@ def test_s3_cache_set_cache_with_ttl(mock_s3_dependencies):
assert "max-age=3600" in call_args[1]["CacheControl"]
def test_s3_cache_get_cache(mock_s3_dependencies):
def test_s3_cache_get_cache_no_expires_info_in_response(mock_s3_dependencies):
"""Test basic get_cache functionality"""
cache = S3Cache("test-bucket")
@ -75,6 +75,54 @@ def test_s3_cache_get_cache(mock_s3_dependencies):
assert result == {"key": "value", "number": 42}
def test_s3_cache_get_cache_with_expires_valid(mock_s3_dependencies):
"""Test get_cache when response contains Expires and cache entry is still valid"""
cache = S3Cache("test-bucket")
# Create a future expiration time (1 hour from now)
future_time = datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(hours=1)
mock_response = {
"Body": MagicMock(),
"Expires": future_time
}
mock_response["Body"].read.return_value = b'{"key": "value", "number": 42}'
cache.s3_client.get_object.return_value = mock_response
result = cache.get_cache("test_key")
cache.s3_client.get_object.assert_called_once_with(
Bucket="test-bucket",
Key="test_key"
)
# Should return the cached value since it's not expired
assert result == {"key": "value", "number": 42}
def test_s3_cache_get_cache_with_expires_expired(mock_s3_dependencies):
"""Test get_cache when response contains Expires and cache entry is no longer valid"""
cache = S3Cache("test-bucket")
# Create a past expiration time (1 hour ago)
past_time = datetime.datetime.now(datetime.timezone.utc) - datetime.timedelta(hours=1)
mock_response = {
"Body": MagicMock(),
"Expires": past_time
}
mock_response["Body"].read.return_value = b'{"key": "value", "number": 42}'
cache.s3_client.get_object.return_value = mock_response
result = cache.get_cache("test_key")
cache.s3_client.get_object.assert_called_once_with(
Bucket="test-bucket",
Key="test_key"
)
# Should return None since the cache entry is expired
assert result is None
def test_s3_cache_get_cache_not_found(mock_s3_dependencies):
"""Test get_cache when key is not found"""
@ -126,7 +174,6 @@ def test_s3_cache_initialization():
cache_with_path = S3Cache("test-bucket", s3_path="my/cache/path")
assert cache_with_path.key_prefix == "my/cache/path/"
# ============================================================================
# ASYNC TESTS
# ============================================================================

View file

@ -10,77 +10,152 @@ import httpx
import pytest
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system-path
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
import litellm
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
OpenAiResponsesToChatCompletionStreamIterator,
)
from litellm.types.llms.openai import Reasoning
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
def test_convert_chat_completion_messages_to_responses_api_image_input():
from litellm.completion_extras.litellm_responses_transformation.transformation import (
LiteLLMResponsesTransformationHandler,
)
class TestLiteLLMResponsesTransformation:
def setup_method(self):
self.handler = LiteLLMResponsesTransformationHandler()
self.model = "responses-api-model"
self.logging_obj = MagicMock()
handler = LiteLLMResponsesTransformationHandler()
def test_transform_request_reasoning_effort(self):
"""
Test that reasoning_effort is mapped to reasoning parameter correctly.
"""
# Case 1: reasoning_effort = "high"
optional_params_high = {"reasoning_effort": "high"}
result_high = self.handler.transform_request(
model=self.model,
messages=[],
optional_params=optional_params_high,
litellm_params={},
headers={},
litellm_logging_obj=self.logging_obj,
)
assert "reasoning" in result_high
assert result_high["reasoning"] == Reasoning(effort="high", summary="detailed")
user_content = "What's in this image?"
user_image = "https://w7.pngwing.com/pngs/666/274/png-transparent-image-pictures-icon-photo-thumbnail.png"
# Case 2: reasoning_effort = "medium"
optional_params_medium = {"reasoning_effort": "medium"}
result_medium = self.handler.transform_request(
model=self.model,
messages=[],
optional_params=optional_params_medium,
litellm_params={},
headers={},
litellm_logging_obj=self.logging_obj,
)
assert "reasoning" in result_medium
assert result_medium["reasoning"] == Reasoning(effort="medium", summary="auto")
messages = [
{
"role": "user",
"content": [
{
"type": "text",
"text": user_content,
},
{
"type": "image_url",
"image_url": {"url": user_image},
},
],
},
]
# Case 3: reasoning_effort = "low"
optional_params_low = {"reasoning_effort": "low"}
result_low = self.handler.transform_request(
model=self.model,
messages=[],
optional_params=optional_params_low,
litellm_params={},
headers={},
litellm_logging_obj=self.logging_obj,
)
assert "reasoning" in result_low
assert result_low["reasoning"] == Reasoning(effort="low", summary="auto")
response, _ = handler.convert_chat_completion_messages_to_responses_api(messages)
# Case 4: no reasoning_effort
optional_params_none = {}
result_none = self.handler.transform_request(
model=self.model,
messages=[],
optional_params=optional_params_none,
litellm_params={},
headers={},
litellm_logging_obj=self.logging_obj,
)
assert "reasoning" in result_none
assert result_none["reasoning"] == Reasoning(summary="auto")
response_str = json.dumps(response)
# Case 5: reasoning_effort = None
optional_params_explicit_none = {"reasoning_effort": None}
result_explicit_none = self.handler.transform_request(
model=self.model,
messages=[],
optional_params=optional_params_explicit_none,
litellm_params={},
headers={},
litellm_logging_obj=self.logging_obj,
)
assert "reasoning" in result_explicit_none
assert result_explicit_none["reasoning"] == Reasoning(summary="auto")
assert user_content in response_str
assert user_image in response_str
def test_convert_chat_completion_messages_to_responses_api_image_input(self):
"""
Test that chat completion messages with image inputs are converted correctly.
"""
user_content = "What's in this image?"
user_image = "https://w7.pngwing.com/pngs/666/274/png-transparent-image-pictures-icon-photo-thumbnail.png"
print("response: ", response)
assert response[0]["content"][1]["image_url"] == user_image
messages = [
{
"role": "user",
"content": [
{
"type": "text",
"text": user_content,
},
{
"type": "image_url",
"image_url": {"url": user_image},
},
],
},
]
response, _ = self.handler.convert_chat_completion_messages_to_responses_api(messages)
def test_openai_responses_chunk_parser_reasoning_summary():
from litellm.completion_extras.litellm_responses_transformation.transformation import (
OpenAiResponsesToChatCompletionStreamIterator,
)
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
response_str = json.dumps(response)
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
assert user_content in response_str
assert user_image in response_str
chunk = {
"delta": "**Compar",
"item_id": "rs_686d544208748198b6912e27b7c299c00e24bd875d35bade",
"output_index": 0,
"sequence_number": 4,
"summary_index": 0,
"type": "response.reasoning_summary_text.delta",
}
print("response: ", response)
assert response[0]["content"][1]["image_url"] == user_image
result = iterator.chunk_parser(chunk)
def test_openai_responses_chunk_parser_reasoning_summary(self):
"""
Test that OpenAI responses chunk parser handles reasoning summary correctly.
"""
iterator = OpenAiResponsesToChatCompletionStreamIterator(
streaming_response=None, sync_stream=True
)
assert isinstance(result, ModelResponseStream)
assert len(result.choices) == 1
choice = result.choices[0]
assert isinstance(choice, StreamingChoices)
assert choice.index == 0
delta = choice.delta
assert isinstance(delta, Delta)
assert delta.content is None
assert delta.reasoning_content == "**Compar"
assert delta.tool_calls is None
assert delta.function_call is None
chunk = {
"delta": "**Compar",
"item_id": "rs_686d544208748198b6912e27b7c299c00e24bd875d35bade",
"output_index": 0,
"sequence_number": 4,
"summary_index": 0,
"type": "response.reasoning_summary_text.delta",
}
result = iterator.chunk_parser(chunk)
assert isinstance(result, ModelResponseStream)
assert len(result.choices) == 1
choice = result.choices[0]
assert isinstance(choice, StreamingChoices)
assert choice.index == 0
delta = choice.delta
assert isinstance(delta, Delta)
assert delta.content is None
assert delta.reasoning_content == "**Compar"
assert delta.tool_calls is None
assert delta.function_call is None

View file

@ -1,7 +1,9 @@
import os
import unittest
from unittest.mock import patch
from datetime import datetime
from unittest.mock import MagicMock, Mock, patch
import litellm
from litellm.integrations.braintrust_logging import BraintrustLogger
class TestBraintrustLogger(unittest.TestCase):
@ -40,4 +42,265 @@ class TestBraintrustLogger(unittest.TestCase):
with patch.dict(os.environ, {}, clear=True):
with self.assertRaises(Exception) as context:
BraintrustLogger(api_key=None)
self.assertIn("Missing keys=['BRAINTRUST_API_KEY']", str(context.exception))
self.assertIn("Missing keys=['BRAINTRUST_API_KEY']", str(context.exception))
@patch('litellm.integrations.braintrust_logging.HTTPHandler')
def test_log_success_event_with_default_span_name(self, MockHTTPHandler):
"""Test log_success_event uses default span name when not provided."""
# Mock HTTP response
mock_response = Mock()
mock_response.json.return_value = {"id": "test-project-id"}
mock_http_handler = Mock()
mock_http_handler.post.return_value = mock_response
MockHTTPHandler.return_value = mock_http_handler
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a mock response object
message_mock = Mock()
message_mock.json = Mock(return_value={"content": "test"})
choice_mock = Mock()
choice_mock.message = message_mock
choice_mock.dict = Mock(return_value={"message": {"content": "test"}})
# Mock the __getitem__ to support response_obj["choices"][0]["message"]
choice_mock.__getitem__ = Mock(return_value=message_mock)
response_obj = Mock(spec=litellm.ModelResponse)
response_obj.choices = [choice_mock]
# Mock the __getitem__ to support response_obj["choices"]
response_obj.__getitem__ = Mock(return_value=[choice_mock])
response_obj.usage = litellm.Usage(
prompt_tokens=10,
completion_tokens=20,
total_tokens=30
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {"metadata": {}},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
}
# Execute
logger.log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Chat Completion')
@patch('litellm.integrations.braintrust_logging.HTTPHandler')
def test_log_success_event_with_custom_span_name(self, MockHTTPHandler):
"""Test log_success_event uses custom span name when provided."""
# Mock HTTP response
mock_response = Mock()
mock_response.json.return_value = {"id": "test-project-id"}
mock_http_handler = Mock()
mock_http_handler.post.return_value = mock_response
MockHTTPHandler.return_value = mock_http_handler
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a mock response object
message_mock = Mock()
message_mock.json = Mock(return_value={"content": "test"})
choice_mock = Mock()
choice_mock.message = message_mock
choice_mock.dict = Mock(return_value={"message": {"content": "test"}})
choice_mock.__getitem__ = Mock(return_value=message_mock)
response_obj = Mock(spec=litellm.ModelResponse)
response_obj.choices = [choice_mock]
response_obj.__getitem__ = Mock(return_value=[choice_mock])
response_obj.usage = litellm.Usage(
prompt_tokens=10,
completion_tokens=20,
total_tokens=30
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {"metadata": {"span_name": "Custom Operation"}},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
}
# Execute
logger.log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Custom Operation')
@patch('litellm.integrations.braintrust_logging.get_async_httpx_client')
async def test_async_log_success_event_with_default_span_name(self, mock_get_http_handler):
"""Test async_log_success_event uses default span name when not provided."""
# Mock async HTTP response
mock_response = Mock()
mock_response.json.return_value = {"id": "test-project-id"}
mock_http_handler = MagicMock()
mock_http_handler.post = MagicMock(return_value=mock_response)
mock_get_http_handler.return_value = mock_http_handler
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a mock response object
message_mock = Mock()
message_mock.json = Mock(return_value={"content": "test"})
choice_mock = Mock()
choice_mock.message = message_mock
choice_mock.dict = Mock(return_value={"message": {"content": "test"}})
choice_mock.__getitem__ = Mock(return_value=message_mock)
response_obj = Mock(spec=litellm.ModelResponse)
response_obj.choices = [choice_mock]
response_obj.__getitem__ = Mock(return_value=[choice_mock])
response_obj.usage = litellm.Usage(
prompt_tokens=10,
completion_tokens=20,
total_tokens=30
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {"metadata": {}},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
}
# Execute
await logger.async_log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Chat Completion')
@patch('litellm.integrations.braintrust_logging.get_async_httpx_client')
async def test_async_log_success_event_with_custom_span_name(self, mock_get_http_handler):
"""Test async_log_success_event uses custom span name when provided."""
# Mock async HTTP response
mock_response = Mock()
mock_response.json.return_value = {"id": "test-project-id"}
mock_http_handler = MagicMock()
mock_http_handler.post = MagicMock(return_value=mock_response)
mock_get_http_handler.return_value = mock_http_handler
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a mock response object
message_mock = Mock()
message_mock.json = Mock(return_value={"content": "test"})
choice_mock = Mock()
choice_mock.message = message_mock
choice_mock.dict = Mock(return_value={"message": {"content": "test"}})
choice_mock.__getitem__ = Mock(return_value=message_mock)
response_obj = Mock(spec=litellm.ModelResponse)
response_obj.choices = [choice_mock]
response_obj.__getitem__ = Mock(return_value=[choice_mock])
response_obj.usage = litellm.Usage(
prompt_tokens=10,
completion_tokens=20,
total_tokens=30
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {"metadata": {"span_name": "Async Custom Operation"}},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
}
# Execute
await logger.async_log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Async Custom Operation')
@patch('litellm.integrations.braintrust_logging.HTTPHandler')
def test_span_name_with_multiple_metadata_fields(self, MockHTTPHandler):
"""Test that span_name works correctly alongside other metadata fields."""
# Mock HTTP response
mock_response = Mock()
mock_response.json.return_value = {"id": "test-project-id"}
mock_http_handler = Mock()
mock_http_handler.post.return_value = mock_response
MockHTTPHandler.return_value = mock_http_handler
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a mock response object
message_mock = Mock()
message_mock.json = Mock(return_value={"content": "test"})
choice_mock = Mock()
choice_mock.message = message_mock
choice_mock.dict = Mock(return_value={"message": {"content": "test"}})
choice_mock.__getitem__ = Mock(return_value=message_mock)
response_obj = Mock(spec=litellm.ModelResponse)
response_obj.choices = [choice_mock]
response_obj.__getitem__ = Mock(return_value=[choice_mock])
response_obj.usage = litellm.Usage(
prompt_tokens=10,
completion_tokens=20,
total_tokens=30
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {
"metadata": {
"span_name": "Multi Metadata Test",
"project_id": "custom-project",
"user_id": "user123",
"session_id": "session456"
}
},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
}
# Execute
logger.log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
# Check span name
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Multi Metadata Test')
# Check that other metadata is preserved
event_metadata = json_data['events'][0]['metadata']
self.assertEqual(event_metadata['user_id'], 'user123')
self.assertEqual(event_metadata['session_id'], 'session456')

View file

@ -0,0 +1,207 @@
import json
import os
import unittest
from datetime import datetime
from unittest.mock import MagicMock, Mock, patch
import litellm
from litellm.integrations.braintrust_logging import BraintrustLogger
class TestBraintrustSpanName(unittest.TestCase):
"""Test custom span_name functionality in Braintrust logging."""
@patch('litellm.integrations.braintrust_logging.HTTPHandler')
def test_default_span_name(self, MockHTTPHandler):
"""Test that default span name is 'Chat Completion' when not provided."""
# Mock HTTP response
mock_http_handler = Mock()
mock_http_handler.post.return_value = Mock()
MockHTTPHandler.return_value = mock_http_handler
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a properly structured mock response
response_obj = litellm.ModelResponse(
id="test-id",
object="chat.completion",
created=1234567890,
model="gpt-3.5-turbo",
choices=[{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop"
}],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {"metadata": {}},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
}
# Execute
logger.log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Chat Completion')
@patch('litellm.integrations.braintrust_logging.HTTPHandler')
def test_custom_span_name(self, MockHTTPHandler):
"""Test that custom span name is used when provided in metadata."""
# Mock HTTP response
mock_http_handler = Mock()
mock_http_handler.post.return_value = Mock()
MockHTTPHandler.return_value = mock_http_handler
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a properly structured mock response
response_obj = litellm.ModelResponse(
id="test-id",
object="chat.completion",
created=1234567890,
model="gpt-3.5-turbo",
choices=[{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop"
}],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {"metadata": {"span_name": "Custom Operation"}},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
}
# Execute
logger.log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Custom Operation')
@patch('litellm.integrations.braintrust_logging.HTTPHandler')
def test_span_name_with_other_metadata(self, MockHTTPHandler):
"""Test that span_name works alongside other metadata fields."""
# Mock HTTP response
mock_http_handler = Mock()
mock_http_handler.post.return_value = Mock()
MockHTTPHandler.return_value = mock_http_handler
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a properly structured mock response
response_obj = litellm.ModelResponse(
id="test-id",
object="chat.completion",
created=1234567890,
model="gpt-3.5-turbo",
choices=[{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop"
}],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {
"metadata": {
"span_name": "Multi Metadata Test",
"project_id": "custom-project",
"user_id": "user123",
"session_id": "session456",
"environment": "production"
}
},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
}
# Execute
logger.log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
# Check span name
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Multi Metadata Test')
# Check that other metadata is preserved (except for filtered keys)
event_metadata = json_data['events'][0]['metadata']
self.assertEqual(event_metadata['user_id'], 'user123')
self.assertEqual(event_metadata['session_id'], 'session456')
self.assertEqual(event_metadata['environment'], 'production')
# Span name should be in span_attributes, not in metadata
self.assertIn('span_name', event_metadata) # span_name is also kept in metadata
@patch('litellm.integrations.braintrust_logging.get_async_httpx_client')
async def test_async_custom_span_name(self, mock_get_http_handler):
"""Test async logging with custom span name."""
# Mock async HTTP response
mock_http_handler = MagicMock()
mock_http_handler.post = MagicMock(return_value=Mock())
mock_get_http_handler.return_value = mock_http_handler
# Setup
logger = BraintrustLogger(api_key="test-key")
logger.default_project_id = "test-project-id"
# Create a properly structured mock response
response_obj = litellm.ModelResponse(
id="test-id",
object="chat.completion",
created=1234567890,
model="gpt-3.5-turbo",
choices=[{
"index": 0,
"message": {"role": "assistant", "content": "test response"},
"finish_reason": "stop"
}],
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30}
)
kwargs = {
"litellm_call_id": "test-call-id",
"messages": [{"role": "user", "content": "test"}],
"litellm_params": {"metadata": {"span_name": "Async Custom Operation"}},
"model": "gpt-3.5-turbo",
"response_cost": 0.001
}
# Execute
await logger.async_log_success_event(kwargs, response_obj, datetime.now(), datetime.now())
# Verify
call_args = mock_http_handler.post.call_args
self.assertIsNotNone(call_args)
json_data = call_args.kwargs['json']
self.assertEqual(json_data['events'][0]['span_attributes']['name'], 'Async Custom Operation')
if __name__ == "__main__":
unittest.main()

View file

@ -229,6 +229,19 @@ class TestLangfuseOtelIntegration:
# Should return an empty dict
assert result == {}
def test_get_langfuse_otel_config_with_otel_host_priority(self):
"""LANGFUSE_OTEL_HOST should take priority over LANGFUSE_HOST."""
with patch.dict(os.environ, {
'LANGFUSE_PUBLIC_KEY': 'test_public_key',
'LANGFUSE_SECRET_KEY': 'test_secret_key',
'LANGFUSE_HOST': 'https://should-not-be-used.com',
'LANGFUSE_OTEL_HOST': 'https://otel-host.com'
}, clear=False):
_ = LangfuseOtelLogger.get_langfuse_otel_config()
assert os.environ.get("OTEL_EXPORTER_OTLP_ENDPOINT") == "https://otel-host.com/api/public/otel"
if __name__ == "__main__":

View file

@ -0,0 +1,43 @@
import pytest
from litellm.litellm_core_utils.get_provider_specific_headers import (
ProviderSpecificHeaderUtils,
)
from litellm.types.utils import ProviderSpecificHeader
class TestProviderSpecificHeaderUtils:
def test_get_provider_specific_headers_matching_provider(self):
"""Test that the method returns extra_headers when custom_llm_provider matches."""
provider_specific_header: ProviderSpecificHeader = {
"custom_llm_provider": "openai",
"extra_headers": {"Authorization": "Bearer token123", "Custom-Header": "value"}
}
custom_llm_provider = "openai"
result = ProviderSpecificHeaderUtils.get_provider_specific_headers(
provider_specific_header, custom_llm_provider
)
expected = {"Authorization": "Bearer token123", "Custom-Header": "value"}
assert result == expected
def test_get_provider_specific_headers_no_match_or_none(self):
"""Test that the method returns empty dict when provider doesn't match or is None."""
# Test case 1: Provider doesn't match
provider_specific_header: ProviderSpecificHeader = {
"custom_llm_provider": "anthropic",
"extra_headers": {"Authorization": "Bearer token123"}
}
custom_llm_provider = "openai"
result = ProviderSpecificHeaderUtils.get_provider_specific_headers(
provider_specific_header, custom_llm_provider
)
assert result == {}
# Test case 2: provider_specific_header is None
result = ProviderSpecificHeaderUtils.get_provider_specific_headers(
None, "openai"
)
assert result == {}

View file

@ -15,7 +15,10 @@ from typing import Optional
import litellm
from litellm.litellm_core_utils.litellm_logging import Logging
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
from litellm.litellm_core_utils.streaming_handler import (
AUDIO_ATTRIBUTE,
CustomStreamWrapper,
)
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
Delta,
@ -813,3 +816,220 @@ def test_optional_combine_thinking_block_with_none_content(
assert final_response.choices[0].delta.content == "</think>The answer is 42"
assert initialized_custom_stream_wrapper.sent_last_thinking_block is True
assert not hasattr(final_response.choices[0].delta, "reasoning_content")
def test_has_special_delta_content(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""Test the _has_special_delta_content helper method"""
# Test empty choices
empty_response = ModelResponseStream(
id="test", created=1742056047, model=None, choices=[]
)
assert not initialized_custom_stream_wrapper._has_special_delta_content(empty_response)
# Test with tool_calls (simulate with mock object)
tool_call_response = ModelResponseStream(
id="test", created=1742056047, model=None,
choices=[
StreamingChoices(
finish_reason=None, index=0,
delta=Delta(content=None, tool_calls=[{"id": "test"}])
)
]
)
assert initialized_custom_stream_wrapper._has_special_delta_content(tool_call_response)
# Test with function_call (simulate with mock object)
function_call_response = ModelResponseStream(
id="test", created=1742056047, model=None,
choices=[
StreamingChoices(
finish_reason=None, index=0,
delta=Delta(content=None, function_call={"name": "test_func"})
)
]
)
assert initialized_custom_stream_wrapper._has_special_delta_content(function_call_response)
# Test with audio (simulate by adding audio attribute)
audio_response = ModelResponseStream(
id="test", created=1742056047, model=None,
choices=[
StreamingChoices(
finish_reason=None, index=0,
delta=Delta(content=None)
)
]
)
# Manually add audio attribute to delta
audio_response.choices[0].delta.audio = {"transcript": "test"}
assert initialized_custom_stream_wrapper._has_special_delta_content(audio_response)
# Test with image (simulate by adding image attribute)
image_response = ModelResponseStream(
id="test", created=1742056047, model=None,
choices=[
StreamingChoices(
finish_reason=None, index=0,
delta=Delta(content=None)
)
]
)
# Manually add image attribute to delta
image_response.choices[0].delta.image = {"url": "test.jpg"}
assert initialized_custom_stream_wrapper._has_special_delta_content(image_response)
# Test with regular content (should return False)
regular_response = ModelResponseStream(
id="test", created=1742056047, model=None,
choices=[
StreamingChoices(
finish_reason=None, index=0,
delta=Delta(content="Hello world")
)
]
)
assert not initialized_custom_stream_wrapper._has_special_delta_content(regular_response)
def test_handle_special_delta_content(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""Test the _handle_special_delta_content helper method"""
test_response = ModelResponseStream(
id="test", created=1742056047, model=None,
choices=[
StreamingChoices(
finish_reason=None, index=0,
delta=Delta(content="test", role="assistant")
)
]
)
# The method should call strip_role_from_delta
result = initialized_custom_stream_wrapper._handle_special_delta_content(test_response)
# Should return the same response object (modified)
assert result is test_response
# Should have set sent_first_chunk to True
assert initialized_custom_stream_wrapper.sent_first_chunk is True
def test_has_any_special_delta_attributes(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""Test the _has_any_special_delta_attributes helper method"""
# Test with delta that has audio attribute
class MockDelta:
def __init__(self):
self.audio = {"transcript": "Hello world"}
audio_delta = MockDelta()
result = initialized_custom_stream_wrapper._has_any_special_delta_attributes(audio_delta)
assert result is True
# Test with delta that has image attribute
class MockDeltaImage:
def __init__(self):
self.image = {"url": "test.jpg"}
image_delta = MockDeltaImage()
result = initialized_custom_stream_wrapper._has_any_special_delta_attributes(image_delta)
assert result is True
# Test with delta that has no special attributes
class MockDeltaRegular:
def __init__(self):
self.content = "regular content"
regular_delta = MockDeltaRegular()
result = initialized_custom_stream_wrapper._has_any_special_delta_attributes(regular_delta)
assert result is False
def test_handle_special_delta_attributes(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""Test the _handle_special_delta_attributes helper method"""
# Create a model response
model_response = ModelResponseStream(
id="test", created=1742056047, model=None,
choices=[
StreamingChoices(
finish_reason=None, index=0,
delta=Delta(content="test")
)
]
)
# Test with delta that has audio attribute
class MockDelta:
def __init__(self):
self.audio = {"transcript": "Hello world"}
audio_delta = MockDelta()
initialized_custom_stream_wrapper._handle_special_delta_attributes(audio_delta, model_response)
# Should copy the audio attribute
assert hasattr(model_response.choices[0].delta, "audio")
assert model_response.choices[0].delta.audio == {"transcript": "Hello world"}
# Test with delta that has image attribute
class MockDeltaImage:
def __init__(self):
self.image = {"url": "test.jpg"}
image_delta = MockDeltaImage()
model_response2 = ModelResponseStream(
id="test", created=1742056047, model=None,
choices=[
StreamingChoices(
finish_reason=None, index=0,
delta=Delta(content="test")
)
]
)
initialized_custom_stream_wrapper._handle_special_delta_attributes(image_delta, model_response2)
# Should copy the image attribute
assert hasattr(model_response2.choices[0].delta, "image")
assert model_response2.choices[0].delta.image == {"url": "test.jpg"}
def test_has_special_delta_attribute(
initialized_custom_stream_wrapper: CustomStreamWrapper,
):
"""Test the _has_special_delta_attribute helper method"""
# Test with None delta
assert not initialized_custom_stream_wrapper._has_special_delta_attribute(None, "audio")
# Test with delta that has the attribute
class MockDelta:
def __init__(self):
self.audio = {"transcript": "test"}
delta_with_audio = MockDelta()
assert initialized_custom_stream_wrapper._has_special_delta_attribute(delta_with_audio, "audio")
# Test with delta that doesn't have the attribute
class MockDeltaNoAudio:
def __init__(self):
self.content = "test"
delta_without_audio = MockDeltaNoAudio()
assert not initialized_custom_stream_wrapper._has_special_delta_attribute(delta_without_audio, "audio")
# Test with delta that has the attribute but it's None
class MockDeltaNone:
def __init__(self):
self.audio = None
delta_with_none = MockDeltaNone()
assert not initialized_custom_stream_wrapper._has_special_delta_attribute(delta_with_none, "audio")

View file

@ -451,6 +451,7 @@ def test_img_url_token_counter(img_url):
def test_token_encode_disallowed_special():
encode(model="gpt-3.5-turbo", text="Hello, world! <|endoftext|>")
token_counter(model="gpt-3.5-turbo", text="Hello, world! <|endoftext|>")
def test_token_counter():

View file

@ -362,3 +362,150 @@ def test_x_initiator_header_system_only_messages():
)
assert headers["X-Initiator"] == "user"
def test_get_supported_openai_params_claude_model():
"""Test that Claude models with extended thinking support have thinking and reasoning parameters."""
config = GithubCopilotConfig()
# Test Claude 4 model supports thinking and reasoning_effort parameters
supported_params = config.get_supported_openai_params("claude-sonnet-4-20250514")
assert "thinking" in supported_params
assert "reasoning_effort" in supported_params
# Test Claude 3-7 model supports thinking and reasoning_effort parameters
supported_params_claude37 = config.get_supported_openai_params("claude-3-7-sonnet-20250219")
assert "thinking" in supported_params_claude37
assert "reasoning_effort" in supported_params_claude37
# Test Claude 3.5 model does NOT support thinking parameters (no extended thinking)
supported_params_claude35 = config.get_supported_openai_params("claude-3.5-sonnet")
assert "thinking" not in supported_params_claude35
assert "reasoning_effort" not in supported_params_claude35
# Test non-Claude model doesn't include thinking parameters but may include reasoning_effort
supported_params_gpt = config.get_supported_openai_params("gpt-4o")
assert "thinking" not in supported_params_gpt
# gpt-4o should NOT have reasoning_effort (not a reasoning model)
assert "reasoning_effort" not in supported_params_gpt
# Test O-series reasoning models include reasoning_effort but not thinking
supported_params_o3 = config.get_supported_openai_params("o3-mini")
assert "thinking" not in supported_params_o3
# o3-mini should have reasoning_effort (it's an O-series reasoning model)
assert "reasoning_effort" in supported_params_o3
def test_get_supported_openai_params_case_insensitive():
"""Test that Claude model detection is case-insensitive for models with extended thinking."""
config = GithubCopilotConfig()
# Test uppercase Claude 4 model with full model name
supported_params_upper = config.get_supported_openai_params("CLAUDE-SONNET-4-20250514")
assert "thinking" in supported_params_upper
assert "reasoning_effort" in supported_params_upper
# Test mixed case Claude 3-7 model (has extended thinking) with full model name
supported_params_mixed = config.get_supported_openai_params("Claude-3-7-Sonnet-20250219")
assert "thinking" in supported_params_mixed
assert "reasoning_effort" in supported_params_mixed
# Test that Claude 3.5 models don't have thinking support (case insensitive)
supported_params_35 = config.get_supported_openai_params("CLAUDE-3.5-SONNET")
assert "thinking" not in supported_params_35
assert "reasoning_effort" not in supported_params_35
def test_copilot_vision_request_header_with_image():
"""Test that Copilot-Vision-Request header is added when messages contain images"""
config = GithubCopilotConfig()
# Mock the authenticator
config.authenticator = MagicMock()
config.authenticator.get_api_key.return_value = "gh.test-key-123"
config.authenticator.get_api_base.return_value = None
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "What's in this image?"},
{
"type": "image_url",
"image_url": {"url": "data:image/jpeg;base64,abc123"}
}
]
}
]
headers = config.validate_environment(
headers={},
model="github_copilot/gpt-4-vision-preview",
messages=messages,
optional_params={},
litellm_params={},
api_key=None,
api_base=None,
)
assert headers["Copilot-Vision-Request"] == "true"
assert headers["X-Initiator"] == "user"
def test_copilot_vision_request_header_text_only():
"""Test that Copilot-Vision-Request header is not added for text-only messages"""
config = GithubCopilotConfig()
# Mock the authenticator
config.authenticator = MagicMock()
config.authenticator.get_api_key.return_value = "gh.test-key-123"
config.authenticator.get_api_base.return_value = None
messages = [
{"role": "user", "content": "Just a text message"},
]
headers = config.validate_environment(
headers={},
model="github_copilot/gpt-4",
messages=messages,
optional_params={},
litellm_params={},
api_key=None,
api_base=None,
)
assert "Copilot-Vision-Request" not in headers
assert headers["X-Initiator"] == "user"
def test_copilot_vision_request_header_with_type_image_url():
"""Test that Copilot-Vision-Request header is added for content with type: image_url"""
config = GithubCopilotConfig()
# Mock the authenticator
config.authenticator = MagicMock()
config.authenticator.get_api_key.return_value = "gh.test-key-123"
config.authenticator.get_api_base.return_value = None
messages = [
{
"role": "user",
"content": [
{"type": "text", "text": "Analyze this image"},
{"type": "image_url", "image_url": {"url": "https://example.com/image.jpg"}}
]
}
]
headers = config.validate_environment(
headers={},
model="github_copilot/gpt-4-vision-preview",
messages=messages,
optional_params={},
litellm_params={},
api_key=None,
api_base=None,
)
assert headers["Copilot-Vision-Request"] == "true"
assert headers["X-Initiator"] == "user"

View file

@ -442,6 +442,82 @@ def test_vertex_ai_map_thinking_param_with_budget_tokens_0():
}
def test_vertex_ai_reasoning_effort_mapping():
"""
Test that reasoning_effort is mapped to thinkingConfig correctly for models that support it.
- A default thinking config is applied if reasoning_effort is not specified.
- reasoning_effort correctly maps to thinkingConfig.
- No thinkingConfig is applied for models that do not support reasoning.
- reasoning_effort is prioritized over thinking param.
"""
v = VertexGeminiConfig()
optional_params = {}
# Case 1: Model supports reasoning, no reasoning_effort provided
# Should apply default thinkingConfig
with patch(
"litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.supports_reasoning",
return_value=True,
):
result_params = v.map_openai_params(
non_default_params={},
optional_params=deepcopy(optional_params),
model="gemini-2.5-pro",
drop_params=False,
)
assert "thinkingConfig" in result_params
assert result_params["thinkingConfig"] == {"includeThoughts": True}
# Case 2: Model supports reasoning, reasoning_effort is 'low'
# Should apply thinkingConfig with budget
with patch(
"litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.supports_reasoning",
return_value=True,
):
result_params_with_effort = v.map_openai_params(
non_default_params={"reasoning_effort": "low"},
optional_params=deepcopy(optional_params),
model="gemini-2.5-pro",
drop_params=False,
)
assert "thinkingConfig" in result_params_with_effort
assert result_params_with_effort["thinkingConfig"]["includeThoughts"] is True
assert "thinkingBudget" in result_params_with_effort["thinkingConfig"]
# Case 3: Model does not support reasoning
# Should not apply thinkingConfig
with patch(
"litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.supports_reasoning",
return_value=False,
):
result_params_no_support = v.map_openai_params(
non_default_params={},
optional_params=deepcopy(optional_params),
model="gemini-pro",
drop_params=False,
)
assert "thinkingConfig" not in result_params_no_support
# Case 4: Model supports reasoning, but reasoning_effort is set, should be prioritized over thinking
with patch(
"litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini.supports_reasoning",
return_value=True,
):
result_params_with_effort = v.map_openai_params(
non_default_params={
"reasoning_effort": "low",
"thinking": {"type": "enabled", "budget_tokens": 1000},
},
optional_params=deepcopy(optional_params),
model="gemini-2.5-pro",
drop_params=False,
)
assert "thinkingConfig" in result_params_with_effort
assert result_params_with_effort["thinkingConfig"]["includeThoughts"] is True
assert "thinkingBudget" in result_params_with_effort["thinkingConfig"]
assert result_params_with_effort["thinkingConfig"]["thinkingBudget"] != 1000
def test_vertex_ai_map_tools():
v = VertexGeminiConfig()
tools = v._map_function(value=[{"code_execution": {}}])
@ -496,6 +572,36 @@ def test_vertex_ai_map_tool_with_anyof():
"anyOf": [{"type": "string", "nullable": True, "title": "Base Branch"}]
}, f"Expected only anyOf field and its contents to be kept, but got {tools[0]['function_declarations'][0]['parameters']['properties']['base_branch']}"
new_value = [
{
"type": "function",
"function": {
"name": "git_create_branch",
"description": "Creates a new branch from an optional base branch",
"parameters": {
"type": "object",
"properties": {
"repo_path": {"title": "Repo Path", "type": "string"},
"branch_name": {"title": "Branch Name", "type": "string"},
"base_branch": {
"anyOf": [{"type": "string"}, {"type": "null"}],
"default": None,
},
},
"required": ["repo_path", "branch_name"],
"title": "GitCreateBranch",
},
},
}
]
new_tools = v._map_function(value=new_value)
assert new_tools[0]["function_declarations"][0]["parameters"]["properties"][
"base_branch"
] == {
"anyOf": [{"type": "string", "nullable": True}]
}, f"Expected only anyOf field and its contents to be kept, but got {new_tools[0]['function_declarations'][0]['parameters']['properties']['base_branch']}"
def test_vertex_ai_streaming_usage_calculation():
"""

View file

@ -1469,3 +1469,37 @@ def test_vertex_parallel_tool_calls_false_single_tool():
parallel_tool_calls=False,
)
assert "tools" in optional_params
from litellm.llms.vertex_ai.gemini.transformation import _transform_request_body
def test_system_prompt_only_adds_blank_user_message():
"""
Test that the system prompt only adds a blank user message when a system message is passed in.
Relevant Issue - https://github.com/BerriAI/litellm/issues/13769
"""
SYSTEM_INSTRUCTION = "System instructions for the model"
data = _transform_request_body(
messages=[{"role": "system", "content": SYSTEM_INSTRUCTION}],
model="gemini-2.5-flash",
optional_params={},
custom_llm_provider="vertex_ai",
litellm_params={},
cached_content=None,
)
print("Final data: ", data)
# validate that a blank user message is added when a system message is passed in
assert len(data["contents"]) == 1
first_content = data["contents"][0]
assert first_content["role"] == "user"
assert len(first_content["parts"]) == 1
#########################################################
# system message was passed in
#########################################################
assert len(data["system_instruction"]) == 1
assert data["system_instruction"]["parts"][0]["text"] == SYSTEM_INSTRUCTION

View file

@ -1,13 +1,14 @@
import json
import os
import sys
from unittest.mock import MagicMock, patch
import httpx
import pytest
from fastapi.testclient import TestClient
from litellm.llms.custom_httpx.http_handler import HTTPHandler
from unittest.mock import MagicMock, patch
import httpx
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
@ -133,4 +134,271 @@ def test_bedrock_non_application_inference_profile_no_encoding():
actual_url = str(call_args.kwargs["url"])
assert "application-inference-profile%2F" not in actual_url
assert "anthropic.claude-3-sonnet-20240229-v1:0" in actual_url
assert response.status_code == 200
def test_update_stream_param_based_on_request_body():
"""
Test _update_stream_param_based_on_request_body handles stream parameter correctly.
"""
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
HttpPassThroughEndpointHelpers,
)
# Test 1: stream in request body should take precedence
parsed_body = {"stream": True, "model": "test-model"}
result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
parsed_body=parsed_body, stream=False
)
assert result is True
# Test 2: no stream in request body should return original stream param
parsed_body = {"model": "test-model"}
result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
parsed_body=parsed_body, stream=False
)
assert result is False
# Test 3: stream=False in request body should return False
parsed_body = {"stream": False, "model": "test-model"}
result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
parsed_body=parsed_body, stream=True
)
assert result is False
# Test 4: no stream param provided, no stream in body
parsed_body = {"model": "test-model"}
result = HttpPassThroughEndpointHelpers._update_stream_param_based_on_request_body(
parsed_body=parsed_body, stream=None
)
assert result is None
@pytest.fixture
def mock_request():
"""Create a mock request with headers"""
from typing import Optional
class QueryParams:
def __init__(self):
self._dict = {}
class MockRequest:
def __init__(
self, headers=None, method="POST", request_body: Optional[dict] = None
):
self.headers = headers or {}
self.query_params = QueryParams()
self.method = method
self.request_body = request_body or {}
# Add url attribute that the actual code expects
self.url = "http://localhost:8000/test"
async def body(self) -> bytes:
return bytes(json.dumps(self.request_body), "utf-8")
return MockRequest
@pytest.fixture
def mock_user_api_key_dict():
"""Create a mock user API key dictionary"""
from litellm.proxy._types import UserAPIKeyAuth
return UserAPIKeyAuth(
api_key="test-key",
user_id="test-user",
team_id="test-team",
end_user_id="test-user",
)
@pytest.mark.asyncio
async def test_pass_through_request_stream_param_override(
mock_request, mock_user_api_key_dict
):
"""
Test that when stream=None is passed as parameter but stream=True
is in request body, the request body value takes precedence and
the eventual POST request uses streaming.
"""
from unittest.mock import AsyncMock, Mock, patch
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
pass_through_request,
)
# Create request body with stream=True
request_body = {
"model": "claude-3-5-sonnet-20241022",
"max_tokens": 256,
"messages": [{"role": "user", "content": "Hello, world"}],
"stream": True # This should override the function parameter
}
# Create a mock streaming response
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "text/event-stream"}
# Mock the streaming response behavior
async def mock_aiter_bytes():
yield b'data: {"content": "Hello"}\n\n'
yield b'data: {"content": "World"}\n\n'
yield b'data: [DONE]\n\n'
mock_response.aiter_bytes = mock_aiter_bytes
# Create mocks for the async client
mock_async_client = AsyncMock()
mock_request_obj = AsyncMock()
# Mock build_request to return a request object (it's a sync method)
mock_async_client.build_request = Mock(return_value=mock_request_obj)
# Mock send to return the streaming response
mock_async_client.send.return_value = mock_response
# Mock get_async_httpx_client to return our mock client
mock_client_obj = Mock()
mock_client_obj.client = mock_async_client
# Create the request
request = mock_request(
headers={}, method="POST", request_body=request_body
)
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client",
return_value=mock_client_obj,
), patch(
"litellm.proxy.proxy_server.proxy_logging_obj.pre_call_hook",
return_value=request_body, # Return the request body unchanged
), patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler",
new=AsyncMock(), # Mock the success handler
):
# Call pass_through_request with stream=False parameter
response = await pass_through_request(
request=request,
target="https://api.anthropic.com/v1/messages",
custom_headers={"Authorization": "Bearer test-key"},
user_api_key_dict=mock_user_api_key_dict,
stream=None, # This should be overridden by request body
)
# Verify that build_request was called (indicating streaming path)
mock_async_client.build_request.assert_called_once_with(
"POST",
httpx.URL("https://api.anthropic.com/v1/messages"),
json=request_body,
params=None,
headers={
"Authorization": "Bearer test-key"
},
)
# Verify that send was called with stream=True
mock_async_client.send.assert_called_once_with(
mock_request_obj,
stream=True # This proves that stream=True from request body was used
)
# Verify that the non-streaming request method was NOT called
mock_async_client.request.assert_not_called()
# Verify response is a StreamingResponse
from fastapi.responses import StreamingResponse
assert isinstance(response, StreamingResponse)
assert response.status_code == 200
@pytest.mark.asyncio
async def test_pass_through_request_stream_param_no_override(
mock_request, mock_user_api_key_dict
):
"""
Test that when stream=False is passed as parameter and no stream
is in request body, the function parameter is used and
the eventual request uses non-streaming.
"""
from unittest.mock import AsyncMock, Mock, patch
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
pass_through_request,
)
# Create request body without stream parameter
request_body = {
"model": "claude-3-5-sonnet-20241022",
"max_tokens": 256,
"messages": [{"role": "user", "content": "Hello, world"}],
# No stream parameter - should use function parameter stream=False
}
# Create a mock non-streaming response
mock_response = AsyncMock()
mock_response.status_code = 200
mock_response.headers = {"content-type": "application/json"}
mock_response._content = b'{"response": "Hello world"}'
async def mock_aread():
return mock_response._content
mock_response.aread = mock_aread
# Create mocks for the async client
mock_async_client = AsyncMock()
# Mock request to return the non-streaming response
mock_async_client.request.return_value = mock_response
# Mock get_async_httpx_client to return our mock client
mock_client_obj = Mock()
mock_client_obj.client = mock_async_client
# Create the request
request = mock_request(
headers={}, method="POST", request_body=request_body
)
with patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client",
return_value=mock_client_obj,
), patch(
"litellm.proxy.proxy_server.proxy_logging_obj.pre_call_hook",
return_value=request_body, # Return the request body unchanged
), patch(
"litellm.proxy.pass_through_endpoints.pass_through_endpoints.pass_through_endpoint_logging.pass_through_async_success_handler",
new=AsyncMock(), # Mock the success handler
):
# Call pass_through_request with stream=False parameter
response = await pass_through_request(
request=request,
target="https://api.anthropic.com/v1/messages",
custom_headers={"Authorization": "Bearer test-key"},
user_api_key_dict=mock_user_api_key_dict,
stream=False, # Should be used since no stream in request body
)
# Verify that build_request was NOT called (no streaming path)
mock_async_client.build_request.assert_not_called()
# Verify that send was NOT called (no streaming path)
mock_async_client.send.assert_not_called()
# Verify that the non-streaming request method WAS called
mock_async_client.request.assert_called_once_with(
method="POST",
url=httpx.URL("https://api.anthropic.com/v1/messages"),
headers={
"Authorization": "Bearer test-key"
},
params=None,
json=request_body,
)
# Verify response is a regular Response (not StreamingResponse)
from fastapi.responses import Response, StreamingResponse
assert not isinstance(response, StreamingResponse)
assert isinstance(response, Response)
assert response.status_code == 200

View file

@ -0,0 +1,498 @@
import os
from unittest.mock import MagicMock, patch
import httpx
import pytest
import litellm
from litellm import ModelResponse
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.guardrails.guardrail_hooks.noma import (
NomaGuardrail,
initialize_guardrail,
)
from litellm.proxy.guardrails.guardrail_hooks.noma.noma import NomaBlockedMessage
from litellm.proxy.guardrails.init_guardrails import init_guardrails_v2
from litellm.types.utils import Choices, Message
@pytest.fixture
def noma_guardrail():
"""Create a NomaGuardrail instance for testing"""
return NomaGuardrail(
api_key="test-api-key",
api_base="https://api.test.noma.security/",
application_id="test-app",
monitor_mode=False,
block_failures=True,
guardrail_name="test-noma-guardrail",
event_hook="pre_call",
default_on=True,
)
@pytest.fixture
def mock_user_api_key_dict():
"""Create a mock UserAPIKeyAuth object"""
return UserAPIKeyAuth(
user_id="test-user-id",
user_email="test@example.com",
key_name="test-key",
key_alias=None,
team_id=None,
team_alias=None,
user_role=None,
api_key="test-api-key",
permissions={},
models=[],
spend=0.0,
max_budget=None,
soft_budget=None,
tpm_limit=None,
rpm_limit=None,
parallel_request_limit=None,
metadata={},
max_parallel_requests=None,
allowed_cache_controls=[],
model_spend={},
model_max_budget={},
)
@pytest.fixture
def mock_request_data():
"""Create mock request data"""
return {
"messages": [
{"role": "system", "content": "You are a helpful assistant"},
{"role": "user", "content": "Hello, how are you?"},
],
"litellm_call_id": "test-call-id",
"metadata": {"requester_ip_address": "192.168.1.1"},
}
class TestNomaGuardrailConfiguration:
"""Test configuration and initialization of Noma guardrail"""
def test_init_with_config(self):
"""Test initializing Noma guardrail via init_guardrails_v2"""
with patch.dict(
os.environ,
{
"NOMA_API_KEY": "test-api-key",
"NOMA_API_BASE": "https://api.test.noma.security/",
},
):
init_guardrails_v2(
all_guardrails=[
{
"guardrail_name": "noma-pre-guard",
"litellm_params": {
"guardrail": "noma",
"mode": "pre_call",
"application_id": "test-app",
"monitor_mode": False,
"block_failures": True,
},
}
],
config_file_path="",
)
def test_init_with_env_vars(self):
"""Test initialization with environment variables"""
with patch.dict(
os.environ,
{
"NOMA_API_KEY": "env-api-key",
"NOMA_API_BASE": "https://env.api.noma.security/",
"NOMA_APPLICATION_ID": "env-app-id",
"NOMA_MONITOR_MODE": "true",
"NOMA_BLOCK_FAILURES": "false",
},
):
guardrail = NomaGuardrail()
assert guardrail.api_key == "env-api-key"
assert guardrail.api_base == "https://env.api.noma.security/"
assert guardrail.application_id == "env-app-id"
assert guardrail.monitor_mode is True
assert guardrail.block_failures is False
def test_init_with_params_override_env(self):
"""Test that constructor params override environment variables"""
with patch.dict(
os.environ,
{
"NOMA_API_KEY": "env-api-key",
"NOMA_MONITOR_MODE": "true",
},
):
guardrail = NomaGuardrail(
api_key="param-api-key",
monitor_mode=False,
)
assert guardrail.api_key == "param-api-key"
assert guardrail.monitor_mode is False
def test_initialize_guardrail_function(self):
"""Test the initialize_guardrail function"""
from litellm.types.guardrails import Guardrail, LitellmParams
litellm_params = LitellmParams(
guardrail="noma",
mode="pre_call",
api_key="test-key",
api_base="https://test.api/",
application_id="test-app",
monitor_mode=True,
block_failures=False,
)
guardrail = Guardrail(
guardrail_name="test-guardrail",
litellm_params=litellm_params,
)
with patch("litellm.logging_callback_manager.add_litellm_callback") as mock_add:
result = initialize_guardrail(litellm_params, guardrail)
assert isinstance(result, NomaGuardrail)
assert result.api_key == "test-key"
assert result.api_base == "https://test.api/"
assert result.application_id == "test-app"
assert result.monitor_mode is True
assert result.block_failures is False
mock_add.assert_called_once_with(result)
class TestNomaBlockedMessage:
"""Test the NomaBlockedMessage exception class"""
def test_blocked_message_basic(self):
"""Test basic blocked message creation"""
response = {
"verdict": False,
"prompt": {
"harmfulContent": {"result": True, "confidence": 0.9},
"code": {"result": False, "confidence": 0.1},
},
}
exception = NomaBlockedMessage(response)
assert exception.status_code == 400
assert exception.detail["error"] == "Request blocked by Noma guardrail"
assert "harmfulContent" in exception.detail["details"]["prompt"]
assert "code" not in exception.detail["details"]["prompt"]
def test_blocked_message_with_sensitive_data(self):
"""Test blocked message with sensitive data detection"""
response = {
"verdict": False,
"prompt": {
"sensitiveData": {
"email": {"result": True, "entities": ["test@example.com"]},
"phone": {"result": False},
},
},
}
exception = NomaBlockedMessage(response)
assert "email" in exception.detail["details"]["prompt"]["sensitiveData"]
assert "phone" not in exception.detail["details"]["prompt"]["sensitiveData"]
def test_blocked_message_with_topics(self):
"""Test blocked message with topic guardrails"""
response = {
"verdict": False,
"prompt": {
"bannedTopics": {
"violence": {"result": True, "confidence": 0.95},
"politics": {"result": False, "confidence": 0.2},
},
},
}
exception = NomaBlockedMessage(response)
assert "violence" in exception.detail["details"]["prompt"]["bannedTopics"]
assert "politics" not in exception.detail["details"]["prompt"]["bannedTopics"]
class TestNomaGuardrailHooks:
"""Test the guardrail hook methods"""
@pytest.mark.asyncio
async def test_pre_call_hook_allowed(
self, noma_guardrail, mock_user_api_key_dict, mock_request_data
):
"""Test pre-call hook when content is allowed"""
mock_response = MagicMock()
mock_response.json.return_value = {"verdict": True}
mock_response.raise_for_status = MagicMock()
with patch.object(
noma_guardrail.async_handler, "post", return_value=mock_response
) as mock_post:
result = await noma_guardrail.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
cache=MagicMock(),
data=mock_request_data,
call_type="completion",
)
assert result == mock_request_data
mock_post.assert_called_once()
# Verify API call details
call_args = mock_post.call_args
assert call_args[0][0].endswith("/ai-dr/v1/prompt/scan/aggregate")
assert call_args[1]["headers"]["X-Noma-AIDR-Application-ID"] == "test-app"
assert call_args[1]["headers"]["Authorization"] == "Bearer test-api-key"
assert call_args[1]["json"]["request"]["text"] == "Hello, how are you?"
@pytest.mark.asyncio
async def test_pre_call_hook_blocked(
self, noma_guardrail, mock_user_api_key_dict, mock_request_data
):
"""Test pre-call hook when content is blocked"""
mock_response = MagicMock()
mock_response.json.return_value = {
"verdict": False,
"originalResponse": {
"prompt": {"harmfulContent": {"result": True, "confidence": 0.9}}
},
}
mock_response.raise_for_status = MagicMock()
with patch.object(
noma_guardrail.async_handler, "post", return_value=mock_response
):
with pytest.raises(NomaBlockedMessage) as exc_info:
await noma_guardrail.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
cache=MagicMock(),
data=mock_request_data,
call_type="completion",
)
assert exc_info.value.status_code == 400
assert "harmfulContent" in exc_info.value.detail["details"]["prompt"]
@pytest.mark.asyncio
async def test_pre_call_hook_monitor_mode(
self, mock_user_api_key_dict, mock_request_data
):
"""Test pre-call hook in monitor mode (logs but doesn't block)"""
guardrail = NomaGuardrail(
api_key="test-key",
monitor_mode=True,
guardrail_name="test-guardrail",
event_hook="pre_call",
default_on=True,
)
mock_response = MagicMock()
mock_response.json.return_value = {
"verdict": False,
"originalResponse": {"prompt": {"harmfulContent": {"result": True}}},
}
mock_response.raise_for_status = MagicMock()
with patch.object(guardrail.async_handler, "post", return_value=mock_response):
# Should not raise exception in monitor mode
result = await guardrail.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
cache=MagicMock(),
data=mock_request_data,
call_type="completion",
)
assert result == mock_request_data
@pytest.mark.asyncio
async def test_post_call_success_hook(
self, noma_guardrail, mock_user_api_key_dict, mock_request_data
):
"""Test post-call success hook"""
# Create a mock ModelResponse
response = ModelResponse(
id="test-response-id",
choices=[
Choices(
finish_reason="stop",
index=0,
message=Message(
content="I'm doing well, thank you!", role="assistant"
),
)
],
created=1234567890,
model="gpt-3.5-turbo",
object="chat.completion",
system_fingerprint=None,
usage={"prompt_tokens": 10, "completion_tokens": 20, "total_tokens": 30},
)
mock_api_response = MagicMock()
mock_api_response.json.return_value = {"verdict": True}
mock_api_response.raise_for_status = MagicMock()
# Update guardrail to use post_call event hook
noma_guardrail.event_hook = "post_call"
with patch.object(
noma_guardrail.async_handler, "post", return_value=mock_api_response
) as mock_post:
result = await noma_guardrail.async_post_call_success_hook(
data=mock_request_data,
user_api_key_dict=mock_user_api_key_dict,
response=response,
)
assert result == response
mock_post.assert_called_once()
# Verify API call details
call_args = mock_post.call_args
assert (
call_args[1]["json"]["response"]["text"] == "I'm doing well, thank you!"
)
assert call_args[1]["json"]["context"]["requestId"] == "test-response-id"
@pytest.mark.asyncio
async def test_moderation_hook(
self, noma_guardrail, mock_user_api_key_dict, mock_request_data
):
"""Test moderation hook (during_call)"""
# Update guardrail to use during_call event hook
noma_guardrail.event_hook = "during_call"
mock_response = MagicMock()
mock_response.json.return_value = {"verdict": True}
mock_response.raise_for_status = MagicMock()
with patch.object(
noma_guardrail.async_handler, "post", return_value=mock_response
):
result = await noma_guardrail.async_moderation_hook(
data=mock_request_data,
user_api_key_dict=mock_user_api_key_dict,
call_type="completion",
)
assert result == mock_request_data
@pytest.mark.asyncio
async def test_api_failure_handling(
self, noma_guardrail, mock_user_api_key_dict, mock_request_data
):
with patch.object(
noma_guardrail.async_handler,
"post",
side_effect=httpx.HTTPStatusError(
"API Error", request=MagicMock(), response=MagicMock(status_code=500)
),
):
with pytest.raises(httpx.HTTPStatusError):
await noma_guardrail.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
cache=MagicMock(),
data=mock_request_data,
call_type="completion",
)
@pytest.mark.asyncio
async def test_api_failure_no_block(
self, mock_user_api_key_dict, mock_request_data
):
guardrail = NomaGuardrail(
api_key="test-key",
block_failures=False,
guardrail_name="test-guardrail",
event_hook="pre_call",
default_on=True,
)
with patch.object(
guardrail.async_handler,
"post",
side_effect=httpx.HTTPStatusError(
"API Error", request=MagicMock(), response=MagicMock(status_code=500)
),
):
result = await guardrail.async_pre_call_hook(
user_api_key_dict=mock_user_api_key_dict,
cache=MagicMock(),
data=mock_request_data,
call_type="completion",
)
assert result == mock_request_data
def test_extract_user_message(self, noma_guardrail):
data = {
"messages": [
{"role": "system", "content": "System prompt"},
{"role": "user", "content": "First user message"},
{"role": "assistant", "content": "Assistant response"},
{"role": "user", "content": "Second user message"},
]
}
import asyncio
message = asyncio.run(noma_guardrail._extract_user_message(data))
assert message == "Second user message"
data = {"messages": [{"role": "system", "content": "System prompt"}]}
message = asyncio.run(noma_guardrail._extract_user_message(data))
assert message is None
data = {"messages": []}
message = asyncio.run(noma_guardrail._extract_user_message(data))
assert message is None
data = {}
message = asyncio.run(noma_guardrail._extract_user_message(data))
assert message is None
class TestIntegration:
@pytest.mark.asyncio
async def test_full_guardrail_flow(self):
"""Test full guardrail flow with multiple hooks"""
with patch.dict(
os.environ,
{
"NOMA_API_KEY": "test-api-key",
"NOMA_API_BASE": "https://api.test.noma.security/",
},
):
init_guardrails_v2(
all_guardrails=[
{
"guardrail_name": "noma-pre-guard",
"litellm_params": {
"guardrail": "noma",
"mode": "pre_call",
"application_id": "test-app",
},
},
{
"guardrail_name": "noma-post-guard",
"litellm_params": {
"guardrail": "noma",
"mode": "post_call",
"application_id": "test-app",
},
},
],
config_file_path="",
)
custom_loggers = (
litellm.logging_callback_manager.get_custom_loggers_for_type(
callback_type=litellm.integrations.custom_guardrail.CustomGuardrail
)
)
assert len(custom_loggers) >= 2

View file

@ -75,6 +75,7 @@ async def test_pangea_ai_guard_request_blocked(pangea_guardrail):
},
]
}
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
with pytest.raises(HTTPException, match="Violated Pangea guardrail policy"):
with patch(
@ -82,9 +83,9 @@ async def test_pangea_ai_guard_request_blocked(pangea_guardrail):
return_value=httpx.Response(
status_code=200,
# Mock only tested part of response
json={"result": {"blocked": True, "prompt_messages": data["messages"]}},
json={"result": {"blocked": True, "transformed": False}},
request=httpx.Request(
method="POST", url=pangea_guardrail.guardrail_endpoint
method="POST", url=guardrail_endpoint,
),
),
) as mock_method:
@ -94,7 +95,52 @@ async def test_pangea_ai_guard_request_blocked(pangea_guardrail):
called_kwargs = mock_method.call_args.kwargs
assert called_kwargs["json"]["recipe"] == "guard_llm_request"
assert called_kwargs["json"]["messages"] == data["messages"]
assert called_kwargs["json"]["input"]["messages"] == data["messages"]
@pytest.mark.asyncio
async def test_pangea_ai_guard_request_transformed(pangea_guardrail):
data = {
"messages": [
{"role": "system", "content": "You are a helpful assistant"},
{
"role": "user",
"content": "Here is an SSN for one my employees: 078-05-1120",
},
]
}
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
return_value=httpx.Response(
status_code=200,
# Mock only tested part of response
json={
"result": {
"blocked": False,
"transformed": True,
"output": {
"messages": [
{"role": "system", "content": "You are a helpful assistant"},
{
"role": "user",
"content": "Here is an SSN for one my employees: <US_SSN>",
},
]
},
},
},
request=httpx.Request(
method="POST", url=guardrail_endpoint,
),
),
):
request = await pangea_guardrail.async_pre_call_hook(
user_api_key_dict=None, cache=None, data=data, call_type="completion"
)
assert request["messages"][1]["content"] == "Here is an SSN for one my employees: <US_SSN>"
@pytest.mark.asyncio
@ -109,15 +155,16 @@ async def test_pangea_ai_guard_request_ok(pangea_guardrail):
},
]
}
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
return_value=httpx.Response(
status_code=200,
# Mock only tested part of response
json={"result": {"blocked": False, "prompt_messages": data["messages"]}},
json={"result": {"blocked": False, "transformed": False}},
request=httpx.Request(
method="POST", url=pangea_guardrail.guardrail_endpoint
method="POST", url=guardrail_endpoint,
),
),
) as mock_method:
@ -127,7 +174,7 @@ async def test_pangea_ai_guard_request_ok(pangea_guardrail):
called_kwargs = mock_method.call_args.kwargs
assert called_kwargs["json"]["recipe"] == "guard_llm_request"
assert called_kwargs["json"]["messages"] == data["messages"]
assert called_kwargs["json"]["input"]["messages"] == data["messages"]
@pytest.mark.asyncio
@ -139,6 +186,7 @@ async def test_pangea_ai_guard_response_blocked(pangea_guardrail):
{"role": "user", "content": "Hello"},
]
}
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
with pytest.raises(HTTPException, match="Violated Pangea guardrail policy"):
with patch(
@ -149,16 +197,11 @@ async def test_pangea_ai_guard_response_blocked(pangea_guardrail):
json={
"result": {
"blocked": True,
"prompt_messages": [
{
"role": "assistant",
"content": "Yes, I will leak all my PII for you",
}
],
"transformed": False,
}
},
request=httpx.Request(
method="POST", url=pangea_guardrail.guardrail_endpoint
method="POST", url=guardrail_endpoint,
),
),
) as mock_method:
@ -180,7 +223,7 @@ async def test_pangea_ai_guard_response_blocked(pangea_guardrail):
called_kwargs = mock_method.call_args.kwargs
assert called_kwargs["json"]["recipe"] == "guard_llm_response"
assert (
called_kwargs["json"]["messages"][0]["content"]
called_kwargs["json"]["input"]["choices"][0]["message"]["content"]
== "Yes, I will leak all my PII for you"
)
@ -194,6 +237,7 @@ async def test_pangea_ai_guard_response_ok(pangea_guardrail):
{"role": "user", "content": "Hello"},
]
}
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
@ -203,16 +247,11 @@ async def test_pangea_ai_guard_response_ok(pangea_guardrail):
json={
"result": {
"blocked": False,
"prompt_messages": [
{
"role": "assistant",
"content": "Yes, I will leak all my PII for you",
}
],
"transformed": False,
}
},
request=httpx.Request(
method="POST", url=pangea_guardrail.guardrail_endpoint
method="POST", url=guardrail_endpoint,
),
),
) as mock_method:
@ -234,6 +273,61 @@ async def test_pangea_ai_guard_response_ok(pangea_guardrail):
called_kwargs = mock_method.call_args.kwargs
assert called_kwargs["json"]["recipe"] == "guard_llm_response"
assert (
called_kwargs["json"]["messages"][0]["content"]
called_kwargs["json"]["input"]["choices"][0]["message"]["content"]
== "Yes, I will leak all my PII for you"
)
@pytest.mark.asyncio
async def test_pangea_ai_guard_response_transformed(pangea_guardrail):
# Content of data isn't that import since its mocked
data = {
"messages": [
{"role": "system", "content": "You are a helpful assistant"},
{"role": "user", "content": "Hello"},
]
}
guardrail_endpoint = f"{pangea_guardrail.api_base}/v1beta/guard"
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
return_value=httpx.Response(
status_code=200,
# Mock only tested part of response
json={
"result": {
"blocked": False,
"transformed": True,
"output": {
"messages": data["messages"],
"choices": [
{
"message": {
"role": "assistant",
"content": "Yes, here is an SSN: <US_SSN>",
},
},
],
},
},
},
request=httpx.Request(
method="POST", url=guardrail_endpoint,
),
),
):
response = await pangea_guardrail.async_post_call_success_hook(
data=data,
user_api_key_dict=None,
response=ModelResponse(
choices=[
{
"message": {
"role": "assistant",
"content": "Yes, here is an SSN: 078-05-1120",
}
}
]
),
)
assert response.choices[0]["message"]["content"] == "Yes, here is an SSN: <US_SSN>"

View file

@ -0,0 +1,147 @@
"""
Test for issue #13995: /batches request throws Internal Server Error when metadata=None
This test verifies that the fix for handling None metadata in batch requests works correctly.
"""
import asyncio
import os
import sys
from unittest.mock import patch, MagicMock, AsyncMock
import pytest
from openai import OpenAI
import litellm
from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup
from litellm.proxy._types import UserAPIKeyAuth
sys.path.insert(
0, os.path.abspath("../../..")
) # Adds the parent directory to the system path
def test_add_key_level_controls_with_none_metadata():
"""
Test that add_key_level_controls handles None metadata gracefully.
This is the core fix for issue #13995.
"""
# Test data
data = {"metadata": {}}
metadata_variable_name = "metadata"
# Test with None key_metadata (this was causing the original error)
result = LiteLLMProxyRequestSetup.add_key_level_controls(
key_metadata=None,
data=data,
_metadata_variable_name=metadata_variable_name
)
# Should return the data unchanged without throwing an error
assert result == data
# Test with empty dict key_metadata (should also work)
result = LiteLLMProxyRequestSetup.add_key_level_controls(
key_metadata={},
data=data,
_metadata_variable_name=metadata_variable_name
)
# Should return the data unchanged
assert result == data
# Test with valid key_metadata containing cache settings
key_metadata_with_cache = {
"cache": {
"ttl": 300,
"s-maxage": 600
}
}
result = LiteLLMProxyRequestSetup.add_key_level_controls(
key_metadata=key_metadata_with_cache,
data=data.copy(),
_metadata_variable_name=metadata_variable_name
)
# Should add cache settings to data
assert "cache" in result
assert result["cache"]["ttl"] == 300
assert result["cache"]["s-maxage"] == 600
def test_add_key_level_controls_simulates_original_issue():
"""
Test that simulates the original issue scenario more directly.
This tests the exact code path that was failing in issue #13995.
"""
# This simulates the scenario where user_api_key_dict.metadata is None
# which was causing the original "'NoneType' object has no attribute 'get'" error
data = {"metadata": {}}
metadata_variable_name = "metadata"
# This is the exact call that was failing before the fix
# user_api_key_dict.metadata was None, causing the error in add_key_level_controls
try:
result = LiteLLMProxyRequestSetup.add_key_level_controls(
key_metadata=None, # This was the root cause of the issue
data=data,
_metadata_variable_name=metadata_variable_name
)
# If we get here, the fix is working
assert result == data
print("✓ Original issue scenario handled correctly - no NoneType error")
except AttributeError as e:
if "'NoneType' object has no attribute 'get'" in str(e):
pytest.fail("The fix for issue #13995 is not working - still getting NoneType error")
else:
# Some other AttributeError, re-raise it
raise
def test_batch_create_with_litellm_sdk():
"""
Test creating a batch using litellm SDK with metadata=None.
This is a more direct test of the original issue.
"""
# Mock the OpenAI batches instance to avoid actual API calls
with patch('litellm.batches.main.openai_batches_instance') as mock_openai_batches:
# Mock the response
mock_response = MagicMock()
mock_response.id = "batch_test123"
mock_openai_batches.create_batch.return_value = mock_response
# This should not raise an exception
try:
response = litellm.create_batch(
completion_window="24h",
endpoint="/v1/chat/completions",
input_file_id="file-test123",
metadata=None, # This was causing the original issue
custom_llm_provider="openai"
)
assert response.id == "batch_test123"
except Exception as e:
if "'NoneType' object has no attribute 'get'" in str(e):
pytest.fail("The fix for issue #13995 is not working - still getting NoneType error")
else:
# Some other exception, re-raise it
raise
if __name__ == "__main__":
# Run the tests
test_add_key_level_controls_with_none_metadata()
print("✓ test_add_key_level_controls_with_none_metadata passed")
test_add_key_level_controls_simulates_original_issue()
print("✓ test_add_key_level_controls_simulates_original_issue passed")
test_batch_create_with_litellm_sdk()
print("✓ test_batch_create_with_litellm_sdk passed")
print("All tests passed! Issue #13995 fix is working correctly.")

View file

@ -541,7 +541,8 @@ class TestFunctionCallTransformation:
result = LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request(
model="gemini/gemini-2.0-flash",
input=test_input,
responses_api_request=responses_api_request
responses_api_request=responses_api_request,
extra_headers={"X-Test-Header": "test-value"}
)
assert "messages" in result
@ -563,6 +564,8 @@ class TestFunctionCallTransformation:
tool_msg = messages[2]
assert tool_msg["role"] == "tool"
assert result["extra_headers"] == {"X-Test-Header": "test-value"}
def test_function_call_without_call_id_fallback_to_id(self):
"""Test that function_call items can use 'id' field when 'call_id' is missing"""
function_call_item = {

View file

@ -5,11 +5,14 @@ This test verifies that the OpenAI-like handler correctly handles
UTF-8 encoded content in streaming responses, specifically fixing
the ASCII encoding error described in issue #12660.
"""
import pytest
import asyncio
from unittest.mock import Mock, AsyncMock
from unittest.mock import AsyncMock, Mock
import pytest
from litellm.llms.openai_like.chat.handler import make_call, make_sync_call
class MockResponse:
"""Mock httpx response for testing UTF-8 handling."""
@ -25,6 +28,14 @@ class MockResponse:
"""Mock aiter_text that yields content with the specified encoding."""
yield self.test_content
def iter_lines(self):
"""Mock iter_lines method for synchronous streaming."""
yield self.test_content
async def aiter_lines(self):
"""Mock aiter_lines method for asynchronous streaming."""
yield self.test_content
def json(self):
return {"choices": [{"delta": {"content": "test"}}]}

View file

@ -1121,4 +1121,119 @@ async def test_retrying() -> None:
model="gpt-4o-mini",
messages=[{"role": "user", "content": "Hello"}],
)
assert mock_request.call_count >= 10, "Expected retrying to be used"
def test_anthropic_disable_url_suffix_env_var():
"""Test that LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX prevents /v1/messages suffix."""
from unittest.mock import patch, MagicMock
import os
from litellm import completion
# Test with environment variable disabled (default behavior)
with patch.dict(os.environ, {"ANTHROPIC_API_BASE": "https://api.example.com"}):
actual_api_base = None
with patch("litellm.main.anthropic_chat_completions") as mock_anthropic:
def capture_completion(**kwargs):
nonlocal actual_api_base
actual_api_base = kwargs.get("api_base")
mock_response = MagicMock()
mock_response.choices = [MagicMock()]
return mock_response
mock_anthropic.completion = capture_completion
# This should append /v1/messages
completion(
model="anthropic/claude-3-sonnet",
messages=[{"role": "user", "content": "test"}],
api_key="test-key"
)
# Verify the api_base has /v1/messages appended
assert actual_api_base.endswith("/v1/messages")
assert actual_api_base == "https://api.example.com/v1/messages"
# Test with environment variable enabled
with patch.dict(os.environ, {
"ANTHROPIC_API_BASE": "https://api.example.com/custom/path",
"LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX": "true"
}):
actual_api_base = None
with patch("litellm.main.anthropic_chat_completions") as mock_anthropic:
def capture_completion(**kwargs):
nonlocal actual_api_base
actual_api_base = kwargs.get("api_base")
mock_response = MagicMock()
mock_response.choices = [MagicMock()]
return mock_response
mock_anthropic.completion = capture_completion
# This should NOT append /v1/messages
completion(
model="anthropic/claude-3-sonnet",
messages=[{"role": "user", "content": "test"}],
api_key="test-key"
)
# Verify the api_base does not have /v1/messages appended
assert actual_api_base == "https://api.example.com/custom/path"
assert not actual_api_base.endswith("/v1/messages")
def test_anthropic_text_disable_url_suffix_env_var():
"""Test that LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX prevents /v1/complete suffix for anthropic_text."""
from unittest.mock import patch, MagicMock
import os
from litellm import completion
# Test with environment variable disabled (default behavior)
with patch.dict(os.environ, {"ANTHROPIC_API_BASE": "https://api.example.com"}):
actual_api_base = None
with patch("litellm.main.base_llm_http_handler") as mock_handler:
def capture_completion(**kwargs):
nonlocal actual_api_base
actual_api_base = kwargs.get("api_base")
return MagicMock()
mock_handler.completion = capture_completion
# This should append /v1/complete
completion(
model="anthropic_text/claude-instant-1",
messages=[{"role": "user", "content": "test"}],
api_key="test-key"
)
# Verify the api_base has /v1/complete appended
assert actual_api_base.endswith("/v1/complete")
assert actual_api_base == "https://api.example.com/v1/complete"
# Test with environment variable enabled
with patch.dict(os.environ, {
"ANTHROPIC_API_BASE": "https://api.example.com/custom/complete",
"LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX": "true"
}):
actual_api_base = None
with patch("litellm.main.base_llm_http_handler") as mock_handler:
def capture_completion(**kwargs):
nonlocal actual_api_base
actual_api_base = kwargs.get("api_base")
return MagicMock()
mock_handler.completion = capture_completion
# This should NOT append /v1/complete
completion(
model="anthropic_text/claude-instant-1",
messages=[{"role": "user", "content": "test"}],
api_key="test-key"
)
# Verify the api_base does not have /v1/complete appended
assert actual_api_base == "https://api.example.com/custom/complete"
assert not actual_api_base.endswith("/v1/complete")

View file

@ -0,0 +1,72 @@
"""
Test for GitHub issue #11267 - System message format issue with Ollama + tools
"""
from unittest.mock import patch
@patch("litellm.add_function_to_prompt", True)
def test_system_message_format_issue_reproduction():
"""
Reproduces the system message format bug from GitHub issue #11267.
"""
from litellm import completion
# Define test data directly from data.jsonl content
model = "ollama/custom_model_name" # Use explicit Ollama model
messages = [
{
"role": "user",
"content": [
{
"type": "text",
"text": "What is the capital of France?"
}
]
},
{
"role": "system",
"content": [
{
"type": "text",
"text": "You are Claude Code, Anthropic's official CLI for Claude.",
"cache_control": {"type": "ephemeral"}
}
]
}
]
temperature = 1
# Add tools to trigger the bug - this is what causes the issue
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get weather for a location",
"parameters": {
"type": "object",
"properties": {
"location": {"type": "string"}
},
"required": ["location"]
}
}
}
]
response = completion(
model=model,
messages=messages,
tools=tools,
temperature=temperature,
mock_response=True
)
assert len(messages[1]["content"]) == 2
if __name__ == "__main__":
print("Testing system message format issue...")
test_system_message_format_issue_reproduction()
print("Tests completed!")

View file

@ -847,6 +847,7 @@ async def test_supports_tool_choice():
or "o3" in model_name
or "mistral" in model_name
or "oci" in model_name
or "openrouter" in model_name
):
continue
@ -957,7 +958,12 @@ def test_get_model_info_shows_supports_computer_use():
def test_pre_process_non_default_params(model, custom_llm_provider):
from pydantic import BaseModel
from litellm.utils import pre_process_non_default_params
from litellm.utils import ProviderConfigManager, pre_process_non_default_params
provider_config = ProviderConfigManager.get_provider_chat_config(
model=model,
provider=LlmProviders(custom_llm_provider)
)
class ResponseFormat(BaseModel):
x: str
@ -974,6 +980,7 @@ def test_pre_process_non_default_params(model, custom_llm_provider):
special_params=special_params,
custom_llm_provider=custom_llm_provider,
additional_drop_params=None,
provider_config=provider_config,
)
print(processed_non_default_params)
assert processed_non_default_params == {

View file

@ -170,8 +170,8 @@ const ChatUI: React.FC<ChatUIProps> = ({
const saved = sessionStorage.getItem('useApiSessionManagement');
return saved ? JSON.parse(saved) : true; // Default to API session management
});
const [uploadedImage, setUploadedImage] = useState<File | null>(null);
const [imagePreviewUrl, setImagePreviewUrl] = useState<string | null>(null);
const [uploadedImages, setUploadedImages] = useState<File[]>([]);
const [imagePreviewUrls, setImagePreviewUrls] = useState<string[]>([]);
const [responsesUploadedImage, setResponsesUploadedImage] = useState<File | null>(null);
const [responsesImagePreviewUrl, setResponsesImagePreviewUrl] = useState<string | null>(null);
const [chatUploadedImage, setChatUploadedImage] = useState<File | null>(null);
@ -450,6 +450,38 @@ const ChatUI: React.FC<ChatUIProps> = ({
]);
};
const updateChatImageUI = (imageUrl: string, model?: string) => {
setChatHistory((prev) => {
const last = prev[prev.length - 1];
// If the last message is from assistant and has content, add image to it
if (last && last.role === "assistant" && !last.isImage) {
const updated = {
...last,
image: {
url: imageUrl,
detail: "auto"
},
model: last.model ?? model
};
return [...prev.slice(0, -1), updated];
} else {
// Otherwise create a new assistant message with just the image
return [
...prev,
{
role: "assistant",
content: "",
model,
image: {
url: imageUrl,
detail: "auto"
}
}
];
}
});
};
const handleKeyDown = (event: React.KeyboardEvent<HTMLTextAreaElement>) => {
if (event.key === 'Enter' && !event.shiftKey) {
event.preventDefault(); // Prevent default to avoid newline
@ -468,18 +500,26 @@ const ChatUI: React.FC<ChatUIProps> = ({
};
const handleImageUpload = (file: File) => {
setUploadedImage(file);
setUploadedImages(prev => [...prev, file]);
const previewUrl = URL.createObjectURL(file);
setImagePreviewUrl(previewUrl);
setImagePreviewUrls(prev => [...prev, previewUrl]);
return false; // Prevent default upload behavior
};
const handleRemoveImage = () => {
if (imagePreviewUrl) {
URL.revokeObjectURL(imagePreviewUrl);
const handleRemoveImage = (index: number) => {
if (imagePreviewUrls[index]) {
URL.revokeObjectURL(imagePreviewUrls[index]);
}
setUploadedImage(null);
setImagePreviewUrl(null);
setUploadedImages(prev => prev.filter((_, i) => i !== index));
setImagePreviewUrls(prev => prev.filter((_, i) => i !== index));
};
const handleRemoveAllImages = () => {
imagePreviewUrls.forEach(url => {
URL.revokeObjectURL(url);
});
setUploadedImages([]);
setImagePreviewUrls([]);
};
const handleResponsesImageUpload = (file: File): false => {
@ -516,8 +556,8 @@ const ChatUI: React.FC<ChatUIProps> = ({
if (inputMessage.trim() === "") return;
// For image edits, require both image and prompt
if (endpointType === EndpointType.IMAGE_EDITS && !uploadedImage) {
NotificationsManager.fromBackend("Please upload an image for editing");
if (endpointType === EndpointType.IMAGE_EDITS && uploadedImages.length === 0) {
NotificationsManager.fromBackend("Please upload at least one image for editing");
return;
}
@ -603,7 +643,8 @@ const ChatUI: React.FC<ChatUIProps> = ({
traceId,
selectedVectorStores.length > 0 ? selectedVectorStores : undefined,
selectedGuardrails.length > 0 ? selectedGuardrails : undefined,
selectedMCPTools // Pass the selected tool directly
selectedMCPTools, // Pass the selected tool directly
updateChatImageUI // Pass the image callback
);
} else if (endpointType === EndpointType.IMAGE) {
// For image generation
@ -617,9 +658,9 @@ const ChatUI: React.FC<ChatUIProps> = ({
);
} else if (endpointType === EndpointType.IMAGE_EDITS) {
// For image edits
if (uploadedImage) {
if (uploadedImages.length > 0) {
await makeOpenAIImageEditsRequest(
uploadedImage,
uploadedImages.length === 1 ? uploadedImages[0] : uploadedImages,
inputMessage,
(imageUrl, model) => updateImageUI(imageUrl, model),
selectedModel,
@ -689,7 +730,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
abortControllerRef.current = null;
// Clear image after successful request for image edits
if (endpointType === EndpointType.IMAGE_EDITS) {
handleRemoveImage();
handleRemoveAllImages();
}
// Clear image after successful request for responses API
if (endpointType === EndpointType.RESPONSES && responsesUploadedImage) {
@ -708,7 +749,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
setChatHistory([]);
setMessageTraceId(null);
setResponsesSessionId(null); // Clear responses session ID
handleRemoveImage(); // Clear any uploaded images for image edits
handleRemoveAllImages(); // Clear any uploaded images for image edits
handleRemoveResponsesImage(); // Clear any uploaded images for responses
handleRemoveChatImage(); // Clear any uploaded images for chat completions
sessionStorage.removeItem('chatHistory');
@ -1049,6 +1090,18 @@ const ChatUI: React.FC<ChatUIProps> = ({
>
{typeof message.content === "string" ? message.content : ""}
</ReactMarkdown>
{/* Show generated image from chat completions */}
{message.image && (
<div className="mt-3">
<img
src={message.image.url}
alt="Generated image"
className="max-w-full rounded-md border border-gray-200 shadow-sm"
style={{ maxHeight: '500px' }}
/>
</div>
)}
</>
)}
@ -1075,7 +1128,7 @@ const ChatUI: React.FC<ChatUIProps> = ({
{/* Image Upload Section for Image Edits */}
{endpointType === EndpointType.IMAGE_EDITS && (
<div className="mb-4">
{!uploadedImage ? (
{uploadedImages.length === 0 ? (
<Dragger
beforeUpload={handleImageUpload}
accept="image/*"
@ -1085,24 +1138,47 @@ const ChatUI: React.FC<ChatUIProps> = ({
<p className="ant-upload-drag-icon">
<PictureOutlined style={{ fontSize: '24px', color: '#666' }} />
</p>
<p className="ant-upload-text text-sm">Click or drag image to upload</p>
<p className="ant-upload-text text-sm">Click or drag images to upload</p>
<p className="ant-upload-hint text-xs text-gray-500">
Support for PNG, JPG, JPEG formats
Support for PNG, JPG, JPEG formats. Multiple images supported.
</p>
</Dragger>
) : (
<div className="relative inline-block">
<img
src={imagePreviewUrl || ''}
alt="Upload preview"
className="max-w-32 max-h-32 rounded-md border border-gray-200 object-cover"
/>
<button
className="absolute top-1 right-1 bg-white shadow-sm border border-gray-200 rounded px-1 py-1 text-red-500 hover:bg-red-50 text-xs"
onClick={handleRemoveImage}
>
<DeleteOutlined />
</button>
<div className="flex flex-wrap gap-2">
{uploadedImages.map((file, index) => (
<div key={index} className="relative inline-block">
<img
src={imagePreviewUrls[index] || ''}
alt={`Upload preview ${index + 1}`}
className="max-w-32 max-h-32 rounded-md border border-gray-200 object-cover"
/>
<button
className="absolute top-1 right-1 bg-white shadow-sm border border-gray-200 rounded px-1 py-1 text-red-500 hover:bg-red-50 text-xs"
onClick={() => handleRemoveImage(index)}
>
<DeleteOutlined />
</button>
</div>
))}
{/* Add more images button */}
<div className="flex items-center justify-center w-32 h-32 border-2 border-dashed border-gray-300 rounded-md hover:border-gray-400 cursor-pointer"
onClick={() => document.getElementById('additional-image-upload')?.click()}>
<div className="text-center">
<PictureOutlined style={{ fontSize: '24px', color: '#666' }} />
<p className="text-xs text-gray-500 mt-1">Add more</p>
</div>
<input
id="additional-image-upload"
type="file"
accept="image/*"
multiple
style={{ display: 'none' }}
onChange={(e) => {
const files = Array.from(e.target.files || []);
files.forEach(file => handleImageUpload(file));
}}
/>
</div>
</div>
)}
</div>

View file

@ -17,7 +17,8 @@ export async function makeOpenAIChatCompletionRequest(
traceId?: string,
vector_store_ids?: string[],
guardrails?: string[],
selectedMCPTool?: string
selectedMCPTool?: string,
onImageGenerated?: (imageUrl: string, model?: string) => void
) {
// base url should be the current base_url
const isLocal = process.env.NODE_ENV === "development";
@ -103,6 +104,12 @@ export async function makeOpenAIChatCompletionRequest(
fullResponseContent += content;
}
// Process image generation if present
if (delta && delta.image && onImageGenerated) {
console.log("Image generated:", delta.image);
onImageGenerated(delta.image.url, chunk.model);
}
// Process reasoning content if present - using type assertion
if (delta && delta.reasoning_content) {
const reasoningContent = delta.reasoning_content;

View file

@ -4,7 +4,7 @@ import { getProxyBaseUrl } from "@/components/networking";
import NotificationManager from "@/components/molecules/notifications_manager";
export async function makeOpenAIImageEditsRequest(
imageFile: File,
imageFiles: File | File[],
prompt: string,
updateUI: (imageUrl: string, model: string) => void,
selectedModel: string,
@ -28,34 +28,60 @@ export async function makeOpenAIImageEditsRequest(
});
try {
const response = await client.images.edit({
model: selectedModel,
image: imageFile,
prompt: prompt,
}, { signal });
console.log(response.data);
// handle single and multiple images
const imagesToProcess = Array.isArray(imageFiles) ? imageFiles : [imageFiles];
if (response.data && response.data[0]) {
// Handle either URL or base64 data from response
if (response.data[0].url) {
// Use the URL directly
updateUI(response.data[0].url, selectedModel);
} else if (response.data[0].b64_json) {
// Convert base64 to data URL format
const base64Data = response.data[0].b64_json;
updateUI(`data:image/png;base64,${base64Data}`, selectedModel);
} else {
throw new Error("No image data found in response");
// For multiple images, we'll make separate API calls for each image
// since OpenAI's edit endpoint processes one image at a time
const results = [];
for (let i = 0; i < imagesToProcess.length; i++) {
const image = imagesToProcess[i];
console.log(`Processing image ${i + 1} of ${imagesToProcess.length}`);
const response = await client.images.edit({
model: selectedModel,
image: image,
prompt: prompt,
}, { signal });
console.log(`Response for image ${i + 1}:`, response.data);
if (response.data && response.data[0]) {
// Handle either URL or base64 data from response
if (response.data[0].url) {
// Use the URL directly
updateUI(response.data[0].url, selectedModel);
results.push(response.data[0].url);
} else if (response.data[0].b64_json) {
// Convert base64 to data URL format
const base64Data = response.data[0].b64_json;
const dataUrl = `data:image/png;base64,${base64Data}`;
updateUI(dataUrl, selectedModel);
results.push(dataUrl);
}
}
} else {
throw new Error("Invalid response format");
}
} catch (error) {
if (results.length > 1) {
NotificationManager.success(`Successfully processed ${results.length} images`);
}
} catch (error: any) {
console.error("Error making image edit request:", error);
if (signal?.aborted) {
console.log("Image edits request was cancelled");
} else {
NotificationManager.fromBackend(`Error occurred while editing image. Please try again. Error: ${error}`);
let errorMessage = "Failed to edit image(s)";
if (error?.error?.message) {
errorMessage = error.error.message;
} else if (error?.message) {
errorMessage = error.message;
}
NotificationManager.fromBackend(`Image edit failed: ${errorMessage}`);
}
throw error; // Re-throw to allow the caller to handle the error
}

View file

@ -7,6 +7,10 @@ export interface Delta {
audio?: any;
refusal?: any;
provider_specific_fields?: any;
image?: {
url: string;
detail: string;
};
}
export interface CompletionTokensDetails {
@ -67,6 +71,10 @@ export interface MessageType {
};
toolName?: string;
imagePreviewUrl?: string; // For storing image preview URL in chat history
image?: {
url: string;
detail: string;
};
}
export interface MultimodalContent {

View file

@ -1383,6 +1383,20 @@ const Teams: React.FC<TeamProps> = ({
>
<TextInput placeholder="e.g., 30d" />
</Form.Item>
<Form.Item
label="Team Member RPM Limit"
name="team_member_rpm_limit"
tooltip="The RPM (Requests Per Minute) limit for individual team members"
>
<NumericalInput step={1} width={400} />
</Form.Item>
<Form.Item
label="Team Member TPM Limit"
name="team_member_tpm_limit"
tooltip="The TPM (Tokens Per Minute) limit for individual team members"
>
<NumericalInput step={1} width={400} />
</Form.Item>
<Form.Item
label="Metadata"
name="metadata"