Merge branch 'main' into litellm_oss_staging_01_26_2026

This commit is contained in:
Sameer Kankute 2026-01-27 17:00:58 +05:30 • committed by GitHub
commit 0214cb04cd
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
126 changed files with 7461 additions and 641 deletions

View file

@ -16,4 +16,5 @@ uvloop==0.21.0
mcp==1.25.0 # for MCP server
semantic_router==0.1.10 # for auto-routing with litellm
fastuuid==0.12.0
responses==0.25.7 # for proxy client tests
responses==0.25.7 # for proxy client tests
pytest-retry==1.6.3 # for automatic test retries

View file

@ -51,12 +51,14 @@ LiteLLM is a unified interface for 100+ LLMs that:
### MAKING CODE CHANGES FOR THE UI (IGNORE FOR BACKEND)
1. **Use Common Components as much as possible**:
1. **Tremor is DEPRECATED, do not use Tremor components in new features/changes**
- The only exception is the Tremor Table component and its required Tremor Table sub components.
2. **Use Common Components as much as possible**:
- These are usually defined in the `common_components` directory
- Use these components as much as possible and avoid building new components unless needed
- Tremor components are deprecated; prefer using Ant Design (AntD) as much as possible
2. **Testing**:
3. **Testing**:
- The codebase uses **Vitest** and **React Testing Library**
- **Query Priority Order**: Use query methods in this order: `getByRole`, `getByLabelText`, `getByPlaceholderText`, `getByText`, `getByTestId`
- **Always use `screen`** instead of destructuring from `render()` (e.g., use `screen.getByText()` not `getByText`)

View file

@ -69,8 +69,8 @@ RUN find /usr/lib -type f -path "*/tornado/test/*" -delete && \
# Convert Windows line endings to Unix and make executable
RUN sed -i 's/\r$//' docker/install_auto_router.sh && chmod +x docker/install_auto_router.sh && ./docker/install_auto_router.sh
# Generate prisma client
RUN prisma generate
# Generate prisma client using the correct schema
RUN prisma generate --schema=./litellm/proxy/schema.prisma
# Convert Windows line endings to Unix for entrypoint scripts
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh
RUN sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh

View file

@ -267,6 +267,7 @@ Support for more providers. Missing a provider or LLM Platform, raise a [feature
<td><img height="60" alt="Greptile" src="https://github.com/user-attachments/assets/0be4bd8a-7cfa-48d3-9090-f415fe948280" /></td>
<td><img height="60" alt="OpenHands" src="https://github.com/user-attachments/assets/a6150c4c-149e-4cae-888b-8b92be6e003f" /></td>
<td><h2>Netflix</h2></td>
<td><img height="60" alt="OpenAI Agents SDK" src="https://github.com/user-attachments/assets/c02f7be0-8c2e-4d27-aea7-7c024bfaebc0" /></td>
</tr>
</table>

View file

@ -68,7 +68,7 @@ Follow [this guide, to add your pydantic ai agent to LiteLLM Agent Gateway](./pr
## Invoking your Agents
Use the [A2A Python SDK](https://pypi.org/project/a2a/) to invoke agents through LiteLLM.
Use the [A2A Python SDK](https://pypi.org/project/a2a-sdk) to invoke agents through LiteLLM.
This example shows how to:
1. **List available agents** - Query `/v1/agents` to see which agents your key can access
@ -193,6 +193,120 @@ The logs show:
style={{width: '100%', display: 'block', margin: '2rem auto'}}
/>
## Forwarding LiteLLM Context Headers
When LiteLLM invokes your A2A agent, it sends special headers that enable:
- **Trace Grouping**: All LLM calls from the same agent execution appear under one trace
- **Agent Spend Tracking**: Costs are attributed to the specific agent
| Header | Purpose |
|--------|---------|
| `X-LiteLLM-Trace-Id` | Links all LLM calls to the same execution flow |
| `X-LiteLLM-Agent-Id` | Attributes spend to the correct agent |
To enable these features, your A2A server must **forward these headers** to any LLM calls it makes back to LiteLLM.
### Implementation Steps
**Step 1: Extract headers from incoming A2A request**
```python def get_litellm_headers(request) -> dict:
"""Extract X-LiteLLM-* headers from incoming A2A request."""
all_headers = request.call_context.state.get('headers', {})
return {
k: v for k, v in all_headers.items()
if k.lower().startswith('x-litellm-')
}
```
**Step 2: Forward headers to your LLM calls**
Pass the extracted headers when making calls back to LiteLLM:
<Tabs>
<TabItem value="openai" label="OpenAI SDK" default>
```python from openai import OpenAI
headers = get_litellm_headers(request)
client = OpenAI(
api_key="sk-your-litellm-key",
base_url="http://localhost:4000",
default_headers=headers, # Forward headers
)
response = client.chat.completions.create(
model="gpt-4o",
messages=[{"role": "user", "content": "Hello"}]
)
```
</TabItem>
<TabItem value="langchain" label="LangChain">
```python
from langchain_openai import ChatOpenAI
headers = get_litellm_headers(request)
llm = ChatOpenAI(
model="gpt-4o",
openai_api_key="sk-your-litellm-key",
base_url="http://localhost:4000",
default_headers=headers, # Forward headers
)
```
</TabItem>
<TabItem value="litellm" label="LiteLLM SDK">
```python
import litellm
headers = get_litellm_headers(request)
response = litellm.completion(
model="gpt-4o",
messages=[{"role": "user", "content": "Hello"}],
api_base="http://localhost:4000",
extra_headers=headers, # Forward headers
)
```
</TabItem>
<TabItem value="requests" label="HTTP (requests/httpx)">
```python
import httpx
headers = get_litellm_headers(request)
headers["Authorization"] = "Bearer sk-your-litellm-key"
response = httpx.post(
"http://localhost:4000/v1/chat/completions",
headers=headers,
json={"model": "gpt-4o", "messages": [{"role": "user", "content": "Hello"}]}
)
```
</TabItem>
</Tabs>
### Result
With header forwarding enabled, you'll see:
**Trace Grouping in Langfuse:**
<Image
img={require('../img/a2a_trace_grouping.png')}
style={{width: '80%', display: 'block', margin: '0', borderRadius: '8px'}}
/>
**Agent Spend Attribution:**
<Image
img={require('../img/a2a_agent_spend.png')}
style={{width: '80%', display: 'block', margin: '0', borderRadius: '8px'}}
/>
## API Reference
### Endpoint

View file

@ -7,6 +7,7 @@ import TabItem from '@theme/TabItem';
LiteLLM Supports logging to the following Datdog Integrations:
- `datadog` [Datadog Logs](https://docs.datadoghq.com/logs/)
- `datadog_llm_observability` [Datadog LLM Observability](https://www.datadoghq.com/product/llm-observability/)
- `datadog_cost_management` [Datadog Cloud Cost Management](#datadog-cloud-cost-management)
- `ddtrace-run` [Datadog Tracing](#datadog-tracing)
## Datadog Logs
@ -73,7 +74,7 @@ Send logs through a local DataDog agent (useful for containerized environments):
```shell
LITELLM_DD_AGENT_HOST="localhost" # hostname or IP of DataDog agent
LITELLM_DD_AGENT_PORT="10518" # [OPTIONAL] port of DataDog agent (default: 10518)
DD_API_KEY="5f2d0f310***********" # [OPTIONAL] your datadog API Key (agent handles auth)
DD_API_KEY="5f2d0f310***********" # [OPTIONAL] your datadog API Key (Agent handles auth for Logs. REQUIRED for LLM Observability)
DD_SOURCE="litellm_dev" # [OPTIONAL] your datadog source
```
@ -84,6 +85,9 @@ When `LITELLM_DD_AGENT_HOST` is set, logs are sent to the agent instead of direc
**Note:** We use `LITELLM_DD_AGENT_HOST` instead of `DD_AGENT_HOST` to avoid conflicts with `ddtrace` which automatically sets `DD_AGENT_HOST` for APM tracing.
> [!IMPORTANT]
> **Datadog LLM Observability**: `DD_API_KEY` is **REQUIRED** even when using the Datadog Agent (`LITELLM_DD_AGENT_HOST`). The agent acts as a proxy but the API key header is mandatory for the LLM Observability endpoint.
**Step 3**: Start the proxy, make a test request
Start proxy
@ -161,6 +165,50 @@ On the Datadog LLM Observability page, you should see that both input messages a
<Image img={require('../../img/dd_llm_obs.png')} />
## Datadog Cloud Cost Management
| Feature | Details |
|---------|---------|
| **What is logged** | Aggregated LLM Costs (FOCUS format) |
| **Events** | Periodic Uploads of Aggregated Cost Data |
| **Product Link** | [Datadog Cloud Cost Management](https://docs.datadoghq.com/cost_management/) |
We will use the `--config` to set `litellm.callbacks = ["datadog_cost_management"]`. This will periodically upload aggregated LLM cost data to Datadog.
**Step 1**: Create a `config.yaml` file and set `litellm_settings`: `success_callback`
```yaml
model_list:
- model_name: gpt-3.5-turbo
litellm_params:
model: gpt-3.5-turbo
litellm_settings:
callbacks: ["datadog_cost_management"]
```
**Step 2**: Set Required env variables
```shell
DD_API_KEY="your-api-key"
DD_APP_KEY="your-app-key" # REQUIRED for Cost Management
DD_SITE="us5.datadoghq.com"
```
**Step 3**: Start the proxy
```shell
litellm --config config.yaml
```
**How it works**
* LiteLLM aggregates costs in-memory by Provider, Model, Date, and Tags.
* Requires `DD_APP_KEY` for the Custom Costs API.
* Costs are uploaded periodically (flushed).
### Datadog Tracing
Use `ddtrace-run` to enable [Datadog Tracing](https://ddtrace.readthedocs.io/en/stable/installation_quickstart.html) on litellm proxy
@ -203,5 +251,5 @@ LiteLLM supports customizing the following Datadog environment variables
| `POD_NAME` | Pod name tag (useful for Kubernetes deployments) | "unknown" | ❌ No |
\* **Required when using Direct API** (default): `DD_API_KEY` and `DD_SITE` are required
\* **Optional when using DataDog Agent**: Set `LITELLM_DD_AGENT_HOST` to use agent mode; `DD_API_KEY` and `DD_SITE` are not required
\* **Optional when using DataDog Agent**: Set `LITELLM_DD_AGENT_HOST` to use agent mode; `DD_API_KEY` and `DD_SITE` are not required for **Datadog Logs**. (**Note: `DD_API_KEY` IS REQUIRED for Datadog LLM Observability**)

View file

@ -11,7 +11,7 @@ import TabItem from '@theme/TabItem';
| Provider Route on LiteLLM | `vercel_ai_gateway/` |
| Link to Provider Doc | [Vercel AI Gateway Documentation ↗](https://vercel.com/docs/ai-gateway) |
| Base URL | `https://ai-gateway.vercel.sh/v1` |
| Supported Operations | `/chat/completions`, `/models` |
| Supported Operations | `/chat/completions`, `/embeddings`, `/models` |
<br />
<br />
@ -73,7 +73,7 @@ messages = [{"content": "Hello, how are you?", "role": "user"}]
# Vercel AI Gateway call with streaming
response = completion(
model="vercel_ai_gateway/openai/gpt-4o",
model="vercel_ai_gateway/openai/gpt-4o",
messages=messages,
stream=True
)
@ -82,6 +82,33 @@ for chunk in response:
print(chunk)
```
### Embeddings
```python showLineNumbers title="Vercel AI Gateway Embeddings"
import os
from litellm import embedding
os.environ["VERCEL_AI_GATEWAY_API_KEY"] = "your-api-key"
# Vercel AI Gateway embedding call
response = embedding(
model="vercel_ai_gateway/openai/text-embedding-3-small",
input="Hello world"
)
print(response.data[0]["embedding"][:5]) # Print first 5 dimensions
```
You can also specify the `dimensions` parameter:
```python showLineNumbers title="Vercel AI Gateway Embeddings with Dimensions"
response = embedding(
model="vercel_ai_gateway/openai/text-embedding-3-small",
input=["Hello world", "Goodbye world"],
dimensions=768
)
```
## Usage - LiteLLM Proxy
Add the following to your LiteLLM Proxy configuration file:
@ -97,6 +124,11 @@ model_list:
litellm_params:
model: vercel_ai_gateway/anthropic/claude-4-sonnet
api_key: os.environ/VERCEL_AI_GATEWAY_API_KEY
- model_name: text-embedding-3-small-gateway
litellm_params:
model: vercel_ai_gateway/openai/text-embedding-3-small
api_key: os.environ/VERCEL_AI_GATEWAY_API_KEY
```
Start your LiteLLM Proxy server:

View file

@ -28,6 +28,37 @@ EXPERIMENTAL_UI_LOGIN="True" litellm --config config.yaml
:::
### Configuration
#### JWT Token Expiration
By default, CLI authentication tokens expire after **24 hours**. You can customize this expiration time by setting the `LITELLM_CLI_JWT_EXPIRATION_HOURS` environment variable when starting your LiteLLM Proxy:
```bash
# Set CLI JWT tokens to expire after 48 hours
export LITELLM_CLI_JWT_EXPIRATION_HOURS=48
export EXPERIMENTAL_UI_LOGIN="True"
litellm --config config.yaml
```
Or in a single command:
```bash
LITELLM_CLI_JWT_EXPIRATION_HOURS=48 EXPERIMENTAL_UI_LOGIN="True" litellm --config config.yaml
```
**Examples:**
- `LITELLM_CLI_JWT_EXPIRATION_HOURS=12` - Tokens expire after 12 hours
- `LITELLM_CLI_JWT_EXPIRATION_HOURS=168` - Tokens expire after 7 days (168 hours)
- `LITELLM_CLI_JWT_EXPIRATION_HOURS=720` - Tokens expire after 30 days (720 hours)
:::tip
You can check your current token's age and expiration status using:
```bash
litellm-proxy whoami
```
:::
### Steps
1. **Install the CLI**

Binary file not shown.

After

Width:  |  Height:  |  Size: 184 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 388 KiB

View file

@ -0,0 +1,423 @@
---
title: "v1.81.3-stable - Performance - 25% CPU Usage Reduction"
slug: "v1-81-3"
date: 2026-01-26T10:00:00
authors:
- name: Krrish Dholakia
title: CEO, LiteLLM
url: https://www.linkedin.com/in/krish-d/
image_url: https://pbs.twimg.com/profile_images/1298587542745358340/DZv3Oj-h_400x400.jpg
- name: Ishaan Jaff
title: CTO, LiteLLM
url: https://www.linkedin.com/in/reffajnaahsi/
image_url: https://pbs.twimg.com/profile_images/1613813310264340481/lz54oEiB_400x400.jpg
hide_table_of_contents: false
---
import Image from '@theme/IdealImage';
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
## Deploy this version
<Tabs>
<TabItem value="docker" label="Docker">
``` showLineNumbers title="docker run litellm"
docker run \
-e STORE_MODEL_IN_DB=True \
-p 4000:4000 \
docker.litellm.ai/berriai/litellm:v1.81.3.rc.2
```
</TabItem>
<TabItem value="pip" label="Pip">
``` showLineNumbers title="pip install litellm"
pip install litellm==1.81.3.rc.2
```
</TabItem>
</Tabs>
---
## New Models / Updated Models
### New Model Support
| Provider | Model | Context Window | Input ($/1M tokens) | Output ($/1M tokens) | Deprecation Date |
| -------- | ----- | -------------- | ------------------- | -------------------- | ---------------- |
| OpenAI | `gpt-audio`, `gpt-audio-2025-08-28` | 128K | $32/1M audio tokens, $2.5/1M text tokens | $64/1M audio tokens, $10/1M text tokens | - |
| OpenAI | `gpt-audio-mini`, `gpt-audio-mini-2025-08-28` | 128K | $10/1M audio tokens, $0.6/1M text tokens | $20/1M audio tokens, $2.4/1M text tokens | - |
| Deepinfra, Vertex AI, Google AI Studio, OpenRouter, Vercel AI Gateway | `gemini-2.0-flash-001`, `gemini-2.0-flash` | - | - | - | 2026-03-31 |
| Groq | `openai/gpt-oss-120b` | 131K | 0.075/1M cache read | 0.6/1M output tokens | - |
| Groq | `groq/openai/gpt-oss-20b` | 131K | 0.0375/1M cache read, $0.075/1M text tokens | 0.3/1M output tokens | - |
| Vertex AI | `gemini-2.5-computer-use-preview-10-2025` | 128K | $1.25 | $10 | - |
| Azure AI | `claude-haiku-4-5` | $1.25/1M cache read, $2/1M cache read above 1 hr, $0.1/1M text tokens | $5/1M output tokens | - |
| Azure AI | `claude-sonnet-4-5` | $3.75/1M cache read, $6/1M cache read above 1 hr, $3/1M text tokens | $15/1M output tokens | - |
| Azure AI | `claude-opus-4-5` | $6.25/1M cache read, $10/1M cache read above 1 hr, $0.5/1M text tokens | $25/1M output tokens | - |
| Azure AI | `claude-opus-4-1` | $18.75/1M cache read, $30/1M cache read above 1 hr, $1.5/1M text tokens | $75/1M output tokens | - |
### Features
- **[OpenAI](../../docs/providers/openai)**
- Add gpt-audio and gpt-audio-mini models to pricing - [PR #19509](https://github.com/BerriAI/litellm/pull/19509)
- correct audio token costs for gpt-4o-audio-preview models - [PR #19500](https://github.com/BerriAI/litellm/pull/19500)
- Limit stop sequence as per openai spec (ensures JetBrains IDE compatibility) - [PR #19562](https://github.com/BerriAI/litellm/pull/19562)
- **[VertexAI](../../docs/providers/vertex)**
- Docs - Google Workload Identity Federation (WIF) support - [PR #19320](https://github.com/BerriAI/litellm/pull/19320)
- **[Agentcore](../../docs/providers/bedrock_agentcore)**
- Fixes streaming issues with AWS Bedrock AgentCore where responses would stop after the first chunk, particularly affecting OAuth-enabled agents - [PR #17141](https://github.com/BerriAI/litellm/pull/17141)
- **[Chatgpt](../../docs/providers/chatgpt)**
- Adds support for calling chatgpt subscription via LiteLLM - [PR #19030](https://github.com/BerriAI/litellm/pull/19030)
- Adds responses API bridge support for chatgpt subscription provider - [PR #19030](https://github.com/BerriAI/litellm/pull/19030)
- **[Bedrock](../../docs/providers/bedrock)**
- support for output format for bedrock invoke via v1/messages - [PR #19560](https://github.com/BerriAI/litellm/pull/19560)
- **[Azure](../../docs/providers/azure/azure)**
- Add support for Azure OpenAI v1 API - [PR #19313](https://github.com/BerriAI/litellm/pull/19313)
- preserve content_policy_violation details for images (#19328) - [PR #19372](https://github.com/BerriAI/litellm/pull/19372)
- Support OpenAI-format nested tool definitions for Responses API - [PR #19526](https://github.com/BerriAI/litellm/pull/19526)
- **Gemini([Vertex AI](../../docs/providers/vertex), [Google AI Studio](../../docs/providers/gemini))**
- use responseJsonSchema for Gemini 2.0+ models - [PR #19314](https://github.com/BerriAI/litellm/pull/19314)
- **[Volcengine](../../docs/providers/volcano)**
- Support Volcengine responses api - [PR #18508](https://github.com/BerriAI/litellm/pull/18508)
- **[Anthropic](../../docs/providers/anthropic)**
- Add Support for calling Claude Code Max subscriptions via LiteLLM - [PR #19453](https://github.com/BerriAI/litellm/pull/19453)
- Add Structured output for /v1/messages with Anthropic API, Azure Anthropic API, Bedrock Converse - [PR #19545](https://github.com/BerriAI/litellm/pull/19545)
- **[Brave Search](../../docs/search/brave)**
- New Search provider - [PR #19433](https://github.com/BerriAI/litellm/pull/19433)
- **Sarvam ai**
- Add support for new sarvam models - [PR #19479](https://github.com/BerriAI/litellm/pull/19479)
- **[GMI](../../docs/providers/gmi)**
- add GMI Cloud provider support - [PR #19376](https://github.com/BerriAI/litellm/pull/19376)
### Bug Fixes
- **[Anthropic](../../docs/providers/anthropic)**
- Fix anthropic-beta sent client side being overridden instead of appended to - [PR #19343](https://github.com/BerriAI/litellm/pull/19343)
- Filter out unsupported fields from JSON schema for Anthropic's output_format API - [PR #19482](https://github.com/BerriAI/litellm/pull/19482)
- **[Bedrock](../../docs/providers/bedrock)**
- Expose stability models via /image_edits endpoint and ensure proper request transformation - [PR #19323](https://github.com/BerriAI/litellm/pull/19323)
- Claude Code x Bedrock Invoke fails with advanced-tool-use-2025-11-20 - [PR #19373](https://github.com/BerriAI/litellm/pull/19373)
- deduplicate tool calls in assistant history - [PR #19324](https://github.com/BerriAI/litellm/pull/19324)
- fix: correct us.anthropic.claude-opus-4-5 In-region pricing - [PR #19310](https://github.com/BerriAI/litellm/pull/19310)
- Fix request validation errors when using Claude 4 via bedrock invoke - [PR #19381](https://github.com/BerriAI/litellm/pull/19381)
- Handle thinking with tool calls for Claude 4 models - [PR #19506](https://github.com/BerriAI/litellm/pull/19506)
- correct streaming choice index for tool calls - [PR #19506](https://github.com/BerriAI/litellm/pull/19506)
- **[Ollama](../../docs/providers/ollama)**
- Fix tool call errors due with improved message extraction - [PR #19369](https://github.com/BerriAI/litellm/pull/19369)
- **[VertexAI](../../docs/providers/vertex)**
- Removed optional vertex_count_tokens_location param before request is sent to vertex - [PR #19359](https://github.com/BerriAI/litellm/pull/19359)
- **Gemini([Vertex AI](../../docs/providers/vertex), [Google AI Studio](../../docs/providers/gemini))**
- Supports setting media_resolution and fps parameters on each video file, when using Gemini video understanding - [PR #19273](https://github.com/BerriAI/litellm/pull/19273)
- handle reasoning_effort as dict from OpenAI Agents SDK - [PR #19419](https://github.com/BerriAI/litellm/pull/19419)
- add file content support in tool results - [PR #19416](https://github.com/BerriAI/litellm/pull/19416)
- **[Azure](../../docs/providers/azure_ai)**
- Fix Azure AI costs for Anthropic models - [PR #19530](https://github.com/BerriAI/litellm/pull/19530)
- **[Giga Chat](../../docs/providers/gigachat)**
- Add tool choice mapping - [PR #19645](https://github.com/BerriAI/litellm/pull/19645)
---
## AI API Endpoints (LLMs, MCP, Agents)
### Features
- **[Files API](../../docs/files_endpoints)**
- Add managed files support when load_balancing is True - [PR #19338](https://github.com/BerriAI/litellm/pull/19338)
- **[Claude Plugin Marketplace](../../docs/tutorials/claude_code_plugin_marketplace)**
- Add self hosted Claude Code Plugin Marketplace - [PR #19378](https://github.com/BerriAI/litellm/pull/19378)
- **[MCP](../../docs/mcp)**
- Add MCP Protocol version 2025-11-25 support - [PR #19379](https://github.com/BerriAI/litellm/pull/19379)
- Log MCP tool calls and list tools in the LiteLLM Spend Logs table for easier debugging - [PR #19469](https://github.com/BerriAI/litellm/pull/19469)
- **[Vertex AI](../../docs/providers/vertex)**
- Ensure only anthropic betas are forwarded down to LLM API (by default) - [PR #19542](https://github.com/BerriAI/litellm/pull/19542)
- Allow overriding to support forwarding incoming headers are forwarded down to target - [PR #19524](https://github.com/BerriAI/litellm/pull/19524)
- **[Chat/Completions](../../docs/completion/input)**
- Add MCP tools response to chat completions - [PR #19552](https://github.com/BerriAI/litellm/pull/19552)
- Add custom vertex ai finish reasons to the output - [PR #19558](https://github.com/BerriAI/litellm/pull/19558)
- Return MCP execution in /chat/completions before model output during streaming - [PR #19623](https://github.com/BerriAI/litellm/pull/19623)
### Bugs
- **[Responses API](../../docs/response_api)**
- Fix duplicate messages during MCP streaming tool execution - [PR #19317](https://github.com/BerriAI/litellm/pull/19317)
- Fix pickle error when using OpenAI's Responses API with stream=True and tool_choice of type allowed_tools (an OpenAI-native parameter) - [PR #17205](https://github.com/BerriAI/litellm/pull/17205)
- stream tool call events for non-openai models - [PR #19368](https://github.com/BerriAI/litellm/pull/19368)
- preserve tool output ordering for gemini in responses bridge - [PR #19360](https://github.com/BerriAI/litellm/pull/19360)
- Add ID caching to prevent ID mismatch text-start and text-delta - [PR #19390](https://github.com/BerriAI/litellm/pull/19390)
- Include output_item, reasoning_summary_Text_done and reasoning_summary_part_done events for non-openai models - [PR #19472](https://github.com/BerriAI/litellm/pull/19472)
- **[Chat/Completions](../../docs/completion/input)**
- fix: drop_params not dropping prompt_cache_key for non-OpenAI providers - [PR #19346](https://github.com/BerriAI/litellm/pull/19346)
- **[Realtime API](../../docs/realtime)**
- disable SSL for ws:// WebSocket connections - [PR #19345](https://github.com/BerriAI/litellm/pull/19345)
- **[Generate Content](../../docs/generateContent)**
- Log actual user input when google genai/vertex endpoints are called client-side - [PR #19156](https://github.com/BerriAI/litellm/pull/19156)
- **[/messages/count_tokens Anthropic Token Counting](../../docs/anthropic_count_tokens)**
- ensure it works for Anthropic, Azure AI Anthropic on AI Gateway - [PR #19432](https://github.com/BerriAI/litellm/pull/19432)
- **[MCP](../../docs/mcp)**
- forward static_headers to MCP servers - [PR #19366](https://github.com/BerriAI/litellm/pull/19366)
- **[Batch API](../../docs/batches)**
- Fix: generation config empty for batch - [PR #19556](https://github.com/BerriAI/litellm/pull/19556)
- **[Pass Through Endpoints](../../docs/proxy/pass_through)**
- Always reupdate registry - [PR #19420](https://github.com/BerriAI/litellm/pull/19420)
---
## Management Endpoints / UI
### Features
- **Cost Estimator**
- Fix model dropdown - [PR #19529](https://github.com/BerriAI/litellm/pull/19529)
- **Claude Code Plugins**
- Allow Adding Claude Code Plugins via UI - [PR #19387](https://github.com/BerriAI/litellm/pull/19387)
- **Guardrails**
- New Policy management UI - [PR #19668](https://github.com/BerriAI/litellm/pull/19668)
- Allow adding policies on Keys/Teams + Viewing on Info panels - [PR #19688](https://github.com/BerriAI/litellm/pull/19688)
- **General**
- respects custom authentication header override - [PR #19276](https://github.com/BerriAI/litellm/pull/19276)
- **Playground**
- Button to Fill Custom API Base - [PR #19440](https://github.com/BerriAI/litellm/pull/19440)
- display mcp output on the play ground - [PR #19553](https://github.com/BerriAI/litellm/pull/19553)
- **Models**
- Paginate /v2/models/info - [PR #19521](https://github.com/BerriAI/litellm/pull/19521)
- All Model Tab Pagination - [PR #19525](https://github.com/BerriAI/litellm/pull/19525)
- Adding Optional scope Param to /models - [PR #19539](https://github.com/BerriAI/litellm/pull/19539)
- Model Search - [PR #19622](https://github.com/BerriAI/litellm/pull/19622)
- Filter by Model ID and Team ID - [PR #19713](https://github.com/BerriAI/litellm/pull/19713)
- **MCP Servers**
- MCP Tools Tab Resetting to Overview - [PR #19468](https://github.com/BerriAI/litellm/pull/19468)
- **Organizations**
- Prevent org admin from creating a new user with proxy_admin permissions - [PR #19296](https://github.com/BerriAI/litellm/pull/19296)
- Edit Page: Reusable Model Select - [PR #19601](https://github.com/BerriAI/litellm/pull/19601)
- **Teams**
- Reusable Model Select - [PR #19543](https://github.com/BerriAI/litellm/pull/19543)
- [Fix] Team Update with Organization having All Proxy Models - [PR #19604](https://github.com/BerriAI/litellm/pull/19604)
- **Logs**
- Include tool arguments in spend logs table - [PR #19640](https://github.com/BerriAI/litellm/pull/19640)
- **Fallbacks / Loadbalancing**
- New fallbacks modal - [PR #19673](https://github.com/BerriAI/litellm/pull/19673)
- Set fallbacks/loadbalancing by team/key - [PR #19686](https://github.com/BerriAI/litellm/pull/19686)
### Bugs
- **Playground**
- increase model selector width in playground Compare view - [PR #19423](https://github.com/BerriAI/litellm/pull/19423)
- **Virtual Keys**
- Sorting Shows Incorrect Entries - [PR #19534](https://github.com/BerriAI/litellm/pull/19534)
- **General**
- UI 404 error when SERVER_ROOT_PATH is set - [PR #19467](https://github.com/BerriAI/litellm/pull/19467)
- Redirect to ui/login on expired JWT - [PR #19687](https://github.com/BerriAI/litellm/pull/19687)
- **SSO**
- Fix SSO user roles not updating for existing users - [PR #19621](https://github.com/BerriAI/litellm/pull/19621)
- **Guardrails**
- ensure guardrail patterns persist on edit and mode toggle - [PR #19265](https://github.com/BerriAI/litellm/pull/19265)
---
## AI Integrations
### Logging
- **General Logging**
- prevent printing duplicate StandardLoggingPayload logs - [PR #19325](https://github.com/BerriAI/litellm/pull/19325)
- Fix: log duplication when json_logs is enabled - [PR #19705](https://github.com/BerriAI/litellm/pull/19705)
- **Langfuse OTEL**
- ignore service logs and fix callback shadowing - [PR #19298](https://github.com/BerriAI/litellm/pull/19298)
- **Langfuse**
- Send litellm_trace_id - [PR #19528](https://github.com/BerriAI/litellm/pull/19528)
- Add Langfuse mock mode for testing without API calls - [PR #19676](https://github.com/BerriAI/litellm/pull/19676)
- **GCS Bucket**
- prevent unbounded queue growth due to slow API calls - [PR #19297](https://github.com/BerriAI/litellm/pull/19297)
- Add GCS mock mode for testing without API calls - [PR #19683](https://github.com/BerriAI/litellm/pull/19683)
- **Responses API Logging**
- Fix pydantic serialization error - [PR #19486](https://github.com/BerriAI/litellm/pull/19486)
- **Arize Phoenix**
- add openinference span kinds to arize phoenix - [PR #19267](https://github.com/BerriAI/litellm/pull/19267)
- **Prometheus**
- Added new prometheus metrics for user count and team count - [PR #19520](https://github.com/BerriAI/litellm/pull/19520)
### Guardrails
- **Bedrock Guardrails**
- Ensure post_call guardrail checks input+output - [PR #19151](https://github.com/BerriAI/litellm/pull/19151)
- **Prompt Security**
- fixing prompt-security's guardrail implementation - [PR #19374](https://github.com/BerriAI/litellm/pull/19374)
- **Presidio**
- Fixes crash in Presidio Guardrail when running in background threads (logging_hook) - [PR #19714](https://github.com/BerriAI/litellm/pull/19714)
- **Pillar Security**
- Migrate Pillar Security to Generic Guardrail API - [PR #19364](https://github.com/BerriAI/litellm/pull/19364)
- **Policy Engine**
- New LiteLLM Policy engine - create policies to manage guardrails, conditions - permissions per Key, Team - [PR #19612](https://github.com/BerriAI/litellm/pull/19612)
- **General**
- add case-insensitive support for guardrail mode and actions - [PR #19480](https://github.com/BerriAI/litellm/pull/19480)
### Prompt Management
- **General**
- fix prompt info lookup and delete using correct IDs - [PR #19358](https://github.com/BerriAI/litellm/pull/19358)
### Secret Manager
- **AWS Secret Manager**
- ensure auto-rotation updates existing AWS secret instead of creating new one - [PR #19455](https://github.com/BerriAI/litellm/pull/19455)
- **Hashicorp Vault**
- Ensure key rotations work with Vault - [PR #19634](https://github.com/BerriAI/litellm/pull/19634)
---
## Spend Tracking, Budgets and Rate Limiting
- **Pricing Updates**
- Add openai/dall-e base pricing entries - [PR #19133](https://github.com/BerriAI/litellm/pull/19133)
- Add `input_cost_per_video_per_second` in ModelInfoBase - [PR #19398](https://github.com/BerriAI/litellm/pull/19398)
---
## Performance / Loadbalancing / Reliability improvements
- **General**
- Fix date overflow/division by zero in proxy utils - [PR #19527](https://github.com/BerriAI/litellm/pull/19527)
- Fix in-flight request termination on SIGTERM when health-check runs in a separate process - [PR #19427](https://github.com/BerriAI/litellm/pull/19427)
- Fix Pass through routes to work with server root path - [PR #19383](https://github.com/BerriAI/litellm/pull/19383)
- Fix logging error for stop iteration - [PR #19649](https://github.com/BerriAI/litellm/pull/19649)
- prevent retrying 4xx client errors - [PR #19275](https://github.com/BerriAI/litellm/pull/19275)
- add better error handling for misconfig on health check - [PR #19441](https://github.com/BerriAI/litellm/pull/19441)
- **Router**
- Fix Azure RPM calculation formula - [PR #19513](https://github.com/BerriAI/litellm/pull/19513)
- Persist scheduler request queue to redis - [PR #19304](https://github.com/BerriAI/litellm/pull/19304)
- pass search_tools to Router during DB-triggered initialization - [PR #19388](https://github.com/BerriAI/litellm/pull/19388)
- Fixed PromptCachingCache to correctly handle messages where cache_control is a sibling key of string content - [PR #19266](https://github.com/BerriAI/litellm/pull/19266)
- **Memory Leaks/OOM**
- prevent OOM with nested $defs in tool schemas - [PR #19112](https://github.com/BerriAI/litellm/pull/19112)
- fix: HTTP client memory leaks in Presidio, OpenAI, and Gemini - [PR #19190](https://github.com/BerriAI/litellm/pull/19190)
- **Non root**
- fix logfile and pidfile of supervisor for non root environment - [PR #17267](https://github.com/BerriAI/litellm/pull/17267)
- resolve Read-only file system error in non-root images - [PR #19449](https://github.com/BerriAI/litellm/pull/19449)
- **Dockerfile**
- Redis Semantic Caching - add missing redisvl dependency to requirements.txt - [PR #19417](https://github.com/BerriAI/litellm/pull/19417)
- Bump OTEL versions to support a2a dependency - resolves modulenotfounderror for Microsoft Agents by @Harshit28j in #18991
- **DB**
- Handle PostgreSQL cached plan errors during rolling deployments - [PR #19424](https://github.com/BerriAI/litellm/pull/19424)
- **Timeouts**
- Fix: total timeout is not respected - [PR #19389](https://github.com/BerriAI/litellm/pull/19389)
- **SDK**
- Field-Existence Checks to Type Classes to Prevent Attribute Errors - [PR #18321](https://github.com/BerriAI/litellm/pull/18321)
- add google-cloud-aiplatform as optional dependency with clear error message - [PR #19437](https://github.com/BerriAI/litellm/pull/19437)
- Make grpc dependency optional - [PR #19447](https://github.com/BerriAI/litellm/pull/19447)
- Add support for retry policies - [PR #19645](https://github.com/BerriAI/litellm/pull/19645)
- **Performance**
- Cut chat_completion latency by ~21% by reducing pre-call processing time - [PR #19535](https://github.com/BerriAI/litellm/pull/19535)
- Optimize strip_trailing_slash with O(1) index check - [PR #19679](https://github.com/BerriAI/litellm/pull/19679)
- Optimize use_custom_pricing_for_model with set intersection - [PR #19677](https://github.com/BerriAI/litellm/pull/19677)
- perf: skip pattern_router.route() for non-wildcard models - [PR #19664](https://github.com/BerriAI/litellm/pull/19664)
- perf: Add LRU caching to get_model_info for faster cost lookups - [PR #19606](https://github.com/BerriAI/litellm/pull/19606)
---
## General Proxy Improvements
### Doc Improvements
- new tutorial for adding MCPs to Cursor via LiteLLM - [PR #19317](https://github.com/BerriAI/litellm/pull/19317)
- fix vertex_region to vertex_location in Vertex AI pass-through docs - [PR #19380](https://github.com/BerriAI/litellm/pull/19380)
- clarify Gemini and Vertex AI model prefix in json file - [PR #19443](https://github.com/BerriAI/litellm/pull/19443)
- update Claude Code integration guides - [PR #19415](https://github.com/BerriAI/litellm/pull/19415)
- adjust opencode tutorial - [PR #19605](https://github.com/BerriAI/litellm/pull/19605)
- add spend-queue-troubleshooting docs - [PR #19659](https://github.com/BerriAI/litellm/pull/19659)
- docs: add litellm-enterprise requirement for managed files - [PR #19689](https://github.com/BerriAI/litellm/pull/19689)
### Helm
- Add support for keda in helm chart - [PR #19337](https://github.com/BerriAI/litellm/pull/19337)
- sync Helm chart version with LiteLLM release version - [PR #19438](https://github.com/BerriAI/litellm/pull/19438)
- Enable PreStop hook configuration in values.yaml - [PR #19613](https://github.com/BerriAI/litellm/pull/19613)
### General
- Add health check scripts and parallel execution support - [PR #19295](https://github.com/BerriAI/litellm/pull/19295)
---
## New Contributors
* @dushyantzz made their first contribution in [PR #19158](https://github.com/BerriAI/litellm/pull/19158)
* @obod-mpw made their first contribution in [PR #19133](https://github.com/BerriAI/litellm/pull/19133)
* @msexxeta made their first contribution in [PR #19030](https://github.com/BerriAI/litellm/pull/19030)
* @rsicart made their first contribution in [PR #19337](https://github.com/BerriAI/litellm/pull/19337)
* @cluebbehusen made their first contribution in [PR #19311](https://github.com/BerriAI/litellm/pull/19311)
* @Lucky-Lodhi2004 made their first contribution in [PR #19315](https://github.com/BerriAI/litellm/pull/19315)
* @binbandit made their first contribution in [PR #19324](https://github.com/BerriAI/litellm/pull/19324)
* @flex-myeonghyeon made their first contribution in [PR #19381](https://github.com/BerriAI/litellm/pull/19381)
* @Lrakotoson made their first contribution in [PR #18321](https://github.com/BerriAI/litellm/pull/18321)
* @bensi94 made their first contribution in [PR #18787](https://github.com/BerriAI/litellm/pull/18787)
* @victorigualada made their first contribution in [PR #19368](https://github.com/BerriAI/litellm/pull/19368)
* @VedantMadane made their first contribution in #19266
* @stiyyagura0901 made their first contribution in #19276
* @kamilio made their first contribution in [PR #19447](https://github.com/BerriAI/litellm/pull/19447)
* @jonathansampson made their first contribution in [PR #19433](https://github.com/BerriAI/litellm/pull/19433)
* @rynecarbone made their first contribution in [PR #19416](https://github.com/BerriAI/litellm/pull/19416)
* @jayy-77 made their first contribution in #19366
* @davida-ps made their first contribution in [PR #19374](https://github.com/BerriAI/litellm/pull/19374)
* @joaodinissf made their first contribution in [PR #19506](https://github.com/BerriAI/litellm/pull/19506)
* @ecao310 made their first contribution in [PR #19520](https://github.com/BerriAI/litellm/pull/19520)
* @mpcusack-altos made their first contribution in [PR #19577](https://github.com/BerriAI/litellm/pull/19577)
* @milan-berri made their first contribution in [PR #19602](https://github.com/BerriAI/litellm/pull/19602)
* @xqe2011 made their first contribution in #19621
---
## Full Changelog
**[View complete changelog on GitHub](https://github.com/BerriAI/litellm/releases/tag/v1.81.3.rc)**

View file

@ -9,7 +9,7 @@ import datetime
from typing import TYPE_CHECKING, Any, AsyncIterator, Coroutine, Dict, Optional, Union
import litellm
from litellm._logging import verbose_logger
from litellm._logging import verbose_logger, verbose_proxy_logger
from litellm.a2a_protocol.streaming_iterator import A2AStreamingIterator
from litellm.a2a_protocol.utils import A2ARequestUtils
from litellm.constants import DEFAULT_A2A_AGENT_TIMEOUT
@ -20,6 +20,7 @@ from litellm.llms.custom_httpx.http_handler import (
)
from litellm.types.agents import LiteLLMSendMessageResponse
from litellm.utils import client
import uuid
if TYPE_CHECKING:
from a2a.client import A2AClient as A2AClientType
@ -225,7 +226,11 @@ async def asend_message(
raise ValueError(
"Either a2a_client or api_base is required for standard A2A flow"
)
a2a_client = await create_a2a_client(base_url=api_base)
trace_id = str(uuid.uuid4())
extra_headers = {"X-LiteLLM-Trace-Id": trace_id}
if agent_id:
extra_headers["X-LiteLLM-Agent-Id"] = agent_id
a2a_client = await create_a2a_client(base_url=api_base, extra_headers=extra_headers)
# Type assertion: a2a_client is guaranteed to be non-None here
assert a2a_client is not None
@ -490,6 +495,10 @@ async def create_a2a_client(
)
httpx_client = http_handler.client
if extra_headers:
httpx_client.headers.update(extra_headers)
verbose_proxy_logger.debug(f"A2A client created with extra_headers={extra_headers}")
# Resolve agent card
resolver = A2ACardResolver(
httpx_client=httpx_client,

View file

@ -1165,6 +1165,7 @@ LITELLM_CLI_SOURCE_IDENTIFIER = "litellm-cli"
LITELLM_CLI_SESSION_TOKEN_PREFIX = "litellm-session-token"
CLI_SSO_SESSION_CACHE_KEY_PREFIX = "cli_sso_session"
CLI_JWT_TOKEN_NAME = "cli-jwt-token"
CLI_JWT_EXPIRATION_HOURS = int(os.getenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", 24))
########################### DB CRON JOB NAMES ###########################
DB_SPEND_UPDATE_JOB_NAME = "db_spend_update_job"

View file

@ -286,7 +286,9 @@ class MCPClient:
return []
async def call_tool(
self, call_tool_request_params: MCPCallToolRequestParams
self,
call_tool_request_params: MCPCallToolRequestParams,
host_progress_callback: Optional[Callable] = None
) -> MCPCallToolResult:
"""
Call an MCP Tool.
@ -295,13 +297,28 @@ class MCPClient:
f"MCP client calling tool '{call_tool_request_params.name}' with arguments: {call_tool_request_params.arguments}"
)
async def on_progress(progress: float, total: float | None, message: str | None):
percentage = (progress / total * 100) if total else 0
verbose_logger.info(
f"MCP Tool '{call_tool_request_params.name}' progress: "
f"{progress}/{total} ({percentage:.0f}%) - {message or ''}"
)
# Forward to Host if callback provided
if host_progress_callback:
try:
await host_progress_callback(progress, total)
except Exception as e:
verbose_logger.warning(f"Failed to forward to Host: {e}")
async def _call_tool_operation(session: ClientSession):
verbose_logger.debug("MCP client sending tool call to session")
return await session.call_tool(
name=call_tool_request_params.name,
arguments=call_tool_request_params.arguments,
)
progress_callback=on_progress,
)
try:
tool_result = await self.run_with_session(_call_tool_operation)
verbose_logger.info(

View file

@ -83,6 +83,33 @@
},
"description": "Datadog Logging Integration"
},
{
"id": "datadog_cost_management",
"displayName": "Datadog Cost Management",
"logo": "datadog.png",
"supports_key_team_logging": false,
"dynamic_params": {
"dd_api_key": {
"type": "password",
"ui_name": "API Key",
"description": "Datadog API key for authentication",
"required": true
},
"dd_app_key": {
"type": "password",
"ui_name": "App Key",
"description": "Datadog Application Key for Cloud Cost Management",
"required": true
},
"dd_site": {
"type": "text",
"ui_name": "Site",
"description": "Datadog site URL (e.g., us5.datadoghq.com)",
"required": true
}
},
"description": "Datadog Cloud Cost Management Integration"
},
{
"id": "lago",
"displayName": "Lago",
@ -407,4 +434,4 @@
},
"description": "SQS Queue (AWS) Logging Integration"
}
]
]

View file

@ -516,7 +516,9 @@ class CustomGuardrail(CustomLogger):
from litellm.types.utils import GuardrailMode
# Use event_type if provided, otherwise fall back to self.event_hook
guardrail_mode: Union[GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks]]
guardrail_mode: Union[
GuardrailEventHooks, GuardrailMode, List[GuardrailEventHooks]
]
if event_type is not None:
guardrail_mode = event_type
elif isinstance(self.event_hook, Mode):
@ -524,11 +526,21 @@ class CustomGuardrail(CustomLogger):
else:
guardrail_mode = self.event_hook # type: ignore[assignment]
from litellm.litellm_core_utils.core_helpers import (
filter_exceptions_from_params,
)
# Sanitize the response to ensure it's JSON serializable and free of circular refs
# This prevents RecursionErrors in downstream loggers (Langfuse, Datadog, etc.)
clean_guardrail_response = filter_exceptions_from_params(
guardrail_json_response
)
slg = StandardLoggingGuardrailInformation(
guardrail_name=self.guardrail_name,
guardrail_provider=guardrail_provider,
guardrail_mode=guardrail_mode,
guardrail_response=guardrail_json_response,
guardrail_response=clean_guardrail_response,
guardrail_status=guardrail_status,
start_time=start_time,
end_time=end_time,

View file

@ -32,6 +32,7 @@ from litellm.integrations.datadog.datadog_handler import (
get_datadog_service,
get_datadog_source,
get_datadog_tags,
get_datadog_base_url_from_env,
)
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.llms.custom_httpx.http_handler import (
@ -100,7 +101,9 @@ class DataDogLogger(
self._configure_dd_direct_api()
# Optional override for testing
self._apply_dd_base_url_override()
dd_base_url = get_datadog_base_url_from_env()
if dd_base_url:
self.intake_url = f"{dd_base_url}/api/v2/logs"
self.sync_client = _get_httpx_client()
asyncio.create_task(self.periodic_flush())
self.flush_lock = asyncio.Lock()
@ -159,18 +162,6 @@ class DataDogLogger(
self.DD_API_KEY = os.getenv("DD_API_KEY")
self.intake_url = f"https://http-intake.logs.{os.getenv('DD_SITE')}/api/v2/logs"
def _apply_dd_base_url_override(self) -> None:
"""
Apply base URL override for testing purposes
"""
dd_base_url: Optional[str] = (
os.getenv("_DATADOG_BASE_URL")
or os.getenv("DATADOG_BASE_URL")
or os.getenv("DD_BASE_URL")
)
if dd_base_url is not None:
self.intake_url = f"{dd_base_url}/api/v2/logs"
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
"""
Async Log success events to Datadog

View file

@ -0,0 +1,204 @@
import asyncio
import os
import time
from datetime import datetime
from typing import Dict, List, Optional, Tuple
from litellm._logging import verbose_logger
from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client,
httpxSpecialProvider,
)
from litellm.types.integrations.datadog_cost_management import (
DatadogFOCUSCostEntry,
)
from litellm.types.utils import StandardLoggingPayload
class DatadogCostManagementLogger(CustomBatchLogger):
def __init__(self, **kwargs):
self.dd_api_key = os.getenv("DD_API_KEY")
self.dd_app_key = os.getenv("DD_APP_KEY")
self.dd_site = os.getenv("DD_SITE", "datadoghq.com")
if not self.dd_api_key or not self.dd_app_key:
verbose_logger.warning(
"Datadog Cost Management: DD_API_KEY and DD_APP_KEY are required. Integration will not work."
)
self.upload_url = f"https://api.{self.dd_site}/api/v2/cost/custom_costs"
self.async_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback
)
# Initialize lock and start periodic flush task
self.flush_lock = asyncio.Lock()
asyncio.create_task(self.periodic_flush())
# Check if flush_lock is already in kwargs to avoid double passing (unlikely but safe)
if "flush_lock" not in kwargs:
kwargs["flush_lock"] = self.flush_lock
super().__init__(**kwargs)
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
try:
standard_logging_object: Optional[StandardLoggingPayload] = kwargs.get(
"standard_logging_object", None
)
if standard_logging_object is None:
return
# Only log if there is a cost associated
if standard_logging_object.get("response_cost", 0) > 0:
self.log_queue.append(standard_logging_object)
if len(self.log_queue) >= self.batch_size:
await self.async_send_batch()
except Exception as e:
verbose_logger.exception(
f"Datadog Cost Management: Error in async_log_success_event: {str(e)}"
)
async def async_send_batch(self):
if not self.log_queue:
return
try:
# Aggregate costs from the batch
aggregated_entries = self._aggregate_costs(self.log_queue)
if not aggregated_entries:
return
# Send to Datadog
await self._upload_to_datadog(aggregated_entries)
# Clear queue only on success (or if we decide to drop on failure)
# CustomBatchLogger clears queue in flush_queue, so we just process here
except Exception as e:
verbose_logger.exception(
f"Datadog Cost Management: Error in async_send_batch: {str(e)}"
)
def _aggregate_costs(
self, logs: List[StandardLoggingPayload]
) -> List[DatadogFOCUSCostEntry]:
"""
Aggregates costs by Provider, Model, and Date.
Returns a list of DatadogFOCUSCostEntry.
"""
aggregator: Dict[Tuple[str, str, str, Tuple[Tuple[str, str], ...]], DatadogFOCUSCostEntry] = {}
for log in logs:
try:
# Extract keys for aggregation
provider = log.get("custom_llm_provider") or "unknown"
model = log.get("model") or "unknown"
cost = log.get("response_cost", 0)
if cost == 0:
continue
# Get date strings (FOCUS format requires specific keys, but for aggregation we group by Day)
# UTC date
# We interpret "ChargePeriod" as the day of the request.
ts = log.get("startTime") or time.time()
dt = datetime.fromtimestamp(ts)
date_str = dt.strftime("%Y-%m-%d")
# ChargePeriodStart and End
# If we want daily granularity, end date is usually same day or next day?
# Datadog Custom Costs usually expects periods.
# "ChargePeriodStart": "2023-01-01", "ChargePeriodEnd": "2023-12-31" in example.
# If we send daily, we can say Start=Date, End=Date.
# Grouping Key: Provider + Model + Date + Tags?
# For simplicity, let's aggregate by Provider + Model + Date first.
# If we handle tags, we need to include them in the key.
tags = self._extract_tags(log)
tags_key = tuple(sorted(tags.items())) if tags else ()
key = (provider, model, date_str, tags_key)
if key not in aggregator:
aggregator[key] = {
"ProviderName": provider,
"ChargeDescription": f"LLM Usage for {model}",
"ChargePeriodStart": date_str,
"ChargePeriodEnd": date_str,
"BilledCost": 0.0,
"BillingCurrency": "USD",
"Tags": tags if tags else None,
}
aggregator[key]["BilledCost"] += cost
except Exception as e:
verbose_logger.warning(
f"Error processing log for cost aggregation: {e}"
)
continue
return list(aggregator.values())
def _extract_tags(self, log: StandardLoggingPayload) -> Dict[str, str]:
from litellm.integrations.datadog.datadog_handler import (
get_datadog_env,
get_datadog_hostname,
get_datadog_pod_name,
get_datadog_service,
)
tags = {
"env": get_datadog_env(),
"service": get_datadog_service(),
"host": get_datadog_hostname(),
"pod_name": get_datadog_pod_name(),
}
# Add metadata as tags
metadata = log.get("metadata", {})
if metadata:
# Add user info
if "user_api_key_alias" in metadata:
tags["user"] = str(metadata["user_api_key_alias"])
if "user_api_key_team_alias" in metadata:
tags["team"] = str(metadata["user_api_key_team_alias"])
# model_group is not in StandardLoggingMetadata TypedDict, so we need to access it via dict.get()
model_group = metadata.get("model_group") # type: ignore[misc]
if model_group:
tags["model_group"] = str(model_group)
return tags
async def _upload_to_datadog(self, payload: List[Dict]):
if not self.dd_api_key or not self.dd_app_key:
return
headers = {
"Content-Type": "application/json",
"DD-API-KEY": self.dd_api_key,
"DD-APPLICATION-KEY": self.dd_app_key,
}
# The API endpoint expects a list of objects directly in the body (file content behavior)
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
data_json = safe_dumps(payload)
response = await self.async_client.put(
self.upload_url, content=data_json, headers=headers
)
response.raise_for_status()
verbose_logger.debug(
f"Datadog Cost Management: Uploaded {len(payload)} cost entries. Status: {response.status_code}"
)

View file

@ -20,6 +20,14 @@ def get_datadog_hostname() -> str:
return os.getenv("HOSTNAME", "")
def get_datadog_base_url_from_env() -> Optional[str]:
"""
Get base URL override from common DD_BASE_URL env var.
This is useful for testing or custom endpoints.
"""
return os.getenv("DD_BASE_URL")
def get_datadog_env() -> str:
return os.getenv("DD_ENV", "unknown")

View file

@ -21,6 +21,7 @@ from litellm.integrations.custom_batch_logger import CustomBatchLogger
from litellm.integrations.datadog.datadog_handler import (
get_datadog_service,
get_datadog_tags,
get_datadog_base_url_from_env,
)
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.prompt_templates.common_utils import (
@ -43,24 +44,22 @@ class DataDogLLMObsLogger(CustomBatchLogger):
def __init__(self, **kwargs):
try:
verbose_logger.debug("DataDogLLMObs: Initializing logger")
if os.getenv("DD_API_KEY", None) is None:
raise Exception("DD_API_KEY is not set, set 'DD_API_KEY=<>'")
if os.getenv("DD_SITE", None) is None:
raise Exception(
"DD_SITE is not set, set 'DD_SITE=<>', example sit = `us5.datadoghq.com`"
)
# Configure DataDog endpoint (Agent or Direct API)
# Use LITELLM_DD_AGENT_HOST to avoid conflicts with ddtrace's DD_AGENT_HOST
dd_agent_host = os.getenv("LITELLM_DD_AGENT_HOST")
self.async_client = get_async_httpx_client(
llm_provider=httpxSpecialProvider.LoggingCallback
)
self.DD_API_KEY = os.getenv("DD_API_KEY")
self.DD_SITE = os.getenv("DD_SITE")
self.intake_url = (
f"https://api.{self.DD_SITE}/api/intake/llm-obs/v1/trace/spans"
)
# testing base url
dd_base_url = os.getenv("DD_BASE_URL")
if dd_agent_host:
self._configure_dd_agent(dd_agent_host=dd_agent_host)
else:
self._configure_dd_direct_api()
# Optional override for testing
dd_base_url = get_datadog_base_url_from_env()
if dd_base_url:
self.intake_url = f"{dd_base_url}/api/intake/llm-obs/v1/trace/spans"
@ -78,6 +77,38 @@ class DataDogLLMObsLogger(CustomBatchLogger):
verbose_logger.exception(f"DataDogLLMObs: Error initializing - {str(e)}")
raise e
def _configure_dd_agent(self, dd_agent_host: str):
"""
Configure the Datadog logger to send traces to the Agent.
"""
# When using the Agent, LLM Observability Intake does NOT require the API Key
# Reference: https://docs.datadoghq.com/llm_observability/setup/sdk/#agent-setup
# Use specific port for LLM Obs (Trace Agent) to avoid conflict with Logs Agent (10518)
agent_port = os.getenv("LITELLM_DD_LLM_OBS_PORT", "8126")
self.DD_SITE = "localhost" # Not used for URL construction in agent mode
self.intake_url = (
f"http://{dd_agent_host}:{agent_port}/api/intake/llm-obs/v1/trace/spans"
)
verbose_logger.debug(f"DataDogLLMObs: Using DD Agent at {self.intake_url}")
def _configure_dd_direct_api(self):
"""
Configure the Datadog logger to send traces directly to the Datadog API.
"""
if not self.DD_API_KEY:
raise Exception("DD_API_KEY is not set, set 'DD_API_KEY=<>'")
self.DD_SITE = os.getenv("DD_SITE")
if not self.DD_SITE:
raise Exception(
"DD_SITE is not set, set 'DD_SITE=<>', example site = `us5.datadoghq.com`"
)
self.intake_url = (
f"https://api.{self.DD_SITE}/api/intake/llm-obs/v1/trace/spans"
)
def _get_datadog_llm_obs_params(self) -> Dict:
"""
Get the datadog_llm_observability_params from litellm.datadog_llm_observability_params
@ -164,13 +195,14 @@ class DataDogLLMObsLogger(CustomBatchLogger):
json_payload = safe_dumps(payload)
headers = {"Content-Type": "application/json"}
if self.DD_API_KEY:
headers["DD-API-KEY"] = self.DD_API_KEY
response = await self.async_client.post(
url=self.intake_url,
content=json_payload,
headers={
"DD-API-KEY": self.DD_API_KEY,
"Content-Type": "application/json",
},
headers=headers,
)
if response.status_code != 202:

View file

@ -23,6 +23,7 @@ from litellm.constants import MAX_LANGFUSE_INITIALIZED_CLIENTS
from litellm.litellm_core_utils.core_helpers import (
safe_deep_copy,
reconstruct_model_name,
filter_exceptions_from_params,
)
from litellm.litellm_core_utils.redact_messages import redact_user_api_key_info
from litellm.integrations.langfuse.langfuse_mock_client import (
@ -75,9 +76,8 @@ def _extract_cache_read_input_tokens(usage_obj) -> int:
# Check prompt_tokens_details.cached_tokens (used by Gemini and other providers)
if hasattr(usage_obj, "prompt_tokens_details"):
prompt_tokens_details = getattr(usage_obj, "prompt_tokens_details", None)
if (
prompt_tokens_details is not None
and hasattr(prompt_tokens_details, "cached_tokens")
if prompt_tokens_details is not None and hasattr(
prompt_tokens_details, "cached_tokens"
):
cached_tokens = getattr(prompt_tokens_details, "cached_tokens", None)
if (
@ -540,7 +540,6 @@ class LangFuseLogger:
verbose_logger.debug("Langfuse Layer Logging - logging to langfuse v2")
try:
metadata = metadata or {}
standard_logging_object: Optional[StandardLoggingPayload] = cast(
Optional[StandardLoggingPayload],
kwargs.get("standard_logging_object", None),
@ -706,9 +705,10 @@ class LangFuseLogger:
clean_metadata["litellm_response_cost"] = cost
if standard_logging_object is not None:
clean_metadata["hidden_params"] = standard_logging_object[
"hidden_params"
]
hidden_params = standard_logging_object.get("hidden_params", {})
clean_metadata["hidden_params"] = filter_exceptions_from_params(
hidden_params
)
if (
litellm.langfuse_default_tags is not None

View file

@ -351,9 +351,9 @@ def filter_exceptions_from_params(data: Any, max_depth: int = 20) -> Any:
# Skip callable objects (functions, methods, lambdas) but not classes (type objects)
if callable(data) and not isinstance(data, type):
return None
# Skip known non-serializable object types (Logging, etc.)
# Skip known non-serializable object types (Logging, Router, etc.)
obj_type_name = type(data).__name__
if obj_type_name in ["Logging", "LiteLLMLoggingObj"]:
if obj_type_name in ["Logging", "LiteLLMLoggingObj", "Router"]:
return None
if isinstance(data, dict):

View file

@ -3304,6 +3304,7 @@ def _get_masked_values(
"token",
"key",
"secret",
"vertex_credentials",
]
return {
k: (

View file

@ -21,11 +21,13 @@ from litellm.types.utils import (
ChatCompletionMessageToolCall,
ChatCompletionRedactedThinkingBlock,
Choices,
CompletionTokensDetailsWrapper,
Delta,
EmbeddingResponse,
Function,
HiddenParams,
ImageResponse,
PromptTokensDetailsWrapper,
)
from litellm.types.utils import Logprobs as TextCompletionLogprobs
from litellm.types.utils import (
@ -304,6 +306,22 @@ class LiteLLMResponseObjectHandler:
"text_tokens": 0,
}
# Map Responses API naming to Chat Completions API naming for cost calculator
if usage.get("prompt_tokens") is None:
usage["prompt_tokens"] = usage.get("input_tokens", 0)
if usage.get("completion_tokens") is None:
usage["completion_tokens"] = usage.get("output_tokens", 0)
# Convert dicts to wrapper objects so getattr() works in cost calculation
if isinstance(usage.get("input_tokens_details"), dict):
usage["prompt_tokens_details"] = PromptTokensDetailsWrapper(
**usage["input_tokens_details"]
)
if isinstance(usage.get("output_tokens_details"), dict):
usage["completion_tokens_details"] = CompletionTokensDetailsWrapper(
**usage["output_tokens_details"]
)
if model_response_object is None:
model_response_object = ImageResponse(**response_object)
return model_response_object

View file

@ -4408,7 +4408,7 @@ def _bedrock_tools_pt(tools: List) -> List[BedrockToolBlock]:
]
"""
"""
Bedrock toolConfig looks like:
Bedrock toolConfig looks like:
"tools": [
{
"toolSpec": {
@ -4436,6 +4436,7 @@ def _bedrock_tools_pt(tools: List) -> List[BedrockToolBlock]:
tool_block_list: List[BedrockToolBlock] = []
for tool in tools:
# Handle regular function tools
parameters = tool.get("function", {}).get(
"parameters", {"type": "object", "properties": {}}
)

View file

@ -110,6 +110,10 @@ class AnthropicMessagesHandler(BaseTranslation):
inputs["tools"] = tools_to_check
if structured_messages:
inputs["structured_messages"] = structured_messages
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=data,
@ -309,6 +313,14 @@ class AnthropicMessagesHandler(BaseTranslation):
inputs["images"] = images_to_check
if tool_calls_to_check:
inputs["tool_calls"] = tool_calls_to_check
# Include model information from the response if available
response_model = None
if isinstance(response, dict):
response_model = response.get("model")
elif hasattr(response, "model"):
response_model = getattr(response, "model", None)
if response_model:
inputs["model"] = response_model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
@ -552,7 +564,7 @@ class AnthropicMessagesHandler(BaseTranslation):
response_content = response.get("content", [])
else:
response_content = getattr(response, "content", None) or []
if not response_content:
return False
for content_block in response_content:

View file

@ -168,6 +168,36 @@ class LiteLLMAnthropicMessagesAdapter:
return provider_specific_fields.get("signature")
return None
def _add_cache_control_if_applicable(
self,
source: Any,
target: Any,
model: Optional[str],
) -> None:
"""
Extract cache_control from source and add to target if it should be preserved.
This method accepts Any type to support both regular dicts and TypedDict objects.
TypedDict objects (like ChatCompletionTextObject, ChatCompletionImageObject, etc.)
are dicts at runtime but have specific types at type-check time. Using Any allows
this method to work with both while maintaining runtime correctness.
Args:
source: Dict or TypedDict containing potential cache_control field
target: Dict or TypedDict to add cache_control to
model: Model name to check if cache_control should be preserved
"""
# TypedDict objects are dicts at runtime, so .get() works
cache_control = source.get("cache_control") if isinstance(source, dict) else getattr(source, "cache_control", None)
if cache_control and model and self.is_anthropic_claude_model(model):
# TypedDict objects support dict operations at runtime
# Use type ignore consistent with codebase pattern (see anthropic/chat/transformation.py:432)
if isinstance(target, dict):
target["cache_control"] = cache_control # type: ignore[typeddict-item]
else:
# Fallback for non-dict objects (shouldn't happen in practice)
cast(Dict[str, Any], target)["cache_control"] = cache_control
def translatable_anthropic_params(self) -> List:
"""
Which anthropic params, we need to translate to the openai format.
@ -205,12 +235,8 @@ class LiteLLMAnthropicMessagesAdapter:
text_obj = ChatCompletionTextObject(
type="text", text=content.get("text", "")
)
# Preserve cache_control if present (for prompt caching)
# Only for Anthropic models that support prompt caching
cache_control = content.get("cache_control")
if cache_control and model and self.is_anthropic_claude_model(model):
text_obj["cache_control"] = cache_control # type: ignore
new_user_content_list.append(text_obj)
self._add_cache_control_if_applicable(content, text_obj, model)
new_user_content_list.append(text_obj) # type: ignore
elif content.get("type") == "image":
# Convert Anthropic image format to OpenAI format
source = content.get("source", {})
@ -225,7 +251,24 @@ class LiteLLMAnthropicMessagesAdapter:
image_obj = ChatCompletionImageObject(
type="image_url", image_url=image_url_obj
)
new_user_content_list.append(image_obj)
self._add_cache_control_if_applicable(content, image_obj, model)
new_user_content_list.append(image_obj) # type: ignore
elif content.get("type") == "document":
# Convert Anthropic document format (PDF, etc.) to OpenAI format
source = content.get("source", {})
openai_image_url = (
self._translate_anthropic_image_to_openai(cast(dict, source))
)
if openai_image_url:
image_url_obj = ChatCompletionImageUrlObject(
url=openai_image_url
)
doc_obj = ChatCompletionImageObject(
type="image_url", image_url=image_url_obj
)
self._add_cache_control_if_applicable(content, doc_obj, model)
new_user_content_list.append(doc_obj) # type: ignore
elif content.get("type") == "tool_result":
if "content" not in content:
tool_result = ChatCompletionToolMessage(
@ -233,14 +276,16 @@ class LiteLLMAnthropicMessagesAdapter:
tool_call_id=content.get("tool_use_id", ""),
content="",
)
tool_message_list.append(tool_result)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result) # type: ignore[arg-type]
elif isinstance(content.get("content"), str):
tool_result = ChatCompletionToolMessage(
role="tool",
tool_call_id=content.get("tool_use_id", ""),
content=str(content.get("content", "")),
)
tool_message_list.append(tool_result)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result) # type: ignore[arg-type]
elif isinstance(content.get("content"), list):
# Combine all content items into a single tool message
# to avoid creating multiple tool_result blocks with the same ID
@ -256,7 +301,8 @@ class LiteLLMAnthropicMessagesAdapter:
tool_call_id=content.get("tool_use_id", ""),
content=c,
)
tool_message_list.append(tool_result)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result) # type: ignore[arg-type]
elif isinstance(c, dict):
if c.get("type") == "text":
tool_result = ChatCompletionToolMessage(
@ -266,7 +312,8 @@ class LiteLLMAnthropicMessagesAdapter:
),
content=c.get("text", ""),
)
tool_message_list.append(tool_result)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result) # type: ignore[arg-type]
elif c.get("type") == "image":
source = c.get("source", {})
openai_image_url = (
@ -282,7 +329,8 @@ class LiteLLMAnthropicMessagesAdapter:
),
content=openai_image_url,
)
tool_message_list.append(tool_result)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result) # type: ignore[arg-type]
else:
# For multiple content items, combine into a single tool message
# with list content to preserve all items while having one tool_use_id
@ -331,7 +379,8 @@ class LiteLLMAnthropicMessagesAdapter:
tool_call_id=content.get("tool_use_id", ""),
content=combined_content_parts, # type: ignore
)
tool_message_list.append(tool_result)
self._add_cache_control_if_applicable(content, tool_result, model)
tool_message_list.append(tool_result) # type: ignore[arg-type]
if len(tool_message_list) > 0:
new_messages.extend(tool_message_list)
@ -344,6 +393,8 @@ class LiteLLMAnthropicMessagesAdapter:
## ASSISTANT MESSAGE ##
assistant_message_str: Optional[str] = None
assistant_content_list: List[Dict[str, Any]] = [] # For content blocks with cache_control
has_cache_control_in_text = False
tool_calls: List[ChatCompletionAssistantToolCall] = []
thinking_blocks: List[
Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]
@ -357,10 +408,14 @@ class LiteLLMAnthropicMessagesAdapter:
assistant_message_str = str(content)
elif isinstance(content, dict):
if content.get("type") == "text":
if assistant_message_str is None:
assistant_message_str = content.get("text", "")
else:
assistant_message_str += content.get("text", "")
text_block: Dict[str, Any] = {
"type": "text",
"text": content.get("text", ""),
}
self._add_cache_control_if_applicable(content, text_block, model)
if "cache_control" in text_block:
has_cache_control_in_text = True
assistant_content_list.append(text_block)
elif content.get("type") == "tool_use":
function_chunk: ChatCompletionToolCallFunctionChunk = {
"name": content.get("name", ""),
@ -384,13 +439,13 @@ class LiteLLMAnthropicMessagesAdapter:
provider_specific_fields
)
tool_calls.append(
ChatCompletionAssistantToolCall(
id=content.get("id", ""),
type="function",
function=function_chunk,
)
tool_call = ChatCompletionAssistantToolCall(
id=content.get("id", ""),
type="function",
function=function_chunk,
)
self._add_cache_control_if_applicable(content, tool_call, model)
tool_calls.append(tool_call)
elif content.get("type") == "thinking":
thinking_block = ChatCompletionThinkingBlock(
type="thinking",
@ -411,18 +466,30 @@ class LiteLLMAnthropicMessagesAdapter:
if (
assistant_message_str is not None
or len(assistant_content_list) > 0
or len(tool_calls) > 0
or len(thinking_blocks) > 0
):
# Use list format if any text block has cache_control, otherwise use string
if has_cache_control_in_text and len(assistant_content_list) > 0:
assistant_content: Any = assistant_content_list
elif len(assistant_content_list) > 0 and not has_cache_control_in_text:
# Concatenate text blocks into string when no cache_control
assistant_content = "".join(
block.get("text", "") for block in assistant_content_list
)
else:
assistant_content = assistant_message_str
assistant_message = ChatCompletionAssistantMessage(
role="assistant",
content=assistant_message_str,
content=assistant_content,
thinking_blocks=(
thinking_blocks if len(thinking_blocks) > 0 else None
),
)
if len(tool_calls) > 0:
assistant_message["tool_calls"] = tool_calls
assistant_message["tool_calls"] = tool_calls # type: ignore
if len(thinking_blocks) > 0:
assistant_message["thinking_blocks"] = thinking_blocks # type: ignore
new_messages.append(assistant_message)
@ -532,10 +599,10 @@ class LiteLLMAnthropicMessagesAdapter:
)
def translate_anthropic_tools_to_openai(
self, tools: List[AllAnthropicToolsValues]
self, tools: List[AllAnthropicToolsValues], model: Optional[str] = None
) -> List[ChatCompletionToolParam]:
new_tools: List[ChatCompletionToolParam] = []
mapped_tool_params = ["name", "input_schema", "description"]
mapped_tool_params = ["name", "input_schema", "description", "cache_control"]
for tool in tools:
function_chunk = ChatCompletionToolParamFunctionChunk(
name=tool["name"],
@ -548,11 +615,11 @@ class LiteLLMAnthropicMessagesAdapter:
for k, v in tool.items():
if k not in mapped_tool_params: # pass additional computer kwargs
function_chunk.setdefault("parameters", {}).update({k: v})
new_tools.append(
ChatCompletionToolParam(type="function", function=function_chunk)
)
tool_param = ChatCompletionToolParam(type="function", function=function_chunk)
self._add_cache_control_if_applicable(tool, tool_param, model)
new_tools.append(tool_param) # type: ignore[arg-type]
return new_tools
return new_tools # type: ignore[return-value]
def translate_anthropic_output_format_to_openai(
self, output_format: Any
@ -590,6 +657,41 @@ class LiteLLMAnthropicMessagesAdapter:
},
}
def _add_system_message_to_messages(
self,
new_messages: List[AllMessageValues],
anthropic_message_request: AnthropicMessagesRequest,
) -> None:
"""Add system message to messages list if present in request."""
if "system" not in anthropic_message_request:
return
system_content = anthropic_message_request["system"]
if not system_content:
return
# Handle system as string or array of content blocks
if isinstance(system_content, str):
new_messages.insert(
0,
ChatCompletionSystemMessage(role="system", content=system_content),
)
elif isinstance(system_content, list):
# Convert Anthropic system content blocks to OpenAI format
openai_system_content: List[Dict[str, Any]] = []
model_name = anthropic_message_request.get("model", "")
for block in system_content:
if isinstance(block, dict) and block.get("type") == "text":
text_block: Dict[str, Any] = {
"type": "text",
"text": block.get("text", ""),
}
self._add_cache_control_if_applicable(block, text_block, model_name)
openai_system_content.append(text_block)
if openai_system_content:
new_messages.insert(
0,
ChatCompletionSystemMessage(role="system", content=openai_system_content), # type: ignore
)
def translate_anthropic_to_openai(
self, anthropic_message_request: AnthropicMessagesRequest
) -> ChatCompletionRequest:
@ -618,13 +720,7 @@ class LiteLLMAnthropicMessagesAdapter:
model=anthropic_message_request.get("model"),
)
## ADD SYSTEM MESSAGE TO MESSAGES
if "system" in anthropic_message_request:
system_content = anthropic_message_request["system"]
if system_content:
new_messages.insert(
0,
ChatCompletionSystemMessage(role="system", content=system_content),
)
self._add_system_message_to_messages(new_messages, anthropic_message_request)
new_kwargs: ChatCompletionRequest = {
"model": anthropic_message_request["model"],
@ -655,7 +751,8 @@ class LiteLLMAnthropicMessagesAdapter:
tools = anthropic_message_request["tools"]
if tools:
new_kwargs["tools"] = self.translate_anthropic_tools_to_openai(
tools=cast(List[AllAnthropicToolsValues], tools)
tools=cast(List[AllAnthropicToolsValues], tools),
model=new_kwargs.get("model"),
)
## CONVERT THINKING
@ -843,7 +940,7 @@ class LiteLLMAnthropicMessagesAdapter:
role="assistant",
model=response.model or "unknown-model",
stop_sequence=None,
usage=anthropic_usage,
usage=anthropic_usage, # type: ignore
content=anthropic_content, # type: ignore
stop_reason=anthropic_finish_reason,
)
@ -992,7 +1089,7 @@ class LiteLLMAnthropicMessagesAdapter:
else:
usage_delta = UsageDelta(input_tokens=0, output_tokens=0)
return MessageBlockDelta(
type="message_delta", delta=delta, usage=usage_delta
type="message_delta", delta=delta, usage=usage_delta # type: ignore
)
(
type_of_content,

View file

@ -298,6 +298,39 @@ class AmazonConverseConfig(BaseConfig):
# Check if the model is specifically Nova Lite 2
return "nova-2-lite" in model_without_region
def _map_web_search_options(
self,
web_search_options: dict,
model: str
) -> Optional[BedrockToolBlock]:
"""
Map web_search_options to Nova grounding systemTool.
Nova grounding (web search) is only supported on Amazon Nova models.
Returns None for non-Nova models.
Args:
web_search_options: The web_search_options dict from the request
model: The model identifier string
Returns:
BedrockToolBlock with systemTool for Nova models, None otherwise
Reference: https://docs.aws.amazon.com/nova/latest/userguide/grounding.html
"""
# Only Nova models support nova_grounding
# Model strings can be like: "amazon.nova-pro-v1:0", "us.amazon.nova-pro-v1:0", etc.
if "nova" not in model.lower():
verbose_logger.debug(
f"web_search_options passed but model {model} is not a Nova model. "
"Nova grounding is only supported on Amazon Nova models."
)
return None
# Nova doesn't support search_context_size or user_location params
# (unlike Anthropic), so we just enable grounding with no options
return BedrockToolBlock(systemTool={"name": "nova_grounding"})
def _transform_reasoning_effort_to_reasoning_config(
self, reasoning_effort: str
) -> dict:
@ -438,6 +471,10 @@ class AmazonConverseConfig(BaseConfig):
):
supported_params.append("tools")
# Nova models support web_search_options (mapped to nova_grounding systemTool)
if base_model.startswith("amazon.nova"):
supported_params.append("web_search_options")
if litellm.utils.supports_tool_choice(
model=model, custom_llm_provider=self.custom_llm_provider
) or litellm.utils.supports_tool_choice(
@ -730,6 +767,13 @@ class AmazonConverseConfig(BaseConfig):
if bedrock_tier in ("default", "flex", "priority"):
optional_params["serviceTier"] = {"type": bedrock_tier}
if param == "web_search_options" and value and isinstance(value, dict):
grounding_tool = self._map_web_search_options(value, model)
if grounding_tool is not None:
optional_params = self._add_tools_to_optional_params(
optional_params=optional_params, tools=[grounding_tool]
)
# Only update thinking tokens for non-GPT-OSS models and non-Nova-Lite-2 models
# Nova Lite 2 handles token budgeting differently through reasoningConfig
if "gpt-oss" not in model and not self._is_nova_lite_2_model(model):
@ -1388,20 +1432,23 @@ class AmazonConverseConfig(BaseConfig):
str,
List[ChatCompletionToolCallChunk],
Optional[List[BedrockConverseReasoningContentBlock]],
Optional[List[CitationsContentBlock]],
]:
"""
Translate the message content to a string and a list of tool calls and reasoning content blocks
Translate the message content to a string and a list of tool calls, reasoning content blocks, and citations.
Returns:
content_str: str
tools: List[ChatCompletionToolCallChunk]
reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]]
citationsContentBlocks: Optional[List[CitationsContentBlock]] - Citations from Nova grounding
"""
content_str = ""
tools: List[ChatCompletionToolCallChunk] = []
reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = (
None
)
citationsContentBlocks: Optional[List[CitationsContentBlock]] = None
for idx, content in enumerate(content_blocks):
"""
- Content is either a tool response or text
@ -1446,10 +1493,15 @@ class AmazonConverseConfig(BaseConfig):
if reasoningContentBlocks is None:
reasoningContentBlocks = []
reasoningContentBlocks.append(content["reasoningContent"])
# Handle Nova grounding citations content
if "citationsContent" in content:
if citationsContentBlocks is None:
citationsContentBlocks = []
citationsContentBlocks.append(content["citationsContent"])
return content_str, tools, reasoningContentBlocks
return content_str, tools, reasoningContentBlocks, citationsContentBlocks
def _transform_response(
def _transform_response( # noqa: PLR0915
self,
model: str,
response: httpx.Response,
@ -1525,18 +1577,27 @@ class AmazonConverseConfig(BaseConfig):
reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = (
None
)
citationsContentBlocks: Optional[List[CitationsContentBlock]] = None
if message is not None:
(
content_str,
tools,
reasoningContentBlocks,
citationsContentBlocks,
) = self._translate_message_content(message["content"])
# Initialize provider_specific_fields if we have any special content blocks
provider_specific_fields: dict = {}
if reasoningContentBlocks is not None:
provider_specific_fields["reasoningContentBlocks"] = reasoningContentBlocks
if citationsContentBlocks is not None:
provider_specific_fields["citationsContent"] = citationsContentBlocks
if provider_specific_fields:
chat_completion_message["provider_specific_fields"] = provider_specific_fields
if reasoningContentBlocks is not None:
chat_completion_message["provider_specific_fields"] = {
"reasoningContentBlocks": reasoningContentBlocks,
}
chat_completion_message["reasoning_content"] = (
self._transform_reasoning_content(reasoningContentBlocks)
)

View file

@ -1476,6 +1476,11 @@ class AWSEventStreamDecoder:
reasoning_content = (
"" # set to non-empty string to ensure consistency with Anthropic
)
elif "citationsContent" in delta_obj:
# Handle Nova grounding citations in streaming responses
provider_specific_fields = {
"citationsContent": delta_obj["citationsContent"],
}
return (
text,
tool_use,

View file

@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Optional
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -49,8 +50,13 @@ class CohereRerankHandler(BaseTranslation):
# Process query only
query = data.get("query")
if query is not None and isinstance(query, str):
inputs = GenericGuardrailAPIInputs(texts=[query])
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": [query]},
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,

View file

@ -87,6 +87,10 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
tools = data.get("tools")
if tools:
inputs["tools"] = tools
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
@ -297,6 +301,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
inputs["images"] = images_to_check
if tool_calls_to_check:
inputs["tool_calls"] = tool_calls_to_check # type: ignore
# Include model information from the response if available
if hasattr(response, "model") and response.model:
inputs["model"] = response.model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
@ -417,6 +424,13 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
if images_to_check:
inputs["images"] = images_to_check
# Include model information from the first response if available
if (
responses_so_far
and hasattr(responses_so_far[0], "model")
and responses_so_far[0].model
):
inputs["model"] = responses_so_far[0].model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=request_data,

View file

@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Optional
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -53,8 +54,13 @@ class OpenAITextCompletionHandler(BaseTranslation):
if isinstance(prompt, str):
# Single string prompt
inputs = GenericGuardrailAPIInputs(texts=[prompt])
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": [prompt]},
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
@ -80,8 +86,13 @@ class OpenAITextCompletionHandler(BaseTranslation):
text_indices.append(idx)
if texts_to_check:
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": texts_to_check},
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
@ -154,8 +165,12 @@ class OpenAITextCompletionHandler(BaseTranslation):
if user_metadata:
request_data["litellm_metadata"] = user_metadata
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
# Include model information from the response if available
if hasattr(response, "model") and response.model:
inputs["model"] = response.model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": texts_to_check},
inputs=inputs,
request_data=request_data,
input_type="response",
logging_obj=litellm_logging_obj,

View file

@ -8,8 +8,7 @@ from typing import Optional
from litellm import verbose_logger
from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token
from litellm.responses.utils import ResponseAPILoggingUtils
from litellm.types.utils import ImageResponse
from litellm.types.utils import ImageResponse, Usage
def cost_calculator(
@ -39,11 +38,18 @@ def cost_calculator(
)
return 0.0
# Transform ImageUsage to Usage using the existing helper
# ImageUsage has the same format as ResponseAPIUsage
chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
# If usage is already a Usage object with completion_tokens_details set,
# use it directly (it was already transformed in convert_to_image_response)
if isinstance(usage, Usage) and usage.completion_tokens_details is not None:
chat_usage = usage
else:
# Transform ImageUsage to Usage using the existing helper
# ImageUsage has the same format as ResponseAPIUsage
from litellm.responses.utils import ResponseAPILoggingUtils
chat_usage = ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(
usage
)
# Use generic_cost_per_token for cost calculation
prompt_cost, completion_cost = generic_cost_per_token(

View file

@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Optional
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -52,8 +53,13 @@ class OpenAIImageGenerationHandler(BaseTranslation):
# Apply guardrail to the prompt
if isinstance(prompt, str):
inputs = GenericGuardrailAPIInputs(texts=[prompt])
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": [prompt]},
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,

View file

@ -105,6 +105,10 @@ class OpenAIResponsesHandler(BaseTranslation):
inputs["tools"] = tools_to_check
if structured_messages:
inputs["structured_messages"] = structured_messages # type: ignore
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
@ -150,6 +154,10 @@ class OpenAIResponsesHandler(BaseTranslation):
inputs["tools"] = tools_to_check
if structured_messages:
inputs["structured_messages"] = structured_messages # type: ignore
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
request_data=data,
@ -344,6 +352,14 @@ class OpenAIResponsesHandler(BaseTranslation):
inputs["images"] = images_to_check
if tool_calls_to_check:
inputs["tool_calls"] = tool_calls_to_check
# Include model information from the response if available
response_model = None
if isinstance(response, dict):
response_model = response.get("model")
elif hasattr(response, "model"):
response_model = getattr(response, "model", None)
if response_model:
inputs["model"] = response_model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=inputs,
@ -388,12 +404,15 @@ class OpenAIResponsesHandler(BaseTranslation):
tool_calls = model_response_stream.choices[0].delta.tool_calls
if tool_calls:
inputs = GenericGuardrailAPIInputs()
inputs["tool_calls"] = cast(
List[ChatCompletionToolCallChunk], tool_calls
)
# Include model information if available
if hasattr(model_response_stream, "model") and model_response_stream.model:
inputs["model"] = model_response_stream.model
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={
"tool_calls": cast(
List[ChatCompletionToolCallChunk], tool_calls
)
},
inputs=inputs,
request_data={},
input_type="response",
logging_obj=litellm_logging_obj,
@ -417,7 +436,11 @@ class OpenAIResponsesHandler(BaseTranslation):
guardrail_inputs["tool_calls"] = cast(
List[ChatCompletionToolCallChunk], tool_calls
)
if tool_calls:
# Include model information from the response if available
response_model = final_chunk.get("response", {}).get("model")
if response_model:
guardrail_inputs["model"] = response_model
if tool_calls or text:
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=guardrail_inputs,
request_data={},
@ -429,8 +452,14 @@ class OpenAIResponsesHandler(BaseTranslation):
# tool_calls = model_response_stream.choices[0].tool_calls
# convert openai response to model response
string_so_far = self.get_streaming_string_so_far(responses_so_far)
inputs = GenericGuardrailAPIInputs(texts=[string_so_far])
# Try to get model from the final chunk if available
if isinstance(final_chunk, dict):
response_model = final_chunk.get("response", {}).get("model") if isinstance(final_chunk.get("response"), dict) else None
if response_model:
inputs["model"] = response_model
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": [string_so_far]},
inputs=inputs,
request_data={},
input_type="response",
logging_obj=litellm_logging_obj,

View file

@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Optional
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -50,8 +51,13 @@ class OpenAITextToSpeechHandler(BaseTranslation):
return data
if isinstance(input_text, str):
inputs = GenericGuardrailAPIInputs(texts=[input_text])
# Include model information if available (voice model)
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": [input_text]},
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,

View file

@ -9,6 +9,7 @@ from typing import TYPE_CHECKING, Any, Optional
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -88,8 +89,12 @@ class OpenAIAudioTranscriptionHandler(BaseTranslation):
if user_metadata:
request_data["litellm_metadata"] = user_metadata
inputs = GenericGuardrailAPIInputs(texts=[original_text])
# Include model information from the response if available
if hasattr(response, "model") and response.model:
inputs["model"] = response.model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": [original_text]},
inputs=inputs,
request_data=request_data,
input_type="response",
logging_obj=litellm_logging_obj,

View file

@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, List, Optional
from litellm._logging import verbose_proxy_logger
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
from litellm.proxy._types import PassThroughGuardrailSettings
from litellm.types.utils import GenericGuardrailAPIInputs
if TYPE_CHECKING:
from litellm.integrations.custom_guardrail import CustomGuardrail
@ -118,8 +119,13 @@ class PassThroughEndpointHandler(BaseTranslation):
return data
# Apply guardrail (pass-through doesn't modify the text, just checks it)
inputs = GenericGuardrailAPIInputs(texts=[text_to_check])
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": [text_to_check]},
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,
@ -178,8 +184,13 @@ class PassThroughEndpointHandler(BaseTranslation):
request_data["litellm_metadata"] = user_metadata
# Apply guardrail (pass-through doesn't modify the text, just checks it)
inputs = GenericGuardrailAPIInputs(texts=[text_to_check])
# Include model information from the response if available
response_model = response.get("model") if isinstance(response, dict) else None
if response_model:
inputs["model"] = response_model
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs={"texts": [text_to_check]},
inputs=inputs,
request_data=request_data,
input_type="response",
logging_obj=litellm_logging_obj,

View file

@ -0,0 +1,176 @@
"""
Vercel AI Gateway Embedding API Configuration.
This module provides the configuration for Vercel AI Gateway's Embedding API.
Vercel AI Gateway is OpenAI-compatible and supports embeddings via the /v1/embeddings endpoint.
Docs: https://vercel.com/docs/ai-gateway/openai-compat/embeddings
"""
from typing import TYPE_CHECKING, Any, Optional
import httpx
from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig
from litellm.secret_managers.main import get_secret_str
from litellm.types.llms.openai import AllEmbeddingInputValues
from litellm.types.utils import EmbeddingResponse
from litellm.utils import convert_to_model_response_object
from ..common_utils import VercelAIGatewayException
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
LiteLLMLoggingObj = _LiteLLMLoggingObj
else:
LiteLLMLoggingObj = Any
class VercelAIGatewayEmbeddingConfig(BaseEmbeddingConfig):
"""
Configuration for Vercel AI Gateway's Embedding API.
Reference: https://vercel.com/docs/ai-gateway/openai-compat/embeddings
"""
def validate_environment(
self,
headers: dict,
model: str,
messages: list,
optional_params: dict,
litellm_params: dict,
api_key: Optional[str] = None,
api_base: Optional[str] = None,
) -> dict:
"""
Validate environment and set up headers for Vercel AI Gateway API.
Vercel AI Gateway requires:
- Authorization header with Bearer token (API key or OIDC token)
"""
vercel_headers = {
"Content-Type": "application/json",
}
# Add Authorization header if api_key is provided
if api_key:
vercel_headers["Authorization"] = f"Bearer {api_key}"
# Merge with existing headers (user's extra_headers take priority)
merged_headers = {**vercel_headers, **headers}
return merged_headers
def get_complete_url(
self,
api_base: Optional[str],
api_key: Optional[str],
model: str,
optional_params: dict,
litellm_params: dict,
stream: Optional[bool] = None,
) -> str:
"""
Get the complete URL for Vercel AI Gateway Embedding API endpoint.
"""
if api_base:
api_base = api_base.rstrip("/")
else:
api_base = (
get_secret_str("VERCEL_AI_GATEWAY_API_BASE")
or "https://ai-gateway.vercel.sh/v1"
)
return f"{api_base}/embeddings"
def transform_embedding_request(
self,
model: str,
input: AllEmbeddingInputValues,
optional_params: dict,
headers: dict,
) -> dict:
"""
Transform embedding request to Vercel AI Gateway format (OpenAI-compatible).
"""
# Ensure input is a list
if isinstance(input, str):
input = [input]
# Strip 'vercel_ai_gateway/' prefix if present
if model.startswith("vercel_ai_gateway/"):
model = model.replace("vercel_ai_gateway/", "", 1)
return {
"model": model,
"input": input,
**optional_params,
}
def transform_embedding_response(
self,
model: str,
raw_response: httpx.Response,
model_response: EmbeddingResponse,
logging_obj: LiteLLMLoggingObj,
api_key: Optional[str],
request_data: dict,
optional_params: dict,
litellm_params: dict,
) -> EmbeddingResponse:
"""
Transform embedding response from Vercel AI Gateway format (OpenAI-compatible).
"""
logging_obj.post_call(original_response=raw_response.text)
# Vercel AI Gateway returns standard OpenAI-compatible embedding response
response_json = raw_response.json()
return convert_to_model_response_object(
response_object=response_json,
model_response_object=model_response,
response_type="embedding",
)
def get_supported_openai_params(self, model: str) -> list:
"""
Get list of supported OpenAI parameters for Vercel AI Gateway embeddings.
Vercel AI Gateway supports the standard OpenAI embeddings parameters
and auto-maps 'dimensions' to each provider's expected field.
"""
return [
"timeout",
"dimensions",
"encoding_format",
"user",
]
def map_openai_params(
self,
non_default_params: dict,
optional_params: dict,
model: str,
drop_params: bool,
) -> dict:
"""
Map OpenAI parameters to Vercel AI Gateway format.
"""
for param, value in non_default_params.items():
if param in self.get_supported_openai_params(model):
optional_params[param] = value
return optional_params
def get_error_class(
self, error_message: str, status_code: int, headers: Any
) -> Any:
"""
Get the error class for Vercel AI Gateway errors.
"""
return VercelAIGatewayException(
message=error_message,
status_code=status_code,
headers=headers,
)

View file

@ -4867,6 +4867,36 @@ def embedding( # noqa: PLR0915
headers = openrouter_headers
response = base_llm_http_handler.embedding(
model=model,
input=input,
custom_llm_provider=custom_llm_provider,
api_base=api_base,
api_key=api_key,
logging_obj=logging,
timeout=timeout,
model_response=EmbeddingResponse(),
optional_params=optional_params,
client=client,
aembedding=aembedding,
litellm_params=litellm_params_dict,
headers=headers,
)
elif custom_llm_provider == "vercel_ai_gateway":
api_base = (
api_base
or litellm.api_base
or get_secret_str("VERCEL_AI_GATEWAY_API_BASE")
or "https://ai-gateway.vercel.sh/v1"
)
api_key = (
api_key
or litellm.api_key
or get_secret_str("VERCEL_AI_GATEWAY_API_KEY")
or get_secret_str("VERCEL_OIDC_TOKEN")
)
response = base_llm_http_handler.embedding(
model=model,
input=input,

View file

@ -46,8 +46,13 @@ class MCPGuardrailTranslationHandler(BaseTranslation):
)
return data
inputs = GenericGuardrailAPIInputs(texts=[content])
# Include model information if available
model = data.get("model")
if model:
inputs["model"] = model
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
inputs=GenericGuardrailAPIInputs(texts=[content]),
inputs=inputs,
request_data=data,
input_type="request",
logging_obj=litellm_logging_obj,

View file

@ -11,7 +11,7 @@ import datetime
import hashlib
import json
import re
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast
from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, cast, Callable
from urllib.parse import urlparse
from fastapi import HTTPException
@ -1842,6 +1842,7 @@ class MCPServerManager:
oauth2_headers: Optional[Dict[str, str]],
raw_headers: Optional[Dict[str, str]],
proxy_logging_obj: Optional[ProxyLogging],
host_progress_callback: Optional[Callable] = None,
) -> CallToolResult:
"""
Call a regular MCP tool using the MCP client.
@ -1926,7 +1927,7 @@ class MCPServerManager:
)
async def _call_tool_via_client(client, params):
return await client.call_tool(params)
return await client.call_tool(params, host_progress_callback=host_progress_callback)
tasks.append(
asyncio.create_task(_call_tool_via_client(client, call_tool_params))
@ -1963,6 +1964,8 @@ class MCPServerManager:
proxy_logging_obj: Optional[ProxyLogging] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
host_progress_callback: Optional[Callable] = None,
) -> CallToolResult:
"""
Call a tool with the given name and arguments
@ -2038,6 +2041,7 @@ class MCPServerManager:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
proxy_logging_obj=proxy_logging_obj,
host_progress_callback=host_progress_callback,
)
# For OpenAPI tools, await outside the client context

View file

@ -8,8 +8,7 @@ import contextlib
from datetime import datetime
import traceback
import uuid
from typing import Any, AsyncIterator, Dict, List, Optional, Tuple, Union, cast
from typing import Any, AsyncIterator, Dict, List, Optional, Tuple, Union, cast, Callable
from fastapi import FastAPI, HTTPException
from pydantic import AnyUrl, ConfigDict
from starlette.types import Receive, Scope, Send
@ -128,8 +127,8 @@ if MCP_AVAILABLE:
session_manager = StreamableHTTPSessionManager(
app=server,
event_store=None,
json_response=True, # Use JSON responses instead of SSE by default
stateless=True,
json_response=False, # enables SSE streaming
stateless=False, # enables session state
)
# Create SSE session manager
@ -282,6 +281,30 @@ if MCP_AVAILABLE:
verbose_logger.debug(
f"MCP mcp_server_tool_call - User API Key Auth from context: {user_api_key_auth}"
)
host_progress_callback = None
try:
host_ctx = server.request_context
if host_ctx and hasattr(host_ctx, 'meta') and host_ctx.meta:
host_token = getattr(host_ctx.meta, 'progressToken', None)
if host_token and hasattr(host_ctx, 'session') and host_ctx.session:
host_session = host_ctx.session
async def forward_progress(progress: float, total: float | None):
"""Forward progress notifications from external MCP to Host"""
try:
await host_session.send_progress_notification(
progress_token=host_token,
progress=progress,
total=total
)
verbose_logger.debug(f"Forwarded progress {progress}/{total} to Host")
except Exception as e:
verbose_logger.error(f"Failed to forward progress to Host: {e}")
host_progress_callback = forward_progress
verbose_logger.debug(f"Host progressToken captured: {host_token[:8]}...")
except Exception as e:
verbose_logger.warning(f"Could not capture host progress context: {e}")
try:
# Create a body date for logging
body_data = {"name": name, "arguments": arguments}
@ -311,6 +334,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers=mcp_server_auth_headers,
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
host_progress_callback=host_progress_callback,
**data, # for logging
)
except BlockedPiiEntityError as e:
@ -1345,6 +1369,7 @@ if MCP_AVAILABLE:
mcp_server_auth_headers: Optional[Dict[str, Dict[str, str]]] = None,
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
host_progress_callback: Optional[Callable] = None,
**kwargs: Any,
) -> CallToolResult:
"""
@ -1442,6 +1467,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
litellm_logging_obj=litellm_logging_obj,
host_progress_callback=host_progress_callback,
)
# Fall back to local tool registry with original name (legacy support)
@ -1689,6 +1715,7 @@ if MCP_AVAILABLE:
oauth2_headers: Optional[Dict[str, str]] = None,
raw_headers: Optional[Dict[str, str]] = None,
litellm_logging_obj: Optional[Any] = None,
host_progress_callback: Optional[Callable] = None,
) -> CallToolResult:
"""Handle tool execution for managed server tools"""
# Import here to avoid circular import
@ -1704,6 +1731,7 @@ if MCP_AVAILABLE:
oauth2_headers=oauth2_headers,
raw_headers=raw_headers,
proxy_logging_obj=proxy_logging_obj,
host_progress_callback=host_progress_callback,
)
verbose_logger.debug("CALL TOOL RESULT: %s", call_tool_result)
return call_tool_result

View file

@ -851,6 +851,13 @@ class GenerateRequestBase(LiteLLMPydanticObjectBase):
aliases: Optional[dict] = {}
object_permission: Optional[LiteLLM_ObjectPermissionBase] = None
@field_validator("max_budget", mode="before")
@classmethod
def check_max_budget(cls, v):
if v == "":
return None
return v
class AllowedVectorStoreIndexItem(LiteLLMPydanticObjectBase):
index_name: str
@ -3421,6 +3428,8 @@ class LitellmMetadataFromRequestHeaders(TypedDict, total=False):
"""
spend_logs_metadata: Optional[dict]
agent_id: Optional[str]
trace_id: Optional[str]
class JWTKeyItem(TypedDict, total=False):

View file

@ -21,6 +21,8 @@ from litellm._logging import verbose_proxy_logger
from litellm.caching.caching import DualCache
from litellm.caching.dual_cache import LimitedSizeOrderedDict
from litellm.constants import (
CLI_JWT_EXPIRATION_HOURS,
CLI_JWT_TOKEN_NAME,
DEFAULT_IN_MEMORY_TTL,
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
DEFAULT_MAX_RECURSE_DEPTH,
@ -1602,7 +1604,10 @@ class ExperimentalUIJWTToken:
user_info: LiteLLM_UserTable, team_id: Optional[str] = None
) -> str:
"""
Generate a JWT token for CLI authentication with 24-hour expiration.
Generate a JWT token for CLI authentication with configurable expiration.
The expiration time can be controlled via the LITELLM_CLI_JWT_EXPIRATION_HOURS
environment variable (defaults to 24 hours).
Args:
user_info: User information from the database
@ -1613,7 +1618,6 @@ class ExperimentalUIJWTToken:
"""
from datetime import timedelta
from litellm.constants import CLI_JWT_TOKEN_NAME
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
encrypt_value_helper,
)
@ -1621,8 +1625,8 @@ class ExperimentalUIJWTToken:
if user_info.user_role is None:
raise Exception("User role is required for CLI JWT login")
# Calculate expiration time (24 hours from now - matching old CLI key behavior)
expiration_time = get_utc_datetime() + timedelta(hours=24)
# Calculate expiration time (configurable via LITELLM_CLI_JWT_EXPIRATION_HOURS env var)
expiration_time = get_utc_datetime() + timedelta(hours=CLI_JWT_EXPIRATION_HOURS)
# Format the expiration time as ISO 8601 string
expires = expiration_time.strftime("%Y-%m-%dT%H:%M:%S.%f")[:-3] + "+00:00"

View file

@ -197,9 +197,9 @@ async def authenticate_user( # noqa: PLR0915
- Login with UI_USERNAME and UI_PASSWORD
- Login with Invite Link `user_email` and `password` combination
"""
if secrets.compare_digest(username, ui_username) and secrets.compare_digest(
password, ui_password
):
if secrets.compare_digest(
username.encode("utf-8"), ui_username.encode("utf-8")
) and secrets.compare_digest(password.encode("utf-8"), ui_password.encode("utf-8")):
# Non SSO -> If user is using UI_USERNAME and UI_PASSWORD they are Proxy admin
user_role = LitellmUserRoles.PROXY_ADMIN
user_id = LITELLM_PROXY_ADMIN_NAME
@ -313,9 +313,9 @@ async def authenticate_user( # noqa: PLR0915
# check if password == _user_row.password
hash_password = hash_token(token=password)
if secrets.compare_digest(password, _password) or secrets.compare_digest(
hash_password, _password
):
if secrets.compare_digest(
password.encode("utf-8"), _password.encode("utf-8")
) or secrets.compare_digest(hash_password.encode("utf-8"), _password.encode("utf-8")):
if os.getenv("DATABASE_URL") is not None:
# Expire any previous UI session tokens for this user
await expire_previous_ui_session_tokens(

View file

@ -11,6 +11,8 @@ import requests
from rich.console import Console
from rich.table import Table
from litellm.constants import CLI_JWT_EXPIRATION_HOURS
# Token storage utilities
def get_token_file_path() -> str:
@ -592,8 +594,8 @@ def whoami():
age_hours = (time.time() - timestamp) / 3600
click.echo(f"Token age: {age_hours:.1f} hours")
if age_hours > 24:
click.echo("⚠️ Warning: Token is more than 24 hours old and may have expired.")
if age_hours > CLI_JWT_EXPIRATION_HOURS:
click.echo(f"⚠️ Warning: Token is more than {CLI_JWT_EXPIRATION_HOURS} hours old and may have expired.")
# Export functions for use by other CLI commands

View file

@ -274,11 +274,20 @@ def initialize_callbacks_on_proxy( # noqa: PLR0915
WebSearchInterceptionLogger,
)
websearch_interception_obj = WebSearchInterceptionLogger.initialize_from_proxy_config(
litellm_settings=litellm_settings,
callback_specific_params=callback_specific_params,
websearch_interception_obj = (
WebSearchInterceptionLogger.initialize_from_proxy_config(
litellm_settings=litellm_settings,
callback_specific_params=callback_specific_params,
)
)
imported_list.append(websearch_interception_obj)
elif isinstance(callback, str) and callback == "datadog_cost_management":
from litellm.integrations.datadog.datadog_cost_management import (
DatadogCostManagementLogger,
)
datadog_cost_management_obj = DatadogCostManagementLogger()
imported_list.append(datadog_cost_management_obj)
elif isinstance(callback, CustomLogger):
imported_list.append(callback)
else:
@ -353,17 +362,17 @@ def get_remaining_tokens_and_requests_from_request_data(data: Dict) -> Dict[str,
remaining_requests_variable_name = f"litellm-key-remaining-requests-{model_group}"
remaining_requests = _metadata.get(remaining_requests_variable_name, None)
if remaining_requests:
headers[f"x-litellm-key-remaining-requests-{h11_model_group_name}"] = (
remaining_requests
)
headers[
f"x-litellm-key-remaining-requests-{h11_model_group_name}"
] = remaining_requests
# Remaining Tokens
remaining_tokens_variable_name = f"litellm-key-remaining-tokens-{model_group}"
remaining_tokens = _metadata.get(remaining_tokens_variable_name, None)
if remaining_tokens:
headers[f"x-litellm-key-remaining-tokens-{h11_model_group_name}"] = (
remaining_tokens
)
headers[
f"x-litellm-key-remaining-tokens-{h11_model_group_name}"
] = remaining_tokens
return headers
@ -438,9 +447,9 @@ def add_guardrail_response_to_standard_logging_object(
):
if litellm_logging_obj is None:
return
standard_logging_object: Optional[StandardLoggingPayload] = (
litellm_logging_obj.model_call_details.get("standard_logging_object")
)
standard_logging_object: Optional[
StandardLoggingPayload
] = litellm_logging_obj.model_call_details.get("standard_logging_object")
if standard_logging_object is None:
return
guardrail_information = standard_logging_object.get("guardrail_information", [])
@ -469,7 +478,9 @@ def get_metadata_variable_name_from_kwargs(
return "litellm_metadata" if "litellm_metadata" in kwargs else "metadata"
def process_callback(_callback: str, callback_type: str, environment_variables: dict) -> dict:
def process_callback(
_callback: str, callback_type: str, environment_variables: dict
) -> dict:
"""Process a single callback and return its data with environment variables"""
env_vars = CustomLogger.get_callback_env_vars(_callback)
@ -481,11 +492,9 @@ def process_callback(_callback: str, callback_type: str, environment_variables:
else:
env_vars_dict[_var] = env_variable
return {
"name": _callback,
"variables": env_vars_dict,
"type": callback_type
}
return {"name": _callback, "variables": env_vars_dict, "type": callback_type}
def normalize_callback_names(callbacks: Iterable[Any]) -> List[Any]:
if callbacks is None:
return []

View file

@ -185,6 +185,7 @@ class GenericGuardrailAPI(CustomGuardrail):
tools = inputs.get("tools")
structured_messages = inputs.get("structured_messages")
tool_calls = inputs.get("tool_calls")
model = inputs.get("model")
# Use provided request_data or create an empty dict
if request_data is None:
@ -215,6 +216,7 @@ class GenericGuardrailAPI(CustomGuardrail):
tool_calls=tool_calls,
additional_provider_specific_params=additional_params,
input_type=input_type,
model=model,
)
# Prepare headers

View file

@ -558,6 +558,16 @@ class LiteLLMProxyRequestSetup:
#########################################################################################
# Finally update the requests metadata with the `metadata_from_headers`
#########################################################################################
agent_id_from_header = headers.get("x-litellm-agent-id")
trace_id_from_header = headers.get("x-litellm-trace-id")
if agent_id_from_header:
metadata_from_headers["agent_id"] = agent_id_from_header
verbose_proxy_logger.debug(f"Extracted agent_id from header: {agent_id_from_header}")
if trace_id_from_header:
metadata_from_headers["trace_id"] = trace_id_from_header
verbose_proxy_logger.debug(f"Extracted trace_id from header: {trace_id_from_header}")
if isinstance(data[_metadata_variable_name], dict):
data[_metadata_variable_name].update(metadata_from_headers)
return data

View file

@ -814,7 +814,9 @@ def _update_internal_user_params(
) -> dict:
non_default_values = {}
for k, v in data_json.items():
if (
if k == "max_budget":
non_default_values[k] = v
elif (
v is not None
and v
not in (

View file

@ -99,24 +99,24 @@ def determine_role_from_groups(
) -> Optional[LitellmUserRoles]:
"""
Determine the highest privilege role for a user based on their groups.
Role hierarchy (highest to lowest):
- proxy_admin
- proxy_admin_viewer
- internal_user
- internal_user_viewer
Args:
user_groups: List of group names from the SSO token
role_mappings: RoleMappings configuration object
Returns:
The highest privilege role found, or default_role if no matches, or None
"""
if not role_mappings.roles:
# No role mappings configured, return default_role
return role_mappings.default_role
# Role hierarchy (highest to lowest)
role_hierarchy = [
LitellmUserRoles.PROXY_ADMIN,
@ -124,20 +124,22 @@ def determine_role_from_groups(
LitellmUserRoles.INTERNAL_USER,
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
]
# Convert user_groups to a set for efficient lookup
user_groups_set = set(user_groups) if isinstance(user_groups, list) else set()
# Find the highest privilege role the user belongs to
for role in role_hierarchy:
if role in role_mappings.roles:
role_groups = role_mappings.roles[role]
if isinstance(role_groups, list) and user_groups_set.intersection(set(role_groups)):
if isinstance(role_groups, list) and user_groups_set.intersection(
set(role_groups)
):
verbose_proxy_logger.debug(
f"User groups {user_groups} matched role '{role.value}' via groups: {role_groups}"
)
return role
# No matching groups found, return default_role
verbose_proxy_logger.debug(
f"User groups {user_groups} did not match any role mappings, using default_role: {role_mappings.default_role}"
@ -326,9 +328,7 @@ def generic_response_convertor(
"GENERIC_USER_PROVIDER_ATTRIBUTE", "provider"
)
generic_user_role_attribute_name = os.getenv(
"GENERIC_USER_ROLE_ATTRIBUTE", "role"
)
generic_user_role_attribute_name = os.getenv("GENERIC_USER_ROLE_ATTRIBUTE", "role")
verbose_proxy_logger.debug(
f" generic_user_id_attribute_name: {generic_user_id_attribute_name}\n generic_user_email_attribute_name: {generic_user_email_attribute_name}"
@ -345,12 +345,15 @@ def generic_response_convertor(
# Determine user role based on role_mappings if available
# Only apply role_mappings for GENERIC SSO provider
user_role: Optional[LitellmUserRoles] = None
if role_mappings is not None and role_mappings.provider.lower() in ["generic", "okta"]:
if role_mappings is not None and role_mappings.provider.lower() in [
"generic",
"okta",
]:
# Use role_mappings to determine role from groups
group_claim = role_mappings.group_claim
user_groups_raw: Any = get_nested_value(response, group_claim)
# Handle different formats: could be a list, string (comma-separated), or single value
user_groups: List[str] = []
if isinstance(user_groups_raw, list):
@ -361,7 +364,7 @@ def generic_response_convertor(
elif user_groups_raw is not None:
# Single value
user_groups = [str(user_groups_raw)]
if user_groups:
user_role = determine_role_from_groups(user_groups, role_mappings)
verbose_proxy_logger.debug(
@ -373,10 +376,12 @@ def generic_response_convertor(
verbose_proxy_logger.debug(
f"No groups found in '{group_claim}', using default_role: {role_mappings.default_role}"
)
# Fallback to existing logic if role_mappings not used
if user_role is None:
user_role_from_sso = get_nested_value(response, generic_user_role_attribute_name)
user_role_from_sso = get_nested_value(
response, generic_user_role_attribute_name
)
if user_role_from_sso is not None:
role = get_litellm_user_role(user_role_from_sso)
if role is not None:
@ -399,7 +404,9 @@ def generic_response_convertor(
)
def _setup_generic_sso_env_vars(generic_client_id: str, redirect_url: str) -> Tuple[str, List[str], str, str, str, bool]:
def _setup_generic_sso_env_vars(
generic_client_id: str, redirect_url: str
) -> Tuple[str, List[str], str, str, str, bool]:
"""Setup and validate Generic SSO environment variables."""
generic_client_secret = os.getenv("GENERIC_CLIENT_SECRET", None)
generic_scope = os.getenv("GENERIC_SCOPE", "openid email profile").split(" ")
@ -492,7 +499,43 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]:
verbose_proxy_logger.debug(
f"Could not load role_mappings from database: {e}. Continuing with existing role logic."
)
generic_role_mappings = os.getenv("GENERIC_ROLE_MAPPINGS_ROLES", None)
generic_role_mappings_group_claim = os.getenv(
"GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", None
)
generic_role_mappoings_default_role = os.getenv(
"GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", None
)
if generic_role_mappings is not None:
verbose_proxy_logger.debug(
"Found role_mappings for generic provider in environment variables"
)
import ast
try:
generic_user_role_mappings_data: Dict[
LitellmUserRoles, List[str]
] = ast.literal_eval(generic_role_mappings)
if isinstance(generic_user_role_mappings_data, dict):
from litellm.types.proxy.management_endpoints.ui_sso import (
RoleMappings,
)
role_mappings_data = {
"provider": "generic",
"group_claim": generic_role_mappings_group_claim,
"default_role": generic_role_mappoings_default_role,
"roles": generic_user_role_mappings_data,
}
role_mappings = RoleMappings(**role_mappings_data)
verbose_proxy_logger.debug(
f"Loaded role_mappings from environments for provider '{role_mappings.provider}'."
)
return role_mappings
except TypeError as e:
verbose_proxy_logger.warning(f"Error decoding role mappings from environment variables: {e}. Continuing with existing role logic.")
return role_mappings
@ -529,7 +572,7 @@ async def get_generic_sso_response(
# Get role_mappings from SSO settings if available
role_mappings = await _setup_role_mappings()
def response_convertor(response, client):
nonlocal received_response # return for user debugging
received_response = response
@ -1156,7 +1199,7 @@ async def cli_poll_key(key_id: str, team_id: Optional[str] = None):
max_budget=litellm.max_ui_session_budget,
)
# Generate CLI JWT on-demand (24hr expiration)
# Generate CLI JWT on-demand (expiration configurable via LITELLM_CLI_JWT_EXPIRATION_HOURS)
# Pass selected team_id to ensure JWT has correct team
jwt_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(
user_info=user_info, team_id=team_id
@ -1217,20 +1260,24 @@ async def insert_sso_user(
role_mappings_configured = False
try:
from litellm.proxy.utils import get_prisma_client_or_throw
prisma_client = get_prisma_client_or_throw(
"Prisma client is None, connect a database to your proxy"
)
# Get SSO config from dedicated table
sso_db_record = await prisma_client.db.litellm_ssoconfig.find_unique(
where={"id": "sso_config"}
)
if sso_db_record and sso_db_record.sso_settings:
sso_settings_dict = dict(sso_db_record.sso_settings)
role_mappings_data = sso_settings_dict.get("role_mappings")
role_mappings_configured = role_mappings_data is not None
generic_user_role_mappings = os.getenv("GENERIC_USER_ROLE_MAPPINGS", None)
if generic_user_role_mappings is not None:
role_mappings_configured = True
except Exception as e:
# If we can't check role_mappings, continue with existing logic
verbose_proxy_logger.debug(
@ -1240,7 +1287,10 @@ async def insert_sso_user(
# Apply default_internal_user_params
if litellm.default_internal_user_params:
# If role_mappings is configured and user_role is already set from SSO, preserve it
if role_mappings_configured and user_defined_values.get("user_role") is not None:
if (
role_mappings_configured
and user_defined_values.get("user_role") is not None
):
# Preserve the SSO-extracted role, but apply other defaults
preserved_role = user_defined_values.get("user_role")
user_defined_values.update(litellm.default_internal_user_params) # type: ignore

View file

@ -825,6 +825,7 @@ app = FastAPI(
title=_title,
description=_description,
version=version,
root_path=server_root_path,
lifespan=proxy_startup_event, # type: ignore[reportGeneralTypeIssues]
)
@ -5419,6 +5420,24 @@ async def chat_completion( # noqa: PLR0915
global general_settings, user_debug, proxy_logging_obj, llm_model_list
global user_temperature, user_request_timeout, user_max_tokens, user_api_base
data = await _read_request_body(request=request)
if user_api_key_dict is not None:
if data.get("metadata") is None:
data["metadata"] = {}
if (
hasattr(user_api_key_dict, "user_id")
and user_api_key_dict.user_id is not None
):
data["metadata"]["user_api_key_user_id"] = user_api_key_dict.user_id
if (
hasattr(user_api_key_dict, "team_id")
and user_api_key_dict.team_id is not None
):
data["metadata"]["user_api_key_team_id"] = user_api_key_dict.team_id
if (
hasattr(user_api_key_dict, "org_id")
and user_api_key_dict.org_id is not None
):
data["metadata"]["user_api_key_org_id"] = user_api_key_dict.org_id
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
try:
result = await base_llm_response_processor.base_process_llm_request(
@ -5570,6 +5589,24 @@ async def completion( # noqa: PLR0915
data = {}
try:
data = await _read_request_body(request=request)
if user_api_key_dict is not None:
if data.get("metadata") is None:
data["metadata"] = {}
if (
hasattr(user_api_key_dict, "user_id")
and user_api_key_dict.user_id is not None
):
data["metadata"]["user_api_key_user_id"] = user_api_key_dict.user_id
if (
hasattr(user_api_key_dict, "team_id")
and user_api_key_dict.team_id is not None
):
data["metadata"]["user_api_key_team_id"] = user_api_key_dict.team_id
if (
hasattr(user_api_key_dict, "org_id")
and user_api_key_dict.org_id is not None
):
data["metadata"]["user_api_key_org_id"] = user_api_key_dict.org_id
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)
return await base_llm_response_processor.base_process_llm_request(
request=request,
@ -5789,6 +5826,25 @@ async def embeddings( # noqa: PLR0915
)
data["input"] = input_list
if user_api_key_dict is not None:
if data.get("metadata") is None:
data["metadata"] = {}
if (
hasattr(user_api_key_dict, "user_id")
and user_api_key_dict.user_id is not None
):
data["metadata"]["user_api_key_user_id"] = user_api_key_dict.user_id
if (
hasattr(user_api_key_dict, "team_id")
and user_api_key_dict.team_id is not None
):
data["metadata"]["user_api_key_team_id"] = user_api_key_dict.team_id
if (
hasattr(user_api_key_dict, "org_id")
and user_api_key_dict.org_id is not None
):
data["metadata"]["user_api_key_org_id"] = user_api_key_dict.org_id
# Use unified request processor (same as chat/completions and responses)
base_llm_response_processor = ProxyBaseLLMRequestProcessing(data=data)

View file

@ -26,6 +26,184 @@ from litellm.proxy.common_utils.http_parsing_utils import (
router = APIRouter()
def _build_file_metadata_entry(
response: Any,
file_data: Optional[Tuple[str, bytes, str]] = None,
file_url: Optional[str] = None,
) -> Dict[str, Any]:
"""
Build a file metadata entry for storing in vector_store_metadata.
Args:
response: The response from litellm.aingest containing file_id
file_data: Optional tuple of (filename, content, content_type)
file_url: Optional URL if file was ingested from URL
Returns:
Dictionary with file metadata (file_id, filename, file_url, ingested_at, etc.)
"""
from datetime import datetime, timezone
# Extract file_id from response
file_id = None
if hasattr(response, "get"):
file_id = response.get("file_id")
elif hasattr(response, "file_id"):
file_id = response.file_id
# Extract file information from file_data tuple
filename = None
file_size = None
content_type = None
if file_data:
filename = file_data[0]
file_size = len(file_data[1]) if len(file_data) > 1 else None
content_type = file_data[2] if len(file_data) > 2 else None
# Build file metadata entry
file_entry = {
"file_id": file_id,
"filename": filename,
"file_url": file_url,
"ingested_at": datetime.now(timezone.utc).isoformat(),
}
# Add optional fields if available
if file_size is not None:
file_entry["file_size"] = file_size
if content_type is not None:
file_entry["content_type"] = content_type
return file_entry
async def _save_vector_store_to_db_from_rag_ingest(
response: Any,
ingest_options: Dict[str, Any],
prisma_client,
user_api_key_dict: UserAPIKeyAuth,
file_data: Optional[Tuple[str, bytes, str]] = None,
file_url: Optional[str] = None,
) -> None:
"""
Helper function to save a newly created vector store from RAG ingest to the database.
This function:
- Extracts vector store ID and config from the ingest response
- Checks if the vector store already exists in the database
- Creates a new database entry if it doesn't exist
- Adds the vector store to the registry
Args:
response: The response from litellm.aingest()
ingest_options: The ingest options containing vector store config
prisma_client: The Prisma database client
user_api_key_dict: User API key authentication info
"""
from litellm.proxy.vector_store_endpoints.management_endpoints import (
create_vector_store_in_db,
)
# Handle both dict and object responses
if hasattr(response, "get"):
vector_store_id = response.get("vector_store_id")
elif hasattr(response, "vector_store_id"):
vector_store_id = response.vector_store_id
else:
verbose_proxy_logger.warning(
f"Unable to extract vector_store_id from response type: {type(response)}"
)
return
if vector_store_id is None or not isinstance(vector_store_id, str):
verbose_proxy_logger.warning(
"Vector store ID is None or not a string, skipping database save"
)
return
vector_store_config = ingest_options.get("vector_store", {})
custom_llm_provider = vector_store_config.get("custom_llm_provider")
# Extract litellm_vector_store_params for custom name and description
litellm_vector_store_params = ingest_options.get("litellm_vector_store_params", {})
custom_vector_store_name = litellm_vector_store_params.get("vector_store_name")
custom_vector_store_description = litellm_vector_store_params.get("vector_store_description")
# Build file metadata entry using helper
file_entry = _build_file_metadata_entry(
response=response,
file_data=file_data,
file_url=file_url,
)
try:
# Check if vector store already exists in database
existing_vector_store = (
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
where={"vector_store_id": vector_store_id}
)
)
# Only create if it doesn't exist
if existing_vector_store is None:
verbose_proxy_logger.info(
f"Saving newly created vector store {vector_store_id} to database"
)
# Initialize metadata with first file
initial_metadata = {
"ingested_files": [file_entry]
}
# Use custom name if provided, otherwise default
vector_store_name = custom_vector_store_name or f"RAG Vector Store - {vector_store_id[:8]}"
vector_store_description = custom_vector_store_description or "Created via RAG ingest endpoint"
await create_vector_store_in_db(
vector_store_id=vector_store_id,
custom_llm_provider=custom_llm_provider or "openai",
prisma_client=prisma_client,
vector_store_name=vector_store_name,
vector_store_description=vector_store_description,
vector_store_metadata=initial_metadata,
)
verbose_proxy_logger.info(
f"Vector store {vector_store_id} saved to database successfully"
)
else:
verbose_proxy_logger.info(
f"Vector store {vector_store_id} already exists, appending file to metadata"
)
# Update existing vector store with new file
existing_metadata = existing_vector_store.vector_store_metadata or {}
if isinstance(existing_metadata, str):
import json
existing_metadata = json.loads(existing_metadata)
ingested_files = existing_metadata.get("ingested_files", [])
ingested_files.append(file_entry)
existing_metadata["ingested_files"] = ingested_files
# Update the vector store
from litellm.proxy.utils import safe_dumps
await prisma_client.db.litellm_managedvectorstorestable.update(
where={"vector_store_id": vector_store_id},
data={"vector_store_metadata": safe_dumps(existing_metadata)}
)
verbose_proxy_logger.info(
f"Added file {file_entry.get('filename') or file_entry.get('file_url', 'Unknown')} to vector store {vector_store_id} metadata"
)
except Exception as db_error:
# Log the error but don't fail the request since ingestion succeeded
verbose_proxy_logger.exception(
f"Failed to save vector store {vector_store_id} to database: {db_error}"
)
async def parse_rag_ingest_request(
request: Request,
) -> Tuple[Dict[str, Any], Optional[Tuple[str, bytes, str]], Optional[str], Optional[str]]:
@ -158,6 +336,7 @@ async def rag_ingest(
add_litellm_data_to_request,
general_settings,
llm_router,
prisma_client,
proxy_config,
version,
)
@ -189,6 +368,25 @@ async def rag_ingest(
**request_data,
)
# Save vector store to database if it was newly created and prisma_client is available
verbose_proxy_logger.debug(
f"RAG Ingest - Checking database save conditions: prisma_client={prisma_client is not None}, response={response is not None}, response_type={type(response)}"
)
if prisma_client is not None and response is not None:
await _save_vector_store_to_db_from_rag_ingest(
response=response,
ingest_options=ingest_options,
prisma_client=prisma_client,
user_api_key_dict=user_api_key_dict,
file_data=file_data,
file_url=file_url,
)
else:
verbose_proxy_logger.warning(
f"Skipping database save: prisma_client={prisma_client is not None}, response={response is not None}"
)
return response
except HTTPException:

View file

@ -396,7 +396,7 @@ def get_logging_payload( # noqa: PLR0915
)
# Extract agent_id for A2A requests (set directly on model_call_details)
agent_id: Optional[str] = kwargs.get("agent_id")
agent_id: Optional[str] = kwargs.get("agent_id") or metadata.get("agent_id")
custom_llm_provider = kwargs.get("custom_llm_provider")
raw_model = cast(str, kwargs.get("model") or "")
model_name = reconstruct_model_name(raw_model, custom_llm_provider, metadata or {})

View file

@ -982,7 +982,9 @@ class ProxyLogging:
try:
# Check if load balancing should be used
if guardrail_name and self._should_use_guardrail_load_balancing(guardrail_name):
if guardrail_name and self._should_use_guardrail_load_balancing(
guardrail_name
):
response = await self._execute_guardrail_with_load_balancing(
guardrail_name=guardrail_name,
hook_type="pre_call",
@ -1017,7 +1019,11 @@ class ProxyLogging:
latency_seconds = guardrail_end_time - guardrail_start_time
# Get guardrail name for metrics (fallback if not set)
metrics_guardrail_name = guardrail_name or getattr(callback, "guardrail_name", callback.__class__.__name__) or "unknown"
metrics_guardrail_name = (
guardrail_name
or getattr(callback, "guardrail_name", callback.__class__.__name__)
or "unknown"
)
# Find PrometheusLogger in callbacks and record metrics
for prom_callback in litellm.callbacks:
@ -1793,9 +1799,11 @@ class ProxyLogging:
#################################################################
for callback in other_callbacks:
await callback.async_post_call_success_hook(
callback_response = await callback.async_post_call_success_hook(
user_api_key_dict=user_api_key_dict, data=data, response=response
)
if callback_response is not None:
response = callback_response
except Exception as e:
raise e
return response
@ -1852,16 +1860,14 @@ class ProxyLogging:
complete_response = str_so_far + response_str
else:
complete_response = response_str
potential_error_response = (
callback_response = (
await _callback.async_post_call_streaming_hook(
user_api_key_dict=user_api_key_dict,
response=complete_response,
)
)
if isinstance(
potential_error_response, str
) and potential_error_response.startswith("data: "):
return potential_error_response
if callback_response is not None:
response = callback_response
except Exception as e:
raise e
return response
@ -2218,17 +2224,17 @@ class PrismaClient:
) -> Optional[dict]:
"""
Execute a query with automatic fallback for PostgreSQL cached plan errors.
This handles the "cached plan must not change result type" error that occurs
during rolling deployments when schema changes are applied while old pods
still have cached query plans expecting the old schema.
Args:
sql_query: SQL query string to execute
Returns:
Query result or None
Raises:
Original exception if not a cached plan error
"""
@ -2241,7 +2247,7 @@ class PrismaClient:
# Add a unique comment to make the query different
sql_query_retry = sql_query.replace(
"SELECT",
f"SELECT /* cache_invalidated_{int(time.time() * 1000)} */"
f"SELECT /* cache_invalidated_{int(time.time() * 1000)} */",
)
verbose_proxy_logger.warning(
"PostgreSQL cached plan error detected for token lookup, "
@ -2583,7 +2589,9 @@ class PrismaClient:
WHERE v.token = '{token}'
"""
response = await self._query_first_with_cached_plan_fallback(sql_query)
response = await self._query_first_with_cached_plan_fallback(
sql_query
)
if response is not None:
if response["team_models"] is None:
@ -4227,7 +4235,7 @@ def get_server_root_path() -> str:
- If SERVER_ROOT_PATH is set, return it.
- Otherwise, default to "/".
"""
return os.getenv("SERVER_ROOT_PATH", "/")
return os.getenv("SERVER_ROOT_PATH", "")
def get_prisma_client_or_throw(message: str):

View file

@ -133,6 +133,112 @@ async def _resolve_embedding_config_from_db(
return None
########################################################
# Helper Functions
########################################################
async def create_vector_store_in_db(
vector_store_id: str,
custom_llm_provider: str,
prisma_client,
vector_store_name: Optional[str] = None,
vector_store_description: Optional[str] = None,
vector_store_metadata: Optional[Dict] = None,
litellm_params: Optional[Dict] = None,
litellm_credential_name: Optional[str] = None,
) -> LiteLLM_ManagedVectorStore:
"""
Helper function to create a vector store in the database.
This function handles:
- Checking if vector store already exists
- Creating the vector store in the database
- Adding it to the vector store registry
Returns:
LiteLLM_ManagedVectorStore: The created vector store object
Raises:
HTTPException: If vector store already exists or database error occurs
"""
from litellm.types.router import GenericLiteLLMParams
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
# Check if vector store already exists
existing_vector_store = (
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
where={"vector_store_id": vector_store_id}
)
)
if existing_vector_store is not None:
raise HTTPException(
status_code=400,
detail=f"Vector store with ID {vector_store_id} already exists",
)
# Prepare data for database
data_to_create: Dict[str, Any] = {
"vector_store_id": vector_store_id,
"custom_llm_provider": custom_llm_provider,
}
if vector_store_name is not None:
data_to_create["vector_store_name"] = vector_store_name
if vector_store_description is not None:
data_to_create["vector_store_description"] = vector_store_description
if vector_store_metadata is not None:
data_to_create["vector_store_metadata"] = safe_dumps(vector_store_metadata)
if litellm_credential_name is not None:
data_to_create["litellm_credential_name"] = litellm_credential_name
# Handle litellm_params - always provide at least an empty dict
if litellm_params:
# Auto-resolve embedding config if embedding model is provided but config is not
embedding_model = litellm_params.get("litellm_embedding_model")
if embedding_model and not litellm_params.get("litellm_embedding_config"):
resolved_config = await _resolve_embedding_config_from_db(
embedding_model=embedding_model,
prisma_client=prisma_client
)
if resolved_config:
litellm_params["litellm_embedding_config"] = resolved_config
verbose_proxy_logger.info(
f"Auto-resolved embedding config for model {embedding_model}"
)
litellm_params_dict = GenericLiteLLMParams(
**litellm_params
).model_dump(exclude_none=True)
data_to_create["litellm_params"] = safe_dumps(litellm_params_dict)
else:
# Provide empty dict if no litellm_params provided
data_to_create["litellm_params"] = safe_dumps({})
# Create in database
_new_vector_store = (
await prisma_client.db.litellm_managedvectorstorestable.create(
data=data_to_create
)
)
new_vector_store: LiteLLM_ManagedVectorStore = LiteLLM_ManagedVectorStore(
**_new_vector_store.model_dump()
)
# Add vector store to registry
if litellm.vector_store_registry is not None:
litellm.vector_store_registry.add_vector_store_to_registry(
vector_store=new_vector_store
)
verbose_proxy_logger.info(
f"Vector store {vector_store_id} created in database successfully"
)
return new_vector_store
########################################################
# Management Endpoints
########################################################
@ -156,71 +262,34 @@ async def new_vector_store(
- vector_store_metadata: Optional[Dict] - Additional metadata for the vector store
"""
from litellm.proxy.proxy_server import prisma_client
from litellm.types.router import GenericLiteLLMParams
if prisma_client is None:
raise HTTPException(status_code=500, detail="Database not connected")
try:
# Check if vector store already exists
existing_vector_store = (
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
where={"vector_store_id": vector_store.get("vector_store_id")}
)
)
if existing_vector_store is not None:
vector_store_id = vector_store.get("vector_store_id")
custom_llm_provider = vector_store.get("custom_llm_provider")
if not vector_store_id or not custom_llm_provider:
raise HTTPException(
status_code=400,
detail=f"Vector store with ID {vector_store.get('vector_store_id')} already exists",
)
if vector_store.get("vector_store_metadata") is not None:
vector_store["vector_store_metadata"] = safe_dumps(
vector_store.get("vector_store_metadata")
)
# Safely handle JSON serialization of litellm_params
litellm_params_json: Optional[str] = None
_input_litellm_params: dict = vector_store.get("litellm_params", {}) or {}
if _input_litellm_params is not None:
# Auto-resolve embedding config if embedding model is provided but config is not
embedding_model = _input_litellm_params.get("litellm_embedding_model")
if embedding_model and not _input_litellm_params.get("litellm_embedding_config"):
resolved_config = await _resolve_embedding_config_from_db(
embedding_model=embedding_model,
prisma_client=prisma_client
)
if resolved_config:
_input_litellm_params["litellm_embedding_config"] = resolved_config
verbose_proxy_logger.info(
f"Auto-resolved embedding config for model {embedding_model}"
)
litellm_params_dict = GenericLiteLLMParams(
**_input_litellm_params
).model_dump(exclude_none=True)
litellm_params_json = safe_dumps(litellm_params_dict)
del vector_store["litellm_params"]
_new_vector_store = (
await prisma_client.db.litellm_managedvectorstorestable.create(
data={
**vector_store,
"litellm_params": litellm_params_json,
}
detail="vector_store_id and custom_llm_provider are required"
)
# Extract and validate metadata
metadata = vector_store.get("vector_store_metadata")
validated_metadata: Optional[Dict] = None
if metadata is not None and isinstance(metadata, dict):
validated_metadata = metadata
new_vector_store = await create_vector_store_in_db(
vector_store_id=vector_store_id,
custom_llm_provider=custom_llm_provider,
prisma_client=prisma_client,
vector_store_name=vector_store.get("vector_store_name"),
vector_store_description=vector_store.get("vector_store_description"),
vector_store_metadata=validated_metadata,
litellm_params=vector_store.get("litellm_params"),
litellm_credential_name=vector_store.get("litellm_credential_name"),
)
new_vector_store: LiteLLM_ManagedVectorStore = LiteLLM_ManagedVectorStore(
**_new_vector_store.model_dump()
)
# Add vector store to registry
if litellm.vector_store_registry is not None:
litellm.vector_store_registry.add_vector_store_to_registry(
vector_store=new_vector_store
)
return {
"status": "success",
"message": f"Vector store {vector_store.get('vector_store_id')} created successfully",

View file

@ -198,6 +198,10 @@ async def _execute_query_pipeline(
"""
Execute the RAG query pipeline.
"""
# Extract router from kwargs - use it for completion if available
# to properly resolve virtual model names
router: Optional["Router"] = kwargs.pop("router", None)
# 1. Extract query from last user message
query_text = RAGQuery.extract_query_from_messages(messages)
if not query_text:
@ -233,12 +237,21 @@ async def _execute_query_pipeline(
context_message = RAGQuery.build_context_message(context_chunks)
modified_messages = messages[:-1] + [context_message] + [messages[-1]]
response = await litellm.acompletion(
model=model,
messages=modified_messages,
stream=stream,
**kwargs,
)
# Use router if available to properly resolve virtual model names
if router is not None:
response = await router.acompletion(
model=model,
messages=modified_messages,
stream=stream,
**kwargs,
)
else:
response = await litellm.acompletion(
model=model,
messages=modified_messages,
stream=stream,
**kwargs,
)
# 5. Attach search results to response
if not stream and isinstance(response, ModelResponse):

View file

@ -1707,8 +1707,11 @@ class Router:
litellm_params = deployment.get("litellm_params", {})
dep_num_retries = litellm_params.get("num_retries")
if dep_num_retries is not None and isinstance(dep_num_retries, int):
exception.num_retries = dep_num_retries # type: ignore
if dep_num_retries is not None:
try:
exception.num_retries = int(dep_num_retries) # type: ignore # Handle both int and str
except (ValueError, TypeError):
pass # Skip if value can't be converted to int
def _update_kwargs_with_default_litellm_params(
self, kwargs: dict, metadata_variable_name: Optional[str] = "metadata"

View file

@ -17,6 +17,22 @@ from litellm.secret_managers.secret_manager_handler import get_secret_from_manag
oidc_cache = DualCache()
def _get_oidc_http_handler(timeout: Optional[httpx.Timeout] = None) -> HTTPHandler:
"""
Factory function to create HTTPHandler for OIDC requests.
This function can be mocked in tests.
Args:
timeout: Optional timeout for HTTP requests. Defaults to 600.0 seconds with 5.0 connect timeout.
Returns:
HTTPHandler instance configured for OIDC requests.
"""
if timeout is None:
timeout = httpx.Timeout(timeout=600.0, connect=5.0)
return HTTPHandler(timeout=timeout)
######### Secret Manager ############################
# checks if user has passed in a secret manager client
# if passed in then checks the secret there
@ -103,7 +119,7 @@ def get_secret( # noqa: PLR0915
if oidc_token is not None:
return oidc_token
oidc_client = HTTPHandler(timeout=httpx.Timeout(timeout=600.0, connect=5.0))
oidc_client = _get_oidc_http_handler()
# https://cloud.google.com/compute/docs/instances/verifying-instance-identity#request_signature
response = oidc_client.get(
"http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity",
@ -141,7 +157,7 @@ def get_secret( # noqa: PLR0915
if oidc_token is not None:
return oidc_token
oidc_client = HTTPHandler(timeout=httpx.Timeout(timeout=600.0, connect=5.0))
oidc_client = _get_oidc_http_handler()
response = oidc_client.get(
actions_id_token_request_url,
params={"audience": oidc_aud},

View file

@ -0,0 +1,27 @@
from typing import Dict, Optional, TypedDict
from litellm.types.integrations.custom_logger import StandardCustomLoggerInitParams
class DatadogCostManagementInitParams(StandardCustomLoggerInitParams):
"""
Init params for Datadog Cost Management
"""
datadog_cost_management_params: Optional[Dict] = None
class DatadogFOCUSCostEntry(TypedDict):
"""
Represents a single cost line item in the FOCUS format.
Ref: https://focus.finops.org/#specification
"""
ProviderName: str
ChargeDescription: str
ChargePeriodStart: str
ChargePeriodEnd: str
BilledCost: float
BillingCurrency: str
Tags: Optional[Dict[str, str]]

View file

@ -93,6 +93,67 @@ class GuardrailConverseContentBlock(TypedDict, total=False):
text: GuardrailConverseTextBlock
class CitationWebLocationBlock(TypedDict, total=False):
"""
Web location block for Nova grounding citations.
Contains the URL and domain from web search results.
Reference: https://docs.aws.amazon.com/nova/latest/userguide/grounding.html
"""
url: str
domain: str
class CitationLocationBlock(TypedDict, total=False):
"""
Location block containing the web location for a citation.
"""
web: CitationWebLocationBlock
class CitationReferenceBlock(TypedDict, total=False):
"""
Citation reference block containing a single citation with its location.
Each citation contains:
- location.web.url: The URL of the source
- location.web.domain: The domain of the source
"""
location: CitationLocationBlock
class CitationsContentBlock(TypedDict, total=False):
"""
Citations content block returned by Nova grounding (web search) tool.
When Nova grounding is enabled via systemTool, the model may return
citationsContent blocks containing web search citation references.
Reference: https://docs.aws.amazon.com/nova/latest/userguide/grounding.html
Example response structure:
{
"citationsContent": {
"citations": [
{
"location": {
"web": {
"url": "https://example.com/article",
"domain": "example.com"
}
}
}
]
}
}
"""
citations: List[CitationReferenceBlock]
class ContentBlock(TypedDict, total=False):
text: str
image: ImageBlock
@ -103,6 +164,7 @@ class ContentBlock(TypedDict, total=False):
cachePoint: CachePointBlock
reasoningContent: BedrockConverseReasoningContentBlock
guardContent: GuardrailConverseContentBlock
citationsContent: CitationsContentBlock
class MessageBlock(TypedDict):
@ -159,8 +221,24 @@ class ToolSpecBlock(TypedDict, total=False):
description: str
class SystemToolBlock(TypedDict, total=False):
"""
System tool block for Nova grounding and other built-in tools.
Example:
{
"systemTool": {
"name": "nova_grounding"
}
}
"""
name: Required[str]
class ToolBlock(TypedDict, total=False):
toolSpec: Optional[ToolSpecBlock]
systemTool: Optional[SystemToolBlock]
cachePoint: Optional[CachePointBlock]
@ -210,11 +288,13 @@ class ContentBlockStartEvent(TypedDict, total=False):
class ContentBlockDeltaEvent(TypedDict, total=False):
"""
Either 'text' or 'toolUse' will be specified for Converse API streaming response.
May also include 'citationsContent' when Nova grounding is enabled.
"""
text: str
toolUse: ToolBlockDeltaEvent
reasoningContent: BedrockConverseReasoningContentBlockDelta
citationsContent: CitationsContentBlock
class PerformanceConfigBlock(TypedDict):
@ -879,3 +959,8 @@ class BedrockGetBatchResponse(TypedDict, total=False):
outputDataConfig: BedrockOutputDataConfig
timeoutDurationInHours: Optional[int]
clientRequestToken: Optional[str]
class BedrockToolBlock(TypedDict, total=False):
toolSpec: Optional[ToolSpecBlock]
systemTool: Optional[SystemToolBlock] # For Nova grounding
cachePoint: Optional[CachePointBlock]

View file

@ -51,19 +51,20 @@ class GenericGuardrailAPIRequest(BaseModel):
"""Request model for the Generic Guardrail API"""
input_type: Literal["request", "response"]
litellm_call_id: Optional[str] # the call id of the individual LLM call
litellm_call_id: Optional[str] = None # the call id of the individual LLM call
litellm_trace_id: Optional[
str
] # the trace id of the LLM call - useful if there are multiple LLM calls for the same conversation
structured_messages: Optional[List[AllMessageValues]]
images: Optional[List[str]]
tools: Optional[List[ChatCompletionToolParam]]
texts: Optional[List[str]]
] = None # the trace id of the LLM call - useful if there are multiple LLM calls for the same conversation
structured_messages: Optional[List[AllMessageValues]] = None
images: Optional[List[str]] = None
tools: Optional[List[ChatCompletionToolParam]] = None
texts: Optional[List[str]] = None
request_data: GenericGuardrailAPIMetadata
additional_provider_specific_params: Optional[Dict[str, Any]]
additional_provider_specific_params: Optional[Dict[str, Any]] = None
tool_calls: Optional[
Union[List[ChatCompletionToolCallChunk], List[ChatCompletionMessageToolCall]]
]
] = None
model: Optional[str] = None # the model being used for the LLM call
class GenericGuardrailAPIResponse:

View file

@ -3431,3 +3431,4 @@ class GenericGuardrailAPIInputs(TypedDict, total=False):
structured_messages: List[
AllMessageValues
] # structured messages sent to the LLM - indicates if text is from system or user
model: Optional[str] # the model being used for the LLM call

View file

@ -8053,6 +8053,12 @@ class ProviderConfigManager:
)
return OpenrouterEmbeddingConfig()
elif litellm.LlmProviders.VERCEL_AI_GATEWAY == provider:
from litellm.llms.vercel_ai_gateway.embedding.transformation import (
VercelAIGatewayEmbeddingConfig,
)
return VercelAIGatewayEmbeddingConfig()
elif litellm.LlmProviders.GIGACHAT == provider:
return litellm.GigaChatEmbeddingConfig()
elif litellm.LlmProviders.SAGEMAKER == provider:

View file

@ -10232,6 +10232,48 @@
"mode": "completion",
"output_cost_per_token": 5e-07
},
"deepseek-v3-2-251201": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",
"max_input_tokens": 98304,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 0.0,
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"glm-4-7-251222": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",
"max_input_tokens": 204800,
"max_output_tokens": 131072,
"max_tokens": 131072,
"mode": "chat",
"output_cost_per_token": 0.0,
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"kimi-k2-thinking-251104": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",
"max_input_tokens": 229376,
"max_output_tokens": 32768,
"max_tokens": 32768,
"mode": "chat",
"output_cost_per_token": 0.0,
"supports_assistant_prefill": true,
"supports_function_calling": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_tool_choice": true
},
"doubao-embedding": {
"input_cost_per_token": 0.0,
"litellm_provider": "volcengine",

View file

@ -1,6 +1,6 @@
[tool.poetry]
name = "litellm"
version = "1.81.3"
version = "1.81.4"
description = "Library to easily interface with LLM API providers"
authors = ["BerriAI"]
license = "MIT"
@ -145,6 +145,7 @@ mypy = "^1.0"
pytest = "^7.4.3"
pytest-mock = "^3.12.0"
pytest-asyncio = "^0.21.1"
pytest-retry = "^1.6.3"
requests-mock = "^1.12.1"
responses = "^0.25.7"
respx = "^0.22.0"
@ -173,7 +174,7 @@ requires = ["poetry-core", "wheel"]
build-backend = "poetry.core.masonry.api"
[tool.commitizen]
version = "1.81.3"
version = "1.81.4"
version_files = [
"pyproject.toml:^version"
]
@ -183,6 +184,8 @@ plugins = "pydantic.mypy"
[tool.pytest.ini_options]
asyncio_mode = "auto"
retries = 20
retry_delay = 5
markers = [
"asyncio: mark test as an asyncio test",
"limit_leaks: mark test with memory limit for leak detection (e.g., '40 MB')",

View file

@ -6,7 +6,7 @@ def get_function_names_from_file(file_path):
"""
Extracts all function names from a given Python file.
"""
with open(file_path, "r") as file:
with open(file_path, "r", encoding="utf-8") as file:
tree = ast.parse(file.read())
function_names = []
@ -45,7 +45,7 @@ def get_all_functions_called_in_tests(base_dir):
if file.endswith(".py") and "router" in file.lower():
print("file: ", file)
file_path = os.path.join(root, file)
with open(file_path, "r") as f:
with open(file_path, "r", encoding="utf-8") as f:
try:
tree = ast.parse(f.read())
except SyntaxError:

View file

@ -3954,3 +3954,288 @@ def test_bedrock_openai_error_handling():
assert exc_info.value.status_code == 422
print("✓ Error handling works correctly")
# ============================================================================
# Nova Grounding (web_search_options) Unit Tests (Mocked)
# ============================================================================
def test_bedrock_nova_grounding_web_search_options_non_streaming():
"""
Unit test for Nova grounding using web_search_options parameter (non-streaming).
This test mocks the HTTP call to verify:
1. web_search_options is correctly mapped to systemTool for Nova models
2. The request structure is correct
Related: https://docs.aws.amazon.com/nova/latest/userguide/grounding.html
"""
from unittest.mock import patch, MagicMock
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
messages = [
{
"role": "user",
"content": "What is the current population of Tokyo, Japan?",
}
]
with patch.object(client, "post") as mock_post:
try:
completion(
model="us.amazon.nova-pro-v1:0", # No bedrock/ prefix when using api_base
messages=messages,
web_search_options={}, # Enables Nova grounding
max_tokens=500,
client=client,
api_base="https://bedrock-runtime.us-east-1.amazonaws.com",
)
except Exception:
pass # Expected - we're just checking the request structure
# Verify the request was made correctly
if mock_post.called:
request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}"))
print(f"Request body: {json.dumps(request_body, indent=2)}")
# Verify toolConfig is present with systemTool
assert "toolConfig" in request_body, "toolConfig should be in request"
tool_config = request_body["toolConfig"]
assert "tools" in tool_config, "tools should be in toolConfig"
# Find the systemTool for nova_grounding
system_tool_found = False
for tool in tool_config["tools"]:
if "systemTool" in tool:
assert tool["systemTool"]["name"] == "nova_grounding"
system_tool_found = True
break
assert system_tool_found, "systemTool with nova_grounding should be present"
print(f"✓ web_search_options correctly transformed to systemTool (non-streaming)")
def test_bedrock_nova_grounding_with_function_tools():
"""
Unit test for Nova grounding combined with regular function tools.
This tests the scenario where users want both web grounding AND
custom function calling capabilities.
"""
from unittest.mock import patch
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
# Regular function tool
tools = [
{
"type": "function",
"function": {
"name": "get_stock_price",
"description": "Get the current stock price for a given ticker symbol",
"parameters": {
"type": "object",
"properties": {
"ticker": {
"type": "string",
"description": "The stock ticker symbol, e.g. AAPL, GOOGL",
}
},
"required": ["ticker"],
},
},
}
]
messages = [
{
"role": "user",
"content": "What is the current market cap of Apple Inc?",
}
]
with patch.object(client, "post") as mock_post:
try:
completion(
model="us.amazon.nova-pro-v1:0", # No bedrock/ prefix when using api_base
messages=messages,
tools=tools,
web_search_options={}, # Also enable web grounding
max_tokens=500,
client=client,
api_base="https://bedrock-runtime.us-east-1.amazonaws.com",
)
except Exception:
pass # Expected - we're just checking the request structure
# Verify the request was made correctly
if mock_post.called:
request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}"))
print(f"Request body: {json.dumps(request_body, indent=2)}")
# Verify toolConfig has both function tool and systemTool
assert "toolConfig" in request_body, "toolConfig should be in request"
tool_config = request_body["toolConfig"]
assert "tools" in tool_config, "tools should be in toolConfig"
tools_in_request = tool_config["tools"]
# Should have both the function tool and the systemTool
function_tool_found = False
system_tool_found = False
for tool in tools_in_request:
if "toolSpec" in tool:
assert tool["toolSpec"]["name"] == "get_stock_price"
function_tool_found = True
if "systemTool" in tool:
assert tool["systemTool"]["name"] == "nova_grounding"
system_tool_found = True
assert function_tool_found, "Function tool (get_stock_price) should be present"
assert system_tool_found, "systemTool (nova_grounding) should be present"
print(f"✓ Both function tools and web_search_options correctly combined")
@pytest.mark.asyncio
async def test_bedrock_nova_grounding_async():
"""
Async unit test for Nova grounding via web_search_options.
This test verifies the request transformation for async calls.
"""
from unittest.mock import patch, AsyncMock
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
client = AsyncHTTPHandler()
messages = [
{
"role": "user",
"content": "What is the weather forecast for New York City today?",
}
]
with patch.object(client, "post", new=AsyncMock()) as mock_post:
try:
await litellm.acompletion(
model="us.amazon.nova-pro-v1:0", # No bedrock/ prefix when using api_base
messages=messages,
web_search_options={},
max_tokens=500,
client=client,
api_base="https://bedrock-runtime.us-east-1.amazonaws.com",
)
except Exception:
pass # Expected - we're just checking the request structure
# Verify the request was made correctly
if mock_post.called:
request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}"))
print(f"Request body: {json.dumps(request_body, indent=2)}")
# Verify toolConfig is present with systemTool
assert "toolConfig" in request_body, "toolConfig should be in request"
tool_config = request_body["toolConfig"]
assert "tools" in tool_config, "tools should be in toolConfig"
# Find the systemTool for nova_grounding
system_tool_found = False
for tool in tool_config["tools"]:
if "systemTool" in tool:
assert tool["systemTool"]["name"] == "nova_grounding"
system_tool_found = True
break
assert system_tool_found, "systemTool with nova_grounding should be present"
print(f"✓ Async web_search_options correctly transformed to systemTool")
def test_bedrock_nova_web_search_options_ignored_for_non_nova():
"""
Test that web_search_options is ignored for non-Nova Bedrock models.
Nova grounding is only supported on Nova models. For other models,
the parameter should be silently ignored.
"""
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
config = AmazonConverseConfig()
# Should return None for non-Nova models
result = config._map_web_search_options({}, "anthropic.claude-3-sonnet-v1")
assert result is None
result = config._map_web_search_options({}, "amazon.titan-text-express-v1")
assert result is None
# Should return systemTool for Nova models
result = config._map_web_search_options({}, "amazon.nova-pro-v1:0")
assert result is not None
system_tool = result.get("systemTool")
assert system_tool is not None
assert system_tool["name"] == "nova_grounding"
result2 = config._map_web_search_options({}, "us.amazon.nova-premier-v1:0")
assert result2 is not None
system_tool2 = result2.get("systemTool")
assert system_tool2 is not None
assert system_tool2["name"] == "nova_grounding"
def test_bedrock_nova_grounding_request_transformation():
"""
Unit test to verify that web_search_options transforms to systemTool in the request.
"""
from unittest.mock import patch, MagicMock
from litellm.llms.custom_httpx.http_handler import HTTPHandler
client = HTTPHandler()
messages = [{"role": "user", "content": "What is the population of Tokyo?"}]
with patch.object(client, "post") as mock_post:
mock_post.return_value = MagicMock(
status_code=200,
json=lambda: {
"output": {"message": {"role": "assistant", "content": [{"text": "Test"}]}},
"stopReason": "end_turn",
"usage": {"inputTokens": 10, "outputTokens": 5}
}
)
try:
response = completion(
model="bedrock/us.amazon.nova-pro-v1:0",
messages=messages,
web_search_options={},
max_tokens=100,
client=client,
)
except Exception:
pass # Expected - we're just checking the request
if mock_post.called:
request_body = json.loads(mock_post.call_args.kwargs.get("data", "{}"))
print(f"Request body: {json.dumps(request_body, indent=2)}")
# Verify toolConfig is present with systemTool
assert "toolConfig" in request_body, "toolConfig should be in request"
tool_config = request_body["toolConfig"]
assert "tools" in tool_config, "tools should be in toolConfig"
tools_in_request = tool_config["tools"]
# Find the systemTool
system_tool_found = False
for tool in tools_in_request:
if "systemTool" in tool:
assert tool["systemTool"]["name"] == "nova_grounding"
system_tool_found = True
break
assert system_tool_found, "systemTool with nova_grounding should be present"
print("✓ web_search_options correctly transformed to systemTool")

View file

@ -5,7 +5,7 @@ import base64
import os
import sys
import pytest
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import AsyncMock, MagicMock, patch, ANY
# Add the project root to the path
sys.path.insert(0, os.path.abspath("../../.."))
@ -183,7 +183,7 @@ class TestMCPClientUnitTests:
assert result == mock_result
mock_session_instance.initialize.assert_called_once()
mock_session_instance.call_tool.assert_called_once_with(
name="test_tool", arguments={"arg1": "value1"}
name="test_tool", arguments={"arg1": "value1"},progress_callback=ANY
)

View file

@ -473,18 +473,41 @@ async def test_get_users_key_count(prisma_client):
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
await litellm.proxy.proxy_server.prisma_client.connect()
# Get initial user list and select the first user
initial_users = await get_users(role=None, page=1, page_size=20)
# Create a test user with no initial keys to ensure deterministic behavior
test_user_id = f"test_user_key_count-{uuid.uuid4()}"
test_user_request = NewUserRequest(
user_id=test_user_id,
user_role=LitellmUserRoles.INTERNAL_USER.value,
auto_create_key=False, # Ensure we start with 0 keys
)
await new_user(
test_user_request,
UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-1234",
user_id="admin",
),
)
# Get initial key count for the test user
initial_users = await get_users(
user_ids=test_user_id,
role=None,
page=1,
page_size=20,
)
print("initial_users", initial_users)
assert len(initial_users["users"]) > 0, "No users found to test with"
assert len(initial_users["users"]) == 1, "Test user should be found"
test_user = initial_users["users"][0]
assert test_user.user_id == test_user_id
initial_key_count = test_user.key_count
assert initial_key_count == 0, f"Expected initial key count to be 0, but got {initial_key_count}"
# Create a new key for the selected user
# Create a new key for the test user
new_key = await generate_key_fn(
data=GenerateKeyRequest(
user_id=test_user.user_id,
user_id=test_user_id,
key_alias=f"test_key_{uuid.uuid4()}",
models=["fake-model"],
),
@ -496,19 +519,26 @@ async def test_get_users_key_count(prisma_client):
)
# Get updated user list and check key count
updated_users = await get_users(role=None, page=1, page_size=20)
updated_users = await get_users(
user_ids=test_user_id,
role=None,
page=1,
page_size=20,
)
print("updated_users", updated_users)
updated_key_count = None
for user in updated_users["users"]:
if user.user_id == test_user.user_id:
updated_key_count = user.key_count
break
assert len(updated_users["users"]) == 1, "Test user should still be found"
updated_user = updated_users["users"][0]
updated_key_count = updated_user.key_count
assert updated_key_count is not None, "Test user not found in updated users list"
assert (
updated_key_count == initial_key_count + 1
), f"Expected key count to increase by 1, but got {updated_key_count} (was {initial_key_count})"
# Clean up test user and keys
await prisma_client.db.litellm_usertable.delete(
where={"user_id": test_user_id}
)
async def cleanup_existing_teams(prisma_client):
all_teams = await prisma_client.db.litellm_teamtable.find_many()

View file

@ -0,0 +1,64 @@
import os
from unittest import mock
from litellm.proxy import utils
# Test the utility function logic
def test_get_server_root_path_unset():
"""
Test that get_server_root_path returns empty string when SERVER_ROOT_PATH is unset
"""
with mock.patch.dict(os.environ, {}, clear=True):
# We need to make sure SERVER_ROOT_PATH is not in env
if "SERVER_ROOT_PATH" in os.environ:
del os.environ["SERVER_ROOT_PATH"]
root_path = utils.get_server_root_path()
assert (
root_path == ""
), "Should return empty string when unset to allow X-Forwarded-Prefix"
def test_get_server_root_path_set():
"""
Test that get_server_root_path returns the value when SERVER_ROOT_PATH is set
"""
with mock.patch.dict(os.environ, {"SERVER_ROOT_PATH": "/my-path"}, clear=True):
root_path = utils.get_server_root_path()
assert root_path == "/my-path", "Should return the set value"
def test_get_server_root_path_empty_string():
"""
Test that get_server_root_path returns empty string when SERVER_ROOT_PATH is explicitly empty
"""
with mock.patch.dict(os.environ, {"SERVER_ROOT_PATH": ""}, clear=True):
root_path = utils.get_server_root_path()
assert (
root_path == ""
), "Should return empty string when explicitly set to empty"
# Integration test simulation for FastAPI app initialization
def test_fastapi_app_initialization_mock():
"""
Simulate how proxy_server.py initializes FastAPI app with the root_path.
We don't import proxy_server because it has global side effects/singletons.
Instead we verify the logic flow.
"""
from fastapi import FastAPI
# CASE 1: Proxy Mode (Unset)
with mock.patch.dict(os.environ, {}, clear=True):
if "SERVER_ROOT_PATH" in os.environ:
del os.environ["SERVER_ROOT_PATH"]
server_root_path = utils.get_server_root_path()
app = FastAPI(root_path=server_root_path)
assert app.root_path == ""
# CASE 2: Direct Mode (Set)
with mock.patch.dict(os.environ, {"SERVER_ROOT_PATH": "/custom-root"}, clear=True):
server_root_path = utils.get_server_root_path()
app = FastAPI(root_path=server_root_path)
assert app.root_path == "/custom-root"

View file

@ -2169,3 +2169,32 @@ def test_resolve_model_name_from_model_id():
result = router.resolve_model_name_from_model_id("gpt-3.5-turbo")
assert result == "gpt-3.5-turbo"
def test_get_valid_args():
"""Test get_valid_args static method returns valid Router.__init__ arguments"""
# Call the static method
valid_args = Router.get_valid_args()
# Verify it returns a list
assert isinstance(valid_args, list)
assert len(valid_args) > 0
# Verify it contains expected Router.__init__ arguments
expected_args = [
"model_list",
"routing_strategy",
"cache_responses",
"num_retries",
"timeout",
"fallbacks",
]
for arg in expected_args:
assert arg in valid_args, f"Expected argument '{arg}' not found in valid_args"
# Verify "self" is not in the list (since it's removed)
assert "self" not in valid_args
# Verify it contains keyword-only arguments too
# These are common Router.__init__ parameters
assert "assistants_config" in valid_args or "search_tools" in valid_args

View file

@ -357,15 +357,15 @@ class TestContainerIntegration:
def test_error_handling_integration(self):
"""Test error handling in the integration flow."""
with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
# Simulate an API error
mock_handler.container_create_handler.side_effect = litellm.APIError(
status_code=400,
message="API Error occurred",
llm_provider="openai",
model=""
)
# Simulate an API error
api_error = litellm.APIError(
status_code=400,
message="API Error occurred",
llm_provider="openai",
model=""
)
with patch.object(litellm.main.base_llm_http_handler, 'container_create_handler', side_effect=api_error):
with pytest.raises(litellm.APIError):
create_container(
name="Error Test Container",
@ -385,12 +385,12 @@ class TestContainerIntegration:
name="Provider Test Container"
)
with patch('litellm.containers.main.base_llm_http_handler') as mock_handler:
mock_handler.container_create_handler.return_value = mock_response
with patch.object(litellm.main.base_llm_http_handler, 'container_create_handler', return_value=mock_response) as mock_handler:
response = create_container(
name="Provider Test Container",
custom_llm_provider=provider
)
assert response.name == "Provider Test Container"
# Verify the mock was actually called (not making real API calls)
mock_handler.assert_called_once()

View file

@ -2,7 +2,9 @@ import os
import sys
import unittest.mock as mock
import httpx
import pytest
import respx
from httpx import Response
sys.path.insert(0, os.path.abspath("../../.."))
@ -11,6 +13,8 @@ from litellm_enterprise.enterprise_callbacks.send_emails.resend_email import (
ResendEmailLogger,
)
# Test file for Resend email integration
@pytest.fixture
def mock_env_vars():
@ -23,21 +27,27 @@ def mock_httpx_client():
with mock.patch(
"litellm_enterprise.enterprise_callbacks.send_emails.resend_email.get_async_httpx_client"
) as mock_client:
# Create a mock response
mock_response = mock.AsyncMock(spec=Response)
mock_response = mock.Mock(spec=Response)
mock_response.status_code = 200
mock_response.json.return_value = {"id": "test_email_id"}
mock_response.raise_for_status.return_value = None
# Create a mock client
mock_async_client = mock.AsyncMock()
mock_async_client.post.return_value = mock_response
mock_client.return_value = mock_async_client
mock_client.return_value = mock_async_client
yield mock_async_client
@pytest.mark.asyncio
@respx.mock
async def test_send_email_success(mock_env_vars, mock_httpx_client):
# Block all HTTP requests at network level to prevent real API calls
respx.post("https://api.resend.com/emails").mock(
return_value=httpx.Response(200, json={"id": "test_email_id"})
)
# Initialize the logger
logger = ResendEmailLogger()
@ -71,7 +81,13 @@ async def test_send_email_success(mock_env_vars, mock_httpx_client):
@pytest.mark.asyncio
@respx.mock
async def test_send_email_missing_api_key(mock_httpx_client):
# Block all HTTP requests at network level to prevent real API calls
respx.post("https://api.resend.com/emails").mock(
return_value=httpx.Response(200, json={"id": "test_email_id"})
)
# Remove the API key from environment before initializing logger
original_key = os.environ.pop("RESEND_API_KEY", None)
@ -86,7 +102,9 @@ async def test_send_email_missing_api_key(mock_httpx_client):
html_body = "<p>Test email body</p>"
# Mock the response to avoid making real HTTP requests
mock_response = mock.AsyncMock(spec=Response)
mock_response = mock.Mock(spec=Response)
mock_response.raise_for_status.return_value = None
mock_response.status_code = 200
mock_response.json.return_value = {"id": "test_email_id"}
mock_httpx_client.post.return_value = mock_response
@ -107,7 +125,13 @@ async def test_send_email_missing_api_key(mock_httpx_client):
@pytest.mark.asyncio
@respx.mock
async def test_send_email_multiple_recipients(mock_env_vars, mock_httpx_client):
# Block all HTTP requests at network level to prevent real API calls
respx.post("https://api.resend.com/emails").mock(
return_value=httpx.Response(200, json={"id": "test_email_id"})
)
# Initialize the logger
logger = ResendEmailLogger()
@ -118,7 +142,9 @@ async def test_send_email_multiple_recipients(mock_env_vars, mock_httpx_client):
html_body = "<p>Test email body</p>"
# Mock the response to avoid making real HTTP requests
mock_response = mock.AsyncMock(spec=Response)
mock_response = mock.Mock(spec=Response)
mock_response.raise_for_status.return_value = None
mock_response.status_code = 200
mock_response.json.return_value = {"id": "test_email_id"}
mock_httpx_client.post.return_value = mock_response

View file

@ -2,7 +2,9 @@ import os
import sys
import unittest.mock as mock
import httpx
import pytest
import respx
from httpx import Response
sys.path.insert(0, os.path.abspath("../../.."))
@ -14,8 +16,26 @@ from litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email import (
@pytest.fixture
def mock_env_vars():
with mock.patch.dict(os.environ, {"SENDGRID_API_KEY": "test_api_key"}):
# Store original values
original_api_key = os.environ.get("SENDGRID_API_KEY")
original_sender_email = os.environ.get("SENDGRID_SENDER_EMAIL")
# Set test API key and remove SENDGRID_SENDER_EMAIL to ensure isolation
os.environ["SENDGRID_API_KEY"] = "test_api_key"
if "SENDGRID_SENDER_EMAIL" in os.environ:
del os.environ["SENDGRID_SENDER_EMAIL"]
try:
yield
finally:
# Restore original values
if original_api_key is not None:
os.environ["SENDGRID_API_KEY"] = original_api_key
elif "SENDGRID_API_KEY" in os.environ:
del os.environ["SENDGRID_API_KEY"]
if original_sender_email is not None:
os.environ["SENDGRID_SENDER_EMAIL"] = original_sender_email
@pytest.fixture
@ -23,14 +43,16 @@ def mock_httpx_client():
with mock.patch(
"litellm_enterprise.enterprise_callbacks.send_emails.sendgrid_email.get_async_httpx_client"
) as mock_client:
mock_response = mock.AsyncMock(spec=Response)
mock_response = mock.Mock(spec=Response)
mock_response.status_code = 202
mock_response.text = "accepted"
mock_response.raise_for_status.return_value = None
mock_async_client = mock.AsyncMock()
mock_async_client.post.return_value = mock_response
mock_client.return_value = mock_async_client
mock_client.return_value = mock_async_client
yield mock_async_client
@ -62,18 +84,12 @@ async def test_send_email_success(mock_env_vars, mock_httpx_client):
@pytest.mark.asyncio
async def test_send_email_missing_api_key(mock_httpx_client):
# Remove the API key from environment before initializing logger
async def test_send_email_missing_api_key():
original_key = os.environ.pop("SENDGRID_API_KEY", None)
try:
logger = SendGridEmailLogger()
# Mock the response to avoid making real HTTP requests
mock_response = mock.AsyncMock(spec=Response)
mock_response.status_code = 401
mock_httpx_client.post.return_value = mock_response
with pytest.raises(ValueError):
await logger.send_email(
from_email="test@example.com",
@ -81,16 +97,19 @@ async def test_send_email_missing_api_key(mock_httpx_client):
subject="Test Subject",
html_body="<p>Test email body</p>",
)
mock_httpx_client.post.assert_not_called()
finally:
# Restore the original key if it existed
if original_key is not None:
os.environ["SENDGRID_API_KEY"] = original_key
@pytest.mark.asyncio
@respx.mock
async def test_send_email_multiple_recipients(mock_env_vars, mock_httpx_client):
# Block all HTTP requests at network level to prevent real API calls
respx.post("https://api.sendgrid.com/v3/mail/send").mock(
return_value=httpx.Response(202, text="accepted")
)
logger = SendGridEmailLogger()
from_email = "test@example.com"
@ -98,10 +117,10 @@ async def test_send_email_multiple_recipients(mock_env_vars, mock_httpx_client):
subject = "Test Subject"
html_body = "<p>Test email body</p>"
# Mock the response to avoid making real HTTP requests
mock_response = mock.AsyncMock(spec=Response)
mock_response = mock.Mock(spec=Response)
mock_response.status_code = 202
mock_response.text = "accepted"
mock_response.raise_for_status.return_value = None
mock_httpx_client.post.return_value = mock_response
await logger.send_email(

View file

@ -0,0 +1,169 @@
import os
import time
from unittest.mock import AsyncMock
import pytest
from httpx import Response
from litellm.integrations.datadog.datadog_cost_management import (
DatadogCostManagementLogger,
)
from litellm.types.utils import StandardLoggingPayload
@pytest.fixture
def clean_env():
# Save original env
original_api_key = os.environ.get("DD_API_KEY")
original_app_key = os.environ.get("DD_APP_KEY")
original_site = os.environ.get("DD_SITE")
# Set test env
os.environ["DD_API_KEY"] = "test_api_key"
os.environ["DD_APP_KEY"] = "test_app_key"
os.environ["DD_SITE"] = "test.datadoghq.com"
yield
# Restore original env
if original_api_key:
os.environ["DD_API_KEY"] = original_api_key
else:
del os.environ["DD_API_KEY"]
if original_app_key:
os.environ["DD_APP_KEY"] = original_app_key
else:
del os.environ["DD_APP_KEY"]
if original_site:
os.environ["DD_SITE"] = original_site
else:
del os.environ["DD_SITE"]
@pytest.mark.asyncio
async def test_init(clean_env):
"""
Test initialization sets up clients and url correctly
"""
logger = DatadogCostManagementLogger()
assert logger.dd_api_key == "test_api_key"
assert logger.dd_app_key == "test_app_key"
assert (
logger.upload_url == "https://api.test.datadoghq.com/api/v2/cost/custom_costs"
)
@pytest.mark.asyncio
async def test_aggregate_costs(clean_env):
"""
Test that costs are correctly aggregated by provider, model, and date
"""
logger = DatadogCostManagementLogger()
# Mock some log payloads
now = time.time()
day_str = time.strftime("%Y-%m-%d", time.localtime(now))
logs = [
StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=now,
metadata={"user_api_key_team_alias": "team-a"},
),
StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.02,
startTime=now,
metadata={"user_api_key_team_alias": "team-a"},
),
StandardLoggingPayload(
custom_llm_provider="anthropic",
model="claude-3",
response_cost=0.05,
startTime=now,
),
]
aggregated = logger._aggregate_costs(logs)
assert len(aggregated) == 2
# Check OpenAI entry
openai_entry = next(e for e in aggregated if e["ProviderName"] == "openai")
assert openai_entry["BilledCost"] == 0.03
assert openai_entry["ChargeDescription"] == "LLM Usage for gpt-4"
assert openai_entry["ChargePeriodStart"] == day_str
assert openai_entry["Tags"]["team"] == "team-a"
assert "env" in openai_entry["Tags"]
assert "service" in openai_entry["Tags"]
# Check Anthropic entry
anthropic_entry = next(e for e in aggregated if e["ProviderName"] == "anthropic")
assert anthropic_entry["BilledCost"] == 0.05
@pytest.mark.asyncio
async def test_async_log_success_event(clean_env):
"""
Test that logs are added to queue
"""
logger = DatadogCostManagementLogger(batch_size=10)
await logger.async_log_success_event(
kwargs={"standard_logging_object": {"response_cost": 0.01}},
response_obj={},
start_time=time.time(),
end_time=time.time(),
)
assert len(logger.log_queue) == 1
assert logger.log_queue[0]["response_cost"] == 0.01
# Test zero cost ignored
await logger.async_log_success_event(
kwargs={"standard_logging_object": {"response_cost": 0.0}},
response_obj={},
start_time=time.time(),
end_time=time.time(),
)
assert len(logger.log_queue) == 1
@pytest.mark.asyncio
async def test_async_send_batch(clean_env):
"""
Test that batch is aggregated and uploaded
"""
logger = DatadogCostManagementLogger()
logger.async_client = AsyncMock()
logger.async_client.put.return_value = Response(202, json={"status": "ok"})
# Add logs directly to queue
logger.log_queue = [
StandardLoggingPayload(
custom_llm_provider="openai",
model="gpt-4",
response_cost=0.01,
startTime=time.time(),
)
]
await logger.async_send_batch()
# Verify API called
assert logger.async_client.put.called
call_args = logger.async_client.put.call_args
assert call_args[0][0] == "https://api.test.datadoghq.com/api/v2/cost/custom_costs"
import json
# Use call_args.kwargs['content']
content = json.loads(call_args[1]["content"])
assert content[0]["ProviderName"] == "openai"
assert content[0]["BilledCost"] == 0.01

View file

@ -0,0 +1,62 @@
import os
from unittest.mock import patch
from litellm.integrations.datadog.datadog_llm_obs import DataDogLLMObsLogger
def test_datadog_llm_obs_agent_configuration():
"""
Test that DataDog LLM Obs logger correctly configures agent endpoint.
"""
test_env = {
"LITELLM_DD_AGENT_HOST": "localhost",
"LITELLM_DD_LLM_OBS_PORT": "10518",
"DD_API_KEY": "test-api-key", # Optional, but checking if it's preserved
}
# Ensure DD_SITE is NOT set to verify we don't need it in agent mode
with patch.dict(os.environ, test_env, clear=True):
with patch("asyncio.create_task"): # Prevent periodic flush task from running
dd_logger = DataDogLLMObsLogger()
expected_url = "http://localhost:10518/api/intake/llm-obs/v1/trace/spans"
assert dd_logger.intake_url == expected_url
assert dd_logger.DD_API_KEY == "test-api-key"
def test_datadog_llm_obs_agent_no_api_key_ok():
"""
Test that agent mode works WITHOUT DD_API_KEY (agent handles auth).
"""
test_env = {
"LITELLM_DD_AGENT_HOST": "localhost",
# No DD_API_KEY
}
with patch.dict(os.environ, test_env, clear=True):
with patch("asyncio.create_task"):
# Should NOT raise exception anymore
dd_logger = DataDogLLMObsLogger()
assert dd_logger.DD_API_KEY is None
# Default port is 8126 if not set
expected_url = "http://localhost:8126/api/intake/llm-obs/v1/trace/spans"
assert dd_logger.intake_url == expected_url
def test_datadog_llm_obs_direct_api_configuration():
"""
Test that direct API configuration still works as expected.
"""
test_env = {
"DD_API_KEY": "direct-api-key",
"DD_SITE": "us5.datadoghq.com",
}
with patch.dict(os.environ, test_env, clear=True):
with patch("asyncio.create_task"):
dd_logger = DataDogLLMObsLogger()
expected_url = "https://api.us5.datadoghq.com/api/intake/llm-obs/v1/trace/spans"
assert dd_logger.intake_url == expected_url
assert dd_logger.DD_API_KEY == "direct-api-key"

View file

@ -0,0 +1,73 @@
import pytest
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.types.guardrails import GuardrailEventHooks
import json
class TestCustomGuardrailRecursion:
"""
Specific tests for the circular reference / RecursionError fix in logging.
"""
def test_log_guardrail_information_handles_circular_references(self):
"""
Test that add_standard_logging method sanitizes input data containing circular references
instead of crashing.
This reproduces the Langfuse crash scenario:
Request -> Metadata -> GuardrailResponse -> DebugContext -> Request
"""
guardrail = CustomGuardrail(
guardrail_name="recursion_test_guardrail",
event_hook=GuardrailEventHooks.pre_call,
)
# 1. Setup Circular Data
request_data = {"user_id": "test_recursive_user"}
metadata = {"session_id": "123"}
request_data["metadata"] = metadata
# Create the danger: Guardrail Response holding a reference back to request_data
dirty_response = {
"flagged": False,
"debug_context": request_data, # <--- ACCESS TO ROOT (Circular Ref)
}
# 2. Invoke the logging method
# If the fix is working, this will NOT raise RecursionError
try:
guardrail.add_standard_logging_guardrail_information_to_request_data(
guardrail_json_response=dirty_response,
request_data=request_data,
guardrail_status="success",
start_time=1.0,
end_time=2.0,
duration=1.0,
masked_entity_count={},
event_type=GuardrailEventHooks.pre_call,
)
except RecursionError:
pytest.fail(
"RecursionError raised! The cyclic reference sanitization failed."
)
# 3. Verify the data stored is safe
stored_info = request_data["metadata"][
"standard_logging_guardrail_information"
][0]
stored_response = stored_info["guardrail_response"]
# Check that we can dump it to JSON without crashing (Ultimate proof)
try:
json.dumps(stored_response)
except Exception as e:
pytest.fail(f"Stored data is not JSON serializable: {e}")
# Check content - keys should be preserved but recursion broken
assert "debug_context" in stored_response
debug_context = stored_response["debug_context"]
# In a sanitized copy, the nested metadata should be a copy, not the original live dict
assert debug_context["user_id"] == "test_recursive_user"
# The 'metadata' inside 'debug_context' would be where recursion stops or is filtered
assert "metadata" in debug_context

View file

@ -1138,6 +1138,73 @@ def test_bedrock_create_bedrock_block_different_document_formats():
assert block["document"]["name"].endswith(f"_{format_type}")
assert block["document"]["format"] == format_type
def test_bedrock_nova_web_search_options_mapping():
"""
Test that web_search_options is correctly mapped to Nova grounding.
This follows the LiteLLM pattern for web search where:
- Vertex AI maps web_search_options to {"googleSearch": {}}
- Anthropic maps web_search_options to {"type": "web_search_20250305", ...}
- Nova should map web_search_options to {"systemTool": {"name": "nova_grounding"}}
"""
from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig
config = AmazonConverseConfig()
# Test basic mapping for Nova model
result = config._map_web_search_options({}, "amazon.nova-pro-v1:0")
assert result is not None
system_tool = result.get("systemTool")
assert system_tool is not None
assert system_tool["name"] == "nova_grounding"
# Test with search_context_size (should be ignored for Nova)
result2 = config._map_web_search_options(
{"search_context_size": "high"},
"us.amazon.nova-premier-v1:0"
)
assert result2 is not None
system_tool2 = result2.get("systemTool")
assert system_tool2 is not None
assert system_tool2["name"] == "nova_grounding"
# Nova doesn't support search_context_size, so it's just ignored
def test_bedrock_tools_pt_does_not_handle_system_tool():
"""
Verify that _bedrock_tools_pt does NOT handle system_tool format.
System tools (nova_grounding) should be added via web_search_options,
not via the tools parameter directly.
"""
from litellm.litellm_core_utils.prompt_templates.factory import _bedrock_tools_pt
# Regular function tools should still work
tools = [
{
"type": "function",
"function": {
"name": "get_weather",
"description": "Get the current weather",
"parameters": {
"type": "object",
"properties": {
"location": {"type": "string"}
},
"required": ["location"]
}
}
}
]
result = _bedrock_tools_pt(tools=tools)
assert len(result) == 1
tool_spec = result[0].get("toolSpec")
assert tool_spec is not None
assert tool_spec["name"] == "get_weather"
def test_convert_to_anthropic_tool_result_image_with_cache_control():
"""
@ -1305,12 +1372,12 @@ def test_convert_to_anthropic_tool_result_image_url_as_http():
assert result["content"][0]["cache_control"]["type"] == "ephemeral"
def test_anthropic_messages_pt_server_tool_use_passthrough():
"""
Test that anthropic_messages_pt passes through server_tool_use and
Test that anthropic_messages_pt passes through server_tool_use and
tool_search_tool_result blocks in assistant message content.
These are Anthropic-native content types used for tool search functionality
that need to be preserved when reconstructing multi-turn conversations.
Fixes: https://github.com/BerriAI/litellm/issues/XXXXX
"""
from litellm.litellm_core_utils.prompt_templates.factory import anthropic_messages_pt
@ -1359,15 +1426,15 @@ def test_anthropic_messages_pt_server_tool_use_passthrough():
# Verify we have 3 messages (user, assistant, user)
assert len(result) == 3
# Verify the assistant message content
assistant_msg = result[1]
assert assistant_msg["role"] == "assistant"
assert isinstance(assistant_msg["content"], list)
# Find the different content block types
content_types = [block.get("type") for block in assistant_msg["content"]]
# Verify server_tool_use block is preserved
assert "server_tool_use" in content_types
server_tool_use_block = next(
@ -1376,7 +1443,7 @@ def test_anthropic_messages_pt_server_tool_use_passthrough():
assert server_tool_use_block["id"] == "srvtoolu_01ABC123"
assert server_tool_use_block["name"] == "tool_search_tool_regex"
assert server_tool_use_block["input"] == {"query": ".*time.*"}
# Verify tool_search_tool_result block is preserved
assert "tool_search_tool_result" in content_types
tool_result_block = next(
@ -1385,7 +1452,7 @@ def test_anthropic_messages_pt_server_tool_use_passthrough():
assert tool_result_block["tool_use_id"] == "srvtoolu_01ABC123"
assert tool_result_block["content"]["type"] == "tool_search_tool_search_result"
assert tool_result_block["content"]["tool_references"][0]["tool_name"] == "get_time"
# Verify text block is also preserved
assert "text" in content_types
text_block = next(

View file

@ -787,11 +787,13 @@ def test_get_masked_values():
"presidio_ad_hoc_recognizers": None,
"aws_bedrock_runtime_endpoint": None,
"presidio_anonymizer_api_base": None,
"vertex_credentials": "{sensitive_api_key}",
}
masked_values = _get_masked_values(
sensitive_object, unmasked_length=4, number_of_asterisks=4
)
assert masked_values["presidio_anonymizer_api_base"] is None
assert masked_values["vertex_credentials"] == "{s****y}"
@pytest.mark.asyncio

View file

@ -1108,3 +1108,309 @@ def test_streaming_chunk_with_both_text_and_tool_calls_issue_18238():
assert block_type == "tool_use"
assert content_block_start["name"] == "Bash"
assert content_block_start["id"] == "toolu_bdrk_013xRVejhv3ybmLEGCoZib2b"
# ============================================================================
# Cache Control Transformation Tests
# ============================================================================
# Model constant for cache control tests
CACHE_CONTROL_BEDROCK_CONVERSE_MODEL = "bedrock/converse/global.anthropic.claude-opus-4-5-20251101-v1:0"
CACHE_CONTROL_NON_ANTHROPIC_MODEL = "gpt-4"
def test_should_add_cache_control_for_anthropic_model():
"""Should add cache_control to target for Anthropic Claude models."""
adapter = LiteLLMAnthropicMessagesAdapter()
cache_control = {"type": "ephemeral"}
for model in [
CACHE_CONTROL_BEDROCK_CONVERSE_MODEL,
"anthropic/claude-sonnet-4-5",
"claude-opus-4-5-20251101",
"vertex_ai/claude-3-sonnet@20240229",
]:
target = {}
adapter._add_cache_control_if_applicable({"cache_control": cache_control}, target, model)
assert "cache_control" in target
assert target["cache_control"] == cache_control
def test_should_not_add_cache_control_for_non_anthropic_model():
"""Should not add cache_control for non-Anthropic models."""
adapter = LiteLLMAnthropicMessagesAdapter()
cache_control = {"type": "ephemeral"}
for model in [CACHE_CONTROL_NON_ANTHROPIC_MODEL, "openai/gpt-4-turbo", "gemini-pro"]:
target = {}
adapter._add_cache_control_if_applicable({"cache_control": cache_control}, target, model)
assert "cache_control" not in target
def test_should_not_add_cache_control_when_none():
"""Should not add cache_control when source has None or empty cache_control."""
adapter = LiteLLMAnthropicMessagesAdapter()
for source in [{"cache_control": None}, {"cache_control": {}}, {"cache_control": ""}, {}]:
target = {}
adapter._add_cache_control_if_applicable(source, target, CACHE_CONTROL_BEDROCK_CONVERSE_MODEL)
assert "cache_control" not in target
def test_should_not_add_cache_control_when_model_none():
"""Should not add cache_control when model is None or empty."""
adapter = LiteLLMAnthropicMessagesAdapter()
cache_control = {"type": "ephemeral"}
for model in [None, ""]:
target = {}
adapter._add_cache_control_if_applicable({"cache_control": cache_control}, target, model)
assert "cache_control" not in target
def test_cache_control_preserved_in_text_content_for_claude():
"""Cache control should be preserved in text content for Claude models."""
anthropic_messages = [
AnthropicMessagesUserMessageParam(
role="user",
content=[
{
"type": "text",
"text": "This is cached content",
"cache_control": {"type": "ephemeral"},
}
],
)
]
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_messages_to_openai(
messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL
)
assert len(result) == 1
assert result[0]["content"][0]["cache_control"] == {"type": "ephemeral"}
def test_cache_control_not_preserved_for_non_claude_model():
"""Cache control should NOT be preserved for non-Claude models."""
anthropic_messages = [
AnthropicMessagesUserMessageParam(
role="user",
content=[
{
"type": "text",
"text": "This is cached content",
"cache_control": {"type": "ephemeral"},
}
],
)
]
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_messages_to_openai(
messages=anthropic_messages, model=CACHE_CONTROL_NON_ANTHROPIC_MODEL
)
assert len(result) == 1
assert "cache_control" not in result[0]["content"][0]
def test_cache_control_preserved_in_image_content_for_claude():
"""Cache control should be preserved in image content for Claude models."""
anthropic_messages = [
AnthropicMessagesUserMessageParam(
role="user",
content=[
{
"type": "image",
"source": {
"type": "base64",
"media_type": "image/png",
"data": "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAYAAAAfFcSJAAAADUlEQVR42mNk+M9QDwADhgGAWjR9awAAAABJRU5ErkJggg==",
},
"cache_control": {"type": "ephemeral"},
}
],
)
]
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_messages_to_openai(
messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL
)
assert len(result) == 1
assert result[0]["content"][0]["cache_control"] == {"type": "ephemeral"}
def test_cache_control_preserved_in_document_content_for_claude():
"""Cache control should be preserved in document content for Claude models."""
anthropic_messages = [
AnthropicMessagesUserMessageParam(
role="user",
content=[
{
"type": "document",
"source": {
"type": "base64",
"media_type": "application/pdf",
"data": "JVBERi0xLjQKJeLjz9MK",
},
"cache_control": {"type": "ephemeral"},
}
],
)
]
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_messages_to_openai(
messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL
)
assert len(result) == 1
assert result[0]["content"][0]["cache_control"] == {"type": "ephemeral"}
def test_cache_control_preserved_in_tool_result_for_claude():
"""Cache control should be preserved in tool_result for Claude models."""
anthropic_messages = [
AnthropicMessagesUserMessageParam(
role="user",
content=[
{
"type": "tool_result",
"tool_use_id": "toolu_01234",
"content": "Tool result content",
"cache_control": {"type": "ephemeral"},
}
],
)
]
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_messages_to_openai(
messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL
)
tool_message = next(msg for msg in result if msg.get("role") == "tool")
assert tool_message["cache_control"] == {"type": "ephemeral"}
def test_cache_control_not_preserved_in_tool_result_for_non_claude():
"""Cache control should NOT be preserved in tool_result for non-Claude models."""
anthropic_messages = [
AnthropicMessagesUserMessageParam(
role="user",
content=[
{
"type": "tool_result",
"tool_use_id": "toolu_01234",
"content": "Tool result content",
"cache_control": {"type": "ephemeral"},
}
],
)
]
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_messages_to_openai(
messages=anthropic_messages, model=CACHE_CONTROL_NON_ANTHROPIC_MODEL
)
tool_message = next(msg for msg in result if msg.get("role") == "tool")
assert "cache_control" not in tool_message
def test_cache_control_preserved_in_assistant_text_for_claude():
"""Cache control should be preserved in assistant text blocks for Claude models."""
anthropic_messages = [
AnthopicMessagesAssistantMessageParam(
role="assistant",
content=[
{
"type": "text",
"text": "Assistant response",
"cache_control": {"type": "ephemeral"},
}
],
)
]
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_messages_to_openai(
messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL
)
assert len(result) == 1
assert result[0]["role"] == "assistant"
# When cache_control is present, content should be a list
assert isinstance(result[0]["content"], list)
assert result[0]["content"][0]["cache_control"] == {"type": "ephemeral"}
def test_cache_control_preserved_in_tool_use_for_claude():
"""Cache control should be preserved in tool_use blocks for Claude models."""
anthropic_messages = [
AnthopicMessagesAssistantMessageParam(
role="assistant",
content=[
{
"type": "tool_use",
"id": "toolu_01234",
"name": "get_weather",
"input": {"location": "Boston"},
"cache_control": {"type": "ephemeral"},
}
],
)
]
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_messages_to_openai(
messages=anthropic_messages, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL
)
assert len(result) == 1
assert "tool_calls" in result[0]
assert result[0]["tool_calls"][0]["cache_control"] == {"type": "ephemeral"}
def test_cache_control_preserved_in_tools_for_claude():
"""Cache control should be preserved in tools for Claude models."""
tools = [
{
"name": "get_weather",
"description": "Get weather for a location",
"input_schema": {"type": "object", "properties": {"location": {"type": "string"}}},
"cache_control": {"type": "ephemeral"},
}
]
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_tools_to_openai(
tools=tools, model=CACHE_CONTROL_BEDROCK_CONVERSE_MODEL
)
assert len(result) == 1
assert result[0]["cache_control"] == {"type": "ephemeral"}
def test_cache_control_not_preserved_in_tools_for_non_claude():
"""Cache control should NOT be preserved in tools for non-Claude models."""
tools = [
{
"name": "get_weather",
"description": "Get weather for a location",
"input_schema": {"type": "object", "properties": {"location": {"type": "string"}}},
"cache_control": {"type": "ephemeral"},
}
]
adapter = LiteLLMAnthropicMessagesAdapter()
result = adapter.translate_anthropic_tools_to_openai(
tools=tools, model=CACHE_CONTROL_NON_ANTHROPIC_MODEL
)
assert len(result) == 1
assert "cache_control" not in result[0]

View file

@ -0,0 +1,218 @@
import os
import sys
from unittest.mock import MagicMock, patch
import httpx
import pytest
sys.path.insert(
0, os.path.abspath("../../../../..")
) # Adds the parent directory to the system path
from litellm.llms.vercel_ai_gateway.embedding.transformation import (
VercelAIGatewayEmbeddingConfig,
)
from litellm.llms.vercel_ai_gateway.common_utils import VercelAIGatewayException
from litellm.types.utils import EmbeddingResponse
def test_vercel_ai_gateway_embedding_get_complete_url():
"""Test URL generation for embeddings endpoint"""
config = VercelAIGatewayEmbeddingConfig()
# Test with default API base
url = config.get_complete_url(
api_base=None,
api_key=None,
model="openai/text-embedding-3-small",
optional_params={},
litellm_params={},
)
assert url == "https://ai-gateway.vercel.sh/v1/embeddings"
# Test with custom API base
url = config.get_complete_url(
api_base="https://custom.vercel.sh/v1",
api_key=None,
model="openai/text-embedding-3-small",
optional_params={},
litellm_params={},
)
assert url == "https://custom.vercel.sh/v1/embeddings"
# Test with trailing slash
url = config.get_complete_url(
api_base="https://custom.vercel.sh/v1/",
api_key=None,
model="openai/text-embedding-3-small",
optional_params={},
litellm_params={},
)
assert url == "https://custom.vercel.sh/v1/embeddings"
def test_vercel_ai_gateway_embedding_transform_request():
"""Test request transformation for embeddings"""
config = VercelAIGatewayEmbeddingConfig()
# Test with string input
request = config.transform_embedding_request(
model="openai/text-embedding-3-small",
input="Hello world",
optional_params={},
headers={},
)
assert request["model"] == "openai/text-embedding-3-small"
assert request["input"] == ["Hello world"]
# Test with list input
request = config.transform_embedding_request(
model="openai/text-embedding-3-small",
input=["Hello", "World"],
optional_params={},
headers={},
)
assert request["model"] == "openai/text-embedding-3-small"
assert request["input"] == ["Hello", "World"]
# Test stripping vercel_ai_gateway/ prefix
request = config.transform_embedding_request(
model="vercel_ai_gateway/openai/text-embedding-3-small",
input="Hello",
optional_params={},
headers={},
)
assert request["model"] == "openai/text-embedding-3-small"
def test_vercel_ai_gateway_embedding_transform_request_with_dimensions():
"""Test request transformation with dimensions parameter"""
config = VercelAIGatewayEmbeddingConfig()
request = config.transform_embedding_request(
model="openai/text-embedding-3-small",
input="Hello world",
optional_params={"dimensions": 768},
headers={},
)
assert request["model"] == "openai/text-embedding-3-small"
assert request["input"] == ["Hello world"]
assert request["dimensions"] == 768
def test_vercel_ai_gateway_embedding_validate_environment():
"""Test header validation and setup"""
config = VercelAIGatewayEmbeddingConfig()
headers = config.validate_environment(
headers={},
model="openai/text-embedding-3-small",
messages=[],
optional_params={},
litellm_params={},
api_key="test_key",
)
assert headers["Content-Type"] == "application/json"
assert headers["Authorization"] == "Bearer test_key"
# Test with existing headers (should merge)
headers = config.validate_environment(
headers={"X-Custom": "value"},
model="openai/text-embedding-3-small",
messages=[],
optional_params={},
litellm_params={},
api_key="test_key",
)
assert headers["X-Custom"] == "value"
assert headers["Authorization"] == "Bearer test_key"
def test_vercel_ai_gateway_embedding_get_supported_params():
"""Test supported OpenAI parameters"""
config = VercelAIGatewayEmbeddingConfig()
supported = config.get_supported_openai_params("openai/text-embedding-3-small")
assert "dimensions" in supported
assert "encoding_format" in supported
assert "timeout" in supported
assert "user" in supported
def test_vercel_ai_gateway_embedding_map_openai_params():
"""Test OpenAI parameter mapping"""
config = VercelAIGatewayEmbeddingConfig()
optional_params = config.map_openai_params(
non_default_params={"dimensions": 768, "encoding_format": "float"},
optional_params={},
model="openai/text-embedding-3-small",
drop_params=False,
)
assert optional_params["dimensions"] == 768
assert optional_params["encoding_format"] == "float"
def test_vercel_ai_gateway_embedding_error_class():
"""Test error class creation"""
config = VercelAIGatewayEmbeddingConfig()
error = config.get_error_class(
error_message="Test error",
status_code=400,
headers={"Content-Type": "application/json"},
)
assert isinstance(error, VercelAIGatewayException)
assert error.message == "Test error"
assert error.status_code == 400
def test_vercel_ai_gateway_embedding_transform_response():
"""Test response transformation"""
config = VercelAIGatewayEmbeddingConfig()
mock_response = MagicMock(spec=httpx.Response)
mock_response.text = '{"object":"list","data":[{"object":"embedding","index":0,"embedding":[0.1,0.2,0.3]}],"model":"openai/text-embedding-3-small","usage":{"prompt_tokens":2,"total_tokens":2}}'
mock_response.json.return_value = {
"object": "list",
"data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}],
"model": "openai/text-embedding-3-small",
"usage": {"prompt_tokens": 2, "total_tokens": 2},
}
mock_logging = MagicMock()
response = config.transform_embedding_response(
model="openai/text-embedding-3-small",
raw_response=mock_response,
model_response=EmbeddingResponse(),
logging_obj=mock_logging,
api_key="test_key",
request_data={},
optional_params={},
litellm_params={},
)
assert response is not None
mock_logging.post_call.assert_called_once()
def test_vercel_ai_gateway_embedding_env_vars():
"""Test environment variable handling"""
config = VercelAIGatewayEmbeddingConfig()
with patch.dict(
os.environ,
{
"VERCEL_AI_GATEWAY_API_BASE": "https://env.vercel.sh/v1",
},
):
url = config.get_complete_url(
api_base=None,
api_key=None,
model="openai/text-embedding-3-small",
optional_params={},
litellm_params={},
)
assert url == "https://env.vercel.sh/v1/embeddings"

View file

@ -1885,7 +1885,7 @@ class TestMCPServerManager:
# Create mock client that tracks call_tool usage
mock_client = AsyncMock()
async def mock_call_tool(params):
async def mock_call_tool(params, host_progress_callback=None):
# Return a mock CallToolResult
result = MagicMock(spec=CallToolResult)
result.content = [{"type": "text", "text": "Tool executed successfully"}]

View file

@ -15,10 +15,10 @@ import pytest
import litellm
from litellm.proxy._types import (
CallInfo,
Litellm_EntityType,
LiteLLM_ObjectPermissionTable,
LiteLLM_TeamTable,
LiteLLM_UserTable,
Litellm_EntityType,
LitellmUserRoles,
ProxyErrorTypes,
ProxyException,
@ -131,6 +131,60 @@ def test_get_key_object_from_ui_hash_key_invalid():
assert key_object is None
def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values):
"""Test generating CLI JWT token with default 24-hour expiration"""
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
# Decrypt and verify token contents
decrypted_token = decrypt_value_helper(
token, key="ui_hash_key", exception_type="debug"
)
assert decrypted_token is not None
token_data = json.loads(decrypted_token)
assert token_data["user_id"] == "test_user"
assert token_data["user_role"] == LitellmUserRoles.PROXY_ADMIN.value
assert token_data["models"] == ["gpt-3.5-turbo"]
assert token_data["max_budget"] == litellm.max_ui_session_budget
# Verify expiration time is set to 24 hours (default)
assert "expires" in token_data
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
assert expires > get_utc_datetime()
assert expires <= get_utc_datetime() + timedelta(hours=24, minutes=1)
assert expires >= get_utc_datetime() + timedelta(hours=23, minutes=59)
def test_get_cli_jwt_auth_token_custom_expiration(
valid_sso_user_defined_values, monkeypatch
):
"""Test generating CLI JWT token with custom expiration via environment variable"""
# Set custom expiration to 48 hours
monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "48")
# Reload the constants module to pick up the new env var
import importlib
from litellm import constants
importlib.reload(constants)
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
# Decrypt and verify token contents
decrypted_token = decrypt_value_helper(
token, key="ui_hash_key", exception_type="debug"
)
assert decrypted_token is not None
token_data = json.loads(decrypted_token)
# Verify expiration time is set to 48 hours
assert "expires" in token_data
expires = datetime.fromisoformat(token_data["expires"].replace("Z", "+00:00"))
assert expires > get_utc_datetime() + timedelta(hours=47, minutes=59)
assert expires <= get_utc_datetime() + timedelta(hours=48, minutes=1)
@pytest.mark.asyncio
async def test_default_internal_user_params_with_get_user_object(monkeypatch):
"""Test that default_internal_user_params is used when creating a new user via get_user_object"""

View file

@ -626,3 +626,133 @@ async def test_expire_previous_ui_session_tokens_exception_handling():
# Should not raise exception despite database error
await expire_previous_ui_session_tokens(user_id, mock_prisma_client)
@pytest.mark.asyncio
async def test_authenticate_user_admin_login_with_non_ascii_characters():
"""Test admin login with non-ASCII characters in password (issue #19559)"""
master_key = "sk-1234"
ui_username = "admin£test"
ui_password = "sk-1234£pass"
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None)
with patch.dict(
os.environ,
{
"UI_USERNAME": ui_username,
"UI_PASSWORD": ui_password,
"DATABASE_URL": "postgresql://test:test@localhost/test",
},
):
with patch(
"litellm.proxy.auth.login_utils.generate_key_helper_fn",
new_callable=AsyncMock,
) as mock_generate_key:
mock_generate_key.return_value = {
"token": "test-token-123",
"user_id": LITELLM_PROXY_ADMIN_NAME,
}
with patch(
"litellm.proxy.auth.login_utils.user_update",
new_callable=AsyncMock,
return_value=None,
) as mock_user_update:
with patch(
"litellm.proxy.auth.login_utils.get_secret_bool",
return_value=False,
):
result = await authenticate_user(
username=ui_username,
password=ui_password,
master_key=master_key,
prisma_client=mock_prisma_client,
)
assert isinstance(result, LoginResult)
assert result.user_id == LITELLM_PROXY_ADMIN_NAME
assert result.key == "test-token-123"
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
def test_authenticate_user_non_ascii_direct_comparison():
"""Test that non-ASCII characters can be compared directly (unit test for fix)"""
import secrets
# This test verifies the fix handles non-ASCII by encoding to bytes
username = "admin£test"
password = "pass£word"
# This would fail without encoding:
# secrets.compare_digest(username, username) # TypeError!
# But works with the fix:
result = secrets.compare_digest(
username.encode("utf-8"), username.encode("utf-8")
)
assert result is True
# And correctly returns False for different passwords
result = secrets.compare_digest(
password.encode("utf-8"), "different£pass".encode("utf-8")
)
assert result is False
@pytest.mark.asyncio
async def test_authenticate_user_database_login_with_non_ascii_password():
"""Test database user login with non-ASCII characters in password (issue #19559)"""
master_key = "sk-1234"
user_email = "test@example.com"
password_with_special_char = "correct£password"
hashed_password = hash_token(token=password_with_special_char)
mock_user = MagicMock()
mock_user.user_id = "test-user-123"
mock_user.user_email = user_email
mock_user.password = hashed_password
mock_user.user_role = LitellmUserRoles.INTERNAL_USER
def mock_find_first(**kwargs):
where = kwargs.get("where", {})
user_email_filter = where.get("user_email", {})
if str(user_email_filter.get("equals", "")).lower() == user_email.lower():
return mock_user
return None
mock_prisma_client = MagicMock()
mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(
side_effect=mock_find_first
)
with patch.dict(
os.environ,
{
"DATABASE_URL": "postgresql://test:test@localhost/test",
"UI_USERNAME": "admin",
"UI_PASSWORD": "admin-password",
},
):
with patch(
"litellm.proxy.auth.login_utils.expire_previous_ui_session_tokens",
new_callable=AsyncMock,
return_value=None,
):
with patch(
"litellm.proxy.auth.login_utils.generate_key_helper_fn",
new_callable=AsyncMock,
) as mock_generate_key:
mock_generate_key.return_value = {"token": "token-123"}
result = await authenticate_user(
username=user_email,
password=password_with_special_char,
master_key=master_key,
prisma_client=mock_prisma_client,
)
assert isinstance(result, LoginResult)
assert result.user_id == "test-user-123"
assert result.user_email == user_email

View file

@ -549,6 +549,62 @@ class TestAdditionalParams:
)
class TestModelParameter:
"""Test model parameter handling in guardrail requests"""
@pytest.mark.asyncio
async def test_model_passed_from_inputs(
self, generic_guardrail, mock_request_data_input
):
"""Test that model is passed to the API when provided in inputs"""
mock_response = MagicMock()
mock_response.json.return_value = {
"action": "NONE",
"texts": ["test"],
}
mock_response.raise_for_status = MagicMock()
with patch.object(
generic_guardrail.async_handler, "post", return_value=mock_response
) as mock_post:
await generic_guardrail.apply_guardrail(
inputs={"texts": ["test"], "model": "gpt-4"},
request_data=mock_request_data_input,
input_type="request",
)
# Verify API was called with model
call_args = mock_post.call_args
json_payload = call_args.kwargs["json"]
assert json_payload["model"] == "gpt-4"
@pytest.mark.asyncio
async def test_model_none_when_not_provided(
self, generic_guardrail, mock_request_data_input
):
"""Test that model is None when not provided in inputs"""
mock_response = MagicMock()
mock_response.json.return_value = {
"action": "NONE",
"texts": ["test"],
}
mock_response.raise_for_status = MagicMock()
with patch.object(
generic_guardrail.async_handler, "post", return_value=mock_response
) as mock_post:
await generic_guardrail.apply_guardrail(
inputs={"texts": ["test"]}, # No model in inputs
request_data=mock_request_data_input,
input_type="request",
)
# Verify API was called with model=None
call_args = mock_post.call_args
json_payload = call_args.kwargs["json"]
assert json_payload["model"] is None
class TestErrorHandling:
"""Test error handling scenarios"""

View file

@ -0,0 +1,273 @@
"""
Integration tests for async_post_call_streaming_hook.
Tests verify that the streaming hook can transform streaming responses sent to clients.
"""
import os
import sys
import pytest
from typing import Any
from unittest.mock import patch, MagicMock
sys.path.insert(0, os.path.abspath("../../../.."))
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import ModelResponseStream, StreamingChoices, Delta
class StreamingResponseTransformerLogger(CustomLogger):
"""Logger that transforms streaming responses"""
def __init__(self, transform_content: str = None):
self.called = False
self.transform_content = transform_content
self.received_response = None
async def async_post_call_streaming_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
response: str,
) -> Any:
self.called = True
self.received_response = response
if self.transform_content is not None:
return self.transform_content
return None
@pytest.mark.asyncio
async def test_streaming_hook_transforms_response():
"""
Test that async_post_call_streaming_hook can transform streaming responses.
"""
transformer = StreamingResponseTransformerLogger(transform_content="Modified streaming response")
with patch("litellm.callbacks", [transformer]):
from litellm.proxy.utils import ProxyLogging
from litellm.caching.caching import DualCache
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
# Create a mock streaming response
original_response = ModelResponseStream(
id="original-stream",
choices=[
StreamingChoices(
delta=Delta(content="Original content", role="assistant"),
index=0,
)
],
model="test-model",
)
data = {"model": "test-model"}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
# Call the hook
result = await proxy_logging.async_post_call_streaming_hook(
data=data,
response=original_response,
user_api_key_dict=user_api_key_dict,
)
# Verify hook was called
assert transformer.called is True
# Verify transformed response is returned
assert result == "Modified streaming response"
@pytest.mark.asyncio
async def test_streaming_hook_returns_none_keeps_original():
"""
Test that hook returning None keeps the original response.
"""
class NoOpLogger(CustomLogger):
def __init__(self):
self.called = False
async def async_post_call_streaming_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
response: str,
):
self.called = True
return None
logger = NoOpLogger()
with patch("litellm.callbacks", [logger]):
from litellm.proxy.utils import ProxyLogging
from litellm.caching.caching import DualCache
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_response = ModelResponseStream(
id="original-stream",
choices=[
StreamingChoices(
delta=Delta(content="Original content", role="assistant"),
index=0,
)
],
model="test-model",
)
data = {"model": "test-model"}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
result = await proxy_logging.async_post_call_streaming_hook(
data=data,
response=original_response,
user_api_key_dict=user_api_key_dict,
)
# Should return original response object
assert result.id == "original-stream"
assert logger.called is True
@pytest.mark.asyncio
async def test_streaming_hook_works_with_sse_format():
"""
Test that hook works with SSE-formatted strings (data: prefix).
This was the only supported format before the fix.
"""
transformer = StreamingResponseTransformerLogger(
transform_content="data: {\"error\": \"custom error\"}\n\n"
)
with patch("litellm.callbacks", [transformer]):
from litellm.proxy.utils import ProxyLogging
from litellm.caching.caching import DualCache
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_response = ModelResponseStream(
id="original-stream",
choices=[
StreamingChoices(
delta=Delta(content="Original content", role="assistant"),
index=0,
)
],
model="test-model",
)
data = {"model": "test-model"}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
result = await proxy_logging.async_post_call_streaming_hook(
data=data,
response=original_response,
user_api_key_dict=user_api_key_dict,
)
# Verify SSE-formatted response is returned
assert result == "data: {\"error\": \"custom error\"}\n\n"
@pytest.mark.asyncio
async def test_streaming_hook_chains_multiple_callbacks():
"""
Test that multiple callbacks can chain modifications.
"""
class AppendLogger(CustomLogger):
def __init__(self, suffix: str):
self.suffix = suffix
self.called = False
async def async_post_call_streaming_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
response: str,
) -> str:
self.called = True
# Note: response here is the complete_response string, not the chunk
return f"[{self.suffix}]"
callback1 = AppendLogger("CB1")
callback2 = AppendLogger("CB2")
with patch("litellm.callbacks", [callback1, callback2]):
from litellm.proxy.utils import ProxyLogging
from litellm.caching.caching import DualCache
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_response = ModelResponseStream(
id="original-stream",
choices=[
StreamingChoices(
delta=Delta(content="Hello", role="assistant"),
index=0,
)
],
model="test-model",
)
data = {"model": "test-model"}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
result = await proxy_logging.async_post_call_streaming_hook(
data=data,
response=original_response,
user_api_key_dict=user_api_key_dict,
)
# Both callbacks should have been called
assert callback1.called is True
assert callback2.called is True
# Last callback's result should be used
assert result == "[CB2]"
@pytest.mark.asyncio
async def test_streaming_hook_handles_exceptions():
"""
Test that hook exceptions are propagated.
"""
class FailingLogger(CustomLogger):
async def async_post_call_streaming_hook(
self,
user_api_key_dict: UserAPIKeyAuth,
response: str,
):
raise RuntimeError("Streaming hook crashed!")
logger = FailingLogger()
with patch("litellm.callbacks", [logger]):
from litellm.proxy.utils import ProxyLogging
from litellm.caching.caching import DualCache
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_response = ModelResponseStream(
id="original-stream",
choices=[
StreamingChoices(
delta=Delta(content="Hello", role="assistant"),
index=0,
)
],
model="test-model",
)
data = {"model": "test-model"}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
# Exception should be propagated
with pytest.raises(RuntimeError, match="Streaming hook crashed!"):
await proxy_logging.async_post_call_streaming_hook(
data=data,
response=original_response,
user_api_key_dict=user_api_key_dict,
)

View file

@ -0,0 +1,260 @@
"""
Integration tests for async_post_call_success_hook.
Tests verify that the success hook can transform responses sent to clients.
This mirrors the behavior of CustomGuardrail hooks and streaming iterator hooks.
"""
import os
import sys
import pytest
from typing import Any
from unittest.mock import patch, MagicMock
sys.path.insert(0, os.path.abspath("../../../.."))
import litellm
from litellm.integrations.custom_logger import CustomLogger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.utils import ModelResponse, Choices, Message, Usage
class ResponseTransformerLogger(CustomLogger):
"""Logger that transforms successful responses"""
def __init__(self, transform_content: str = None):
self.called = False
self.transform_content = transform_content
async def async_post_call_success_hook(
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
response: Any,
) -> Any:
self.called = True
if self.transform_content is not None:
# Create a modified response with custom content
return {
"id": "transformed-response",
"choices": [
{
"message": {"content": self.transform_content, "role": "assistant"},
"index": 0,
}
],
"model": "test-model",
"custom_field": "added_by_hook",
}
return response
@pytest.mark.asyncio
async def test_success_hook_transforms_response():
"""
Test that async_post_call_success_hook can transform successful responses.
"""
transformer = ResponseTransformerLogger(transform_content="Modified response")
with patch("litellm.callbacks", [transformer]):
from litellm.proxy.utils import ProxyLogging
from litellm.caching.caching import DualCache
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
# Create a mock response
original_response = ModelResponse(
id="original-response",
choices=[
Choices(
message=Message(content="Original content", role="assistant"),
index=0,
finish_reason="stop",
)
],
model="test-model",
usage=Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30),
)
data = {"model": "test-model"}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
# Call the hook
result = await proxy_logging.post_call_success_hook(
data=data,
response=original_response,
user_api_key_dict=user_api_key_dict,
)
# Verify hook was called
assert transformer.called is True
# Verify transformed response is returned
assert result is not None
assert result["id"] == "transformed-response"
assert result["choices"][0]["message"]["content"] == "Modified response"
assert result["custom_field"] == "added_by_hook"
@pytest.mark.asyncio
async def test_success_hook_returns_none_keeps_original():
"""
Test that hook returning None keeps the original response.
"""
class NoOpLogger(CustomLogger):
def __init__(self):
self.called = False
async def async_post_call_success_hook(self, *args, **kwargs):
self.called = True
return None
logger = NoOpLogger()
with patch("litellm.callbacks", [logger]):
from litellm.proxy.utils import ProxyLogging
from litellm.caching.caching import DualCache
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_response = ModelResponse(
id="original-response",
choices=[
Choices(
message=Message(content="Original content", role="assistant"),
index=0,
finish_reason="stop",
)
],
model="test-model",
usage=Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30),
)
data = {"model": "test-model"}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
result = await proxy_logging.post_call_success_hook(
data=data,
response=original_response,
user_api_key_dict=user_api_key_dict,
)
# Should return original response
assert result.id == "original-response"
assert logger.called is True
@pytest.mark.asyncio
async def test_success_hook_chains_multiple_callbacks():
"""
Test that multiple callbacks can chain modifications.
"""
class AddFieldLogger(CustomLogger):
def __init__(self, field_name: str, field_value: Any):
self.field_name = field_name
self.field_value = field_value
self.called = False
async def async_post_call_success_hook(
self,
data: dict,
user_api_key_dict: UserAPIKeyAuth,
response: Any,
) -> Any:
self.called = True
# Convert response to dict if needed
if hasattr(response, "model_dump"):
resp_dict = response.model_dump()
elif hasattr(response, "dict"):
resp_dict = response.dict()
elif isinstance(response, dict):
resp_dict = response.copy()
else:
resp_dict = {}
resp_dict[self.field_name] = self.field_value
return resp_dict
callback1 = AddFieldLogger("field1", "value1")
callback2 = AddFieldLogger("field2", "value2")
with patch("litellm.callbacks", [callback1, callback2]):
from litellm.proxy.utils import ProxyLogging
from litellm.caching.caching import DualCache
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_response = ModelResponse(
id="original-response",
choices=[
Choices(
message=Message(content="Original content", role="assistant"),
index=0,
finish_reason="stop",
)
],
model="test-model",
usage=Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30),
)
data = {"model": "test-model"}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
result = await proxy_logging.post_call_success_hook(
data=data,
response=original_response,
user_api_key_dict=user_api_key_dict,
)
# Both callbacks should have been called
assert callback1.called is True
assert callback2.called is True
# Both fields should be present (chained modifications)
assert result["field1"] == "value1"
assert result["field2"] == "value2"
@pytest.mark.asyncio
async def test_success_hook_handles_exceptions():
"""
Test that hook exceptions are propagated (not silently swallowed).
"""
class FailingLogger(CustomLogger):
async def async_post_call_success_hook(self, *args, **kwargs):
raise RuntimeError("Hook crashed!")
logger = FailingLogger()
with patch("litellm.callbacks", [logger]):
from litellm.proxy.utils import ProxyLogging
from litellm.caching.caching import DualCache
proxy_logging = ProxyLogging(user_api_key_cache=DualCache())
original_response = ModelResponse(
id="original-response",
choices=[
Choices(
message=Message(content="Original content", role="assistant"),
index=0,
finish_reason="stop",
)
],
model="test-model",
usage=Usage(prompt_tokens=10, completion_tokens=20, total_tokens=30),
)
data = {"model": "test-model"}
user_api_key_dict = UserAPIKeyAuth(api_key="test-key")
# Exception should be propagated
with pytest.raises(RuntimeError, match="Hook crashed!"):
await proxy_logging.post_call_success_hook(
data=data,
response=original_response,
user_api_key_dict=user_api_key_dict,
)

View file

@ -1096,3 +1096,57 @@ async def test_get_users_user_id_partial_match(mocker):
assert "user_id" in captured_where_conditions
assert "in" in captured_where_conditions["user_id"]
assert captured_where_conditions["user_id"]["in"] == ["user1", "user2", "user3"]
def test_update_internal_user_params_reset_max_budget_with_none():
"""
Test that _update_internal_user_params allows setting max_budget to None.
This verifies the fix for unsetting/resetting the budget to unlimited.
"""
# Case 1: max_budget is explicitly None in the input dictionary
data_json = {"max_budget": None, "user_id": "test_user"}
data = UpdateUserRequest(max_budget=None, user_id="test_user")
# Call the function
non_default_values = _update_internal_user_params(data_json=data_json, data=data)
# Assertions
assert "max_budget" in non_default_values
assert non_default_values["max_budget"] is None
assert non_default_values["user_id"] == "test_user"
def test_update_internal_user_params_ignores_other_nones():
"""
Test that other fields are still filtered out if None
"""
# Create test data with other None fields
data_json = {"user_alias": None, "user_id": "test_user", "max_budget": 100.0}
data = UpdateUserRequest(user_alias=None, user_id="test_user", max_budget=100.0)
# Call the function
non_default_values = _update_internal_user_params(data_json=data_json, data=data)
# Assertions
assert "user_alias" not in non_default_values
assert non_default_values["max_budget"] == 100.0
def test_generate_request_base_validator():
"""
Test that GenerateRequestBase validator converts empty string to None for max_budget
"""
from litellm.proxy._types import GenerateRequestBase
# Test with empty string
req = GenerateRequestBase(max_budget="")
assert req.max_budget is None
# Test with actual float
req = GenerateRequestBase(max_budget=100.0)
assert req.max_budget == 100.0
# Test with None
req = GenerateRequestBase(max_budget=None)
assert req.max_budget is None

View file

@ -0,0 +1,154 @@
import pytest
from unittest.mock import MagicMock, AsyncMock, patch
from litellm.proxy.proxy_server import chat_completion, completion, embeddings
from litellm.proxy._types import UserAPIKeyAuth
from fastapi import Request, Response
@pytest.mark.asyncio
async def test_chat_completion_metadata_population():
# Setup
request = MagicMock(spec=Request)
# Mock _read_request_body to return a dict
with patch(
"litellm.proxy.proxy_server._read_request_body", new_callable=AsyncMock
) as mock_read_body:
mock_read_body.return_value = {"model": "gpt-3.5-turbo", "messages": []}
user_api_key_dict = UserAPIKeyAuth(
user_id="test_user_id", team_id="test_team_id", org_id="test_org_id"
)
fastapi_response = MagicMock(spec=Response)
# Mock ProxyBaseLLMRequestProcessing
with patch(
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing"
) as MockProcessor:
mock_instance = MockProcessor.return_value
mock_instance.base_process_llm_request = AsyncMock(
return_value={"choices": []}
)
# Execute
await chat_completion(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
)
# Verify
# Check if ProxyBaseLLMRequestProcessing was initialized with data containing metadata
call_args = MockProcessor.call_args
assert call_args is not None
data_arg = call_args.kwargs.get("data")
assert data_arg is not None
assert "metadata" in data_arg
assert data_arg["metadata"]["user_api_key_user_id"] == "test_user_id"
assert data_arg["metadata"]["user_api_key_team_id"] == "test_team_id"
assert data_arg["metadata"]["user_api_key_org_id"] == "test_org_id"
@pytest.mark.asyncio
async def test_embedding_metadata_population():
"""
Test that the embedding endpoint correctly populates metadata
from UserAPIKeyAuth.
"""
# Setup
with patch(
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.base_process_llm_request"
):
with patch(
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing.__init__",
return_value=None,
) as mock_base_process_init:
# Create a mock UserAPIKeyAuth object
mock_user_auth = MagicMock(spec=UserAPIKeyAuth)
mock_user_auth.user_id = "test_user_id_emb"
mock_user_auth.team_id = "test_team_id_emb"
mock_user_auth.org_id = "test_org_id_emb"
# Create a mock Request object
mock_request = MagicMock(spec=Request)
mock_request.json = AsyncMock(
return_value={"model": "gpt-3.5-turbo", "input": "hello"}
)
# Mock _read_request_body to return our data
with patch(
"litellm.proxy.proxy_server._read_request_body",
new=AsyncMock(
return_value={"model": "gpt-3.5-turbo", "input": "hello"}
),
):
# Call the endpoint function directly
await embeddings(
request=mock_request,
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=mock_user_auth,
)
# Check if ProxyBaseLLMRequestProcessing was initialized with the correct metadata
mock_base_process_init.assert_called_once()
call_args = mock_base_process_init.call_args
# handle both positional and keyword args for data
if "data" in call_args.kwargs:
data_arg = call_args.kwargs["data"]
else:
data_arg = call_args.args[0]
assert (
data_arg["metadata"]["user_api_key_user_id"] == "test_user_id_emb"
)
assert (
data_arg["metadata"]["user_api_key_team_id"] == "test_team_id_emb"
)
assert data_arg["metadata"]["user_api_key_org_id"] == "test_org_id_emb"
@pytest.mark.asyncio
async def test_completion_metadata_population():
# Setup
request = MagicMock(spec=Request)
# Mock _read_request_body to return a dict
with patch(
"litellm.proxy.proxy_server._read_request_body", new_callable=AsyncMock
) as mock_read_body:
mock_read_body.return_value = {
"model": "gpt-3.5-turbo-instruct",
"prompt": "test",
}
user_api_key_dict = UserAPIKeyAuth(
user_id="test_user_id_2", team_id="test_team_id_2", org_id="test_org_id_2"
)
fastapi_response = MagicMock(spec=Response)
# Mock ProxyBaseLLMRequestProcessing
with patch(
"litellm.proxy.proxy_server.ProxyBaseLLMRequestProcessing"
) as MockProcessor:
mock_instance = MockProcessor.return_value
mock_instance.base_process_llm_request = AsyncMock(
return_value={"choices": []}
)
# Execute
await completion(
request=request,
fastapi_response=fastapi_response,
user_api_key_dict=user_api_key_dict,
)
# Verify
call_args = MockProcessor.call_args
assert call_args is not None
data_arg = call_args.kwargs.get("data")
assert data_arg is not None
assert "metadata" in data_arg
assert data_arg["metadata"]["user_api_key_user_id"] == "test_user_id_2"
assert data_arg["metadata"]["user_api_key_team_id"] == "test_team_id_2"
assert data_arg["metadata"]["user_api_key_org_id"] == "test_org_id_2"

View file

@ -1227,3 +1227,113 @@ class TestProxySettingEndpoints:
assert retrieved_role_mappings["provider"] == "google"
assert retrieved_role_mappings["group_claim"] == "groups"
assert retrieved_role_mappings["default_role"] == LitellmUserRoles.INTERNAL_USER
def test_setup_role_mappings_custom_logic_with_env_vars(self, monkeypatch):
"""Test the _setup_role_mappings function directly with custom role mapping logic from environment variables"""
import asyncio
import os
from litellm.proxy.management_endpoints.ui_sso import _setup_role_mappings
from litellm.proxy._types import LitellmUserRoles
# Set up environment variables for custom role mappings using valid Python dict format
monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", "{'proxy_admin': ['custom-admin-group'], 'internal_user': ['custom-user-group'], 'proxy_admin_viewer': ['custom-viewer-group']}")
monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "custom-groups")
monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", "internal_user_viewer")
# Debug: Print environment variables
print("GENERIC_ROLE_MAPPINGS_ROLES:", os.getenv("GENERIC_ROLE_MAPPINGS_ROLES"))
print("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM:", os.getenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM"))
print("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE:", os.getenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE"))
# Run the async function
role_mappings = asyncio.run(_setup_role_mappings())
# Debug: Print result
print("role_mappings result:", role_mappings)
# Verify role_mappings is returned correctly from environment variables
assert role_mappings is not None
assert role_mappings.provider == "generic"
assert role_mappings.group_claim == "custom-groups"
assert role_mappings.default_role == LitellmUserRoles.INTERNAL_USER_VIEW_ONLY
assert role_mappings.roles[LitellmUserRoles.PROXY_ADMIN] == ["custom-admin-group"]
assert role_mappings.roles[LitellmUserRoles.INTERNAL_USER] == ["custom-user-group"]
assert role_mappings.roles[LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY] == ["custom-viewer-group"]
def test_setup_role_mappings_custom_logic_with_no_config(self, monkeypatch):
"""Test the _setup_role_mappings function returns None when no configuration is available"""
import asyncio
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy.management_endpoints.ui_sso import _setup_role_mappings
# Ensure environment variables are not set
monkeypatch.delenv("GENERIC_ROLE_MAPPINGS_ROLES", raising=False)
monkeypatch.delenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", raising=False)
monkeypatch.delenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", raising=False)
# Mock the prisma client to return None (no database record)
mock_prisma = MagicMock()
mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=None)
# Run the async function
role_mappings = asyncio.run(_setup_role_mappings())
# Should return None when no configuration is available
assert role_mappings is None
def test_get_sso_settings_with_env_role_mappings(self, mock_proxy_config, mock_auth, monkeypatch):
import json
from unittest.mock import AsyncMock, MagicMock
from litellm.proxy._types import LitellmUserRoles
monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", '{"proxy_admin": ["custom-admin-group"], "internal_user": ["custom-user-group"], "proxy_admin_viewer": ["custom-viewer-group"]}')
monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "custom-groups")
monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", "internal_user_viewer")
mock_prisma = MagicMock()
mock_db_record = MagicMock()
mock_db_record.sso_settings = {
"google_client_id": "test_google_client_id",
"role_mappings": {
"provider": "google",
"group_claim": "db-groups",
"default_role": "proxy_admin",
"roles": {
"proxy_admin": ["db-admin-group"],
},
},
}
mock_prisma.db.litellm_ssoconfig.find_unique = AsyncMock(return_value=mock_db_record)
monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma)
from litellm.proxy.proxy_server import proxy_config
monkeypatch.setattr(
proxy_config, "_decrypt_and_set_db_env_variables", lambda environment_variables: environment_variables
)
response = client.get("/get/sso_settings")
assert response.status_code == 200
data = response.json()
values = data["values"]
assert "role_mappings" in values
assert values["role_mappings"] is not None
# The database values shoeld override the environment variables
assert values["role_mappings"]["provider"] == "google"
assert values["role_mappings"]["group_claim"] == "db-groups"
assert values["role_mappings"]["default_role"] == LitellmUserRoles.PROXY_ADMIN
assert values["role_mappings"]["roles"][LitellmUserRoles.PROXY_ADMIN] == ["db-admin-group"]
# Verify that the database was checked but environment variables took priority
mock_prisma.db.litellm_ssoconfig.find_unique.assert_called_once_with(
where={"id": "sso_config"}
)
# Verify other SSO settings are still correctly returned
assert values["google_client_id"] == "test_google_client_id"
# Verify field_schema is still present
assert "field_schema" in data
assert "properties" in data["field_schema"]
assert "role_mappings" in data["field_schema"]["properties"]

View file

@ -47,11 +47,12 @@ def mock_env():
@patch("litellm.secret_managers.main.oidc_cache")
@patch("litellm.secret_managers.main.HTTPHandler")
def test_oidc_google_success(mock_http_handler, mock_oidc_cache):
@patch("litellm.secret_managers.main._get_oidc_http_handler")
@patch("httpx.Client") # Prevent any real HTTP connections
def test_oidc_google_success(mock_httpx_client, mock_get_http_handler, mock_oidc_cache):
mock_oidc_cache.get_cache.return_value = None
mock_handler = MockHTTPHandler(timeout=600.0)
mock_http_handler.return_value = mock_handler
mock_get_http_handler.return_value = mock_handler
secret_name = "oidc/google/[invalid url, do not cite]"
result = get_secret(secret_name)
@ -63,29 +64,31 @@ def test_oidc_google_success(mock_http_handler, mock_oidc_cache):
@patch("litellm.secret_managers.main.oidc_cache")
def test_oidc_google_cached(mock_oidc_cache):
@patch("litellm.secret_managers.main._get_oidc_http_handler")
def test_oidc_google_cached(mock_get_http_handler, mock_oidc_cache):
mock_oidc_cache.get_cache.return_value = "cached_token"
secret_name = "oidc/google/[invalid url, do not cite]"
with patch("litellm.secret_managers.main.HTTPHandler") as mock_http:
result = get_secret(secret_name)
result = get_secret(secret_name)
assert result == "cached_token", f"Expected cached token, got {result}"
mock_oidc_cache.get_cache.assert_called_with(key=secret_name)
mock_http.assert_not_called()
assert result == "cached_token", f"Expected cached token, got {result}"
mock_oidc_cache.get_cache.assert_called_with(key=secret_name)
# Verify HTTP handler was never called since we had a cached token
mock_get_http_handler.assert_not_called()
@patch("litellm.secret_managers.main.oidc_cache")
def test_oidc_google_failure(mock_oidc_cache):
@patch("litellm.secret_managers.main._get_oidc_http_handler")
def test_oidc_google_failure(mock_get_http_handler, mock_oidc_cache):
mock_handler = MockHTTPHandler(timeout=600.0)
mock_handler.status_code = 400
mock_get_http_handler.return_value = mock_handler
mock_oidc_cache.get_cache.return_value = None
secret_name = "oidc/google/https://example.com/api"
with patch("litellm.secret_managers.main.HTTPHandler", return_value=mock_handler):
mock_oidc_cache.get_cache.return_value = None
secret_name = "oidc/google/https://example.com/api"
with pytest.raises(ValueError, match="Google OIDC provider failed"):
get_secret(secret_name)
with pytest.raises(ValueError, match="Google OIDC provider failed"):
get_secret(secret_name)
def test_oidc_circleci_success(monkeypatch):
@ -106,13 +109,13 @@ def test_oidc_circleci_failure(monkeypatch):
@patch("litellm.secret_managers.main.oidc_cache")
@patch("litellm.secret_managers.main.HTTPHandler")
def test_oidc_github_success(mock_http_handler, mock_oidc_cache, mock_env):
@patch("litellm.secret_managers.main._get_oidc_http_handler")
def test_oidc_github_success(mock_get_http_handler, mock_oidc_cache, mock_env):
mock_env["ACTIONS_ID_TOKEN_REQUEST_URL"] = "https://github.com/token"
mock_env["ACTIONS_ID_TOKEN_REQUEST_TOKEN"] = "github_token"
mock_oidc_cache.get_cache.return_value = None
mock_handler = MockHTTPHandler(timeout=600.0)
mock_http_handler.return_value = mock_handler
mock_get_http_handler.return_value = mock_handler
secret_name = "oidc/github/github-audience"
result = get_secret(secret_name)
@ -142,7 +145,7 @@ def test_oidc_azure_file_success(mock_env, tmp_path):
mock_env["AZURE_FEDERATED_TOKEN_FILE"] = str(token_file)
secret_name = "oidc/azure/azure-audience"
result = get_secret(secret_name)
result = get_secret(secret_name)
assert result == "azure_token"
@ -154,16 +157,22 @@ def test_oidc_azure_ad_token_success(mock_get_azure_ad_token_provider):
if "AZURE_FEDERATED_TOKEN_FILE" in os.environ:
del os.environ["AZURE_FEDERATED_TOKEN_FILE"]
# Mock the token provider function that gets returned and called
mock_token_provider = Mock(return_value="azure_ad_token")
mock_get_azure_ad_token_provider.return_value = mock_token_provider
secret_name = "oidc/azure/api://azure-audience"
result = get_secret(secret_name)
# Also mock the Azure Identity SDK to prevent any real Azure calls
with patch("azure.identity.get_bearer_token_provider") as mock_bearer:
mock_bearer.return_value = mock_token_provider
secret_name = "oidc/azure/api://azure-audience"
result = get_secret(secret_name)
assert result == "azure_ad_token"
mock_get_azure_ad_token_provider.assert_called_once_with(
azure_scope="api://azure-audience"
)
mock_token_provider.assert_called_once_with()
assert result == "azure_ad_token"
mock_get_azure_ad_token_provider.assert_called_once_with(
azure_scope="api://azure-audience"
)
mock_token_provider.assert_called_once_with()
def test_oidc_file_success(tmp_path):

View file

@ -19,10 +19,13 @@ import pytest
import litellm
from litellm.types.utils import (
CompletionTokensDetailsWrapper,
ImageResponse,
ImageObject,
ImageUsage,
ImageUsageInputTokensDetails,
PromptTokensDetailsWrapper,
Usage,
)
@ -202,6 +205,71 @@ class TestGPTImageCostRouting:
assert cost >= 0
class TestGPTImage15OutputImageTokens:
"""
Test for GitHub issue #19508:
Image usage calculation does not include image tokens in gpt-image-1.5
gpt-image-1.5 returns output_tokens_details with separate image_tokens and text_tokens,
and these must be correctly included in cost calculation.
"""
def test_gpt_image_15_output_image_tokens_cost(self):
"""
Test that output image tokens are correctly included in cost calculation.
This tests the fix for issue #19508 where output_tokens_details.image_tokens
were not being included in the cost calculation, causing costs to be
underreported (e.g., $0.046 instead of $0.14).
"""
# Simulate gpt-image-1.5 response with output_tokens_details
# This is what the API returns and what convert_to_image_response transforms
usage = Usage(
prompt_tokens=169,
completion_tokens=4599,
total_tokens=4768,
prompt_tokens_details=PromptTokensDetailsWrapper(
text_tokens=169,
image_tokens=0,
),
completion_tokens_details=CompletionTokensDetailsWrapper(
text_tokens=439,
image_tokens=4160,
),
)
image_response = ImageResponse(
created=1234567890,
data=[ImageObject(b64_json="test")],
)
image_response.usage = usage
image_response._hidden_params = {"custom_llm_provider": "openai"}
cost = litellm.completion_cost(
completion_response=image_response,
model="gpt-image-1.5",
call_type="image_generation",
custom_llm_provider="openai",
)
# gpt-image-1.5 pricing:
# - input_cost_per_token: 5e-06 ($5/1M for text input)
# - output_cost_per_token: 1e-05 ($10/1M for text output)
# - output_cost_per_image_token: 3.2e-05 ($32/1M for image output)
#
# Expected cost:
# Input text: 169 * $5/1M = $0.000845
# Output text: 439 * $10/1M = $0.00439
# Output image: 4160 * $32/1M = $0.13312
# Total: $0.138355
expected_cost = 169 * 5e-06 + 439 * 1e-05 + 4160 * 3.2e-05
assert abs(cost - expected_cost) < 1e-6, (
f"Expected {expected_cost}, got {cost}. "
f"Image tokens may not be included in cost calculation."
)
class TestCompletionCostIntegration:
"""Test the full completion_cost integration for gpt-image-1"""

View file

@ -477,6 +477,12 @@ async def test_openai_env_base(
respx_mock: respx.MockRouter, env_base, openai_api_response, monkeypatch
):
"This tests OpenAI env variables are honored, including legacy OPENAI_API_BASE"
# Clear cache to ensure no cached clients from previous tests interfere
# This prevents cache pollution where a previous test cached a client with
# aiohttp transport, which would bypass respx mocks
if hasattr(litellm, "in_memory_llm_clients_cache"):
litellm.in_memory_llm_clients_cache.flush_cache()
# Ensure aiohttp transport is disabled to use httpx which respx can mock
litellm.disable_aiohttp_transport = True

View file

@ -32,17 +32,17 @@ class TestPerDeploymentNumRetries:
)
deployment = router.model_list[0]
# Create a mock exception without num_retries
class MockException(Exception):
pass
exc = MockException("test error")
assert not hasattr(exc, "num_retries") or exc.num_retries is None
# Call the helper
router._set_deployment_num_retries_on_exception(exc, deployment)
# Verify num_retries was set from deployment
assert exc.num_retries == 5
@ -66,16 +66,16 @@ class TestPerDeploymentNumRetries:
)
deployment = router.model_list[0]
# Create an exception that already has num_retries
class MockException(Exception):
num_retries = 10 # Already set
exc = MockException("test error")
# Call the helper
router._set_deployment_num_retries_on_exception(exc, deployment)
# Verify num_retries was NOT overridden
assert exc.num_retries == 10
@ -99,15 +99,15 @@ class TestPerDeploymentNumRetries:
)
deployment = router.model_list[0]
class MockException(Exception):
pass
exc = MockException("test error")
# Call the helper
router._set_deployment_num_retries_on_exception(exc, deployment)
# Verify num_retries was not set (deployment has no num_retries)
assert not hasattr(exc, "num_retries") or exc.num_retries is None
@ -155,3 +155,36 @@ class TestPerDeploymentNumRetries:
kwargs = {}
router._update_kwargs_before_fallbacks(model="test-model", kwargs=kwargs)
assert kwargs["num_retries"] == 7 # Uses global
def test_set_deployment_num_retries_with_string_value(self):
"""
Test that _set_deployment_num_retries_on_exception handles string values
from environment variables correctly.
GitHub Issue: #19481
"""
router = Router(
model_list=[
{
"model_name": "test-model",
"litellm_params": {
"model": "openai/gpt-4",
"api_key": "test-key",
"num_retries": "6", # String value (as from env var)
},
},
],
num_retries=0, # Global setting
)
deployment = router.model_list[0]
class MockException(Exception):
pass
exc = MockException("test error")
# Call the helper
router._set_deployment_num_retries_on_exception(exc, deployment)
# Verify num_retries was converted from string to int
assert exc.num_retries == 6

View file

@ -119,8 +119,12 @@ def test_add_vector_store_to_registry():
@respx.mock
def test_search_uses_registry_credentials():
"""search() should pull credentials from vector_store_registry when available"""
# Block all HTTP requests at the network level to prevent real API calls
respx.route().mock(return_value=httpx.Response(200, json={"object": "list", "data": []}))
vector_store = LiteLLM_ManagedVectorStore(
vector_store_id="vs1",
custom_llm_provider="bedrock",

View file

@ -0,0 +1,35 @@
// hooks/useDisableShowPrompts.ts
import { useSyncExternalStore } from "react";
import { getLocalStorageItem } from "@/utils/localStorageUtils";
import { LOCAL_STORAGE_EVENT } from "@/utils/localStorageUtils";
function subscribe(callback: () => void) {
const onStorage = (e: StorageEvent) => {
if (e.key === "disableShowPrompts") {
callback();
}
};
const onCustom = (e: Event) => {
const { key } = (e as CustomEvent).detail;
if (key === "disableShowPrompts") {
callback();
}
};
window.addEventListener("storage", onStorage);
window.addEventListener(LOCAL_STORAGE_EVENT, onCustom);
return () => {
window.removeEventListener("storage", onStorage);
window.removeEventListener(LOCAL_STORAGE_EVENT, onCustom);
};
}
function getSnapshot() {
return getLocalStorageItem("disableShowPrompts") === "true";
}
export function useDisableShowPrompts() {
return useSyncExternalStore(subscribe, getSnapshot);
}

Some files were not shown because too many files have changed in this diff Show more