Merge pull request #24211 from BerriAI/litellm_dev_sameer_16_march_week

Litellm dev sameer 16 march week
This commit is contained in:
yuneng-jiang 2026-03-21 15:19:15 -07:00 • committed by GitHub
commit 1986f1034e
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
44 changed files with 5762 additions and 529 deletions

View file

@ -0,0 +1,106 @@
---
slug: gpt_5_4_mini_nano
title: "Day 0 Support: GPT-5.4-mini and GPT-5.4-nano"
date: 2026-03-17T10:00:00
authors:
- name: Sameer Kankute
title: SWE @ LiteLLM (LLM Translation)
url: https://www.linkedin.com/in/sameer-kankute/
image_url: https://pbs.twimg.com/profile_images/2001352686994907136/ONgNuSk5_400x400.jpg
- 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
description: "GPT-5.4-mini and GPT-5.4-nano model support in LiteLLM"
tags: [openai, gpt-5.4-mini, gpt-5.4-nano, completion]
hide_table_of_contents: false
---
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
LiteLLM now supports GPT-5.4-mini and GPT-5.4-nano — cost-effective models for simple completions and high-throughput workloads.
:::note
If you're on **v1.82.3-stable** or above, you don't need any update to use these models.
:::
## Usage
<Tabs>
<TabItem value="proxy" label="LiteLLM Proxy">
**1. Setup config.yaml**
```yaml
model_list:
- model_name: gpt-5.4-mini
litellm_params:
model: openai/gpt-5.4-mini
api_key: os.environ/OPENAI_API_KEY
- model_name: gpt-5.4-nano
litellm_params:
model: openai/gpt-5.4-nano
api_key: os.environ/OPENAI_API_KEY
```
**2. Start the proxy**
```bash
litellm --config /path/to/config.yaml
```
**3. Test it**
```bash
# GPT-5.4-mini
curl -X POST "http://localhost:4000/v1/chat/completions" \
-H "Content-Type: application/json" \
-H "Authorization: Bearer $LITELLM_KEY" \
-d '{
"model": "gpt-5.4-mini",
"messages": [{"role": "user", "content": "What is the capital of France?"}]
}'
# GPT-5.4-nano
curl -X POST "http://localhost:4000/v1/chat/completions" \
-H "Content-Type: application/json" \
-H "Authorization: Bearer $LITELLM_KEY" \
-d '{
"model": "gpt-5.4-nano",
"messages": [{"role": "user", "content": "What is 2 + 2?"}]
}'
```
</TabItem>
<TabItem value="sdk" label="LiteLLM SDK">
```python
from litellm import completion
# GPT-5.4-mini
response = completion(
model="openai/gpt-5.4-mini",
messages=[{"role": "user", "content": "What is the capital of France?"}],
)
print(response.choices[0].message.content)
# GPT-5.4-nano
response = completion(
model="openai/gpt-5.4-nano",
messages=[{"role": "user", "content": "What is 2 + 2?"}],
)
print(response.choices[0].message.content)
```
</TabItem>
</Tabs>
## Notes
- Both models support function calling, vision, and tool-use — see the [OpenAI provider docs](../../docs/providers/openai) for advanced usage.
- GPT-5.4-nano is the most cost-effective option for simple tasks; GPT-5.4-mini offers a balance of speed and capability.

View file

@ -0,0 +1,48 @@
---
title: Prompt Management with Responses API
---
# Prompt Management with Responses API
Use LiteLLM Prompt Management with `/v1/responses` by passing `prompt_id` and optional `prompt_variables`.
## Basic Usage
```bash
curl -X POST "http://localhost:4000/v1/responses" \
-H "Authorization: Bearer sk-1234" \
-H "Content-Type: application/json" \
-d '{
"model": "gpt-4o",
"prompt_id": "my-responses-prompt",
"prompt_variables": {"topic": "large language models"},
"input": []
}'
```
## Multi-turn Follow-up in `input`
To send follow-up turns in one request, pass message history in `input`.
```bash
curl -X POST "http://localhost:4000/v1/responses" \
-H "Authorization: Bearer sk-1234" \
-H "Content-Type: application/json" \
-d '{
"model": "gpt-4o",
"prompt_id": "my-responses-prompt",
"prompt_variables": {"topic": "large language models"},
"input": [
{"role": "user", "content": "Topic is LLMs. Start short."},
{"role": "assistant", "content": "Sure, go ahead."},
{"role": "user", "content": "Now give me 3 bullets and include pricing caveat."}
]
}'
```
## Notes
- Prompt template messages are merged with your `input` messages.
- Prompt variable substitution applies to prompt message content.
- Tool call payload fields are not substituted by prompt variables.
- For follow-ups with `previous_response_id`, include `prompt_id` again if you want prompt management applied on that turn.

View file

@ -361,8 +361,9 @@ router_settings:
| redis_url | str | URL for Redis server. **Known performance issue with Redis URL.** |
| cache_responses | boolean | Flag to enable caching LLM Responses, if cache set under `router_settings`. If true, caches responses. Defaults to False. |
| router_general_settings | RouterGeneralSettings | [SDK-Only] Router general settings - contains optimizations like 'async_only_mode'. [Docs](../routing.md#router-general-settings) |
| optional_pre_call_checks | List[str] | List of pre-call checks to add to the router. Supported: `router_budget_limiting`, `prompt_caching`, `responses_api_deployment_check`, `encrypted_content_affinity`, `deployment_affinity`, `session_affinity`, `forward_client_headers_by_model_group` |
| optional_pre_call_checks | List[str] | List of pre-call checks to add to the router. Supported: `router_budget_limiting`, `prompt_caching`, `responses_api_deployment_check`, `encrypted_content_affinity` (requires LiteLLM >= 1.82.3), `deployment_affinity`, `session_affinity`, `forward_client_headers_by_model_group` |
| deployment_affinity_ttl_seconds | int | TTL (seconds) for user-key → deployment affinity mapping when `deployment_affinity` is enabled (configured at Router init / proxy startup). Defaults to `3600` (1 hour). |
| model_group_affinity_config | Dict[str, List[str]] | Per-model-group affinity flags. Keys are model group names; values are lists of checks to enable (`deployment_affinity`, `responses_api_deployment_check`, `session_affinity`). Groups not listed fall back to the global `optional_pre_call_checks`. [Docs](../response_api.md#per-model-group-affinity-configuration) |
| ignore_invalid_deployments | boolean | If true, ignores invalid deployments. Default for proxy is True - to prevent invalid models from blocking other models from being loaded. |
| search_tools | List[SearchToolTypedDict] | List of search tool configurations for Search API integration. Each tool specifies a search_tool_name and litellm_params with search_provider, api_key, api_base, etc. [Further Docs](../search/index.md) |
| guardrail_list | List[GuardrailTypedDict] | List of guardrail configurations for guardrail load balancing. Enables load balancing across multiple guardrail deployments with the same guardrail_name. [Further Docs](./guardrails/guardrail_load_balancing.md) |

View file

@ -352,7 +352,7 @@ If `order=1` deployment is unavailable (e.g., rate-limited), the router falls ba
When load balancing OpenAI's Responses API across deployments with **different API keys** (e.g., different Azure regions or organizations), encrypted content items (like `rs_...` reasoning items) can only be decrypted by the originating API key.
**Solution:** Use the `encrypted_content_affinity` pre-call check to automatically route follow-up requests containing encrypted items to the correct deployment:
**Solution:** Use the `encrypted_content_affinity` pre-call check (requires LiteLLM >= 1.82.3) to automatically route follow-up requests containing encrypted items to the correct deployment:
```yaml
model_list:

View file

@ -311,7 +311,7 @@ litellm_settings:
1. **At Startup**: When the proxy starts, it reads the `prompts` field from `config.yaml`
2. **Initialization**: Each prompt is initialized based on its `prompt_integration` type
3. **In-Memory Storage**: Prompts are stored in the `IN_MEMORY_PROMPT_REGISTRY`
4. **Access**: Use these prompts via the `/v1/chat/completions` endpoint with `prompt_id` in the request
4. **Access**: Use these prompts via `/v1/chat/completions` or `/v1/responses` with `prompt_id` in the request
### Using Config-Loaded Prompts
@ -331,6 +331,23 @@ curl -L -X POST 'http://0.0.0.0:4000/v1/chat/completions' \
}'
```
You can also use the same `prompt_id` with the Responses API:
```bash
curl -L -X POST 'http://0.0.0.0:4000/v1/responses' \
-H 'Content-Type: application/json' \
-H 'Authorization: Bearer sk-1234' \
-d '{
"model": "gpt-4o",
"prompt_id": "coding_assistant",
"prompt_variables": {
"language": "python",
"task": "create a web scraper"
},
"input": []
}'
```
### Prompt Schema Reference
Each prompt in the `prompts` list requires:

View file

@ -1160,12 +1160,12 @@ follow_up = await router.aresponses(
To enable session continuity for Responses API in your LiteLLM proxy, set `optional_pre_call_checks` in your proxy config.yaml.
- `responses_api_deployment_check`: high priority routing when `previous_response_id` is provided
- `encrypted_content_affinity`: **[Recommended]** content-aware routing for encrypted items (e.g., `rs_...` reasoning items)
- `encrypted_content_affinity`: **[Recommended]** content-aware routing for encrypted items (e.g., `rs_...` reasoning items) (**requires LiteLLM >= 1.82.3**)
- `session_affinity`: sticky sessions based on session id (takes priority over `deployment_affinity`)
- `deployment_affinity`: sticky sessions based on user key (applies even without `previous_response_id`)
:::tip Recommended: Use `encrypted_content_affinity`
For Responses API with load balancing across deployments with **different API keys**, use `encrypted_content_affinity` instead of `deployment_affinity`. It only pins requests that contain encrypted content, avoiding quota reduction while preventing `invalid_encrypted_content` errors.
For Responses API with load balancing across deployments with **different API keys**, use `encrypted_content_affinity` instead of `deployment_affinity`. It only pins requests that contain encrypted content, avoiding quota reduction while preventing `invalid_encrypted_content` errors. (Requires LiteLLM >= 1.82.3.)
:::
Notes:
@ -1364,6 +1364,85 @@ litellm --config config.yaml
| `deployment_affinity` | Simple sticky sessions | All requests from same API key | ❌ Reduces quota by # of users |
## Per-Model-Group Affinity Configuration
By default, `optional_pre_call_checks` applies globally to all model groups. Use `model_group_affinity_config` when you want different affinity behavior per model group — for example, enabling stickiness only for models spread across providers (Azure + Bedrock) while leaving single-provider groups free to load-balance.
Groups not listed fall back to the global `optional_pre_call_checks` settings.
<Tabs>
<TabItem value="python-sdk" label="Python SDK">
```python
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {"model": "azure/gpt-4", "api_key": "...", "api_base": "https://endpoint1.openai.azure.com"},
},
{
"model_name": "gpt-4",
"litellm_params": {"model": "bedrock/anthropic.claude-v2", "aws_region_name": "us-east-1"},
},
{
"model_name": "text-embedding-ada-002",
"litellm_params": {"model": "azure/text-embedding-ada-002", "api_key": "...", "api_base": "https://endpoint1.openai.azure.com"},
},
{
"model_name": "text-embedding-ada-002",
"litellm_params": {"model": "azure/text-embedding-ada-002", "api_key": "...", "api_base": "https://endpoint2.openai.azure.com"},
},
],
# gpt-4: cross-provider (Azure + Bedrock) — enable deployment affinity
# text-embedding-ada-002: same provider — no affinity, let it load balance freely
model_group_affinity_config={
"gpt-4": ["deployment_affinity", "responses_api_deployment_check"],
},
)
```
</TabItem>
<TabItem value="proxy-server" label="Proxy Server">
```yaml title="config.yaml"
model_list:
- model_name: gpt-4
litellm_params:
model: azure/gpt-4
api_key: os.environ/AZURE_API_KEY_1
api_base: https://endpoint1.openai.azure.com
- model_name: gpt-4
litellm_params:
model: bedrock/anthropic.claude-v2
aws_region_name: us-east-1
- model_name: text-embedding-ada-002
litellm_params:
model: azure/text-embedding-ada-002
api_key: os.environ/AZURE_API_KEY_1
api_base: https://endpoint1.openai.azure.com
- model_name: text-embedding-ada-002
litellm_params:
model: azure/text-embedding-ada-002
api_key: os.environ/AZURE_API_KEY_2
api_base: https://endpoint2.openai.azure.com
router_settings:
# gpt-4: cross-provider — enable stickiness
# text-embedding-ada-002: not listed — load balances freely
model_group_affinity_config:
"gpt-4":
- deployment_affinity
- responses_api_deployment_check
```
</TabItem>
</Tabs>
**Supported values:** `deployment_affinity`, `responses_api_deployment_check`, `session_affinity`
## Calling non-Responses API endpoints (`/responses` to `/chat/completions` Bridge)
LiteLLM allows you to call non-Responses API models via a bridge to LiteLLM's `/chat/completions` endpoint. This is useful for calling Anthropic, Gemini and even non-Responses API OpenAI models.
@ -1556,6 +1635,12 @@ curl -X POST "http://localhost:4000/v1/responses" \
}'
```
## File Search (Vector Stores)
For full `file_search` usage (native + emulated fallback), SDK/Proxy examples, architecture diagram, and Q&A, see:
- [`File Search in the Responses API — E2E Testing Guide`](/docs/tutorials/file_search_responses_api)
## Session Management
LiteLLM Proxy supports session management for all supported models. This allows you to store and fetch conversation history (state) in LiteLLM Proxy.

View file

@ -0,0 +1,241 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# File Search in the Responses API
LiteLLM now supports `file_search` in the Responses API across both:
- providers that support it natively (like OpenAI / Azure), and
- providers that do not (like Anthropic, Bedrock, and other non-native providers) via emulation.
## What this is
`file_search` lets models retrieve grounded context from your vector stores and answer with citations.
LiteLLM keeps one OpenAI-compatible output shape while routing requests through either native passthrough or an emulated fallback.
Two paths are covered:
| Path | When it runs | What LiteLLM does |
| --- | --- | --- |
| **Native passthrough** | Provider natively supports `file_search` (OpenAI, Azure) | Decodes unified vector store ID → forwards to provider as-is |
| **Emulated fallback** | Provider doesn't support `file_search` (Anthropic, Bedrock, etc.) | Converts to a function tool → intercepts tool call → runs vector search → synthesizes OpenAI-format output |
In `tools[].vector_store_ids`, LiteLLM accepts both provider-native IDs (e.g. `vs_...`) **and** **managed vector store unified IDs** (URL-safe base64 strings from the proxy managed-vector flow), e.g. `litellm.responses(..., tools=[{"type": "file_search", "vector_store_ids": ["bGl0ZWxsbV9wcm94eT..."]}])`.
## Usage
<Tabs>
<TabItem value="proxy" label="LiteLLM Proxy" default>
### 1. Setup `config.yaml`
```yaml title="config.yaml"
model_list:
- model_name: gpt-4.1
litellm_params:
model: openai/gpt-4.1
api_key: os.environ/OPENAI_API_KEY
- model_name: claude-sonnet
litellm_params:
model: anthropic/claude-sonnet-4-5
api_key: os.environ/ANTHROPIC_API_KEY
```
### 2. Start the proxy
```bash
litellm --config config.yaml
```
### 3. Call Responses API with `file_search`
```python title="Proxy call"
from openai import OpenAI
client = OpenAI(base_url="http://localhost:4000", api_key="sk-your-proxy-key")
response = client.responses.create(
model="claude-sonnet", # swap to "gpt-4.1" for native path
input="What does LiteLLM support?",
tools=[{
"type": "file_search",
"vector_store_ids": ["vs_abc123"]
}],
include=["file_search_call.results"],
)
print(response.output)
```
</TabItem>
<TabItem value="sdk" label="LiteLLM SDK">
### 1. Install + set keys
```bash
pip install litellm
export OPENAI_API_KEY="sk-..."
export ANTHROPIC_API_KEY="sk-ant-..."
```
### 2. Call Responses API with `file_search`
```python title="SDK call"
import litellm
response = litellm.responses(
model="anthropic/claude-sonnet-4-5", # swap to openai/gpt-4.1 for native path
input="What does LiteLLM support?",
tools=[{
"type": "file_search",
"vector_store_ids": ["vs_abc123"]
}],
include=["file_search_call.results"],
)
print(response.output)
```
</TabItem>
</Tabs>
### Behavior Matrix
| Path | SDK model | Proxy model | Behavior |
| --- | --- | --- | --- |
| Native passthrough | `openai/gpt-4.1` | `gpt-4.1` | Provider executes native `file_search` |
| Emulated fallback | `anthropic/claude-sonnet-4-5` | `claude-sonnet` | LiteLLM converts to function tool and synthesizes OpenAI-format output |
## Architecture Diagram
```mermaid
flowchart TD
A[Client SDK or Proxy Caller] --> B[LiteLLM Responses API]
B --> C{Provider supports native file_search?}
C -->|Yes| D[Native passthrough path]
D --> D1[Decode unified vector_store_id if needed]
D1 --> D2[Forward request to provider unchanged]
D2 --> D3[Provider performs file_search]
D3 --> Z[OpenAI-compatible output]
C -->|No| E[Emulated fallback path]
E --> E1[Convert file_search to litellm_file_search function tool]
E1 --> E2[First model call returns tool call with one or more queries]
E2 --> E3[LiteLLM executes vector search for each query]
E3 --> E4[Second model call with tool_result context]
E4 --> E5[Synthesize file_search_call + message + citations]
E5 --> Z[OpenAI-compatible output]
```
## Prerequisites
```bash
pip install 'litellm[proxy]'
export OPENAI_API_KEY="sk-..." # for native path
export ANTHROPIC_API_KEY="sk-ant-..." # for emulated path
```
## Example response shape
## Validating the Output Format
Regardless of which path ran, the response always follows the OpenAI Responses API format:
```json
{
"output": [
{
"type": "file_search_call",
"id": "fs_abc123",
"status": "completed",
"queries": ["What does LiteLLM support?"],
"search_results": null
},
{
"type": "message",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": "LiteLLM is a unified interface...",
"annotations": [
{
"type": "file_citation",
"index": 150,
"file_id": "file-xxxx",
"filename": "knowledge.txt"
}
]
}
]
}
]
}
```
**Validation script:**
```python showLineNumbers title="Validate response structure"
def validate_file_search_response(response):
"""Assert that response follows OpenAI file_search output format."""
output = response.output
assert len(output) >= 2, "Expected at least 2 output items"
# First item: file_search_call
fs_call = output[0]
fs_type = fs_call["type"] if isinstance(fs_call, dict) else fs_call.type
assert fs_type == "file_search_call", f"Expected file_search_call, got {fs_type}"
fs_status = fs_call["status"] if isinstance(fs_call, dict) else fs_call.status
assert fs_status == "completed"
# Second item: message
msg = output[1]
msg_type = msg["type"] if isinstance(msg, dict) else msg.type
assert msg_type == "message"
content = msg["content"] if isinstance(msg, dict) else msg.content
assert len(content) > 0
text_block = content[0]
text = text_block["text"] if isinstance(text_block, dict) else text_block.text
assert isinstance(text, str) and len(text) > 0
print("✅ Response structure valid")
print(f" Queries: {fs_call['queries'] if isinstance(fs_call, dict) else fs_call.queries}")
print(f" Answer length: {len(text)} chars")
annotations = text_block["annotations"] if isinstance(text_block, dict) else text_block.annotations
print(f" Citations: {len(annotations)}")
validate_file_search_response(response)
```
## Q&A
- **Why do I see `UnsupportedParamsError`?** This usually means `file_search` was passed to a provider that does not support it natively and emulation could not route correctly. Check:
- The model string is valid (for example, `anthropic/claude-sonnet-4-5`).
- `custom_llm_provider` resolves correctly so LiteLLM can load the provider config.
- **Why does vector search return no results?** Common causes:
- The vector store ID is wrong or has no files attached.
- In LiteLLM-managed stores, file ingestion is not complete (`status != completed`).
- The query is too narrow; try a broader query.
- **Why am I getting `403 Access denied` on vector store calls?** The caller does not have access to that vector store.
- The store may belong to another team.
- Use an admin/proxy key if your setup requires cross-team access.
- **Why are `annotations` empty in emulated mode?** `file_citation` annotations require `file_id` metadata in search results. If your vector backend does not return file-level metadata, the answer text is still generated but citations can be empty.
## What to check next
- [File Search reference in Responses API docs](/docs/response_api#file-search-vector-stores) — full API reference
- [Vector Store management](/docs/vector_store_files) — create and manage vector stores
- [Managed vector stores](/docs/providers/bedrock_vector_store) — provider-specific setup

View file

@ -0,0 +1,151 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Vertex AI PayGo and Priority
## Priority PayGo
LiteLLM supports Priority PayGo.
Send a priority header, get priority queueing, and pay priority token rates.
:::info Which models support Priority PayGo?
As of this writing: `gemini/gemini-2.5-pro`, `vertex_ai/gemini-3-pro-preview`, `vertex_ai/gemini-3.1-pro-preview`, `vertex_ai/gemini-3-flash-preview`, and their variants.
Check `supports_service_tier: true` in LiteLLM's [model pricing JSON](https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json).
:::
### Send a priority request
Use this header:
`X-Vertex-AI-LLM-Shared-Request-Type: priority`
<Tabs>
<TabItem value="litellm-sdk" label="LiteLLM SDK">
```python
import litellm
response = litellm.completion(
model="vertex_ai/gemini-3-pro-preview",
messages=[{"role": "user", "content": "Summarize the Gettysburg Address."}],
vertex_project="YOUR_PROJECT_ID",
vertex_location="us-central1",
extra_headers={"X-Vertex-AI-LLM-Shared-Request-Type": "priority"},
)
print(response.choices[0].message.content)
```
</TabItem>
<TabItem value="proxy-config" label="Proxy config">
```yaml title="config.yaml"
model_list:
- model_name: gemini-priority
litellm_params:
model: vertex_ai/gemini-3-pro-preview
vertex_project: "YOUR_PROJECT_ID"
vertex_location: "us-central1"
vertex_credentials: os.environ/GOOGLE_APPLICATION_CREDENTIALS
extra_headers:
X-Vertex-AI-LLM-Shared-Request-Type: priority
```
```bash
curl http://localhost:4000/v1/chat/completions \
-H "Authorization: Bearer sk-your-key" \
-H "Content-Type: application/json" \
-d '{"model": "gemini-priority", "messages": [{"role": "user", "content": "Hello"}]}'
```
</TabItem>
<TabItem value="pass-through" label="Pass-through mode">
Use `x-pass-` so LiteLLM forwards provider-specific headers.
```bash
MODEL_ID="gemini-3-pro-preview-0325"
PROJECT_ID="YOUR_PROJECT_ID"
curl -X POST \
"${LITELLM_PROXY_BASE_URL}/vertex_ai/v1/projects/${PROJECT_ID}/locations/global/publishers/google/models/${MODEL_ID}:generateContent" \
-H "Authorization: Bearer sk-your-litellm-key" \
-H "Content-Type: application/json" \
-H "x-pass-X-Vertex-AI-LLM-Shared-Request-Type: priority" \
-d '{"contents": [{"role": "user", "parts": [{"text": "Hello!"}]}]}'
```
</TabItem>
</Tabs>
### How cost tracking works
![Vertex AI Priority PayGo Cost Tracking Flow](/img/vertex_cost_tracking_flow.svg)
**`trafficType` → `service_tier` mapping**
| `usageMetadata.trafficType` | `service_tier` | Pricing keys used |
|---|---|---|
| `ON_DEMAND` | `None` | `input_cost_per_token` |
| `ON_DEMAND_PRIORITY` | `"priority"` | `input_cost_per_token_priority` |
| `FLEX` / `BATCH` | `"flex"` | `input_cost_per_token_flex` |
If a tier-specific key is missing, LiteLLM falls back to standard pricing keys.
---
## Standard PayGo vs Provisioned Throughput
This is a different header from priority routing:
| Header value | Behavior |
|---|---|
| `X-Vertex-AI-LLM-Request-Type: shared` | Force standard PayGo (bypass PT) |
| `X-Vertex-AI-LLM-Request-Type: dedicated` | Force Provisioned Throughput only (`429` if exhausted) |
### Native route example
```python
import litellm
response = litellm.completion(
model="vertex_ai/gemini-2.0-flash",
messages=[{"role": "user", "content": "Hello!"}],
vertex_project="YOUR_PROJECT_ID",
vertex_location="us-central1",
extra_headers={"X-Vertex-AI-LLM-Request-Type": "shared"},
)
```
### Pass-through example
```bash
MODEL_ID="gemini-2.0-flash-001"
PROJECT_ID="YOUR_PROJECT_ID"
curl -X POST \
"${LITELLM_PROXY_BASE_URL}/vertex_ai/v1/projects/${PROJECT_ID}/locations/global/publishers/google/models/${MODEL_ID}:generateContent" \
-H "Authorization: Bearer sk-your-litellm-key" \
-H "Content-Type: application/json" \
-H "x-pass-X-Vertex-AI-LLM-Request-Type: shared" \
-d '{
"contents": [{"role": "user", "parts": [{"text": "Hello!"}]}]
}'
```
---
## Troubleshooting
**Q: What does `403 Permission denied` or `IAM_PERMISSION_DENIED` mean?**
A: The service account or Application Default Credentials (ADC) user does not have the `roles/aiplatform.user` role. To resolve this, re-run the `gcloud projects add-iam-policy-binding`.
**Q: What should I do if I get a `429 Quota exceeded` error?**
A: This means you've hit the per-region QPM (queries per minute) or TPM (tokens per minute) quota. You can:
- Request a quota increase from the [GCP Quotas console](https://console.cloud.google.com/iam-admin/quotas)
- Add more regions to your LiteLLM configuration for load balancing
- Upgrade to [Provisioned Throughput](https://cloud.google.com/vertex-ai/generative-ai/docs/provisioned-throughput) for guaranteed capacity
**Q: How do I fix the `VERTEXAI_PROJECT not set` error?**
A: Either pass the `vertex_project` parameter explicitly in your LiteLLM call, or set the `VERTEXAI_PROJECT` environment variable before running your code.

View file

@ -584,6 +584,7 @@ const sidebars = {
label: "Spend Tracking",
items: [
"proxy/cost_tracking",
"tutorials/vertex_ai_pay_go",
"proxy/request_tags",
"proxy/custom_pricing",
"proxy/pricing_calculator",
@ -737,6 +738,7 @@ const sidebars = {
"proxy/realtime_webrtc",
"rerank",
"response_api",
"prompt_management",
"response_api_compact",
{
type: "category",
@ -1433,6 +1435,7 @@ const learnSidebar = {
},
items: [
"tutorials/prompt_caching",
"tutorials/file_search_responses_api",
"tutorials/anthropic_file_usage",
"tutorials/gemini_realtime_with_audio",
"tutorials/litellm_proxy_aporia",

View file

@ -0,0 +1,63 @@
<svg width="100%" viewBox="0 0 680 560" xmlns="http://www.w3.org/2000/svg">
<defs>
<marker id="arrow" viewBox="0 0 10 10" refX="8" refY="5" markerWidth="6" markerHeight="6" orient="auto-start-reverse">
<path d="M2 1L8 5L2 9" fill="none" stroke="context-stroke" stroke-width="1.5" stroke-linecap="round" stroke-linejoin="round"/>
</marker>
</defs>
<!-- Step 1: HTTP Request -->
<g style="fill:rgb(0, 0, 0);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto">
<rect x="190" y="30" width="300" height="56" rx="8" stroke-width="0.5" style="fill:rgb(12, 68, 124);stroke:rgb(133, 183, 235);color:rgb(255, 255, 255);stroke-width:0.5px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto"/>
<text x="340" y="52" text-anchor="middle" dominant-baseline="central" style="fill:rgb(181, 212, 244);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:14px;font-weight:500;text-anchor:middle;dominant-baseline:central">HTTP request</text>
<text x="340" y="70" text-anchor="middle" dominant-baseline="central" style="fill:rgb(133, 183, 235);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:12px;font-weight:400;text-anchor:middle;dominant-baseline:central">X-Vertex-AI-LLM-Shared-Request-Type: priority</text>
</g>
<!-- Arrow 1 -->
<line x1="340" y1="86" x2="340" y2="120" marker-end="url(#arrow)" style="fill:none;stroke:rgb(156, 154, 146);color:rgb(255, 255, 255);stroke-width:1.5px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto"/>
<text x="356" y="108" dominant-baseline="central" style="fill:rgb(194, 192, 182);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:12px;font-weight:400;text-anchor:start;dominant-baseline:central">Vertex AI</text>
<!-- Step 2: Vertex response -->
<g style="fill:rgb(0, 0, 0);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto">
<rect x="190" y="120" width="300" height="56" rx="8" stroke-width="0.5" style="fill:rgb(8, 80, 65);stroke:rgb(93, 202, 165);color:rgb(255, 255, 255);stroke-width:0.5px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto"/>
<text x="340" y="142" text-anchor="middle" dominant-baseline="central" style="fill:rgb(159, 225, 203);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:14px;font-weight:500;text-anchor:middle;dominant-baseline:central">Vertex response</text>
<text x="340" y="160" text-anchor="middle" dominant-baseline="central" style="fill:rgb(93, 202, 165);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:12px;font-weight:400;text-anchor:middle;dominant-baseline:central">usageMetadata.trafficType = ON_DEMAND_PRIORITY</text>
</g>
<!-- Arrow 2 -->
<line x1="340" y1="176" x2="340" y2="210" marker-end="url(#arrow)" style="fill:none;stroke:rgb(156, 154, 146);color:rgb(255, 255, 255);stroke-width:1.5px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto"/>
<!-- Step 3: LiteLLM hidden params -->
<g style="fill:rgb(0, 0, 0);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto">
<rect x="190" y="210" width="300" height="56" rx="8" stroke-width="0.5" style="fill:rgb(60, 52, 137);stroke:rgb(175, 169, 236);color:rgb(255, 255, 255);stroke-width:0.5px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto"/>
<text x="340" y="232" text-anchor="middle" dominant-baseline="central" style="fill:rgb(206, 203, 246);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:14px;font-weight:500;text-anchor:middle;dominant-baseline:central">LiteLLM stores it</text>
<text x="340" y="250" text-anchor="middle" dominant-baseline="central" style="fill:rgb(175, 169, 236);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:12px;font-weight:400;text-anchor:middle;dominant-baseline:central">_hidden_params.provider_specific_fields.traffic_type</text>
</g>
<!-- Arrow 3 -->
<line x1="340" y1="266" x2="340" y2="300" marker-end="url(#arrow)" style="fill:none;stroke:rgb(156, 154, 146);color:rgb(255, 255, 255);stroke-width:1.5px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto"/>
<!-- Step 4: completion_cost() -->
<g style="fill:rgb(0, 0, 0);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto">
<rect x="190" y="300" width="300" height="56" rx="8" stroke-width="0.5" style="fill:rgb(99, 56, 6);stroke:rgb(239, 159, 39);color:rgb(255, 255, 255);stroke-width:0.5px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto"/>
<text x="340" y="322" text-anchor="middle" dominant-baseline="central" style="fill:rgb(250, 199, 117);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:14px;font-weight:500;text-anchor:middle;dominant-baseline:central">completion_cost()</text>
<text x="340" y="340" text-anchor="middle" dominant-baseline="central" style="fill:rgb(239, 159, 39);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:12px;font-weight:400;text-anchor:middle;dominant-baseline:central">Maps traffic_type → service_tier = "priority"</text>
</g>
<!-- Arrow 4 -->
<line x1="340" y1="356" x2="340" y2="390" marker-end="url(#arrow)" style="fill:none;stroke:rgb(156, 154, 146);color:rgb(255, 255, 255);stroke-width:1.5px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto"/>
<!-- Step 5: Pricing lookup -->
<g style="fill:rgb(0, 0, 0);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto">
<rect x="190" y="390" width="300" height="56" rx="8" stroke-width="0.5" style="fill:rgb(113, 43, 19);stroke:rgb(240, 153, 123);color:rgb(255, 255, 255);stroke-width:0.5px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:16px;font-weight:400;text-anchor:start;dominant-baseline:auto"/>
<text x="340" y="412" text-anchor="middle" dominant-baseline="central" style="fill:rgb(245, 196, 179);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:14px;font-weight:500;text-anchor:middle;dominant-baseline:central">Pricing lookup</text>
<text x="340" y="430" text-anchor="middle" dominant-baseline="central" style="fill:rgb(240, 153, 123);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:12px;font-weight:400;text-anchor:middle;dominant-baseline:central">input/output_cost_per_token_priority</text>
</g>
<!-- Step numbers in left margin -->
<text x="172" y="58" text-anchor="end" dominant-baseline="central" style="fill:rgb(194, 192, 182);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:12px;font-weight:400;text-anchor:end;dominant-baseline:central">①</text>
<text x="172" y="148" text-anchor="end" dominant-baseline="central" style="fill:rgb(194, 192, 182);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:12px;font-weight:400;text-anchor:end;dominant-baseline:central">②</text>
<text x="172" y="238" text-anchor="end" dominant-baseline="central" style="fill:rgb(194, 192, 182);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:12px;font-weight:400;text-anchor:end;dominant-baseline:central">③</text>
<text x="172" y="328" text-anchor="end" dominant-baseline="central" style="fill:rgb(194, 192, 182);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:12px;font-weight:400;text-anchor:end;dominant-baseline:central">④</text>
<text x="172" y="418" text-anchor="end" dominant-baseline="central" style="fill:rgb(194, 192, 182);stroke:none;color:rgb(255, 255, 255);stroke-width:1px;stroke-linecap:butt;stroke-linejoin:miter;opacity:1;font-family:&quot;Anthropic Sans&quot;, -apple-system, &quot;system-ui&quot;, &quot;Segoe UI&quot;, sans-serif;font-size:12px;font-weight:400;text-anchor:end;dominant-baseline:central">⑤</text>
</svg>

After

Width:  |  Height:  |  Size: 12 KiB

View file

@ -29,7 +29,7 @@ from litellm.proxy.openai_files_endpoints.common_utils import (
get_models_from_unified_file_id,
normalize_mime_type_for_provider,
)
from litellm.types.llms.openai import (
from litellm.types.llms.openai import ( # pyright: ignore[reportAttributeAccessIssue]
AllMessageValues,
AsyncCursorPage,
ChatCompletionFileObject,
@ -442,25 +442,33 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
elif call_type == CallTypes.aresponses.value or call_type == CallTypes.responses.value:
# Handle managed files in responses API input and tools
file_ids = []
# Extract file IDs from input parameter
input_data = data.get("input")
if input_data:
file_ids.extend(self.get_file_ids_from_responses_input(input_data))
# Extract file IDs from tools parameter (e.g., code_interpreter container)
tools = data.get("tools")
if tools:
file_ids.extend(self.get_file_ids_from_responses_tools(tools))
if file_ids:
# Check user has access to all managed files
await self.check_file_ids_access(file_ids, user_api_key_dict)
model_file_id_mapping = await self.get_model_file_id_mapping(
file_ids, user_api_key_dict.parent_otel_span
)
data["model_file_id_mapping"] = model_file_id_mapping
# Check access for file_search vector_store_ids
if tools:
unified_vs_ids = self.get_vector_store_ids_from_file_search_tools(tools)
if unified_vs_ids:
await self.check_vector_store_ids_access(
unified_vs_ids, user_api_key_dict
)
elif call_type == CallTypes.afile_content.value:
retrieve_file_id = cast(Optional[str], data.get("file_id"))
potential_file_id = (
@ -704,6 +712,101 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
return file_ids
def get_vector_store_ids_from_file_search_tools(
self, tools: List[Dict[str, Any]]
) -> List[str]:
"""
Extract unified vector_store_ids from file_search tools.
Only returns IDs that are LiteLLM-managed (base64 unified IDs).
Native provider IDs are skipped — they have no LiteLLM access record.
"""
from litellm.llms.base_llm.managed_resources.utils import (
is_base64_encoded_unified_id,
)
vs_ids: List[str] = []
if not isinstance(tools, list):
return vs_ids
for tool in tools:
if not isinstance(tool, dict) or tool.get("type") != "file_search":
continue
vector_store_ids = tool.get("vector_store_ids")
if not isinstance(vector_store_ids, list):
continue
for vs_id in vector_store_ids:
if isinstance(vs_id, str) and is_base64_encoded_unified_id(vs_id):
vs_ids.append(vs_id)
return vs_ids
async def check_vector_store_ids_access(
self,
vector_store_ids: List[str],
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""
Verify the caller's team can access each LiteLLM-managed vector store.
Batch-fetches vector stores from DB and checks team_id.
Raises HTTPException(403) on the first access violation.
Non-managed (native) IDs should already be filtered out before calling this.
"""
from litellm.llms.base_llm.managed_resources.utils import (
extract_unified_uuid_from_unified_id,
)
from litellm.proxy.auth.auth_checks import (
get_managed_vector_store_rows_by_uuids,
)
from litellm.proxy.proxy_server import (
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if not vector_store_ids or prisma_client is None:
return
# Map each unified ID to its internal UUID for a single batch DB fetch
uuid_to_unified: Dict[str, str] = {}
for vs_id in vector_store_ids:
uuid = extract_unified_uuid_from_unified_id(vs_id)
if uuid:
uuid_to_unified[uuid] = vs_id
if not uuid_to_unified:
return
rows = await get_managed_vector_store_rows_by_uuids(
uuids=list(uuid_to_unified.keys()),
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
found_uuids = {row.vector_store_id for row in rows}
for uuid, original_id in uuid_to_unified.items():
if uuid not in found_uuids:
raise HTTPException(
status_code=403,
detail=f"Vector store '{original_id}' not found or access denied.",
)
caller_team_id = user_api_key_dict.team_id
for row in rows:
vs_team_id = getattr(row, "team_id", None)
if vs_team_id is not None and vs_team_id != caller_team_id:
raise HTTPException(
status_code=403,
detail=(
f"Team '{caller_team_id}' does not have access to vector "
f"store '{row.vector_store_id}'. The store belongs to team "
f"'{vs_team_id}'."
),
)
async def get_model_file_id_mapping(
self, file_ids: List[str], litellm_parent_otel_span: Span
) -> dict:
@ -954,7 +1057,7 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints):
)
else:
file_object = await litellm.afile_retrieve(
custom_llm_provider=model_name.split("/")[0] if model_name and "/" in model_name else "openai",
custom_llm_provider=model_name.split("/")[0] if model_name and "/" in model_name else "openai", # type: ignore[arg-type]
file_id=original_file_id,
)
verbose_logger.debug(

View file

@ -885,7 +885,7 @@ def list_batches(
async def acancel_batch(
batch_id: str,
model: Optional[str] = None,
custom_llm_provider: Literal["openai", "azure"] = "openai",
custom_llm_provider: Literal["openai", "azure", "vertex_ai"] = "openai",
metadata: Optional[Dict[str, str]] = None,
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
@ -931,7 +931,7 @@ async def acancel_batch(
def cancel_batch(
batch_id: str,
model: Optional[str] = None,
custom_llm_provider: Union[Literal["openai", "azure"], str] = "openai",
custom_llm_provider: Union[Literal["openai", "azure", "vertex_ai"], str] = "openai",
metadata: Optional[Dict[str, str]] = None,
extra_headers: Optional[Dict[str, str]] = None,
extra_body: Optional[Dict[str, str]] = None,
@ -1048,9 +1048,35 @@ def cancel_batch(
cancel_batch_data=_cancel_batch_request,
litellm_params=litellm_params,
)
elif custom_llm_provider == "vertex_ai":
api_base = optional_params.api_base or None
vertex_ai_project = (
optional_params.vertex_project
or litellm.vertex_project
or get_secret_str("VERTEXAI_PROJECT")
)
vertex_ai_location = (
optional_params.vertex_location
or litellm.vertex_location
or get_secret_str("VERTEXAI_LOCATION")
)
vertex_credentials = optional_params.vertex_credentials or get_secret_str(
"VERTEXAI_CREDENTIALS"
)
response = vertex_ai_batches_instance.cancel_batch(
_is_async=_is_async,
batch_id=batch_id,
api_base=api_base,
vertex_project=vertex_ai_project,
vertex_location=vertex_ai_location,
vertex_credentials=vertex_credentials,
timeout=timeout,
max_retries=optional_params.max_retries,
)
else:
raise litellm.exceptions.BadRequestError(
message="LiteLLM doesn't support {} for 'cancel_batch'. Only 'openai' and 'azure' are supported.".format(
message="LiteLLM doesn't support {} for 'cancel_batch'. Only 'openai', 'azure', and 'vertex_ai' are supported.".format(
custom_llm_provider
),
model="n/a",

View file

@ -257,6 +257,15 @@ def detect_first_expected_role(
return None
def _counts_for_alternation(message: AllMessageValues) -> bool:
role = message.get("role")
if role == "user":
return True
if role == "assistant":
return not bool(message.get("tool_calls"))
return False
def _insert_user_continue_message(
messages: List[AllMessageValues],
user_continue_message: Optional[ChatCompletionUserMessage],
@ -269,8 +278,8 @@ def _insert_user_continue_message(
2. Final assistant message
3. Consecutive assistant messages
Only inserts messages between consecutive assistant messages,
ignoring all other role types.
Skips tool messages and assistant messages with tool calls in the
alternation check, matching strict templates like llama.cpp.
"""
if not messages:
return messages
@ -278,25 +287,42 @@ def _insert_user_continue_message(
result_messages = messages.copy() # Don't modify the input list
continue_message = user_continue_message or DEFAULT_USER_CONTINUE_MESSAGE
# Handle first message if it's an assistant message
# Handle first message if it's an assistant message — always prepend
# user_continue regardless of tool_calls, to preserve backward compatibility.
if result_messages[0]["role"] == "assistant":
result_messages.insert(0, continue_message)
# Handle consecutive assistant messages and final message
i = 1 # Start from second message since we handled first message
# Handle consecutive assistant messages in the counted sequence
i = 1
while i < len(result_messages):
curr_message = result_messages[i]
prev_message = result_messages[i - 1]
# Only check for consecutive assistant messages
# Ignore all other role types
if curr_message["role"] == "assistant" and prev_message["role"] == "assistant":
result_messages.insert(i, continue_message)
i += 2 # Skip over the message we just inserted
else:
inserted_continue_message = False
if (
_counts_for_alternation(curr_message)
and curr_message["role"] == "assistant"
):
# Preserve old behavior for malformed adjacent assistant sequences like
# assistant(tool_calls) -> assistant(no-tool-calls) with no tool message.
if i > 0 and result_messages[i - 1].get("role") == "assistant":
result_messages.insert(i, continue_message)
i += 2
inserted_continue_message = True
else:
j = i - 1
while j >= 0:
previous_message = result_messages[j]
if _counts_for_alternation(previous_message):
if previous_message["role"] == "assistant":
result_messages.insert(i, continue_message)
i += 2
inserted_continue_message = True
break
j -= 1
if not inserted_continue_message:
i += 1
# Handle final message
# Handle final message — append user_continue after any trailing assistant,
# including ones with tool_calls, to preserve backward compatibility.
if result_messages[-1]["role"] == "assistant" and ensure_alternating_roles:
result_messages.append(continue_message)
@ -311,34 +337,24 @@ def _insert_assistant_continue_message(
"""
Add assistant continuation messages between consecutive user messages.
Args:
messages: List of message dictionaries
assistant_continue_message: Optional custom assistant message
ensure_alternating_roles: Whether to enforce alternating roles
Returns:
Modified list of messages with inserted assistant messages
Only checks directly adjacent messages to preserve backward compatibility.
"""
if not ensure_alternating_roles or len(messages) <= 1:
return messages
# Create a new list to store modified messages
continue_message = assistant_continue_message or DEFAULT_ASSISTANT_CONTINUE_MESSAGE
modified_messages: List[AllMessageValues] = []
for i, message in enumerate(messages):
modified_messages.append(message)
# Check if we need to insert an assistant message
if (
i < len(messages) - 1 # Not the last message
and message.get("role") == "user" # Current is user
i < len(messages) - 1
and message.get("role") == "user"
and messages[i + 1].get("role") == "user"
): # Next is user
# Insert assistant message
continue_message = (
assistant_continue_message or DEFAULT_ASSISTANT_CONTINUE_MESSAGE
)
):
modified_messages.append(message)
modified_messages.append(continue_message)
else:
modified_messages.append(message)
return modified_messages
@ -536,6 +552,61 @@ def update_responses_input_with_model_file_ids(
return updated_input
def _decode_vector_store_ids_in_tools(
tools: Optional[List[Dict[str, Any]]],
) -> Optional[List[Dict[str, Any]]]:
"""
Decodes unified (LiteLLM-managed) vector_store_ids in file_search tools to
provider-native IDs. Non-unified IDs are passed through unchanged.
This runs unconditionally — no file-ID mapping is required.
"""
if not tools or not isinstance(tools, list):
return tools
from litellm.llms.base_llm.managed_resources.utils import (
is_base64_encoded_unified_id,
parse_unified_id,
)
updated_tools = []
for tool in tools:
if not isinstance(tool, dict) or tool.get("type") != "file_search":
updated_tools.append(tool)
continue
vector_store_ids = tool.get("vector_store_ids")
if not isinstance(vector_store_ids, list):
updated_tools.append(tool)
continue
decoded_ids = []
for vs_id in vector_store_ids:
if not isinstance(vs_id, str) or not is_base64_encoded_unified_id(vs_id):
decoded_ids.append(vs_id)
continue
parsed = parse_unified_id(vs_id)
provider_resource_id = (
parsed.get("provider_resource_id") if parsed else None
)
if not provider_resource_id:
verbose_logger.warning(
"file_search tool contains unified vector_store_id '%s' that could "
"not be decoded to a provider resource ID — passing original ID. "
"Ensure the vector store was created via LiteLLM.",
vs_id,
)
decoded_ids.append(vs_id)
else:
decoded_ids.append(provider_resource_id)
updated_tools.append({**tool, "vector_store_ids": decoded_ids})
return updated_tools
def update_responses_tools_with_model_file_ids(
tools: Optional[List[Dict[str, Any]]],
model_id: Optional[str] = None,
@ -544,7 +615,8 @@ def update_responses_tools_with_model_file_ids(
"""
Updates responses API tools with provider-specific file IDs.
Handles code_interpreter tools with container.file_ids.
Pass 1 (always): decode unified vector_store_ids in file_search tools.
Pass 2 (needs mapping): map code_interpreter container file_ids to provider IDs.
Args:
tools: The responses API tools parameter
@ -555,6 +627,10 @@ def update_responses_tools_with_model_file_ids(
if not tools or not isinstance(tools, list):
return tools
# Pass 1: decode unified vector_store_ids (no mapping needed)
tools = _decode_vector_store_ids_in_tools(tools) or tools
# Pass 2: map code_interpreter file IDs (requires mapping)
if not model_file_id_mapping or not model_id:
return tools

View file

@ -97,14 +97,59 @@ class AzureAIAgentsHandler:
# -------------------------------------------------------------------------
# Response Helpers
# -------------------------------------------------------------------------
def _extract_content_from_messages(self, messages_data: dict) -> str:
"""Extract assistant content from the messages response."""
def _extract_content_from_messages(
self, messages_data: dict
) -> Tuple[str, Optional[List[Dict[str, Any]]]]:
"""Extract assistant content and annotations from the messages response.
Returns (content, annotations) where annotations is a list of
OpenAI-compatible ChatCompletionAnnotation dicts, or None.
"""
for msg in messages_data.get("data", []):
if msg.get("role") == "assistant":
for content_item in msg.get("content", []):
if content_item.get("type") == "text":
return content_item.get("text", {}).get("value", "")
return ""
text_obj = content_item.get("text", {})
content = text_obj.get("value", "")
raw_annotations = text_obj.get("annotations")
annotations = self._transform_annotations(raw_annotations)
return content, annotations
return "", None
def _transform_annotations(
self,
raw_annotations: Optional[List[Dict[str, Any]]],
) -> Optional[List[Dict[str, Any]]]:
"""Transform Azure AI Foundry annotations to OpenAI-compatible format.
Azure AI returns annotations like:
{"type": "url_citation", "text": "[1]", "start_index": 10,
"end_index": 13, "url_citation": {"url": "...", "title": "..."}}
OpenAI expects:
{"type": "url_citation", "url_citation": {"url": "...", "title": "...",
"start_index": 10, "end_index": 13}}
"""
if not raw_annotations:
return None
result: List[Dict[str, Any]] = []
for ann in raw_annotations:
ann_type = ann.get("type")
if ann_type == "url_citation":
url_citation = dict(ann.get("url_citation", {}))
# Azure puts start/end_index at annotation level; OpenAI
# expects them inside url_citation
if "start_index" in ann and "start_index" not in url_citation:
url_citation["start_index"] = ann["start_index"]
if "end_index" in ann and "end_index" not in url_citation:
url_citation["end_index"] = ann["end_index"]
result.append({"type": "url_citation", "url_citation": url_citation})
else:
# Pass through unknown annotation types as-is
result.append(ann)
return result if result else None
def _build_model_response(
self,
@ -113,15 +158,23 @@ class AzureAIAgentsHandler:
model_response: ModelResponse,
thread_id: str,
messages: List[Dict[str, Any]],
annotations: Optional[List[Dict[str, Any]]] = None,
) -> ModelResponse:
"""Build the ModelResponse from agent output."""
from litellm.types.utils import Choices, Message, Usage
message_kwargs: Dict[str, Any] = {
"content": content,
"role": "assistant",
}
if annotations:
message_kwargs["annotations"] = annotations
model_response.choices = [
Choices(
finish_reason="stop",
index=0,
message=Message(content=content, role="assistant"),
message=Message(**message_kwargs),
)
]
model_response.model = model
@ -250,7 +303,7 @@ class AzureAIAgentsHandler:
)
# Execute the agent flow
thread_id, content = self._execute_agent_flow_sync(
thread_id, content, annotations = self._execute_agent_flow_sync(
make_request=make_request,
api_base=api_base,
api_version=api_version,
@ -261,7 +314,7 @@ class AzureAIAgentsHandler:
)
return self._build_model_response(
model, content, model_response, thread_id, messages
model, content, model_response, thread_id, messages, annotations
)
def _execute_agent_flow_sync(
@ -273,8 +326,8 @@ class AzureAIAgentsHandler:
thread_id: Optional[str],
messages: List[Dict[str, Any]],
optional_params: dict,
) -> Tuple[str, str]:
"""Execute the agent flow synchronously. Returns (thread_id, content)."""
) -> Tuple[str, str, Optional[List[Dict[str, Any]]]]:
"""Execute the agent flow synchronously. Returns (thread_id, content, annotations)."""
# Step 1: Create thread if not provided
if not thread_id:
@ -347,8 +400,8 @@ class AzureAIAgentsHandler:
)
self._check_response(response, [200], "Failed to get messages")
content = self._extract_content_from_messages(response.json())
return thread_id, content
content, annotations = self._extract_content_from_messages(response.json())
return thread_id, content, annotations
# -------------------------------------------------------------------------
# Async Completion
@ -399,7 +452,7 @@ class AzureAIAgentsHandler:
)
# Execute the agent flow
thread_id, content = await self._execute_agent_flow_async(
thread_id, content, annotations = await self._execute_agent_flow_async(
make_request=make_request,
api_base=api_base,
api_version=api_version,
@ -410,7 +463,7 @@ class AzureAIAgentsHandler:
)
return self._build_model_response(
model, content, model_response, thread_id, messages
model, content, model_response, thread_id, messages, annotations
)
async def _execute_agent_flow_async(
@ -422,8 +475,8 @@ class AzureAIAgentsHandler:
thread_id: Optional[str],
messages: List[Dict[str, Any]],
optional_params: dict,
) -> Tuple[str, str]:
"""Execute the agent flow asynchronously. Returns (thread_id, content)."""
) -> Tuple[str, str, Optional[List[Dict[str, Any]]]]:
"""Execute the agent flow asynchronously. Returns (thread_id, content, annotations)."""
# Step 1: Create thread if not provided
if not thread_id:
@ -496,8 +549,8 @@ class AzureAIAgentsHandler:
)
self._check_response(response, [200], "Failed to get messages")
content = self._extract_content_from_messages(response.json())
return thread_id, content
content, annotations = self._extract_content_from_messages(response.json())
return thread_id, content, annotations
# -------------------------------------------------------------------------
# Streaming Completion (Native SSE)
@ -585,6 +638,7 @@ class AzureAIAgentsHandler:
response_id = f"chatcmpl-{uuid.uuid4().hex[:8]}"
created = int(time.time())
thread_id = None
collected_annotations: Optional[List[Dict[str, Any]]] = None
current_event = None
@ -600,6 +654,9 @@ class AzureAIAgentsHandler:
if data_str == "[DONE]":
# Send final chunk with finish_reason
final_delta_kwargs: Dict[str, Any] = {"content": None}
if collected_annotations:
final_delta_kwargs["annotations"] = collected_annotations
final_chunk = ModelResponseStream(
id=response_id,
created=created,
@ -609,7 +666,7 @@ class AzureAIAgentsHandler:
StreamingChoices(
finish_reason="stop",
index=0,
delta=Delta(content=None),
delta=Delta(**final_delta_kwargs),
)
],
)
@ -628,6 +685,19 @@ class AzureAIAgentsHandler:
thread_id = data["id"]
verbose_logger.debug(f"Stream created thread: {thread_id}")
# Extract annotations from completed message
if current_event == "thread.message.completed":
for content_item in data.get("content", []):
if content_item.get("type") == "text":
raw_annotations = content_item.get("text", {}).get(
"annotations"
)
transformed = self._transform_annotations(raw_annotations)
if transformed:
if collected_annotations is None:
collected_annotations = []
collected_annotations.extend(transformed)
# Process message deltas - this is where the actual content comes
if current_event == "thread.message.delta":
delta_content = data.get("delta", {}).get("content", [])

View file

@ -54,6 +54,14 @@ class BaseResponsesAPIConfig(ABC):
and v is not None
}
def supports_native_file_search(self) -> bool:
"""Return True if this provider handles the file_search tool natively.
Override in provider subclasses that support file_search without
LiteLLM emulation (e.g. OpenAI, Azure OpenAI).
"""
return False
@abstractmethod
def get_supported_openai_params(self, model: str) -> list:
pass

View file

@ -3,7 +3,7 @@
from typing import Optional, Union
import litellm
from litellm.utils import _supports_factory
from litellm.utils import _is_explicitly_disabled_factory, _supports_factory
from .gpt_transformation import OpenAIGPTConfig
@ -113,6 +113,25 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
key=f"supports_{level}_reasoning_effort",
)
@classmethod
def _is_reasoning_effort_level_explicitly_disabled(
cls, model: str, level: str
) -> bool:
"""Return True only when the model map explicitly sets the capability to False.
Unlike ``_supports_reasoning_effort_level`` (which requires an explicit True),
this method returns True only when ``supports_{level}_reasoning_effort`` is
explicitly set to ``False`` in the model map. A missing key is treated as
supported (i.e. this method returns False = not disabled).
Use this for opt-out checks where unknown models should be allowed through.
"""
return _is_explicitly_disabled_factory(
model=model,
custom_llm_provider=None,
key=f"supports_{level}_reasoning_effort",
)
def get_supported_openai_params(self, model: str) -> list:
if self.is_model_gpt_5_search_model(model):
return [
@ -200,14 +219,32 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
if "reasoning_effort" in optional_params:
optional_params["reasoning_effort"] = normalized
if effective_effort is not None and effective_effort == "xhigh":
if not self._supports_reasoning_effort_level(model, "xhigh"):
if effective_effort == "xhigh":
# xhigh is an opt-in capability: only allow if model explicitly supports it.
if not self._supports_reasoning_effort_level(model, effective_effort):
if litellm.drop_params or drop_params:
non_default_params.pop("reasoning_effort", None)
optional_params.pop("reasoning_effort", None)
else:
raise litellm.utils.UnsupportedParamsError(
message=(
"reasoning_effort='xhigh' is only supported for gpt-5.1-codex-max, gpt-5.2, and gpt-5.4+ models."
f"reasoning_effort={effective_effort} is not supported for this model."
),
status_code=400,
)
elif effective_effort == "minimal":
# minimal is opt-out: unknown models pass through; only block when
# the model map explicitly sets supports_minimal_reasoning_effort=false.
if self._is_reasoning_effort_level_explicitly_disabled(
model, effective_effort
):
if litellm.drop_params or drop_params:
non_default_params.pop("reasoning_effort", None)
optional_params.pop("reasoning_effort", None)
else:
raise litellm.utils.UnsupportedParamsError(
message=(
f"reasoning_effort={effective_effort} is not supported for this model."
),
status_code=400,
)

View file

@ -32,6 +32,9 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
def custom_llm_provider(self) -> LlmProviders:
return LlmProviders.OPENAI
def supports_native_file_search(self) -> bool:
return True
def get_supported_openai_params(self, model: str) -> list:
"""
All OpenAI Responses API params are supported

View file

@ -376,3 +376,148 @@ class VertexAIBatchPrediction(VertexLLM):
response=_json_response
)
return vertex_batch_response
def cancel_batch(
self,
_is_async: bool,
batch_id: str,
api_base: Optional[str],
vertex_credentials: Optional[VERTEX_CREDENTIALS_TYPES],
vertex_project: Optional[str],
vertex_location: Optional[str],
timeout: Union[float, httpx.Timeout],
max_retries: Optional[int],
) -> Union[LiteLLMBatch, Coroutine[Any, Any, LiteLLMBatch]]:
access_token, project_id = self._ensure_access_token(
credentials=vertex_credentials,
project_id=vertex_project,
custom_llm_provider="vertex_ai",
)
default_api_base = self.create_vertex_batch_url(
vertex_location=vertex_location or "us-central1",
vertex_project=vertex_project or project_id,
)
retrieve_api_base_default = f"{default_api_base}/{batch_id}"
cancel_api_base_default = f"{retrieve_api_base_default}:cancel"
_, api_base = self._check_custom_proxy(
api_base=api_base,
custom_llm_provider="vertex_ai",
gemini_api_key=None,
endpoint="cancel",
stream=None,
auth_header=None,
url=cancel_api_base_default,
model=None,
vertex_project=vertex_project or project_id,
vertex_location=vertex_location or "us-central1",
vertex_api_version="v1",
)
if api_base.endswith(":cancel"):
retrieve_api_base = api_base.removesuffix(":cancel")
else:
retrieve_api_base = api_base.rsplit(":cancel", 1)[0].rstrip("/")
headers = {
"Content-Type": "application/json; charset=utf-8",
"Authorization": f"Bearer {access_token}",
}
if _is_async is True:
return self._async_cancel_batch(
api_base=api_base,
retrieve_api_base=retrieve_api_base,
headers=headers,
timeout=timeout,
)
sync_handler = _get_httpx_client()
try:
response = sync_handler.post(
url=api_base,
headers=headers,
data=json.dumps({}),
timeout=timeout,
)
except httpx.HTTPStatusError as e:
litellm.verbose_logger.error(
"Vertex AI batch cancel failed: status=%s, body=%s",
e.response.status_code,
e.response.text[:1000],
)
raise
if response.status_code != 200:
raise Exception(f"Error: {response.status_code} {response.text}")
# HTTPHandler.get() does not accept a timeout parameter
retrieve_response = sync_handler.get(
url=retrieve_api_base,
headers=headers,
)
if retrieve_response.status_code != 200:
litellm.verbose_logger.error(
"Vertex AI batch retrieve-after-cancel failed: status=%s, body=%s",
retrieve_response.status_code,
retrieve_response.text[:1000],
)
raise Exception(
f"Error: {retrieve_response.status_code} {retrieve_response.text}"
)
_json_response = retrieve_response.json()
vertex_batch_response = VertexAIBatchTransformation.transform_vertex_ai_batch_response_to_openai_batch_response(
response=_json_response
)
return vertex_batch_response
async def _async_cancel_batch(
self,
api_base: str,
retrieve_api_base: str,
headers: Dict[str, str],
timeout: Union[float, httpx.Timeout] = 600.0,
) -> LiteLLMBatch:
client = get_async_httpx_client(
llm_provider=litellm.LlmProviders.VERTEX_AI,
)
try:
response = await client.post(
url=api_base,
headers=headers,
data=json.dumps({}),
timeout=timeout,
)
except httpx.HTTPStatusError as e:
litellm.verbose_logger.error(
"Vertex AI batch cancel failed: status=%s, body=%s",
e.response.status_code,
e.response.text[:1000],
)
raise
if response.status_code != 200:
raise Exception(f"Error: {response.status_code} {response.text}")
# AsyncHTTPHandler.get() does not accept a timeout parameter
retrieve_response = await client.get(
url=retrieve_api_base,
headers=headers,
)
if retrieve_response.status_code != 200:
litellm.verbose_logger.error(
"Vertex AI batch retrieve-after-cancel failed: status=%s, body=%s",
retrieve_response.status_code,
retrieve_response.text[:1000],
)
raise Exception(
f"Error: {retrieve_response.status_code} {retrieve_response.text}"
)
_json_response = retrieve_response.json()
vertex_batch_response = VertexAIBatchTransformation.transform_vertex_ai_batch_response_to_openai_batch_response(
response=_json_response
)
return vertex_batch_response

View file

@ -3435,7 +3435,8 @@
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_none_reasoning_effort": true
"supports_none_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"azure/gpt-5.1-chat-2025-11-13": {
"cache_read_input_token_cost": 1.25e-07,
@ -18305,7 +18306,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1": {
"cache_read_input_token_cost": 1.25e-07,
@ -18344,7 +18346,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1-2025-11-13": {
"cache_read_input_token_cost": 1.25e-07,
@ -18383,7 +18386,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1-chat-latest": {
"cache_read_input_token_cost": 1.25e-07,
@ -18421,7 +18425,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2": {
"cache_read_input_token_cost": 1.75e-07,
@ -18461,7 +18466,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2-2025-12-11": {
"cache_read_input_token_cost": 1.75e-07,
@ -18501,7 +18507,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2-chat-latest": {
"cache_read_input_token_cost": 1.75e-07,
@ -18538,7 +18545,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.3-chat-latest": {
"cache_read_input_token_cost": 1.75e-07,
@ -18575,7 +18583,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2-pro": {
"input_cost_per_token": 2.1e-05,
@ -18608,7 +18617,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2-pro-2025-12-11": {
"input_cost_per_token": 2.1e-05,
@ -18641,7 +18651,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.4": {
"cache_read_input_token_cost": 2.5e-07,
@ -18690,7 +18701,8 @@
"supports_service_tier": true,
"supports_vision": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.4-2026-03-05": {
"cache_read_input_token_cost": 2.5e-07,
@ -18785,7 +18797,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.4-pro-2026-03-05": {
"cache_read_input_token_cost": 3e-06,
@ -18833,7 +18846,94 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.4-mini": {
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 1e-08,
"cache_read_input_token_cost_batches": 3.8e-08,
"input_cost_per_token": 7.5e-07,
"input_cost_per_token_flex": 3.75e-07,
"input_cost_per_token_batches": 3.75e-07,
"litellm_provider": "openai",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.5e-06,
"output_cost_per_token_flex": 2.25e-06,
"output_cost_per_token_batches": 2.25e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"gpt-5.4-nano": {
"cache_read_input_token_cost": 2e-08,
"cache_read_input_token_cost_flex": 1e-08,
"cache_read_input_token_cost_batches": 1e-08,
"input_cost_per_token": 2e-07,
"input_cost_per_token_flex": 1e-07,
"input_cost_per_token_batches": 1e-07,
"litellm_provider": "openai",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_flex": 6.25e-07,
"output_cost_per_token_batches": 6.25e-07,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"gpt-5-pro": {
"input_cost_per_token": 1.5e-05,
@ -18868,7 +18968,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-pro-2025-10-06": {
"input_cost_per_token": 1.5e-05,
@ -18903,7 +19004,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-2025-08-07": {
"cache_read_input_token_cost": 1.25e-07,
@ -18945,7 +19047,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-chat": {
"cache_read_input_token_cost": 1.25e-07,
@ -18979,7 +19082,8 @@
"supports_tool_choice": false,
"supports_vision": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-chat-latest": {
"cache_read_input_token_cost": 1.25e-07,
@ -19013,7 +19117,8 @@
"supports_tool_choice": false,
"supports_vision": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-codex": {
"cache_read_input_token_cost": 1.25e-07,
@ -19046,7 +19151,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1-codex": {
"cache_read_input_token_cost": 1.25e-07,
@ -19082,7 +19188,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1-codex-max": {
"cache_read_input_token_cost": 1.25e-07,
@ -19115,7 +19222,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1-codex-mini": {
"cache_read_input_token_cost": 2.5e-08,
@ -19151,7 +19259,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2-codex": {
"cache_read_input_token_cost": 1.75e-07,
@ -19187,7 +19296,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.3-codex": {
"cache_read_input_token_cost": 1.75e-07,
@ -19223,7 +19333,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-mini": {
"cache_read_input_token_cost": 2.5e-08,
@ -19265,7 +19376,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-mini-2025-08-07": {
"cache_read_input_token_cost": 2.5e-08,
@ -19307,7 +19419,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-nano": {
"cache_read_input_token_cost": 5e-09,
@ -19346,7 +19459,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-nano-2025-08-07": {
"cache_read_input_token_cost": 5e-09,
@ -19384,7 +19498,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-image-1": {
"cache_read_input_image_token_cost": 2.5e-06,
@ -36408,7 +36523,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-search-api-2025-10-14": {
"cache_read_input_token_cost": 1.25e-07,

View file

@ -39,6 +39,7 @@ from litellm.proxy._types import (
LiteLLM_EndUserTable,
Litellm_EntityType,
LiteLLM_JWTAuth,
LiteLLM_ManagedVectorStoresTable,
LiteLLM_ObjectPermissionTable,
LiteLLM_OrganizationMembershipTable,
LiteLLM_OrganizationTable,
@ -2294,6 +2295,71 @@ async def get_object_permission(
return None
@log_db_metrics
async def get_managed_vector_store_rows_by_uuids(
uuids: List[str],
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
parent_otel_span: Optional[Span] = None,
proxy_logging_obj: Optional[ProxyLogging] = None,
) -> List[LiteLLM_ManagedVectorStoresTable]:
"""
Fetch managed vector store rows by their internal UUIDs.
Follows the get_team_object / get_key_object / get_object_permission pattern:
cache-first lookup (in-memory / Redis), DB fallback only on cache miss.
Critical-path DB access must go through this helper to avoid raw Prisma
calls on the hot request path.
"""
if not uuids or prisma_client is None:
return []
result: List[LiteLLM_ManagedVectorStoresTable] = []
cache_misses: List[str] = []
for uuid in uuids:
key = "managed_vector_store_id:{}".format(uuid)
cached = await user_api_key_cache.async_get_cache(key=key)
if cached is not None:
if isinstance(cached, dict):
result.append(LiteLLM_ManagedVectorStoresTable(**cached))
elif isinstance(cached, LiteLLM_ManagedVectorStoresTable):
result.append(cached)
else:
cache_misses.append(uuid)
else:
cache_misses.append(uuid)
if not cache_misses:
return result
rows = await prisma_client.db.litellm_managedvectorstorestable.find_many(
where={"vector_store_id": {"in": cache_misses}},
take=len(cache_misses),
)
for row in rows:
row_dict = (
row.model_dump()
if hasattr(row, "model_dump")
else (row.dict() if hasattr(row, "dict") else None)
)
if not isinstance(row_dict, dict) or not row_dict:
row_dict = dict(row) if hasattr(row, "__dict__") else {}
if not row_dict:
continue
cached_obj = LiteLLM_ManagedVectorStoresTable(**row_dict)
key = "managed_vector_store_id:{}".format(cached_obj.vector_store_id)
await user_api_key_cache.async_set_cache(
key=key,
value=row_dict,
ttl=DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL,
)
result.append(cached_obj)
return result
@log_db_metrics
async def get_org_object(
org_id: str,

View file

@ -5,7 +5,7 @@
######################################################################
import asyncio
from typing import Dict, Optional, cast
from typing import Any, Dict, Optional, cast
from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response
@ -655,7 +655,7 @@ async def list_batches(
managed_files_obj, "list_user_batches"
):
verbose_proxy_logger.debug("Using managed objects table for batch listing")
response = await managed_files_obj.list_user_batches(
response = await cast(Any, managed_files_obj).list_user_batches(
user_api_key_dict=user_api_key_dict,
limit=limit,
after=after,
@ -686,8 +686,9 @@ async def list_batches(
# Encode batch IDs in the list response so clients can use
# them for retrieve/cancel/file downloads through the proxy.
if response and hasattr(response, "data") and response.data:
for batch in response.data:
response_data = getattr(response, "data", None)
if response_data:
for batch in response_data:
encode_batch_response_ids(batch, model=model_param)
verbose_proxy_logger.debug(f"Listed batches using model: {model_param}")
@ -897,7 +898,11 @@ async def cancel_batch(
# SCENARIO 3: Fallback to custom_llm_provider (uses env variables)
else:
custom_llm_provider = (
provider or data.pop("custom_llm_provider", None) or "openai"
provider
or data.pop("custom_llm_provider", None)
or get_custom_llm_provider_from_request_headers(request=request)
or get_custom_llm_provider_from_request_query(request=request)
or "openai"
)
# Extract batch_id from data to avoid "multiple values for keyword argument" error
# data was cast from CancelBatchRequest which already contains batch_id

View file

@ -30,6 +30,7 @@ from litellm.constants import (
MAX_PAYLOAD_SIZE_FOR_DEBUG_LOG,
STREAM_SSE_DATA_PREFIX,
)
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.litellm_core_utils.dd_tracing import tracer
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.llm_response_utils.get_headers import (
@ -45,7 +46,6 @@ from litellm.proxy.common_utils.callback_utils import (
from litellm.proxy.dd_span_tagger import DDSpanTagger
from litellm.proxy.route_llm_request import route_request
from litellm.proxy.utils import ProxyLogging
from litellm.integrations.custom_guardrail import CustomGuardrail
from litellm.router import Router
from litellm.types.guardrails import GuardrailEventHooks
from litellm.types.utils import ServerToolUse
@ -902,6 +902,7 @@ class ProxyBaseLLMRequestProcessing:
version: Optional[str] = None,
is_streaming_request: Optional[bool] = False,
contents: Optional[list] = None, # Add contents parameter
skip_pre_call_logic: bool = False,
) -> Any:
"""
Common request processing logic for both chat completions and responses API endpoints
@ -911,22 +912,30 @@ class ProxyBaseLLMRequestProcessing:
)
self._debug_log_request_payload()
self.data, logging_obj = await self.common_processing_pre_call_logic(
request=request,
general_settings=general_settings,
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
version=version,
proxy_config=proxy_config,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
model=model,
route_type=route_type,
llm_router=llm_router,
)
if skip_pre_call_logic:
logging_obj = self.data.get("litellm_logging_obj")
if logging_obj is None:
raise ValueError(
"skip_pre_call_logic=True requires litellm_logging_obj to be set in data. "
"Ensure common_processing_pre_call_logic was called before using this parameter."
)
else:
self.data, logging_obj = await self.common_processing_pre_call_logic(
request=request,
general_settings=general_settings,
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
version=version,
proxy_config=proxy_config,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
model=model,
route_type=route_type,
llm_router=llm_router,
)
# Defer async logging when post-call guardrails are configured so the
# StandardLoggingPayload is built after guardrails write to metadata.
@ -1082,7 +1091,7 @@ class ProxyBaseLLMRequestProcessing:
cache_hit=cache_hit,
)
logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete # type: ignore[attr-defined]
logging_obj._on_deferred_stream_complete = _on_deferred_stream_complete # type: ignore[union-attr]
if route_type == "allm_passthrough_route":
# Check if response is an async generator
@ -1095,15 +1104,15 @@ class ProxyBaseLLMRequestProcessing:
# For passthrough routes, stream directly without error parsing
# since we're dealing with raw binary data (e.g., AWS event streams)
return StreamingResponse(
content=generator,
content=generator, # type: ignore[arg-type]
status_code=status.HTTP_200_OK,
headers=custom_headers,
)
else:
# Traditional HTTP response with aiter_bytes
return StreamingResponse(
content=response.aiter_bytes(),
status_code=response.status_code,
content=response.aiter_bytes(), # type: ignore[union-attr]
status_code=response.status_code, # type: ignore[union-attr]
headers=custom_headers,
)
elif route_type == "anthropic_messages":
@ -1144,9 +1153,11 @@ class ProxyBaseLLMRequestProcessing:
# Clear the closure so guardrails run inline as before — this
# preserves blocking behavior and avoids double invocation.
if getattr(logging_obj, "_on_deferred_stream_complete", None):
logging_obj._on_deferred_stream_complete = None # type: ignore[attr-defined]
logging_obj._on_deferred_stream_complete = None # type: ignore[union-attr]
response = await proxy_logging_obj.post_call_success_hook(
data=self.data, user_api_key_dict=user_api_key_dict, response=response
data=self.data,
user_api_key_dict=user_api_key_dict,
response=response, # type: ignore[arg-type]
)
except Exception:
_exception_raised = True
@ -1159,7 +1170,7 @@ class ProxyBaseLLMRequestProcessing:
# returns before the deferred block), so _enqueue_fn is None — no-op.
_enqueue_fn = getattr(logging_obj, "_enqueue_deferred_logging", None)
if _enqueue_fn is not None:
logging_obj._enqueue_deferred_logging = None # type: ignore[attr-defined]
logging_obj._enqueue_deferred_logging = None # type: ignore[union-attr]
try:
_enqueue_fn()
except Exception as e:
@ -1180,7 +1191,7 @@ class ProxyBaseLLMRequestProcessing:
logging_obj, "_on_deferred_stream_complete", None
)
if _deferred_fn is not None:
logging_obj._on_deferred_stream_complete = None # type: ignore[attr-defined]
logging_obj._on_deferred_stream_complete = None # type: ignore[union-attr]
try:
asyncio.create_task(
logging_obj.async_success_handler(
@ -1350,18 +1361,16 @@ class ProxyBaseLLMRequestProcessing:
@staticmethod
def _has_post_call_guardrails() -> bool:
"""
Check if any registered callback is a post-call guardrail.
Uses the global litellm.callbacks list rather than per-request
should_run_guardrail() — intentionally conservative so that the
check is simple and stateless. The deferral path produces
identical logging output, just fires it slightly later, so
false-positives are harmless.
True when a guardrail explicitly registers post_call. event_hook=None
matches all hooks in should_run_guardrail but must not defer async logging
on non-streaming /chat/completions (no post_call_success_hook flush path).
"""
for cb in litellm.callbacks:
if isinstance(cb, CustomGuardrail) and cb._event_hook_is_event_type(
GuardrailEventHooks.post_call
):
if not isinstance(cb, CustomGuardrail):
continue
if cb.event_hook is None:
continue
if cb._event_hook_is_event_type(GuardrailEventHooks.post_call):
return True
return False
@ -1395,8 +1404,8 @@ class ProxyBaseLLMRequestProcessing:
from litellm.proxy.proxy_server import llm_router as _global_llm_router
from litellm.proxy.utils import (
_check_and_merge_model_level_guardrails,
unified_guardrail as _unified_guardrail,
)
from litellm.proxy.utils import unified_guardrail as _unified_guardrail
guardrail_data = _check_and_merge_model_level_guardrails(
data=captured_data, llm_router=_global_llm_router

View file

@ -40,6 +40,7 @@ from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.passthrough import BasePassthroughUtils
from litellm.proxy._types import (
CommonProxyErrors,
ConfigFieldInfo,
ConfigFieldUpdate,
LiteLLMRoutes,
@ -651,6 +652,7 @@ async def pass_through_request( # noqa: PLR0915
_parsed_body: Optional[dict] = None
# kwargs for pass through endpoint, contains metadata, litellm_params, call_type, litellm_call_id, passthrough_logging_payload
kwargs: Optional[dict] = None
logging_obj: Optional[Logging] = None
#########################################################
try:
@ -2021,13 +2023,24 @@ class InitPassThroughEndpointHelpers:
@staticmethod
def remove_endpoint_routes(endpoint_id: str):
"""Remove all routes for a specific endpoint ID from the registry"""
"""Remove all routes for a specific endpoint ID from the registry
and clean up corresponding entries from LiteLLMRoutes.openai_routes."""
keys_to_remove = [
key
for key, value in _registered_pass_through_routes.items()
if value["endpoint_id"] == endpoint_id
]
for key in keys_to_remove:
route_info = _registered_pass_through_routes[key]
path = route_info.get("path")
if isinstance(path, str):
openai_routes = LiteLLMRoutes.openai_routes.value
if path in openai_routes:
openai_routes.remove(path)
if route_info.get("type") == "subpath":
wildcard_path = path.rstrip("/") + "/*"
if wildcard_path in openai_routes:
openai_routes.remove(wildcard_path)
del _registered_pass_through_routes[key]
verbose_proxy_logger.debug(
"Removed pass-through route from registry: %s", key
@ -2143,6 +2156,102 @@ def _get_combined_pass_through_endpoints(
return pass_through_endpoints + config_pass_through_endpoints
async def _register_pass_through_endpoint(
endpoint: Union[Dict[str, Any], PassThroughGenericEndpoint],
app: FastAPI,
premium_user: bool,
visited_endpoints: set[str],
) -> None:
endpoint_data: Dict[str, Any]
if isinstance(endpoint, PassThroughGenericEndpoint):
endpoint_data = endpoint.model_dump()
else:
endpoint_data = endpoint
if endpoint_data.get("id") is None:
endpoint_data["id"] = str(uuid.uuid4())
endpoint_id = cast(str, endpoint_data["id"])
target = endpoint_data.get("target")
path = endpoint_data.get("path")
if path is None:
raise ValueError("Path is required for pass-through endpoint")
custom_headers = await set_env_variables_in_header(
custom_headers=endpoint_data.get("headers")
)
forward_headers = endpoint_data.get("forward_headers")
merge_query_params = endpoint_data.get("merge_query_params")
default_query_params = endpoint_data.get("default_query_params")
auth = endpoint_data.get("auth")
dependencies = None
if auth is not None and str(auth).lower() == "true":
if premium_user is not True:
raise ValueError(
"Error Setting Authentication on Pass Through Endpoint: {}".format(
CommonProxyErrors.not_premium_user.value
)
)
dependencies = [Depends(user_api_key_auth)]
if path not in LiteLLMRoutes.openai_routes.value:
LiteLLMRoutes.openai_routes.value.append(path)
if target is None:
return
guardrails = endpoint_data.get("guardrails")
methods = endpoint_data.get("methods")
cost_per_request = endpoint_data.get("cost_per_request")
verbose_proxy_logger.debug(
"Initializing pass through endpoint: %s (ID: %s)", path, endpoint_id
)
InitPassThroughEndpointHelpers.add_exact_path_route(
app=app,
path=path,
target=target,
custom_headers=custom_headers,
forward_headers=forward_headers,
merge_query_params=merge_query_params,
dependencies=dependencies,
cost_per_request=cost_per_request,
endpoint_id=endpoint_id,
guardrails=guardrails,
methods=methods,
default_query_params=default_query_params,
)
methods_for_key = methods if methods else ["GET", "POST", "PUT", "DELETE", "PATCH"]
methods_str = ",".join(sorted(methods_for_key))
visited_endpoints.add(f"{endpoint_id}:exact:{path}:{methods_str}")
if endpoint_data.get("include_subpath", False) is True:
if auth is not None and str(auth).lower() == "true":
wildcard_path = path.rstrip("/") + "/*"
if wildcard_path not in LiteLLMRoutes.openai_routes.value:
LiteLLMRoutes.openai_routes.value.append(wildcard_path)
InitPassThroughEndpointHelpers.add_subpath_route(
app=app,
path=path,
target=target,
custom_headers=custom_headers,
forward_headers=forward_headers,
merge_query_params=merge_query_params,
dependencies=dependencies,
cost_per_request=cost_per_request,
endpoint_id=endpoint_id,
guardrails=guardrails,
methods=methods,
default_query_params=default_query_params,
)
visited_endpoints.add(f"{endpoint_id}:subpath:{path}:{methods_str}")
verbose_proxy_logger.debug(
"Added new pass through endpoint: %s (ID: %s)", path, endpoint_id
)
async def initialize_pass_through_endpoints(
pass_through_endpoints: Union[List[Dict], List[PassThroughGenericEndpoint]],
):
@ -2159,10 +2268,7 @@ async def initialize_pass_through_endpoints(
Returns:
None
"""
from litellm._uuid import uuid
verbose_proxy_logger.debug("initializing pass through endpoints")
from litellm.proxy._types import CommonProxyErrors, LiteLLMRoutes
from litellm.proxy.proxy_server import (
app,
config_passthrough_endpoints,
@ -2189,98 +2295,14 @@ async def initialize_pass_through_endpoints(
InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes()
)
visited_endpoints = set()
visited_endpoints: set[str] = set()
for endpoint in combined_pass_through_endpoints:
if isinstance(endpoint, PassThroughGenericEndpoint):
endpoint = endpoint.model_dump()
# Auto-generate ID for backwards compatibility if not present
if endpoint.get("id") is None:
endpoint["id"] = str(uuid.uuid4())
# Get the endpoint_id as a string (guaranteed to be set at this point)
endpoint_id: str = endpoint["id"]
_target = endpoint.get("target", None)
_path: Optional[str] = endpoint.get("path", None)
if _path is None:
raise ValueError("Path is required for pass-through endpoint")
_custom_headers = endpoint.get("headers", None)
_custom_headers = await set_env_variables_in_header(
custom_headers=_custom_headers
)
_forward_headers = endpoint.get("forward_headers", None)
_merge_query_params = endpoint.get("merge_query_params", None)
_default_query_params = endpoint.get("default_query_params", None)
_auth = endpoint.get("auth", None)
_dependencies = None
if _auth is not None and str(_auth).lower() == "true":
if premium_user is not True:
raise ValueError(
"Error Setting Authentication on Pass Through Endpoint: {}".format(
CommonProxyErrors.not_premium_user.value
)
)
_dependencies = [Depends(user_api_key_auth)]
LiteLLMRoutes.openai_routes.value.append(_path)
if _target is None:
continue
# Get guardrails config if present
_guardrails = endpoint.get("guardrails", None)
# Get methods list if present (None means all methods for backward compatibility)
_methods = endpoint.get("methods", None)
# Add exact path route
verbose_proxy_logger.debug(
"Initializing pass through endpoint: %s (ID: %s)", _path, endpoint_id
)
InitPassThroughEndpointHelpers.add_exact_path_route(
await _register_pass_through_endpoint(
endpoint=endpoint,
app=app,
path=_path,
target=_target,
custom_headers=_custom_headers,
forward_headers=_forward_headers,
merge_query_params=_merge_query_params,
dependencies=_dependencies,
cost_per_request=endpoint.get("cost_per_request", None),
endpoint_id=endpoint_id,
guardrails=_guardrails,
methods=_methods,
default_query_params=_default_query_params,
)
# Generate route key with methods for tracking
methods_for_key = (
_methods if _methods else ["GET", "POST", "PUT", "DELETE", "PATCH"]
)
methods_str = ",".join(sorted(methods_for_key))
visited_endpoints.add(f"{endpoint_id}:exact:{_path}:{methods_str}")
# Add wildcard route for sub-paths
if endpoint.get("include_subpath", False) is True:
InitPassThroughEndpointHelpers.add_subpath_route(
app=app,
path=_path,
target=_target,
custom_headers=_custom_headers,
forward_headers=_forward_headers,
merge_query_params=_merge_query_params,
dependencies=_dependencies,
cost_per_request=endpoint.get("cost_per_request", None),
endpoint_id=endpoint_id,
guardrails=_guardrails,
methods=_methods,
default_query_params=_default_query_params,
)
visited_endpoints.add(f"{endpoint_id}:subpath:{_path}:{methods_str}")
verbose_proxy_logger.debug(
"Added new pass through endpoint: %s (ID: %s)", _path, endpoint_id
premium_user=premium_user,
visited_endpoints=visited_endpoints,
)
# remove the ones that are not visited from the list

View file

@ -119,6 +119,35 @@ async def responses_api(
f"Starting background response with polling for model={data.get('model')}"
)
# Run pre-call checks (rate limits, guardrails, budget) BEFORE creating
# polling ID. This ensures rate-limited requests get a synchronous 429
# instead of a polling ID that immediately fails in the background task.
processor = ProxyBaseLLMRequestProcessing(data=data)
try:
data, _logging_obj = await processor.common_processing_pre_call_logic(
request=request,
general_settings=general_settings,
proxy_logging_obj=proxy_logging_obj,
user_api_key_dict=user_api_key_dict,
version=version,
proxy_config=proxy_config,
user_model=user_model,
user_temperature=user_temperature,
user_request_timeout=user_request_timeout,
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
model=None,
route_type="aresponses",
llm_router=llm_router,
)
except Exception as e:
raise await processor._handle_llm_api_exception(
e=e,
user_api_key_dict=user_api_key_dict,
proxy_logging_obj=proxy_logging_obj,
version=version,
)
# Initialize polling handler with configured TTL (from global config)
polling_handler = ResponsePollingHandler(
redis_cache=redis_usage_cache,
@ -134,7 +163,9 @@ async def responses_api(
request_data=data,
)
# Start background task to stream and update cache
# Start background task to stream and update cache.
# Pass pre-processed data so the background task skips pre-call logic
# (rate limits, guardrails already checked above).
asyncio.create_task(
background_streaming_task(
polling_id=polling_id,

View file

@ -65,7 +65,9 @@ async def background_streaming_task( # noqa: PLR0915
# Create processor
processor = ProxyBaseLLMRequestProcessing(data=data)
# Make streaming request
# Make streaming request.
# Pre-call checks (rate limits, guardrails, budget) were already run
# before polling ID creation, so skip them here to avoid double-counting.
response = await processor.base_process_llm_request(
request=request,
fastapi_response=fastapi_response,
@ -83,6 +85,7 @@ async def background_streaming_task( # noqa: PLR0915
user_max_tokens=user_max_tokens,
user_api_base=user_api_base,
version=version,
skip_pre_call_logic=True,
)
# Process streaming response following OpenAI events format

View file

@ -0,0 +1,592 @@
"""
Emulated file_search for providers that don't support the tool natively.
Flow:
1. Convert file_search tools to a single function tool definition.
2. Call the provider with the function tool.
3. If the provider issues a file_search function_call, execute vector search
via litellm.vector_stores.main.asearch().
4. Feed results back and get the final answer.
5. Wrap everything in OpenAI Responses-API format:
[file_search_call output item] + [message output item with file_citation annotations]
"""
import json
import time
import uuid
from typing import Any, Dict, Iterable, List, Optional, Tuple, Union, cast
from litellm._logging import verbose_logger
from litellm.types.llms.openai import ResponseOutputItem, ResponsesAPIResponse
from litellm.types.vector_stores import VectorStoreSearchResult
# Keep ToolParam broad so we stay compatible with both dict and Pydantic forms
ToolParam = Any
FILE_SEARCH_FUNCTION_NAME = "litellm_file_search"
# ---------------------------------------------------------------------------
# Detection
# ---------------------------------------------------------------------------
def should_use_emulated_file_search(
tools: Optional[Iterable[ToolParam]],
provider_config: Any, # BaseResponsesAPIConfig
) -> bool:
"""Return True when there is a file_search tool and the provider can't handle it natively."""
if not tools:
return False
has_fs = any(isinstance(t, dict) and t.get("type") == "file_search" for t in tools)
if not has_fs:
return False
return provider_config is None or not provider_config.supports_native_file_search()
# ---------------------------------------------------------------------------
# Tool conversion
# ---------------------------------------------------------------------------
def _build_function_tool(vector_store_ids: List[str]) -> Dict[str, Any]:
"""
Create a Responses API function-tool definition that describes file search.
The function accepts one or more natural-language queries (like OpenAI's native
file_search); LiteLLM runs the actual vector search against the configured
vector stores.
Note: Uses Responses API format (name/description/parameters at top level),
NOT Chat Completion format (nested under "function"), so that the
LiteLLMCompletionResponsesConfig transformation picks up name and description.
"""
return {
"type": "function",
"name": FILE_SEARCH_FUNCTION_NAME,
"description": (
"Search the knowledge base for information relevant to the query. "
"Use this whenever you need to look up specific facts, documents, "
"or content from the vector store. You can provide multiple queries "
"to search for different aspects of the information."
),
"parameters": {
"type": "object",
"properties": {
"queries": {
"type": "array",
"items": {"type": "string"},
"description": (
"One or more search queries to look up in the vector store. "
"Multiple queries help find comprehensive information from "
"different angles."
),
},
"vector_store_id": {
"type": "string",
"description": "ID of the vector store to search.",
"enum": vector_store_ids,
},
},
"required": ["queries"],
},
}
def _replace_file_search_tools(
tools: Optional[Iterable[ToolParam]],
) -> Tuple[List[Dict[str, Any]], List[str]]:
"""
Replace all file_search tools with a single function tool.
Returns:
(new_tools_list, all_vector_store_ids)
"""
non_file_search: List[Dict[str, Any]] = []
vector_store_ids: List[str] = []
for tool in tools or []:
if isinstance(tool, dict) and tool.get("type") == "file_search":
ids = tool.get("vector_store_ids") or []
vector_store_ids.extend(ids)
else:
non_file_search.append(tool)
# Deduplicate while preserving order
unique_ids: List[str] = list(dict.fromkeys(vector_store_ids))
if unique_ids:
non_file_search.append(_build_function_tool(unique_ids))
return non_file_search, unique_ids
# ---------------------------------------------------------------------------
# Search execution
# ---------------------------------------------------------------------------
async def _run_vector_searches(
queries: List[str],
vector_store_ids: List[str],
) -> Tuple[List[str], List[VectorStoreSearchResult]]:
"""
Run `asearch` against all vector stores for all queries and collect results.
Args:
queries: List of search queries to execute (like OpenAI's multi-query approach)
vector_store_ids: Vector store IDs to search
Returns:
(queries_list, combined_results)
"""
import litellm.vector_stores.main as vs_main
all_results: List[VectorStoreSearchResult] = []
ids_to_search = vector_store_ids
# Execute each query against all vector stores
for query in queries:
for vs_id in ids_to_search:
try:
response = await vs_main.asearch(
vector_store_id=vs_id,
query=query,
)
results_data = (
response.get("data")
if isinstance(response, dict)
else getattr(response, "data", None)
)
if results_data:
all_results.extend(results_data)
except Exception as exc:
verbose_logger.warning(
"file_search emulated: search failed for query='%s', vector_store_id='%s': %s",
query,
vs_id,
exc,
)
return queries, all_results
# ---------------------------------------------------------------------------
# Result formatting
# ---------------------------------------------------------------------------
def _get_field(result: Any, key: str, default: Any = None) -> Any:
"""Read a field from either a dict/TypedDict or an attribute-based object."""
if isinstance(result, dict):
return result.get(key, default)
return getattr(result, key, default)
def _format_search_results_as_tool_output(
results: List[VectorStoreSearchResult],
) -> str:
"""Serialize search results into a string to pass back as the tool's output."""
if not results:
return "No results found in the vector store."
parts: List[str] = []
for i, result in enumerate(results, 1):
score = _get_field(result, "score")
file_id = _get_field(result, "file_id")
filename = _get_field(result, "filename")
content_items = _get_field(result, "content") or []
text_chunks = [
c.get("text", "") if isinstance(c, dict) else getattr(c, "text", "")
for c in content_items
]
text = " ".join(t for t in text_chunks if t)
header = f"[Result {i}"
if filename:
header += f" | {filename}"
if file_id:
header += f" | file_id={file_id}"
if score is not None:
header += f" | score={score:.3f}"
header += "]"
parts.append(f"{header}\n{text}")
return "\n\n".join(parts)
def _build_search_results_for_include(
results: List[VectorStoreSearchResult],
) -> List[Dict[str, Any]]:
"""
Convert VectorStoreSearchResult objects to the format expected in
file_search_call.search_results (mirrors OpenAI's include= format).
All chunks are returned — no deduplication by file_id — matching the
behaviour of OpenAI's native file_search which surfaces every relevant
chunk even when multiple chunks originate from the same document.
"""
formatted: List[Dict[str, Any]] = []
for result in results:
file_id = _get_field(result, "file_id") or ""
content_items = _get_field(result, "content") or []
text_chunks = [
c.get("text", "") if isinstance(c, dict) else getattr(c, "text", "")
for c in content_items
]
text = " ".join(t for t in text_chunks if t)
formatted.append(
{
"file_id": file_id,
"filename": _get_field(result, "filename") or "",
"score": _get_field(result, "score"),
"text": text,
"attributes": _get_field(result, "attributes") or {},
}
)
return formatted
def _build_file_search_call_output(
call_id: str,
queries: List[str],
results: Optional[List[VectorStoreSearchResult]] = None,
include_search_results: bool = False,
) -> Dict[str, Any]:
"""Build the file_search_call output item (mirrors OpenAI's format).
Args:
call_id: Unique ID for this file_search call.
queries: List of search queries used.
results: The raw search results (used when include_search_results=True).
include_search_results: Populate search_results when the caller passed
``include=["file_search_call.results"]``.
"""
search_results = None
if include_search_results and results:
search_results = _build_search_results_for_include(results)
return {
"type": "file_search_call",
"id": call_id,
"status": "completed",
"queries": queries,
"search_results": search_results,
}
def _build_file_citation_annotations(
results: List[VectorStoreSearchResult],
text: str,
) -> List[Dict[str, Any]]:
"""
Build file_citation annotations for the text.
Each result with a file_id gets a citation at the end of the text.
"""
annotations: List[Dict[str, Any]] = []
index = len(text) # cite at end of text block
seen_file_ids: set = set()
for result in results:
file_id = _get_field(result, "file_id")
filename = _get_field(result, "filename")
if not file_id or file_id in seen_file_ids:
continue
seen_file_ids.add(file_id)
annotations.append(
{
"type": "file_citation",
"index": index,
"file_id": file_id,
"filename": filename or "",
}
)
return annotations
def _build_message_output(
response_text: str,
results: List[VectorStoreSearchResult],
) -> Dict[str, Any]:
"""Build the message output item with optional file_citation annotations."""
annotations = _build_file_citation_annotations(results, response_text)
return {
"type": "message",
"role": "assistant",
"content": [
{
"type": "output_text",
"text": response_text,
"annotations": annotations,
}
],
}
def _extract_text_from_responses_output(response: ResponsesAPIResponse) -> str:
"""Pull the assistant's text from the provider's response."""
for item in response.output:
item_type = (
item.get("type") if isinstance(item, dict) else getattr(item, "type", None)
)
if item_type == "message":
content = (
item.get("content")
if isinstance(item, dict)
else getattr(item, "content", [])
)
for block in content or []:
block_type = (
block.get("type")
if isinstance(block, dict)
else getattr(block, "type", None)
)
if block_type == "output_text":
raw = (
block.get("text")
if isinstance(block, dict)
else getattr(block, "text", "")
)
return str(raw) if raw is not None else ""
return ""
def _synthesize_responses_api_response(
original_response: ResponsesAPIResponse,
file_search_call_output: Dict[str, Any],
message_output: Dict[str, Any],
first_response: Optional[ResponsesAPIResponse] = None,
) -> ResponsesAPIResponse:
"""
Return a new ResponsesAPIResponse with:
output[0] = file_search_call item
output[1] = message item (with citations)
When first_response is provided, its response_cost is accumulated into the
synthesized _hidden_params so that billing callbacks see the total cost of
both provider calls that the emulated flow makes.
"""
synthesized_output: List[Dict[str, Any]] = [file_search_call_output, message_output]
synthesized = ResponsesAPIResponse(
id=getattr(original_response, "id", f"resp_{uuid.uuid4().hex}"),
object="response",
created_at=getattr(original_response, "created_at", int(time.time())),
status="completed",
model=getattr(original_response, "model", ""),
output=cast(
List[Union[ResponseOutputItem, Dict[str, Any]]], synthesized_output
),
usage=getattr(original_response, "usage", None),
error=None,
)
if hasattr(original_response, "_hidden_params"):
hidden = dict(getattr(original_response, "_hidden_params") or {})
if first_response is not None and hasattr(first_response, "_hidden_params"):
first_hidden = getattr(first_response, "_hidden_params") or {}
first_cost = (
first_hidden.get("response_cost")
if isinstance(first_hidden, dict)
else getattr(first_hidden, "response_cost", None)
)
if first_cost is not None:
current_cost = (
hidden.get("response_cost") if isinstance(hidden, dict) else 0
)
hidden["response_cost"] = (current_cost or 0) + first_cost
synthesized._hidden_params = hidden
return synthesized
# ---------------------------------------------------------------------------
# Main entry point
# ---------------------------------------------------------------------------
async def _call_aresponses(
input, model, tools, **kwargs
): # pragma: no cover – thin wrapper for patching in tests
from litellm.responses.main import aresponses
return await aresponses(input=input, model=model, tools=tools, **kwargs)
def _prepare_emulated_file_search_call(
kwargs: Dict[str, Any],
) -> Tuple[bool, Dict[str, Any]]:
include_items: List[str] = list(kwargs.get("include") or [])
include_search_results = "file_search_call.results" in include_items
original_stream = kwargs.get("stream")
updated_kwargs = kwargs
if original_stream:
verbose_logger.debug(
"Streaming is not yet supported for emulated file_search. "
"Disabling stream for this request."
)
updated_kwargs = {**kwargs, "stream": False}
return include_search_results, updated_kwargs
async def aresponses_with_emulated_file_search(
input: Any,
model: str,
tools: Optional[Iterable[ToolParam]] = None,
# Pass-through params — forwarded as-is to the underlying aresponses call
**kwargs: Any,
) -> ResponsesAPIResponse:
"""
Emulated file_search for providers that don't support it natively.
Replaces file_search tools with a function tool, intercepts the tool call,
runs vector search, and synthesizes an OpenAI-format response.
"""
# Determine whether caller wants search_results populated in the output.
_include_search_results, kwargs = _prepare_emulated_file_search_call(kwargs=kwargs)
# 1. Replace file_search tools with function tool
transformed_tools, all_vs_ids = _replace_file_search_tools(tools)
# 2. First provider call — provider will call the file_search function.
# Mark as an internal sub-call so wrapper_async skips billing callbacks;
# the parent litellm_logging_obj (propagated via kwargs) fires once at the end.
first_response: ResponsesAPIResponse = cast(
ResponsesAPIResponse,
await _call_aresponses(
input=input,
model=model,
tools=transformed_tools or None,
**{**kwargs, "_is_litellm_internal_call": True},
),
)
# 3. Look for a file_search function_call in the output
file_search_calls = [
item
for item in first_response.output
if (
isinstance(item, dict)
and item.get("type") == "function_call"
and item.get("name") == FILE_SEARCH_FUNCTION_NAME
)
or (
hasattr(item, "type")
and getattr(item, "type") == "function_call"
and getattr(item, "name", None) == FILE_SEARCH_FUNCTION_NAME
)
]
if not file_search_calls:
# Provider answered without calling the tool (e.g. it had enough context).
# Return as-is wrapped in OpenAI format.
call_id = f"fs_{uuid.uuid4().hex[:24]}"
response_text = _extract_text_from_responses_output(first_response)
return _synthesize_responses_api_response(
original_response=first_response,
file_search_call_output=_build_file_search_call_output(
call_id=call_id,
queries=[str(input)],
results=None,
include_search_results=False,
),
message_output=_build_message_output(response_text, []),
)
# 4. Execute each file_search tool call
tool_results: List[Dict[str, Any]] = []
all_queries: List[str] = []
all_results: List[VectorStoreSearchResult] = []
file_search_call_id = f"fs_{uuid.uuid4().hex[:24]}"
for tool_call in file_search_calls:
if isinstance(tool_call, dict):
call_id = str(
tool_call.get("call_id") or tool_call.get("id") or file_search_call_id
)
raw_args = tool_call.get("arguments") or "{}"
else:
raw_call_id = (
getattr(tool_call, "call_id", None)
or getattr(tool_call, "id", None)
or file_search_call_id
)
call_id = str(raw_call_id)
raw_args = getattr(tool_call, "arguments", "{}") or "{}"
try:
args = json.loads(raw_args) if isinstance(raw_args, str) else raw_args
except json.JSONDecodeError:
args = {}
# Extract queries array (OpenAI-style multi-query support)
queries_from_call = args.get("queries")
if not queries_from_call:
# Fallback: check for single "query" field (backward compat)
single_query = args.get("query")
queries_from_call = [single_query] if single_query else [str(input)]
elif not isinstance(queries_from_call, list):
queries_from_call = [str(queries_from_call)]
vs_id_arg = args.get("vector_store_id")
vs_ids_for_call = [vs_id_arg] if vs_id_arg else all_vs_ids
queries, results = await _run_vector_searches(
queries=queries_from_call,
vector_store_ids=vs_ids_for_call,
)
all_queries.extend(queries)
all_results.extend(results)
tool_results.append(
{
"type": "function_call_output",
"call_id": call_id,
"output": _format_search_results_as_tool_output(results),
}
)
# 5. Build follow-up input: original messages + ALL first-response output items + tool results
# Including all output items (text blocks, reasoning, non-file-search calls) ensures providers
# like Anthropic that emit text before the tool call have complete conversation context.
# Serialize Pydantic model instances to plain dicts so the transformation layer can call .get().
original_input_items = (
list(input)
if isinstance(input, (list, tuple))
else [{"role": "user", "content": str(input)}]
)
first_response_output_items: List[Any] = []
for _item in first_response.output:
if isinstance(_item, dict):
first_response_output_items.append(_item)
elif hasattr(_item, "model_dump"):
first_response_output_items.append(_item.model_dump(exclude_none=True)) # type: ignore[union-attr]
else:
first_response_output_items.append(_item)
follow_up_input = original_input_items + first_response_output_items + tool_results
# 6. Follow-up call — provider writes the final answer given search results.
# Also an internal sub-call; billing is suppressed so the outer call fires once.
final_response: ResponsesAPIResponse = cast(
ResponsesAPIResponse,
await _call_aresponses(
input=follow_up_input,
model=model,
tools=None, # no tools needed for the answer step
**{**kwargs, "_is_litellm_internal_call": True},
),
)
# 7. Synthesize OpenAI-format output
response_text = _extract_text_from_responses_output(final_response)
return _synthesize_responses_api_response(
original_response=final_response,
file_search_call_output=_build_file_search_call_output(
call_id=file_search_call_id,
queries=all_queries or [str(input)],
results=all_results,
include_search_results=_include_search_results,
),
message_output=_build_message_output(response_text, all_results),
first_response=first_response,
)

View file

@ -37,6 +37,7 @@ from litellm.responses.litellm_completion_transformation.handler import (
)
from litellm.responses.utils import ResponsesAPIRequestUtils
from litellm.types.llms.openai import (
AllMessageValues,
PromptObject,
Reasoning,
ResponseIncludable,
@ -72,6 +73,13 @@ litellm_completion_transformation_handler = LiteLLMCompletionTransformationHandl
#################################################
def _has_file_search_tool(tools: Optional[Any]) -> bool:
"""Return True if any tool in the list has type 'file_search'."""
if not tools:
return False
return any(isinstance(t, dict) and t.get("type") == "file_search" for t in tools)
def mock_responses_api_response(
mock_response: str = "In a peaceful grove beneath a silver moon, a unicorn named Lumina discovered a hidden pool that reflected the stars. As she dipped her horn into the water, the pool began to shimmer, revealing a pathway to a magical realm of endless night skies. Filled with wonder, Lumina whispered a wish for all who dream to find their own hidden magic, and as she glanced back, her hoofprints sparkled like stardust.",
):
@ -463,6 +471,53 @@ async def aresponses(
# Update local_vars with detected provider (fixes #19782)
local_vars["custom_llm_provider"] = custom_llm_provider
#########################################################
# ASYNC PROMPT MANAGEMENT
# Run the async hook here so async-only prompt loggers are honoured.
# Then pop prompt_id from kwargs so the sync responses() path does NOT
# re-run the hook (which would double-prepend template messages).
# Pass merged_optional_params via an internal kwarg so responses()
# can apply them to local_vars without re-invoking the hook.
#########################################################
litellm_logging_obj = kwargs.get("litellm_logging_obj", None)
prompt_id = cast(Optional[str], kwargs.get("prompt_id", None))
prompt_variables = cast(Optional[dict], kwargs.get("prompt_variables", None))
original_model = model
if isinstance(
litellm_logging_obj, LiteLLMLoggingObj
) and litellm_logging_obj.should_run_prompt_management_hooks(
prompt_id=prompt_id, non_default_params=kwargs
):
if isinstance(input, str):
client_input: List[AllMessageValues] = [
{"role": "user", "content": input}
]
else:
client_input = [
item # type: ignore[misc]
for item in input
if isinstance(item, dict) and "role" in item
]
(
model,
merged_input,
merged_optional_params,
) = await litellm_logging_obj.async_get_chat_completion_prompt(
model=model,
messages=client_input,
non_default_params=kwargs,
prompt_id=prompt_id,
prompt_variables=prompt_variables,
prompt_label=kwargs.get("prompt_label", None),
prompt_version=kwargs.get("prompt_version", None),
)
input = cast(Union[str, ResponseInputParam], merged_input)
if model != original_model:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
kwargs.pop("prompt_id", None)
kwargs["_async_prompt_merged_params"] = merged_optional_params
func = partial(
responses,
input=input,
@ -531,6 +586,125 @@ async def aresponses(
)
def _apply_prompt_management_to_responses_call(
input: Union[str, ResponseInputParam],
model: str,
custom_llm_provider: Optional[str],
litellm_logging_obj: Optional[LiteLLMLoggingObj],
kwargs: Dict[str, Any],
local_vars: Dict[str, Any],
) -> tuple[Union[str, ResponseInputParam], str, Optional[str]]:
async_merged = kwargs.pop("_async_prompt_merged_params", None)
if async_merged is not None:
for key, value in async_merged.items():
local_vars[key] = value
return input, model, custom_llm_provider
prompt_id = cast(Optional[str], kwargs.get("prompt_id", None))
prompt_variables = cast(Optional[dict], kwargs.get("prompt_variables", None))
original_model = model
if isinstance(input, str):
client_input: List[AllMessageValues] = [{"role": "user", "content": input}]
else:
client_input = [
item # type: ignore[misc]
for item in input
if isinstance(item, dict) and "role" in item
]
if isinstance(
litellm_logging_obj, LiteLLMLoggingObj
) and litellm_logging_obj.should_run_prompt_management_hooks(
prompt_id=prompt_id, non_default_params=kwargs
):
(
model,
merged_input,
merged_optional_params,
) = litellm_logging_obj.get_chat_completion_prompt(
model=model,
messages=client_input,
non_default_params=kwargs,
prompt_id=prompt_id,
prompt_variables=prompt_variables,
prompt_label=kwargs.get("prompt_label", None),
prompt_version=kwargs.get("prompt_version", None),
)
input = cast(Union[str, ResponseInputParam], merged_input)
local_vars["input"] = input
local_vars["model"] = model
if model != original_model:
_, custom_llm_provider, _, _ = litellm.get_llm_provider(model=model)
local_vars["custom_llm_provider"] = custom_llm_provider
for key, value in merged_optional_params.items():
local_vars[key] = value
return input, model, custom_llm_provider
def _resolve_model_provider_for_responses(
model: str,
custom_llm_provider: Optional[str],
litellm_params: GenericLiteLLMParams,
local_vars: Dict[str, Any],
) -> tuple[str, Optional[str]]:
(
model,
custom_llm_provider,
dynamic_api_key,
dynamic_api_base,
) = litellm.get_llm_provider(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
)
local_vars["custom_llm_provider"] = custom_llm_provider
if dynamic_api_key is not None:
litellm_params.api_key = dynamic_api_key
if dynamic_api_base is not None:
litellm_params.api_base = dynamic_api_base
return model, custom_llm_provider
def _apply_managed_file_id_mapping(
input: Union[str, ResponseInputParam],
tools: Optional[Iterable[ToolParam]],
kwargs: Dict[str, Any],
local_vars: Dict[str, Any],
) -> tuple[Union[str, ResponseInputParam], Optional[Iterable[ToolParam]]]:
model_file_id_mapping = kwargs.get("model_file_id_mapping")
model_info_id = (
kwargs.get("model_info", {}).get("id")
if isinstance(kwargs.get("model_info"), dict)
else None
)
input = cast(
Union[str, ResponseInputParam],
update_responses_input_with_model_file_ids(
input=input,
model_id=model_info_id,
model_file_id_mapping=model_file_id_mapping,
),
)
local_vars["input"] = input
if tools:
tools = cast(
Optional[Iterable[ToolParam]],
update_responses_tools_with_model_file_ids(
tools=cast(Optional[List[Dict[str, Any]]], tools),
model_id=model_info_id,
model_file_id_mapping=model_file_id_mapping,
),
)
local_vars["tools"] = tools
return input, tools
@client
def responses(
input: Union[str, ResponseInputParam],
@ -602,59 +776,35 @@ def responses(
mock_response=litellm_params.mock_response
)
(
model,
custom_llm_provider,
dynamic_api_key,
dynamic_api_base,
) = litellm.get_llm_provider(
model, custom_llm_provider = _resolve_model_provider_for_responses(
model=model,
custom_llm_provider=custom_llm_provider,
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
litellm_params=litellm_params,
local_vars=local_vars,
)
# Update local_vars with detected provider (fixes #19782)
local_vars["custom_llm_provider"] = custom_llm_provider
# Use dynamic credentials from get_llm_provider (e.g., when use_litellm_proxy=True)
if dynamic_api_key is not None:
litellm_params.api_key = dynamic_api_key
if dynamic_api_base is not None:
litellm_params.api_base = dynamic_api_base
#########################################################
# PROMPT MANAGEMENT
# If aresponses() already ran the async hook, it pops prompt_id and
# passes the result via _async_prompt_merged_params — apply those
# directly and skip the sync hook to avoid double-merging.
#########################################################
input, model, custom_llm_provider = _apply_prompt_management_to_responses_call(
input=input,
model=model,
custom_llm_provider=custom_llm_provider,
litellm_logging_obj=litellm_logging_obj,
kwargs=kwargs,
local_vars=local_vars,
)
#########################################################
# Update input and tools with provider-specific file IDs if managed files are used
#########################################################
model_file_id_mapping = kwargs.get("model_file_id_mapping")
model_info_id = (
kwargs.get("model_info", {}).get("id")
if isinstance(kwargs.get("model_info"), dict)
else None
input, tools = _apply_managed_file_id_mapping(
input=input, tools=tools, kwargs=kwargs, local_vars=local_vars
)
input = cast(
Union[str, ResponseInputParam],
update_responses_input_with_model_file_ids(
input=input,
model_id=model_info_id,
model_file_id_mapping=model_file_id_mapping,
),
)
local_vars["input"] = input
# Update tools with provider-specific file IDs if needed
if tools:
tools = cast(
Optional[Iterable[ToolParam]],
update_responses_tools_with_model_file_ids(
tools=cast(Optional[List[Dict[str, Any]]], tools),
model_id=model_info_id,
model_file_id_mapping=model_file_id_mapping,
),
)
local_vars["tools"] = tools
#########################################################
# Native MCP Responses API
#########################################################
@ -692,12 +842,16 @@ def responses(
return run_async_function(aresponses_api_with_mcp, **mcp_call_kwargs)
# get provider config
responses_api_provider_config: Optional[
BaseResponsesAPIConfig
] = ProviderConfigManager.get_provider_responses_api_config(
model=model,
provider=custom_llm_provider,
)
responses_api_provider_config: Optional[BaseResponsesAPIConfig]
if custom_llm_provider is None:
responses_api_provider_config = None
else:
responses_api_provider_config = (
ProviderConfigManager.get_provider_responses_api_config(
model=model,
provider=custom_llm_provider,
)
)
local_vars.update(kwargs)
# Map reasoning_effort (from litellm_params/proxy config) to reasoning when not set
@ -715,6 +869,56 @@ def responses(
)
)
if _has_file_search_tool(tools) and (
responses_api_provider_config is None
or not responses_api_provider_config.supports_native_file_search()
):
from litellm.responses.file_search.emulated_handler import (
aresponses_with_emulated_file_search,
)
_internal_skip = {"litellm_call_id", "aresponses"}
emulated_kwargs = {
"include": include,
"instructions": instructions,
"max_output_tokens": max_output_tokens,
"prompt": prompt,
"metadata": metadata,
"parallel_tool_calls": parallel_tool_calls,
"previous_response_id": previous_response_id,
"reasoning": reasoning,
"store": store,
"background": background,
"stream": stream,
"temperature": temperature,
"text": text,
"tool_choice": tool_choice,
"top_p": top_p,
"truncation": truncation,
"user": user,
"service_tier": service_tier,
"safety_identifier": safety_identifier,
"text_format": text_format,
"allowed_openai_params": allowed_openai_params,
"extra_headers": extra_headers,
"extra_query": extra_query,
"extra_body": extra_body,
"timeout": timeout,
"custom_llm_provider": custom_llm_provider,
**{k: v for k, v in kwargs.items() if k not in _internal_skip},
}
if _is_async:
return aresponses_with_emulated_file_search(
input=input, model=model, tools=tools, **emulated_kwargs
)
return run_async_function(
aresponses_with_emulated_file_search,
input=input,
model=model,
tools=tools,
**emulated_kwargs,
)
if responses_api_provider_config is None:
return litellm_completion_transformation_handler.response_api_handler(
model=model,
@ -758,6 +962,9 @@ def responses(
)
# Call the handler with _is_async flag instead of directly calling the async handler
if custom_llm_provider is None:
raise ValueError("custom_llm_provider is required but passed as None")
response = base_llm_http_handler.response_api_handler(
model=model,
input=input,

View file

@ -301,6 +301,7 @@ class Router:
RouterGeneralSettings
] = RouterGeneralSettings(),
deployment_affinity_ttl_seconds: int = 3600,
model_group_affinity_config: Optional[Dict[str, List[str]]] = None,
ignore_invalid_deployments: bool = False,
) -> None:
"""
@ -641,6 +642,9 @@ class Router:
self.model_group_retry_policy: Optional[
Dict[str, RetryPolicy]
] = model_group_retry_policy
self.model_group_affinity_config: Optional[
Dict[str, List[str]]
] = model_group_affinity_config
self.allowed_fails_policy: Optional[AllowedFailsPolicy] = None
if allowed_fails_policy is not None:
@ -661,6 +665,26 @@ class Router:
if optional_pre_call_checks is not None:
self.add_optional_pre_call_checks(optional_pre_call_checks)
# If model_group_affinity_config is set but no global affinity checks were
# enabled, we still need the DeploymentAffinityCheck callback (with global
# flags all False) so per-group config can activate affinity per model group.
if self.model_group_affinity_config and not any(
isinstance(cb, DeploymentAffinityCheck)
for cb in (self.optional_callbacks or [])
):
if self.optional_callbacks is None:
self.optional_callbacks = []
affinity_callback = DeploymentAffinityCheck(
cache=self.cache,
ttl_seconds=self.deployment_affinity_ttl_seconds,
enable_user_key_affinity=False,
enable_responses_api_affinity=False,
enable_session_id_affinity=False,
model_group_affinity_config=self.model_group_affinity_config,
)
self.optional_callbacks.append(affinity_callback)
litellm.logging_callback_manager.add_litellm_callback(affinity_callback)
if self.alerting_config is not None:
self._initialize_alerting()
@ -1311,6 +1335,10 @@ class Router:
existing_affinity_callback.ttl_seconds = (
self.deployment_affinity_ttl_seconds
)
if self.model_group_affinity_config:
existing_affinity_callback.model_group_affinity_config = (
self.model_group_affinity_config
)
else:
affinity_callback = DeploymentAffinityCheck(
cache=self.cache,
@ -1318,6 +1346,7 @@ class Router:
enable_user_key_affinity=enable_user_key_affinity,
enable_responses_api_affinity=enable_responses_api_affinity,
enable_session_id_affinity=enable_session_id_affinity,
model_group_affinity_config=self.model_group_affinity_config,
)
self.optional_callbacks.append(affinity_callback)
litellm.logging_callback_manager.add_litellm_callback(affinity_callback)

View file

@ -13,7 +13,7 @@ where routing to a consistent deployment is still beneficial.
"""
import hashlib
from typing import Any, Dict, List, Optional, cast
from typing import Any, Dict, List, Optional, Tuple, cast
from typing_extensions import TypedDict
@ -38,6 +38,9 @@ class DeploymentAffinityCheck(CustomLogger):
"""
CACHE_KEY_PREFIX = "deployment_affinity:v1"
VALID_FLAGS = frozenset(
{"deployment_affinity", "responses_api_deployment_check", "session_affinity"}
)
def __init__(
self,
@ -46,6 +49,7 @@ class DeploymentAffinityCheck(CustomLogger):
enable_user_key_affinity: bool,
enable_responses_api_affinity: bool,
enable_session_id_affinity: bool = False,
model_group_affinity_config: Optional[Dict[str, List[str]]] = None,
):
super().__init__()
self.cache = cache
@ -53,6 +57,39 @@ class DeploymentAffinityCheck(CustomLogger):
self.enable_user_key_affinity = enable_user_key_affinity
self.enable_responses_api_affinity = enable_responses_api_affinity
self.enable_session_id_affinity = enable_session_id_affinity
self.model_group_affinity_config: Dict[str, List[str]] = (
model_group_affinity_config or {}
)
for group, flags in self.model_group_affinity_config.items():
unknown = set(flags) - self.VALID_FLAGS
if unknown:
verbose_router_logger.warning(
"DeploymentAffinityCheck: unknown flag(s) %s for model group '%s'; will be ignored. Valid flags: %s",
unknown,
group,
self.VALID_FLAGS,
)
def _get_effective_flags(self, model_group: str) -> Tuple[bool, bool, bool]:
"""
Return (enable_user_key_affinity, enable_responses_api_affinity, enable_session_id_affinity)
for the given model group.
If the model group has an explicit entry in model_group_affinity_config, use it.
Otherwise fall back to the global instance flags.
"""
group_checks = self.model_group_affinity_config.get(model_group)
if group_checks is not None:
return (
"deployment_affinity" in group_checks,
"responses_api_deployment_check" in group_checks,
"session_affinity" in group_checks,
)
return (
self.enable_user_key_affinity,
self.enable_responses_api_affinity,
self.enable_session_id_affinity,
)
@staticmethod
def _looks_like_sha256_hex(value: str) -> bool:
@ -277,8 +314,14 @@ class DeploymentAffinityCheck(CustomLogger):
request_kwargs = request_kwargs or {}
typed_healthy_deployments = cast(List[dict], healthy_deployments)
(
enable_user_key,
enable_responses_api,
enable_session_id,
) = self._get_effective_flags(model)
# 1) Responses API continuity (high priority)
if self.enable_responses_api_affinity:
if enable_responses_api:
previous_response_id = request_kwargs.get("previous_response_id")
if previous_response_id is not None:
responses_model_id = (
@ -305,7 +348,7 @@ class DeploymentAffinityCheck(CustomLogger):
return typed_healthy_deployments
# 2) Session-id -> deployment affinity
if self.enable_session_id_affinity:
if enable_session_id:
session_id = self._get_session_id_from_request_kwargs(
request_kwargs=request_kwargs
)
@ -344,7 +387,7 @@ class DeploymentAffinityCheck(CustomLogger):
)
# 3) User key -> deployment affinity
if not self.enable_user_key_affinity:
if not enable_user_key:
return typed_healthy_deployments
user_key = self._get_user_key_from_request_kwargs(request_kwargs=request_kwargs)
@ -394,22 +437,47 @@ class DeploymentAffinityCheck(CustomLogger):
- LiteLLM runs async success callbacks via a background logging worker for performance.
- We want affinity to be immediately available for subsequent requests.
"""
if not self.enable_user_key_affinity and not self.enable_session_id_affinity:
metadata_dicts = self._iter_metadata_dicts(kwargs)
# Extract deployment_model_name first — needed for both per-group flag resolution
# and cache key scoping.
deployment_model_name: Optional[str] = None
for metadata in metadata_dicts:
maybe_deployment_model_name = metadata.get("deployment_model_name")
if (
isinstance(maybe_deployment_model_name, str)
and maybe_deployment_model_name
):
deployment_model_name = maybe_deployment_model_name
break
if not deployment_model_name:
verbose_router_logger.debug(
"DeploymentAffinityCheck: deployment_model_name missing in metadata; skipping affinity cache update."
)
return None
# Resolve effective flags for this model group
(
enable_user_key,
_enable_responses_api,
enable_session_id,
) = self._get_effective_flags(deployment_model_name)
if not enable_user_key and not enable_session_id:
return None
user_key = None
if self.enable_user_key_affinity:
if enable_user_key:
user_key = self._get_user_key_from_request_kwargs(request_kwargs=kwargs)
session_id = None
if self.enable_session_id_affinity:
if enable_session_id:
session_id = self._get_session_id_from_request_kwargs(request_kwargs=kwargs)
if user_key is None and session_id is None:
return None
metadata_dicts = self._iter_metadata_dicts(kwargs)
model_info = kwargs.get("model_info")
if not isinstance(model_info, dict):
model_info = None
@ -433,25 +501,6 @@ class DeploymentAffinityCheck(CustomLogger):
)
return None
# Scope affinity by the Router deployment model name (alias-safe, consistent across
# heterogeneous providers, and matches standard logging's `model_map_key`).
deployment_model_name: Optional[str] = None
for metadata in metadata_dicts:
maybe_deployment_model_name = metadata.get("deployment_model_name")
if (
isinstance(maybe_deployment_model_name, str)
and maybe_deployment_model_name
):
deployment_model_name = maybe_deployment_model_name
break
if not deployment_model_name:
verbose_router_logger.warning(
"DeploymentAffinityCheck: deployment_model_name missing; skipping affinity cache update. model_id=%s",
model_id,
)
return None
if user_key is not None:
try:
cache_key = self.get_affinity_cache_key(

View file

@ -77,6 +77,7 @@ class UpdateRouterConfig(BaseModel):
routing_strategy_args: Optional[dict] = None
routing_strategy: Optional[str] = None
model_group_retry_policy: Optional[dict] = None
model_group_affinity_config: Optional[Dict[str, List[str]]] = None
allowed_fails: Optional[int] = None
cooldown_time: Optional[float] = None
num_retries: Optional[int] = None

View file

@ -1798,6 +1798,7 @@ def client(original_function): # noqa: PLR0915
model: Optional[str] = args[0] if len(args) > 0 else kwargs.get("model", None)
is_completion_with_fallbacks = kwargs.get("fallbacks") is not None
_is_litellm_internal_call = kwargs.pop("_is_litellm_internal_call", False)
try:
if logging_obj is None:
@ -1944,15 +1945,26 @@ def client(original_function): # noqa: PLR0915
)
# LOG SUCCESS - handle streaming success logging in the _next_ object
# Internal sub-calls (e.g. emulated file-search steps) share the
# parent's logging obj; skip async logging here so only the outer call bills once.
# NOTE: streaming requests return early (before this point) via
# CustomStreamWrapper, so this block is non-streaming only.
if getattr(logging_obj, "_defer_async_logging", False):
# Proxy has post-call guardrails that must complete before the
# SLP is built. Store a closure the proxy will call after
# post_call_success_hook so guardrail_information is in metadata.
# Only create_task is deferred; sync callbacks fire immediately
# (below, outside the if/else) for billing/rate-limiting.
def _enqueue_deferred_logging() -> None:
if not _is_litellm_internal_call:
if getattr(logging_obj, "_defer_async_logging", False):
def _enqueue_deferred_logging() -> None:
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
)
)
logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging # type: ignore
else:
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
@ -1963,19 +1975,6 @@ def client(original_function): # noqa: PLR0915
)
)
logging_obj._enqueue_deferred_logging = _enqueue_deferred_logging # type: ignore
else:
asyncio.create_task(
_client_async_logging_helper(
logging_obj=logging_obj,
result=result,
start_time=start_time,
end_time=end_time,
is_completion_with_fallbacks=is_completion_with_fallbacks,
)
)
# Sync callbacks always fire immediately regardless of deferral
logging_obj.handle_sync_success_callbacks_for_async_calls(
result=result,
start_time=start_time,
@ -2008,7 +2007,7 @@ def client(original_function): # noqa: PLR0915
except Exception as e:
traceback_exception = traceback.format_exc()
end_time = datetime.datetime.now()
if logging_obj:
if logging_obj and not _is_litellm_internal_call:
try:
logging_obj.failure_handler(
e, traceback_exception, start_time, end_time
@ -2599,6 +2598,47 @@ def _supports_factory(model: str, custom_llm_provider: Optional[str], key: str)
return False
def _is_explicitly_disabled_factory(
model: str, custom_llm_provider: Optional[str], key: str
) -> bool:
"""Return True only when the model map explicitly sets *key* to ``False``.
This is the opt-out mirror of :func:`_supports_factory`. Where
``_supports_factory`` requires an explicit ``True`` to return ``True``,
this function requires an explicit ``False``. A missing key (``None``)
is treated as *not* disabled so that unknown or newly-added models are
allowed through without any model-map entry.
Uses the same ``get_llm_provider`` → ``_get_model_info_helper`` chain as
``_supports_factory`` so caching, fallback, and normalisation improvements
apply here automatically.
"""
try:
model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model, custom_llm_provider=custom_llm_provider
)
model_info = _get_model_info_helper(
model=model, custom_llm_provider=custom_llm_provider
)
val = model_info.get(key)
if val is False:
return True
if val is None:
bare_model_key = _get_model_cost_key(model)
if bare_model_key is not None:
bare_entry = litellm.model_cost.get(bare_model_key) or {}
if bare_entry.get(key) is False:
return True
return False
except Exception as e:
verbose_logger.debug(
f"Model not found or error in checking {key} disabled state. "
f"You passed model={model}, custom_llm_provider={custom_llm_provider}. "
f"Error: {str(e)}"
)
return False
def supports_audio_input(model: str, custom_llm_provider: Optional[str] = None) -> bool:
"""Check if a given model supports audio input in a chat completion call"""
return _supports_factory(

View file

@ -3435,7 +3435,8 @@
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_none_reasoning_effort": true
"supports_none_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"azure/gpt-5.1-chat-2025-11-13": {
"cache_read_input_token_cost": 1.25e-07,
@ -18305,7 +18306,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1": {
"cache_read_input_token_cost": 1.25e-07,
@ -18344,7 +18346,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1-2025-11-13": {
"cache_read_input_token_cost": 1.25e-07,
@ -18383,7 +18386,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1-chat-latest": {
"cache_read_input_token_cost": 1.25e-07,
@ -18421,7 +18425,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2": {
"cache_read_input_token_cost": 1.75e-07,
@ -18461,7 +18466,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2-2025-12-11": {
"cache_read_input_token_cost": 1.75e-07,
@ -18501,7 +18507,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2-chat-latest": {
"cache_read_input_token_cost": 1.75e-07,
@ -18538,7 +18545,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.3-chat-latest": {
"cache_read_input_token_cost": 1.75e-07,
@ -18575,7 +18583,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2-pro": {
"input_cost_per_token": 2.1e-05,
@ -18608,7 +18617,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2-pro-2025-12-11": {
"input_cost_per_token": 2.1e-05,
@ -18641,7 +18651,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.4": {
"cache_read_input_token_cost": 2.5e-07,
@ -18690,7 +18701,8 @@
"supports_service_tier": true,
"supports_vision": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.4-2026-03-05": {
"cache_read_input_token_cost": 2.5e-07,
@ -18785,7 +18797,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.4-pro-2026-03-05": {
"cache_read_input_token_cost": 3e-06,
@ -18833,7 +18846,94 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.4-mini": {
"cache_read_input_token_cost": 7.5e-08,
"cache_read_input_token_cost_flex": 1e-08,
"cache_read_input_token_cost_batches": 3.8e-08,
"input_cost_per_token": 7.5e-07,
"input_cost_per_token_flex": 3.75e-07,
"input_cost_per_token_batches": 3.75e-07,
"litellm_provider": "openai",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 4.5e-06,
"output_cost_per_token_flex": 2.25e-06,
"output_cost_per_token_batches": 2.25e-06,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"gpt-5.4-nano": {
"cache_read_input_token_cost": 2e-08,
"cache_read_input_token_cost_flex": 1e-08,
"cache_read_input_token_cost_batches": 1e-08,
"input_cost_per_token": 2e-07,
"input_cost_per_token_flex": 1e-07,
"input_cost_per_token_batches": 1e-07,
"litellm_provider": "openai",
"max_input_tokens": 272000,
"max_output_tokens": 128000,
"max_tokens": 128000,
"mode": "chat",
"output_cost_per_token": 1.25e-06,
"output_cost_per_token_flex": 6.25e-07,
"output_cost_per_token_batches": 6.25e-07,
"supported_endpoints": [
"/v1/chat/completions",
"/v1/batch",
"/v1/responses"
],
"supported_modalities": [
"text",
"image"
],
"supported_output_modalities": [
"text"
],
"supports_function_calling": true,
"supports_native_streaming": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": true,
"supports_service_tier": true,
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": true,
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": false
},
"gpt-5-pro": {
"input_cost_per_token": 1.5e-05,
@ -18868,7 +18968,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-pro-2025-10-06": {
"input_cost_per_token": 1.5e-05,
@ -18903,7 +19004,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-2025-08-07": {
"cache_read_input_token_cost": 1.25e-07,
@ -18945,7 +19047,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-chat": {
"cache_read_input_token_cost": 1.25e-07,
@ -18979,7 +19082,8 @@
"supports_tool_choice": false,
"supports_vision": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-chat-latest": {
"cache_read_input_token_cost": 1.25e-07,
@ -19013,7 +19117,8 @@
"supports_tool_choice": false,
"supports_vision": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-codex": {
"cache_read_input_token_cost": 1.25e-07,
@ -19046,7 +19151,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1-codex": {
"cache_read_input_token_cost": 1.25e-07,
@ -19082,7 +19188,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1-codex-max": {
"cache_read_input_token_cost": 1.25e-07,
@ -19115,7 +19222,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.1-codex-mini": {
"cache_read_input_token_cost": 2.5e-08,
@ -19151,7 +19259,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5.2-codex": {
"cache_read_input_token_cost": 1.75e-07,
@ -19187,7 +19296,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": true
"supports_xhigh_reasoning_effort": true,
"supports_minimal_reasoning_effort": true
},
"gpt-5.3-codex": {
"cache_read_input_token_cost": 1.75e-07,
@ -19223,7 +19333,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-mini": {
"cache_read_input_token_cost": 2.5e-08,
@ -19265,7 +19376,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-mini-2025-08-07": {
"cache_read_input_token_cost": 2.5e-08,
@ -19307,7 +19419,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-nano": {
"cache_read_input_token_cost": 5e-09,
@ -19346,7 +19459,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-nano-2025-08-07": {
"cache_read_input_token_cost": 5e-09,
@ -19384,7 +19498,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-image-1": {
"cache_read_input_image_token_cost": 2.5e-06,
@ -36408,7 +36523,8 @@
"supports_vision": true,
"supports_web_search": true,
"supports_none_reasoning_effort": false,
"supports_xhigh_reasoning_effort": false
"supports_xhigh_reasoning_effort": false,
"supports_minimal_reasoning_effort": true
},
"gpt-5-search-api-2025-10-14": {
"cache_read_input_token_cost": 1.25e-07,

View file

@ -23,12 +23,14 @@ Example environment variables:
See: https://learn.microsoft.com/en-us/azure/ai-foundry/agents/quickstart
"""
import json
import os
import sys
sys.path.insert(0, os.path.abspath("../.."))
import pytest
from unittest.mock import MagicMock
import litellm
@ -343,13 +345,286 @@ def test_azure_ai_agents_extract_content_from_messages():
]
}
content = handler._extract_content_from_messages(messages_data)
content, annotations = handler._extract_content_from_messages(messages_data)
assert content == "The answer is 100."
assert annotations is None
# Test empty response
empty_data = {"data": []}
content = handler._extract_content_from_messages(empty_data)
content, annotations = handler._extract_content_from_messages(empty_data)
assert content == ""
assert annotations is None
def test_azure_ai_agents_extract_content_with_annotations():
"""
Test that annotations (e.g., Bing Search citations) are extracted from
Azure Agents message responses and transformed to OpenAI-compatible format.
Ref: https://github.com/BerriAI/litellm/issues/19126
"""
from litellm.llms.azure_ai.agents.handler import AzureAIAgentsHandler
handler = AzureAIAgentsHandler()
messages_data = {
"data": [
{
"id": "msg_abc",
"role": "assistant",
"content": [
{
"type": "text",
"text": {
"value": "According to sources [1], the answer is yes.",
"annotations": [
{
"type": "url_citation",
"text": "[1]",
"start_index": 22,
"end_index": 25,
"url_citation": {
"url": "https://example.com/source",
"title": "Example Source"
}
}
]
}
}
]
}
]
}
content, annotations = handler._extract_content_from_messages(messages_data)
assert content == "According to sources [1], the answer is yes."
assert annotations is not None
assert len(annotations) == 1
assert annotations[0]["type"] == "url_citation"
assert annotations[0]["url_citation"]["url"] == "https://example.com/source"
assert annotations[0]["url_citation"]["title"] == "Example Source"
# start/end_index should be moved into url_citation for OpenAI compatibility
assert annotations[0]["url_citation"]["start_index"] == 22
assert annotations[0]["url_citation"]["end_index"] == 25
def test_azure_ai_agents_build_model_response_with_annotations():
"""
Test that _build_model_response includes annotations in the Message object.
"""
from litellm.llms.azure_ai.agents.handler import AzureAIAgentsHandler
from litellm.types.utils import ModelResponse
handler = AzureAIAgentsHandler()
model_response = ModelResponse()
annotations = [
{
"type": "url_citation",
"url_citation": {
"url": "https://example.com",
"title": "Example",
"start_index": 0,
"end_index": 5,
},
}
]
result = handler._build_model_response(
model="azure_ai/agents/asst_123",
content="Hello [1]",
model_response=model_response,
thread_id="thread_abc",
messages=[{"role": "user", "content": "test"}],
annotations=annotations,
)
assert result.choices[0].message.content == "Hello [1]"
assert result.choices[0].message.annotations is not None
assert len(result.choices[0].message.annotations) == 1
assert result.choices[0].message.annotations[0]["type"] == "url_citation"
def test_azure_ai_agents_build_model_response_without_annotations():
"""
Test that _build_model_response works correctly without annotations.
"""
from litellm.llms.azure_ai.agents.handler import AzureAIAgentsHandler
from litellm.types.utils import ModelResponse
handler = AzureAIAgentsHandler()
model_response = ModelResponse()
result = handler._build_model_response(
model="azure_ai/agents/asst_123",
content="Hello",
model_response=model_response,
thread_id="thread_abc",
messages=[{"role": "user", "content": "test"}],
)
assert result.choices[0].message.content == "Hello"
assert getattr(result.choices[0].message, "annotations", None) is None
@pytest.mark.asyncio
async def test_azure_ai_agents_streaming_annotations_from_completed_message():
"""
Test that annotations from thread.message.completed SSE events are collected
and attached to the final chunk's delta.
Ref: https://github.com/BerriAI/litellm/issues/19126
"""
from litellm.llms.azure_ai.agents.handler import AzureAIAgentsHandler
handler = AzureAIAgentsHandler()
# SSE lines simulating a stream with annotations in thread.message.completed
completed_data = {
"content": [
{
"type": "text",
"text": {
"value": "According to [1], the answer is 42.",
"annotations": [
{
"type": "url_citation",
"text": "[1]",
"start_index": 12,
"end_index": 15,
"url_citation": {
"url": "https://example.com/citation",
"title": "Citation Source",
},
}
],
},
}
]
}
sse_lines = [
"event: thread.created",
"",
'data: {"id": "thread_stream_123"}',
"",
"event: thread.message.delta",
"",
'data: {"delta": {"content": [{"type": "text", "text": {"value": "According to [1], the answer is 42."}}]}}',
"",
"event: thread.message.completed",
"",
f"data: {json.dumps(completed_data)}",
"",
"data: [DONE]",
]
async def mock_aiter_lines():
for line in sse_lines:
yield line
mock_response = MagicMock()
mock_response.aiter_lines = MagicMock(return_value=mock_aiter_lines())
chunks = []
async for chunk in handler._process_sse_stream(mock_response, "azure_ai/agents/asst_123"):
chunks.append(chunk)
# Should have content chunks + final [DONE] chunk
assert len(chunks) >= 1
final_chunk = chunks[-1]
assert final_chunk.choices[0].finish_reason == "stop"
assert final_chunk.choices[0].delta.annotations is not None
assert len(final_chunk.choices[0].delta.annotations) == 1
ann = final_chunk.choices[0].delta.annotations[0]
assert ann["type"] == "url_citation"
assert ann["url_citation"]["url"] == "https://example.com/citation"
assert ann["url_citation"]["title"] == "Citation Source"
@pytest.mark.asyncio
async def test_azure_ai_agents_streaming_accumulates_annotations_from_multiple_text_items():
"""
Test that annotations from multiple text content items in thread.message.completed
are accumulated (not overwritten).
Ref: Greptile review on PR #23849
"""
from litellm.llms.azure_ai.agents.handler import AzureAIAgentsHandler
handler = AzureAIAgentsHandler()
# Two text blocks, each with distinct citations
completed_data = {
"content": [
{
"type": "text",
"text": {
"value": "First source [1].",
"annotations": [
{
"type": "url_citation",
"text": "[1]",
"start_index": 12,
"end_index": 15,
"url_citation": {
"url": "https://example.com/first",
"title": "First",
},
}
],
},
},
{
"type": "text",
"text": {
"value": "Second source [2].",
"annotations": [
{
"type": "url_citation",
"text": "[2]",
"start_index": 13,
"end_index": 16,
"url_citation": {
"url": "https://example.com/second",
"title": "Second",
},
}
],
},
},
]
}
sse_lines = [
"event: thread.created",
"",
'data: {"id": "thread_multi"}',
"",
"event: thread.message.completed",
"",
f"data: {json.dumps(completed_data)}",
"",
"data: [DONE]",
]
async def mock_aiter_lines():
for line in sse_lines:
yield line
mock_response = MagicMock()
mock_response.aiter_lines = MagicMock(return_value=mock_aiter_lines())
chunks = []
async for chunk in handler._process_sse_stream(mock_response, "azure_ai/agents/asst_123"):
chunks.append(chunk)
final_chunk = chunks[-1]
assert final_chunk.choices[0].delta.annotations is not None
assert len(final_chunk.choices[0].delta.annotations) == 2
urls = [a["url_citation"]["url"] for a in final_chunk.choices[0].delta.annotations]
assert "https://example.com/first" in urls
assert "https://example.com/second" in urls
@pytest.mark.asyncio

View file

@ -775,6 +775,228 @@ def test_ensure_alternating_roles(
assert messages == expected_messages
def test_ensure_alternating_roles_with_tool_calls():
"""Fixes Regression in #18685 """
messages = [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_123",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "NYC"}',
},
}
],
},
{"role": "tool", "tool_call_id": "call_123", "content": "72F, sunny"},
{"role": "assistant", "content": "It's 72F and sunny in NYC."},
{"role": "user", "content": "What about tomorrow?"},
{"role": "user", "content": "And the day after?"},
{"role": "user", "content": "What about next week?"},
]
transformed_messages = get_completion_messages(
messages=messages,
assistant_continue_message=None,
user_continue_message=None,
ensure_alternating_roles=True,
)
assert transformed_messages == [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_123",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "NYC"}',
},
}
],
},
{"role": "tool", "tool_call_id": "call_123", "content": "72F, sunny"},
{"role": "assistant", "content": "It's 72F and sunny in NYC."},
{"role": "user", "content": "What about tomorrow?"},
{"role": "assistant", "content": "Please continue."},
{"role": "user", "content": "And the day after?"},
{"role": "assistant", "content": "Please continue."},
{"role": "user", "content": "What about next week?"},
]
def test_ensure_alternating_roles_three_consecutive_assistants():
messages = [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "A1"},
{"role": "assistant", "content": "A2"},
{"role": "assistant", "content": "A3"},
]
transformed_messages = get_completion_messages(
messages=messages,
assistant_continue_message=None,
user_continue_message=None,
ensure_alternating_roles=True,
)
assert transformed_messages == [
{"role": "user", "content": "Hello"},
{"role": "assistant", "content": "A1"},
{"role": "user", "content": "Please continue."},
{"role": "assistant", "content": "A2"},
{"role": "user", "content": "Please continue."},
{"role": "assistant", "content": "A3"},
{"role": "user", "content": "Please continue."},
]
def test_ensure_alternating_roles_does_not_split_tool_call_chain():
"""Tool-call chains [user, assistant(tc), tool, user] are preserved as-is."""
messages = [
{"role": "user", "content": "Search for X"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "c1",
"type": "function",
"function": {"name": "search", "arguments": "{}"},
}
],
},
{"role": "tool", "tool_call_id": "c1", "content": "results"},
{"role": "user", "content": "Thanks, now do Y"},
]
transformed_messages = get_completion_messages(
messages=messages,
assistant_continue_message=None,
user_continue_message=None,
ensure_alternating_roles=True,
)
assert transformed_messages == [
{"role": "user", "content": "Search for X"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "c1",
"type": "function",
"function": {"name": "search", "arguments": "{}"},
}
],
},
{"role": "tool", "tool_call_id": "c1", "content": "results"},
{"role": "user", "content": "Thanks, now do Y"},
]
def test_ensure_alternating_roles_assistant_tool_call_then_assistant():
"""
Preserve old behavior for malformed adjacent assistant turns:
[assistant(tool_calls), assistant(no-tool-calls), user] should insert
user_continue between assistant messages.
"""
messages = [
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "c1",
"type": "function",
"function": {"name": "search", "arguments": "{}"},
}
],
},
{"role": "assistant", "content": "Here's what I found."},
{"role": "user", "content": "Thanks"},
]
transformed_messages = get_completion_messages(
messages=messages,
assistant_continue_message=None,
user_continue_message=None,
ensure_alternating_roles=True,
)
assert transformed_messages == [
{"role": "user", "content": "Please continue."},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "c1",
"type": "function",
"function": {"name": "search", "arguments": "{}"},
}
],
},
{"role": "user", "content": "Please continue."},
{"role": "assistant", "content": "Here's what I found."},
{"role": "user", "content": "Thanks"},
]
def test_ensure_alternating_roles_trailing_tool_call_assistant():
messages = [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "NYC"}',
},
}
],
},
]
transformed_messages = get_completion_messages(
messages=messages,
assistant_continue_message=None,
user_continue_message=None,
ensure_alternating_roles=True,
)
assert transformed_messages == [
{"role": "user", "content": "What's the weather?"},
{
"role": "assistant",
"content": None,
"tool_calls": [
{
"id": "call_abc",
"type": "function",
"function": {
"name": "get_weather",
"arguments": '{"location": "NYC"}',
},
}
],
},
{"role": "user", "content": "Please continue."},
]
def test_alternating_roles_e2e():
from litellm.llms.custom_httpx.http_handler import HTTPHandler
import json

View file

@ -0,0 +1,182 @@
"""
Unit tests for pre-call checks running before polling ID creation.
Tests that rate limits, guardrails, and budget checks are enforced
BEFORE a polling ID is created, so rate-limited requests get a
synchronous error instead of a polling ID that immediately fails.
"""
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from fastapi import HTTPException, Request, Response
sys.path.insert(0, os.path.abspath("../.."))
import litellm
from litellm.proxy._types import UserAPIKeyAuth
from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing
class TestSkipPreCallLogic:
"""Test that skip_pre_call_logic parameter works correctly"""
@pytest.mark.asyncio
async def test_skip_pre_call_logic_skips_common_processing(self):
"""When skip_pre_call_logic=True, common_processing_pre_call_logic should not be called"""
mock_logging_obj = MagicMock()
data = {
"model": "gpt-4",
"stream": True,
"litellm_logging_obj": mock_logging_obj,
}
processor = ProxyBaseLLMRequestProcessing(data=data)
mock_proxy_logging = AsyncMock()
mock_proxy_logging.during_call_hook = AsyncMock()
with (
patch.object(
processor, "common_processing_pre_call_logic", new_callable=AsyncMock
) as mock_pre_call,
patch(
"litellm.proxy.common_request_processing.route_request",
new_callable=AsyncMock,
return_value=MagicMock(),
),
):
try:
await processor.base_process_llm_request(
request=MagicMock(spec=Request),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
route_type="aresponses",
proxy_logging_obj=mock_proxy_logging,
llm_router=MagicMock(),
general_settings={},
proxy_config=MagicMock(),
skip_pre_call_logic=True,
)
except Exception:
pass # We only care that common_processing_pre_call_logic was not called
mock_pre_call.assert_not_called()
@pytest.mark.asyncio
async def test_without_skip_runs_common_processing(self):
"""When skip_pre_call_logic=False (default), common_processing_pre_call_logic should be called"""
data = {"model": "gpt-4"}
processor = ProxyBaseLLMRequestProcessing(data=data)
mock_logging_obj = MagicMock()
mock_proxy_logging = AsyncMock()
mock_proxy_logging.during_call_hook = AsyncMock()
with (
patch.object(
processor,
"common_processing_pre_call_logic",
new_callable=AsyncMock,
return_value=(data, mock_logging_obj),
) as mock_pre_call,
patch(
"litellm.proxy.common_request_processing.route_request",
new_callable=AsyncMock,
),
):
try:
await processor.base_process_llm_request(
request=MagicMock(spec=Request),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
route_type="aresponses",
proxy_logging_obj=mock_proxy_logging,
llm_router=MagicMock(),
general_settings={},
proxy_config=MagicMock(),
)
except Exception:
pass
mock_pre_call.assert_called_once()
class TestPollingEndpointPreCallGuard:
"""Test that the polling endpoint enforces pre-call checks before polling ID creation"""
@pytest.mark.asyncio
async def test_rate_limit_error_prevents_polling_id_creation(self):
"""responses_api() must raise 429 and never call generate_polling_id when rate-limited"""
from litellm.proxy.response_api_endpoints.endpoints import responses_api
from litellm.proxy.response_polling.polling_handler import ResponsePollingHandler
rate_limit_exc = litellm.RateLimitError(
message="TPM limit exceeded",
llm_provider="",
model="gpt-4",
)
generate_polling_id_mock = MagicMock(return_value="litellm_poll_test")
proxy_server_patches = {
"litellm.proxy.proxy_server._read_request_body": AsyncMock(
return_value={"model": "gpt-4", "background": True}
),
"litellm.proxy.proxy_server.general_settings": {},
"litellm.proxy.proxy_server.llm_router": MagicMock(),
"litellm.proxy.proxy_server.native_background_mode": None,
"litellm.proxy.proxy_server.polling_cache_ttl": 3600,
"litellm.proxy.proxy_server.polling_via_cache_enabled": True,
"litellm.proxy.proxy_server.proxy_config": MagicMock(),
"litellm.proxy.proxy_server.proxy_logging_obj": AsyncMock(),
"litellm.proxy.proxy_server.redis_usage_cache": AsyncMock(),
"litellm.proxy.proxy_server.select_data_generator": None,
"litellm.proxy.proxy_server.user_api_base": None,
"litellm.proxy.proxy_server.user_max_tokens": None,
"litellm.proxy.proxy_server.user_model": None,
"litellm.proxy.proxy_server.user_request_timeout": None,
"litellm.proxy.proxy_server.user_temperature": None,
"litellm.proxy.proxy_server.version": "1.0.0",
}
with (
patch.multiple("litellm.proxy.proxy_server", **{
k.split(".")[-1]: v for k, v in proxy_server_patches.items()
}),
patch(
"litellm.proxy.response_polling.polling_handler.should_use_polling_for_request",
return_value=True,
),
patch.object(
ProxyBaseLLMRequestProcessing,
"common_processing_pre_call_logic",
new_callable=AsyncMock,
side_effect=rate_limit_exc,
),
patch.object(
ProxyBaseLLMRequestProcessing,
"_handle_llm_api_exception",
new_callable=AsyncMock,
return_value=HTTPException(status_code=429, detail="Rate limit exceeded"),
),
patch.object(ResponsePollingHandler, "generate_polling_id", generate_polling_id_mock),
# Prevent background task from running (avoids noise from incomplete mocks)
patch("asyncio.create_task"),
patch.object(
ResponsePollingHandler,
"create_initial_state",
new_callable=AsyncMock,
return_value=MagicMock(),
),
):
with pytest.raises(HTTPException) as exc_info:
await responses_api(
request=MagicMock(spec=Request),
fastapi_response=MagicMock(spec=Response),
user_api_key_dict=MagicMock(spec=UserAPIKeyAuth),
)
assert exc_info.value.status_code == 429
generate_polling_id_mock.assert_not_called()

View file

@ -1,8 +1,10 @@
import pytest
import litellm
from litellm.litellm_core_utils.get_model_cost_map import get_model_cost_map
from litellm.llms.openai.chat.gpt_5_transformation import OpenAIGPT5Config
from litellm.llms.openai.openai import OpenAIConfig
from litellm.utils import _is_explicitly_disabled_factory
@pytest.fixture()
@ -15,15 +17,23 @@ def gpt5_config() -> OpenAIGPT5Config:
return OpenAIGPT5Config()
@pytest.fixture(autouse=True)
def use_local_model_cost_map(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
monkeypatch.setattr(
litellm, "model_cost", get_model_cost_map(url=litellm.model_cost_map_url)
)
litellm.add_known_models(model_cost_map=litellm.model_cost)
def test_gpt5_supports_reasoning_effort(config: OpenAIConfig):
assert "reasoning_effort" in config.get_supported_openai_params(model="gpt-5")
assert "reasoning_effort" in config.get_supported_openai_params(model="gpt-5-mini")
def test_gpt5_chat_does_not_support_reasoning_effort(config: OpenAIConfig):
assert (
"reasoning_effort"
not in config.get_supported_openai_params(model="gpt-5-chat-latest")
assert "reasoning_effort" not in config.get_supported_openai_params(
model="gpt-5-chat-latest"
)
@ -132,7 +142,6 @@ def test_gpt5_codex_temperature_error(config: OpenAIConfig):
)
def test_gpt5_codex_temperature_one_allowed(config: OpenAIConfig):
"""Test that GPT-5-Codex allows temperature=1."""
params = config.map_openai_params(
@ -198,6 +207,8 @@ def test_gpt5_verbosity_parameter(config: OpenAIConfig):
drop_params=False,
)
assert params["verbosity"] == "low"
def test_gpt5_1_reasoning_effort_none(config: OpenAIConfig):
"""Test that GPT-5.1 supports reasoning_effort='none' parameter.
@ -270,7 +281,9 @@ def test_gpt5_1_model_detection(gpt5_config: OpenAIGPT5Config):
# codex/pro/chat variants do not support none
assert not gpt5_config._supports_reasoning_effort_level("gpt-5.1-codex", "none")
assert not gpt5_config._supports_reasoning_effort_level("gpt-5.1-codex-max", "none")
assert not gpt5_config._supports_reasoning_effort_level("gpt-5.2-chat-latest", "none")
assert not gpt5_config._supports_reasoning_effort_level(
"gpt-5.2-chat-latest", "none"
)
assert not gpt5_config._supports_reasoning_effort_level("gpt-5.2-pro", "none")
assert not gpt5_config._supports_reasoning_effort_level("gpt-5", "none")
assert not gpt5_config._supports_reasoning_effort_level("gpt-5-mini", "none")
@ -324,10 +337,211 @@ def test_gpt5_4_pro_allows_reasoning_effort_xhigh(config: OpenAIConfig):
assert params["reasoning_effort"] == "xhigh"
def test_gpt5_4_mini_allows_reasoning_effort_xhigh(config: OpenAIConfig):
"""gpt-5.4-mini supports reasoning_effort='xhigh'."""
params = config.map_openai_params(
non_default_params={"reasoning_effort": "xhigh"},
optional_params={},
model="gpt-5.4-mini",
drop_params=False,
)
assert params["reasoning_effort"] == "xhigh"
def test_gpt5_4_nano_allows_reasoning_effort_xhigh(config: OpenAIConfig):
"""gpt-5.4-nano supports reasoning_effort='xhigh'."""
params = config.map_openai_params(
non_default_params={"reasoning_effort": "xhigh"},
optional_params={},
model="gpt-5.4-nano",
drop_params=False,
)
assert params["reasoning_effort"] == "xhigh"
def test_gpt5_4_nano_allows_reasoning_effort_none(config: OpenAIConfig):
"""gpt-5.4-nano supports reasoning_effort='none'."""
params = config.map_openai_params(
non_default_params={"reasoning_effort": "none"},
optional_params={},
model="gpt-5.4-nano",
drop_params=False,
)
assert params["reasoning_effort"] == "none"
def test_gpt5_4_mini_allows_reasoning_effort_none(config: OpenAIConfig):
"""gpt-5.4-mini supports reasoning_effort='none'."""
params = config.map_openai_params(
non_default_params={"reasoning_effort": "none"},
optional_params={},
model="gpt-5.4-mini",
drop_params=False,
)
assert params["reasoning_effort"] == "none"
def test_gpt5_4_allows_reasoning_effort_minimal(config: OpenAIConfig):
"""gpt-5.4 supports reasoning_effort='minimal'."""
params = config.map_openai_params(
non_default_params={"reasoning_effort": "minimal"},
optional_params={},
model="gpt-5.4",
drop_params=False,
)
assert params["reasoning_effort"] == "minimal"
def test_gpt5_4_pro_allows_reasoning_effort_minimal(config: OpenAIConfig):
"""gpt-5.4-pro supports reasoning_effort='minimal'."""
params = config.map_openai_params(
non_default_params={"reasoning_effort": "minimal"},
optional_params={},
model="gpt-5.4-pro",
drop_params=False,
)
assert params["reasoning_effort"] == "minimal"
def test_gpt5_4_mini_rejects_reasoning_effort_minimal(config: OpenAIConfig):
"""gpt-5.4-mini does not support reasoning_effort='minimal'."""
with pytest.raises(litellm.utils.UnsupportedParamsError):
config.map_openai_params(
non_default_params={"reasoning_effort": "minimal"},
optional_params={},
model="gpt-5.4-mini",
drop_params=False,
)
def test_gpt5_4_nano_rejects_reasoning_effort_minimal(config: OpenAIConfig):
"""gpt-5.4-nano does not support reasoning_effort='minimal'."""
with pytest.raises(litellm.utils.UnsupportedParamsError):
config.map_openai_params(
non_default_params={"reasoning_effort": "minimal"},
optional_params={},
model="gpt-5.4-nano",
drop_params=False,
)
def test_gpt5_4_mini_provider_prefixed_rejects_minimal(config: OpenAIConfig):
"""openai/gpt-5.4-mini correctly rejects minimal (model lookup normalizes prefix)."""
with pytest.raises(litellm.utils.UnsupportedParamsError):
config.map_openai_params(
non_default_params={"reasoning_effort": "minimal"},
optional_params={},
model="openai/gpt-5.4-mini",
drop_params=False,
)
def test_gpt5_drops_reasoning_effort_minimal_when_requested(config: OpenAIConfig):
"""reasoning_effort='minimal' is dropped for unsupported models when drop_params=True."""
params = config.map_openai_params(
non_default_params={"reasoning_effort": "minimal"},
optional_params={},
model="gpt-5.4-mini",
drop_params=True,
)
assert "reasoning_effort" not in params
def test_gpt5_minimal_dict_triggers_validation(config: OpenAIConfig):
"""Dict with effort='minimal' triggers minimal model-support validation."""
with pytest.raises(litellm.utils.UnsupportedParamsError):
config.map_openai_params(
non_default_params={
"reasoning_effort": {"effort": "minimal", "summary": "detailed"}
},
optional_params={},
model="gpt-5.4-mini",
drop_params=False,
)
def test_gpt5_minimal_dict_accepted_for_supported_model(config: OpenAIConfig):
"""Dict with effort='minimal' passes through for gpt-5.4+."""
params = config.map_openai_params(
non_default_params={
"reasoning_effort": {"effort": "minimal", "summary": "detailed"}
},
optional_params={},
model="gpt-5.4",
drop_params=False,
)
assert params["reasoning_effort"] == "minimal"
def test_gpt5_supports_reasoning_effort_level_minimal(gpt5_config: OpenAIGPT5Config):
"""Test that _supports_reasoning_effort_level correctly identifies minimal support."""
assert gpt5_config._supports_reasoning_effort_level("gpt-5.4", "minimal")
assert gpt5_config._supports_reasoning_effort_level("gpt-5.4-pro", "minimal")
assert not gpt5_config._supports_reasoning_effort_level("gpt-5.4-mini", "minimal")
assert not gpt5_config._supports_reasoning_effort_level("gpt-5.4-nano", "minimal")
def test_gpt5_minimal_explicitly_disabled_check(gpt5_config: OpenAIGPT5Config):
"""_is_reasoning_effort_level_explicitly_disabled returns True only for explicit False entries.
Models with supports_minimal_reasoning_effort=false → disabled.
Models with supports_minimal_reasoning_effort=true (or missing) → not disabled.
Provider-prefixed models (openai/gpt-5.4-mini) are normalized before lookup.
"""
assert gpt5_config._is_reasoning_effort_level_explicitly_disabled(
"gpt-5.4-mini", "minimal"
)
assert gpt5_config._is_reasoning_effort_level_explicitly_disabled(
"gpt-5.4-nano", "minimal"
)
assert gpt5_config._is_reasoning_effort_level_explicitly_disabled(
"openai/gpt-5.4-mini", "minimal"
)
assert not gpt5_config._is_reasoning_effort_level_explicitly_disabled(
"gpt-5.4", "minimal"
)
assert not gpt5_config._is_reasoning_effort_level_explicitly_disabled(
"gpt-5.4-pro", "minimal"
)
def test_is_explicitly_disabled_factory_minimal():
"""_is_explicitly_disabled_factory returns True only for explicit False entries.
Verifies the shared helper used by _is_reasoning_effort_level_explicitly_disabled
directly — so future changes to the helper are caught without going through the
method wrapper.
"""
key = "supports_minimal_reasoning_effort"
assert _is_explicitly_disabled_factory("gpt-5.4-mini", None, key)
assert _is_explicitly_disabled_factory("gpt-5.4-nano", None, key)
assert _is_explicitly_disabled_factory("openai/gpt-5.4-mini", None, key)
assert not _is_explicitly_disabled_factory("gpt-5.4", None, key)
assert not _is_explicitly_disabled_factory("gpt-5.4-pro", None, key)
assert not _is_explicitly_disabled_factory("gpt-5.4-turbo-preview", None, key)
def test_gpt5_unknown_model_passes_through_minimal(config: OpenAIConfig):
"""Unknown/unlisted gpt-5 models should pass reasoning_effort='minimal' through.
Missing supports_minimal_reasoning_effort key is treated as supported,
not as unsupported, to avoid breaking custom or newly-announced models.
"""
params = config.map_openai_params(
non_default_params={"reasoning_effort": "minimal"},
optional_params={},
model="gpt-5.4-turbo-preview",
drop_params=False,
)
assert params["reasoning_effort"] == "minimal"
def test_gpt5_normalizes_reasoning_effort_dict_with_summary(config: OpenAIConfig):
"""Dict with summary/generate_summary is normalized for chat completions."""
params = config.map_openai_params(
non_default_params={"reasoning_effort": {"effort": "high", "summary": "detailed"}},
non_default_params={
"reasoning_effort": {"effort": "high", "summary": "detailed"}
},
optional_params={},
model="gpt-5.4",
drop_params=False,
@ -343,7 +557,9 @@ def test_gpt5_xhigh_dict_triggers_validation(config: OpenAIConfig):
"""
with pytest.raises(litellm.utils.UnsupportedParamsError):
config.map_openai_params(
non_default_params={"reasoning_effort": {"effort": "xhigh", "summary": "detailed"}},
non_default_params={
"reasoning_effort": {"effort": "xhigh", "summary": "detailed"}
},
optional_params={},
model="gpt-5.1",
drop_params=False,
@ -353,7 +569,9 @@ def test_gpt5_xhigh_dict_triggers_validation(config: OpenAIConfig):
def test_gpt5_xhigh_dict_accepted_for_supported_model(config: OpenAIConfig):
"""Dict with effort='xhigh' passes through for gpt-5.4+."""
params = config.map_openai_params(
non_default_params={"reasoning_effort": {"effort": "xhigh", "summary": "detailed"}},
non_default_params={
"reasoning_effort": {"effort": "xhigh", "summary": "detailed"}
},
optional_params={},
model="gpt-5.4",
drop_params=False,
@ -369,7 +587,10 @@ def test_gpt5_none_dict_with_tools_no_tool_drop(config: OpenAIConfig):
"""
tools = [{"type": "function", "function": {"name": "test", "description": "test"}}]
params = config.map_openai_params(
non_default_params={"reasoning_effort": {"effort": "none", "summary": "detailed"}, "tools": tools},
non_default_params={
"reasoning_effort": {"effort": "none", "summary": "detailed"},
"tools": tools,
},
optional_params={},
model="gpt-5.4",
drop_params=False,
@ -399,11 +620,15 @@ def test_gpt5_none_dict_with_sampling_params_allowed(config: OpenAIConfig):
assert params["top_p"] == 0.9
def test_gpt5_normalizes_reasoning_effort_dict_with_summary_from_optional_params(config: OpenAIConfig):
def test_gpt5_normalizes_reasoning_effort_dict_with_summary_from_optional_params(
config: OpenAIConfig,
):
"""reasoning_effort dict with summary in optional_params is normalized."""
params = config.map_openai_params(
non_default_params={},
optional_params={"reasoning_effort": {"effort": "medium", "summary": "detailed"}},
optional_params={
"reasoning_effort": {"effort": "medium", "summary": "detailed"}
},
model="gpt-5.4",
drop_params=False,
)
@ -476,7 +701,7 @@ def test_gpt5_4_pro_rejects_non_default_temperature(config: OpenAIConfig):
def test_gpt5_1_temperature_without_reasoning_effort(config: OpenAIConfig):
"""Test that GPT-5.1 supports any temperature when reasoning_effort is not specified.
When reasoning_effort is not provided, it defaults to "none" for gpt-5.1,
so temperature should be allowed.
"""
@ -502,7 +727,7 @@ def test_gpt5_1_temperature_with_reasoning_effort_other_values(config: OpenAICon
model="gpt-5.1",
drop_params=False,
)
# Test that temperature=1 is allowed with other reasoning_effort values
for effort in ["low", "medium", "high"]:
params = config.map_openai_params(
@ -515,7 +740,9 @@ def test_gpt5_1_temperature_with_reasoning_effort_other_values(config: OpenAICon
assert params["reasoning_effort"] == effort
def test_gpt5_1_temperature_with_reasoning_effort_in_optional_params(config: OpenAIConfig):
def test_gpt5_1_temperature_with_reasoning_effort_in_optional_params(
config: OpenAIConfig,
):
"""Test that reasoning_effort can be in optional_params and still work correctly."""
# Test with reasoning_effort="none" in optional_params
params = config.map_openai_params(
@ -525,7 +752,7 @@ def test_gpt5_1_temperature_with_reasoning_effort_in_optional_params(config: Ope
drop_params=False,
)
assert params["temperature"] == 0.5
# Test with reasoning_effort="low" in optional_params (should only allow temp=1)
with pytest.raises(litellm.utils.UnsupportedParamsError):
config.map_openai_params(
@ -535,6 +762,7 @@ def test_gpt5_1_temperature_with_reasoning_effort_in_optional_params(config: Ope
drop_params=False,
)
def test_gpt5_1_temperature_drop_when_not_none(config: OpenAIConfig):
"""Test that GPT-5.1 drops temperature when reasoning_effort != 'none' and drop_params=True."""
params = config.map_openai_params(
@ -557,7 +785,7 @@ def test_gpt5_temperature_still_restricted(config: OpenAIConfig):
model="gpt-5",
drop_params=False,
)
# temperature=1 should still work for gpt-5
params = config.map_openai_params(
non_default_params={"temperature": 1.0},
@ -650,7 +878,9 @@ def test_gpt5_search_supported_params(gpt5_config: OpenAIGPT5Config):
"reasoning_effort",
]
for param in rejected:
assert param not in supported, f"{param} should not be supported for search models"
assert (
param not in supported
), f"{param} should not be supported for search models"
def test_gpt5_search_has_expected_params(gpt5_config: OpenAIGPT5Config):
@ -688,7 +918,11 @@ def test_gpt5_search_maps_max_tokens(config: OpenAIConfig):
def test_gpt5_search_drops_unsupported_params(config: OpenAIConfig):
"""Test that search models drop unsupported params via map_openai_params."""
params = config.map_openai_params(
non_default_params={"n": 2, "temperature": 0.7, "tools": [{"type": "function"}]},
non_default_params={
"n": 2,
"temperature": 0.7,
"tools": [{"type": "function"}],
},
optional_params={},
model="gpt-5-search-api",
drop_params=True,
@ -696,6 +930,8 @@ def test_gpt5_search_drops_unsupported_params(config: OpenAIConfig):
assert "n" not in params
assert "temperature" not in params
assert "tools" not in params
# GPT-5 unsupported params audit (validated via direct API calls)
def test_gpt5_rejects_params_unsupported_by_openai(config: OpenAIConfig):
"""Params that OpenAI rejects for all GPT-5 reasoning models."""
@ -709,9 +945,9 @@ def test_gpt5_rejects_params_unsupported_by_openai(config: OpenAIConfig):
for model in ["gpt-5", "gpt-5-mini", "gpt-5-codex", "gpt-5.1", "gpt-5.2"]:
supported = config.get_supported_openai_params(model=model)
for param in rejected_params:
assert param not in supported, (
f"{param} should not be supported for {model}"
)
assert (
param not in supported
), f"{param} should not be supported for {model}"
def test_gpt5_1_supports_logprobs_top_p(config: OpenAIConfig):
@ -720,16 +956,22 @@ def test_gpt5_1_supports_logprobs_top_p(config: OpenAIConfig):
supported = config.get_supported_openai_params(model=model)
assert "logprobs" in supported, f"logprobs should be supported for {model}"
assert "top_p" in supported, f"top_p should be supported for {model}"
assert "top_logprobs" in supported, f"top_logprobs should be supported for {model}"
assert (
"top_logprobs" in supported
), f"top_logprobs should be supported for {model}"
def test_gpt5_base_does_not_support_logprobs_top_p(config: OpenAIConfig):
"""Base gpt-5/gpt-5-mini do NOT support logprobs, top_p, top_logprobs."""
for model in ["gpt-5", "gpt-5-mini", "gpt-5-codex"]:
supported = config.get_supported_openai_params(model=model)
assert "logprobs" not in supported, f"logprobs should not be supported for {model}"
assert (
"logprobs" not in supported
), f"logprobs should not be supported for {model}"
assert "top_p" not in supported, f"top_p should not be supported for {model}"
assert "top_logprobs" not in supported, f"top_logprobs should not be supported for {model}"
assert (
"top_logprobs" not in supported
), f"top_logprobs should not be supported for {model}"
def test_gpt5_1_logprobs_passthrough(config: OpenAIConfig):
@ -788,4 +1030,4 @@ def test_gpt5_1_logprobs_dropped_with_reasoning_effort(config: OpenAIConfig):
)
assert "logprobs" not in params
assert "top_p" not in params
assert params["reasoning_effort"] == "high"
assert params["reasoning_effort"] == "high"

View file

@ -0,0 +1,892 @@
"""
Unit tests for file_search / vector_store support in the Responses API.
Coverage:
A1-A7 _decode_vector_store_ids_in_tools()
B1-B3 update_responses_tools_with_model_file_ids()
C1,D1 supports_native_file_search()
E1-E4 file_search guard in responses/main.py
F1-F6 ManagedFiles hook access control
G1-G3 get_vector_store_ids_from_file_search_tools()
H1-H14 emulated_handler unit tests
"""
import base64
from typing import Any, Dict, List, Optional
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.litellm_core_utils.prompt_templates.common_utils import (
_decode_vector_store_ids_in_tools,
update_responses_tools_with_model_file_ids,
)
from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_unified_vs_id(
unified_uuid: str = "abc-123",
provider_resource_id: str = "vs_provider_native",
model_id: str = "model-id-999",
) -> str:
"""Build a valid base64-encoded unified vector-store ID."""
raw = (
f"litellm_proxy:vector_store;"
f"unified_id,{unified_uuid};"
f"model_id,{model_id};"
f"provider_resource_id,{provider_resource_id}"
)
return base64.urlsafe_b64encode(raw.encode()).decode().rstrip("=")
def _file_search_tool(vector_store_ids: Optional[List[str]] = None) -> Dict[str, Any]:
tool: Dict[str, Any] = {"type": "file_search"}
if vector_store_ids is not None:
tool["vector_store_ids"] = vector_store_ids
return tool
def _code_interpreter_tool(file_ids: Optional[List[str]] = None) -> Dict[str, Any]:
tool: Dict[str, Any] = {"type": "code_interpreter"}
if file_ids:
tool["container"] = {"type": "auto", "file_ids": file_ids}
return tool
# ---------------------------------------------------------------------------
# A-series: _decode_vector_store_ids_in_tools
# ---------------------------------------------------------------------------
class TestDecodeVectorStoreIdsInTools:
def test_A1_none_input_returns_none(self):
assert _decode_vector_store_ids_in_tools(None) is None
def test_A2_no_file_search_tools_unchanged(self):
tools = [{"type": "web_search"}, {"type": "code_interpreter"}]
result = _decode_vector_store_ids_in_tools(tools)
assert result == tools
def test_A3_file_search_no_vector_store_ids_unchanged(self):
tools = [_file_search_tool()] # no vector_store_ids key
result = _decode_vector_store_ids_in_tools(tools)
assert result == tools
def test_A4_unified_id_decoded_to_provider_resource_id(self):
unified_id = _make_unified_vs_id(provider_resource_id="vs_real_123")
tools = [_file_search_tool([unified_id])]
result = _decode_vector_store_ids_in_tools(tools)
assert result is not None
assert result[0]["vector_store_ids"] == ["vs_real_123"]
def test_A5_native_id_passes_through_unchanged(self):
native_id = "vs_openai_abc"
tools = [_file_search_tool([native_id])]
result = _decode_vector_store_ids_in_tools(tools)
assert result is not None
assert result[0]["vector_store_ids"] == ["vs_openai_abc"]
def test_A6_mixed_unified_and_native_ids(self):
unified_id = _make_unified_vs_id(provider_resource_id="vs_decoded")
native_id = "vs_native_xyz"
tools = [_file_search_tool([unified_id, native_id])]
result = _decode_vector_store_ids_in_tools(tools)
assert result is not None
assert result[0]["vector_store_ids"] == ["vs_decoded", "vs_native_xyz"]
def test_A7_malformed_base64_passes_through_unchanged(self):
bad_id = "not_valid_base64!!!"
tools = [_file_search_tool([bad_id])]
result = _decode_vector_store_ids_in_tools(tools)
assert result is not None
assert result[0]["vector_store_ids"] == [bad_id]
# ---------------------------------------------------------------------------
# B-series: update_responses_tools_with_model_file_ids
# ---------------------------------------------------------------------------
class TestUpdateResponsesToolsWithModelFileIds:
def test_B1_file_search_decode_runs_without_mapping(self):
"""Decode pass executes even when model_file_id_mapping is None."""
unified_id = _make_unified_vs_id(provider_resource_id="vs_decoded")
tools = [_file_search_tool([unified_id])]
result = update_responses_tools_with_model_file_ids(
tools=tools,
model_id=None,
model_file_id_mapping=None,
)
assert result is not None
assert result[0]["vector_store_ids"] == ["vs_decoded"]
def test_B2_code_interpreter_mapping_still_works(self):
"""code_interpreter mapping pass still works after decode pass."""
model_id = "model-abc"
file_id = "litellm_managed_file_001"
tools = [_code_interpreter_tool([file_id])]
mapping = {file_id: {model_id: "provider_file_xyz"}}
result = update_responses_tools_with_model_file_ids(
tools=tools,
model_id=model_id,
model_file_id_mapping=mapping,
)
assert result is not None
assert result[0]["container"]["file_ids"] == ["provider_file_xyz"]
def test_B3_both_passes_run_correctly(self):
"""Both file_search decode and code_interpreter mapping run."""
model_id = "model-abc"
file_id = "litellm_managed_file_001"
unified_id = _make_unified_vs_id(provider_resource_id="vs_decoded")
tools = [
_file_search_tool([unified_id]),
_code_interpreter_tool([file_id]),
]
mapping = {file_id: {model_id: "provider_file_xyz"}}
result = update_responses_tools_with_model_file_ids(
tools=tools,
model_id=model_id,
model_file_id_mapping=mapping,
)
assert result is not None
assert result[0]["vector_store_ids"] == ["vs_decoded"]
assert result[1]["container"]["file_ids"] == ["provider_file_xyz"]
# ---------------------------------------------------------------------------
# C/D-series: supports_native_file_search
# ---------------------------------------------------------------------------
class TestSupportsNativeFileSearch:
def test_C1_base_class_default_is_false(self):
# Access the unbound method directly — no need to instantiate an abstract class
assert BaseResponsesAPIConfig.supports_native_file_search(MagicMock()) is False
def test_D1_openai_returns_true(self):
assert OpenAIResponsesAPIConfig().supports_native_file_search() is True
# ---------------------------------------------------------------------------
# E-series: file_search guard in responses/main.py
# ---------------------------------------------------------------------------
class TestFileSearchGuardInResponsesMain:
"""Tests for _has_file_search_tool helper and emulated routing guard."""
def test_has_file_search_tool_true(self):
from litellm.responses.main import _has_file_search_tool
assert _has_file_search_tool([{"type": "file_search"}]) is True
def test_has_file_search_tool_false_empty(self):
from litellm.responses.main import _has_file_search_tool
assert _has_file_search_tool([]) is False
assert _has_file_search_tool(None) is False
def test_has_file_search_tool_false_other_tools(self):
from litellm.responses.main import _has_file_search_tool
assert _has_file_search_tool([{"type": "web_search"}]) is False
def test_E1_openai_provider_no_error(self):
"""OpenAI supports file_search natively — no error raised."""
from litellm.llms.openai.responses.transformation import (
OpenAIResponsesAPIConfig,
)
from litellm.responses.main import _has_file_search_tool
config = OpenAIResponsesAPIConfig()
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}]
assert _has_file_search_tool(tools)
assert config.supports_native_file_search()
# No exception expected — the guard would pass.
def test_E2_no_provider_config_routes_to_emulated_handler(self):
"""Provider config None + file_search should route to emulated handler."""
from litellm.responses.main import responses
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}]
logging_obj = MagicMock()
expected = {"ok": True}
with (
patch(
"litellm.responses.main.litellm.get_llm_provider",
return_value=("claude-sonnet-4-5", "anthropic", None, None),
),
patch(
"litellm.responses.main.update_responses_input_with_model_file_ids",
return_value="hello",
),
patch(
"litellm.responses.main.update_responses_tools_with_model_file_ids",
return_value=tools,
),
patch(
"litellm.responses.main.ProviderConfigManager.get_provider_responses_api_config",
return_value=None,
),
patch(
"litellm.responses.main.ResponsesAPIRequestUtils.get_requested_response_api_optional_param",
return_value={},
),
patch("litellm.responses.main.run_async_function", return_value=expected) as run_async_mock,
):
result = responses(
input="hello",
model="anthropic/claude-sonnet-4-5",
tools=tools,
litellm_logging_obj=logging_obj,
litellm_call_id="call-123",
)
assert result == expected
assert run_async_mock.called
routed_func = run_async_mock.call_args.args[0]
assert routed_func.__name__ == "aresponses_with_emulated_file_search"
def test_E3_non_native_provider_config_routes_to_emulated_handler(self):
"""Non-native provider config + file_search should route to emulated handler."""
from litellm.llms.base_llm.responses.transformation import (
BaseResponsesAPIConfig,
)
from litellm.responses.main import responses
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}]
logging_obj = MagicMock()
expected = {"ok": True}
mock_config = MagicMock(spec=BaseResponsesAPIConfig)
mock_config.supports_native_file_search.return_value = False
with (
patch(
"litellm.responses.main.litellm.get_llm_provider",
return_value=("claude-sonnet-4-5", "anthropic", None, None),
),
patch(
"litellm.responses.main.update_responses_input_with_model_file_ids",
return_value="hello",
),
patch(
"litellm.responses.main.update_responses_tools_with_model_file_ids",
return_value=tools,
),
patch(
"litellm.responses.main.ProviderConfigManager.get_provider_responses_api_config",
return_value=mock_config,
),
patch(
"litellm.responses.main.ResponsesAPIRequestUtils.get_requested_response_api_optional_param",
return_value={},
),
patch("litellm.responses.main.run_async_function", return_value=expected) as run_async_mock,
):
result = responses(
input="hello",
model="anthropic/claude-sonnet-4-5",
tools=tools,
litellm_logging_obj=logging_obj,
litellm_call_id="call-123",
)
assert result == expected
assert run_async_mock.called
routed_func = run_async_mock.call_args.args[0]
assert routed_func.__name__ == "aresponses_with_emulated_file_search"
def test_E4_no_file_search_tools_no_error(self):
"""No file_search tool in request → guard never fires."""
from litellm.responses.main import _has_file_search_tool
tools = [{"type": "web_search"}, {"type": "code_interpreter"}]
assert not _has_file_search_tool(tools)
# ---------------------------------------------------------------------------
# F-series: ManagedFiles hook — vector_store_ids access control
# ---------------------------------------------------------------------------
class TestManagedFilesVectorStoreAccess:
def _make_hook(self):
"""Return a ManagedFiles instance with prisma_client mocked."""
from enterprise.litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles as ManagedFiles,
)
hook = ManagedFiles.__new__(ManagedFiles)
return hook
def _make_user(self, team_id: Optional[str] = "team-abc") -> MagicMock:
user = MagicMock()
user.team_id = team_id
user.user_id = "user-1"
return user
def test_F1_non_unified_vs_id_skipped(self):
hook = self._make_hook()
result = hook.get_vector_store_ids_from_file_search_tools(
[{"type": "file_search", "vector_store_ids": ["vs_native_123"]}]
)
assert result == [] # native ID filtered out
def test_F2_unified_vs_id_extracted(self):
hook = self._make_hook()
unified_id = _make_unified_vs_id()
result = hook.get_vector_store_ids_from_file_search_tools(
[{"type": "file_search", "vector_store_ids": [unified_id]}]
)
assert result == [unified_id]
def _make_vs_row(self, vector_store_id: str, team_id: Optional[str]) -> Any:
"""Build a row compatible with get_managed_vector_store_rows_by_uuids (Prisma model_dump)."""
from litellm.proxy._types import LiteLLM_ManagedVectorStoresTable
return LiteLLM_ManagedVectorStoresTable(
vector_store_id=vector_store_id,
custom_llm_provider="openai",
vector_store_name=None,
vector_store_description=None,
vector_store_metadata=None,
created_at=None,
updated_at=None,
litellm_credential_name=None,
litellm_params=None,
team_id=team_id,
user_id=None,
)
@pytest.mark.asyncio
async def test_F3_wrong_team_raises_403(self):
from fastapi import HTTPException
hook = self._make_hook()
unified_id = _make_unified_vs_id(unified_uuid="uuid-001")
mock_row = self._make_vs_row(vector_store_id="uuid-001", team_id="team-other")
async def mock_get_rows(uuids, prisma_client, user_api_key_cache, proxy_logging_obj=None):
return [mock_row]
with patch(
"litellm.proxy.proxy_server.prisma_client",
MagicMock(),
), patch(
"litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids",
side_effect=mock_get_rows,
):
with pytest.raises(HTTPException) as exc_info:
await hook.check_vector_store_ids_access(
[unified_id], self._make_user(team_id="team-caller")
)
assert exc_info.value.status_code == 403
@pytest.mark.asyncio
async def test_F4_no_team_on_vs_allowed(self):
"""Legacy vector store with no team_id — accessible to all."""
hook = self._make_hook()
unified_id = _make_unified_vs_id(unified_uuid="uuid-002")
mock_row = self._make_vs_row(vector_store_id="uuid-002", team_id=None)
async def mock_get_rows(uuids, prisma_client, user_api_key_cache, proxy_logging_obj=None):
return [mock_row]
with patch(
"litellm.proxy.proxy_server.prisma_client",
MagicMock(),
), patch(
"litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids",
side_effect=mock_get_rows,
):
await hook.check_vector_store_ids_access(
[unified_id], self._make_user(team_id="team-caller")
)
@pytest.mark.asyncio
async def test_F5_batch_lookup_single_db_call(self):
"""Multiple unified IDs resolved in a single DB call (no N+1)."""
hook = self._make_hook()
ids = [
_make_unified_vs_id(unified_uuid=f"uuid-{i}", provider_resource_id=f"vs_{i}")
for i in range(3)
]
rows = [
self._make_vs_row(vector_store_id=f"uuid-{i}", team_id="team-abc")
for i in range(3)
]
get_rows_mock = AsyncMock(return_value=rows)
with patch(
"litellm.proxy.proxy_server.prisma_client",
MagicMock(),
), patch(
"litellm.proxy.auth.auth_checks.get_managed_vector_store_rows_by_uuids",
get_rows_mock,
):
await hook.check_vector_store_ids_access(ids, self._make_user("team-abc"))
get_rows_mock.assert_called_once()
call_args = get_rows_mock.call_args
assert set(call_args.kwargs["uuids"] or call_args.args[0]) == {"uuid-0", "uuid-1", "uuid-2"}
@pytest.mark.asyncio
async def test_F6_non_responses_call_type_skipped(self):
"""Access check only runs for aresponses/responses call types."""
from enterprise.litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles as ManagedFiles,
)
from litellm.proxy._types import CallTypes
# If call_type is acompletion, the vector_store check branch isn't reached.
# Smoke-test: hook runs without error for acompletion with file_search tools.
hook = MagicMock(spec=ManagedFiles)
hook.async_pre_call_hook = AsyncMock(return_value=None)
await hook.async_pre_call_hook(
user_api_key_dict=self._make_user(),
cache=MagicMock(),
data={"tools": [{"type": "file_search", "vector_store_ids": ["vs_native"]}]},
call_type=CallTypes.acompletion.value,
)
hook.async_pre_call_hook.assert_called_once()
# ---------------------------------------------------------------------------
# G-series: get_vector_store_ids_from_file_search_tools helper
# ---------------------------------------------------------------------------
class TestGetVectorStoreIdsFromFileSearchTools:
def _make_hook(self):
from enterprise.litellm_enterprise.proxy.hooks.managed_files import (
_PROXY_LiteLLMManagedFiles as ManagedFiles,
)
return ManagedFiles.__new__(ManagedFiles)
def test_G1_tools_none_returns_empty(self):
hook = self._make_hook()
assert hook.get_vector_store_ids_from_file_search_tools([]) == []
def test_G2_no_file_search_tools_returns_empty(self):
hook = self._make_hook()
tools = [{"type": "code_interpreter"}, {"type": "web_search"}]
assert hook.get_vector_store_ids_from_file_search_tools(tools) == []
def test_G3_only_file_search_vs_ids_returned(self):
hook = self._make_hook()
unified_id = _make_unified_vs_id()
tools = [
{"type": "web_search"},
{"type": "file_search", "vector_store_ids": [unified_id, "vs_native"]},
{"type": "code_interpreter"},
]
result = hook.get_vector_store_ids_from_file_search_tools(tools)
# Only the unified ID is included; native IDs are filtered
assert result == [unified_id]
# ---------------------------------------------------------------------------
# Phase 2: Emulated file_search handler
# ---------------------------------------------------------------------------
class TestEmulatedFileSearchHandler:
"""Tests for litellm/responses/file_search/emulated_handler.py"""
def _make_mock_responses_api_response(
self,
text: str = "The answer is 42.",
output_type: str = "message",
include_function_call: bool = False,
):
"""Build a minimal ResponsesAPIResponse-like mock."""
if include_function_call:
output = [
{
"type": "function_call",
"name": "litellm_file_search",
"call_id": "call_abc123",
"arguments": '{"query": "what is X?", "vector_store_id": "vs_001"}',
}
]
else:
output = [
{
"type": "message",
"role": "assistant",
"content": [{"type": "output_text", "text": text}],
}
]
resp = MagicMock()
resp.output = output
resp.id = "resp_test123"
resp.created_at = 1700000000
resp.model = "claude-3-5-sonnet"
resp.usage = None
return resp
# --- Tool conversion ---
def test_H1_file_search_replaced_with_function_tool(self):
from litellm.responses.file_search.emulated_handler import (
_replace_file_search_tools,
)
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc", "vs_def"]}]
new_tools, vs_ids = _replace_file_search_tools(tools)
assert vs_ids == ["vs_abc", "vs_def"]
assert len(new_tools) == 1
assert new_tools[0]["type"] == "function"
assert new_tools[0]["name"] == "litellm_file_search"
# Both store IDs appear in the enum
enum_ids = new_tools[0]["parameters"]["properties"]["vector_store_id"]["enum"]
assert "vs_abc" in enum_ids
assert "vs_def" in enum_ids
def test_H2_non_file_search_tools_preserved(self):
from litellm.responses.file_search.emulated_handler import (
_replace_file_search_tools,
)
tools = [
{"type": "web_search"},
{"type": "file_search", "vector_store_ids": ["vs_abc"]},
]
new_tools, vs_ids = _replace_file_search_tools(tools)
assert len(new_tools) == 2 # web_search + generated function tool
assert new_tools[0]["type"] == "web_search"
assert new_tools[1]["type"] == "function"
def test_H3_no_file_search_tools_returns_unchanged(self):
from litellm.responses.file_search.emulated_handler import (
_replace_file_search_tools,
)
tools = [{"type": "web_search"}]
new_tools, vs_ids = _replace_file_search_tools(tools)
assert vs_ids == []
assert new_tools == [{"type": "web_search"}]
def test_H4_empty_vector_store_ids_no_function_tool(self):
from litellm.responses.file_search.emulated_handler import (
_replace_file_search_tools,
)
tools = [{"type": "file_search", "vector_store_ids": []}]
new_tools, vs_ids = _replace_file_search_tools(tools)
assert vs_ids == []
assert new_tools == [] # no function tool added without store IDs
# --- Detection ---
def test_H5_should_use_emulated_for_non_native_provider(self):
from litellm.responses.file_search.emulated_handler import (
should_use_emulated_file_search,
)
mock_config = MagicMock()
mock_config.supports_native_file_search.return_value = False
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}]
assert should_use_emulated_file_search(tools, mock_config) is True
def test_H6_should_not_emulate_for_native_provider(self):
from litellm.llms.openai.responses.transformation import (
OpenAIResponsesAPIConfig,
)
from litellm.responses.file_search.emulated_handler import (
should_use_emulated_file_search,
)
config = OpenAIResponsesAPIConfig()
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}]
assert should_use_emulated_file_search(tools, config) is False
def test_H7_should_not_emulate_without_file_search_tools(self):
from litellm.responses.file_search.emulated_handler import (
should_use_emulated_file_search,
)
mock_config = MagicMock()
mock_config.supports_native_file_search.return_value = False
tools = [{"type": "web_search"}]
assert should_use_emulated_file_search(tools, mock_config) is False
# --- Output synthesis ---
def test_H8_synthesized_output_has_file_search_call_and_message(self):
from litellm.responses.file_search.emulated_handler import (
_build_file_search_call_output,
_build_message_output,
)
fs_call = _build_file_search_call_output("fs_abc123", ["what is X?"])
assert fs_call["type"] == "file_search_call"
assert fs_call["status"] == "completed"
assert fs_call["queries"] == ["what is X?"]
msg = _build_message_output("The answer is 42.", [])
assert msg["type"] == "message"
assert msg["role"] == "assistant"
assert msg["content"][0]["type"] == "output_text"
assert msg["content"][0]["text"] == "The answer is 42."
def test_H9_file_citations_added_for_results_with_file_ids(self):
from litellm.responses.file_search.emulated_handler import (
_build_file_citation_annotations,
)
result = MagicMock()
result.file_id = "file-abc"
result.filename = "doc.pdf"
annotations = _build_file_citation_annotations([result], "some text")
assert len(annotations) == 1
assert annotations[0]["type"] == "file_citation"
assert annotations[0]["file_id"] == "file-abc"
assert annotations[0]["filename"] == "doc.pdf"
def test_H10_no_duplicate_citations_for_same_file(self):
from litellm.responses.file_search.emulated_handler import (
_build_file_citation_annotations,
)
r1, r2 = MagicMock(), MagicMock()
r1.file_id = "file-abc"
r1.filename = "doc.pdf"
r2.file_id = "file-abc" # same file
r2.filename = "doc.pdf"
annotations = _build_file_citation_annotations([r1, r2], "text")
assert len(annotations) == 1
def test_H14_include_search_results_returns_all_chunks(self):
"""All chunks are returned even when they originate from the same file,
matching OpenAI native file_search behaviour."""
from litellm.responses.file_search.emulated_handler import (
_build_search_results_for_include,
)
r1, r2 = MagicMock(), MagicMock()
r1.file_id = "file-abc"
r1.filename = "doc.pdf"
r1.score = 0.9
r1.attributes = {}
r1.content = [{"type": "text", "text": "first hit"}]
r2.file_id = "file-abc" # same file, different chunk from a second query
r2.filename = "doc.pdf"
r2.score = 0.85
r2.attributes = {}
r2.content = [{"type": "text", "text": "second hit"}]
search_results = _build_search_results_for_include([r1, r2])
assert len(search_results) == 2, "Both chunks should be returned, not deduplicated"
assert search_results[0]["text"] == "first hit"
assert search_results[1]["text"] == "second hit"
# --- End-to-end (mocked) ---
@pytest.mark.asyncio
async def test_H11_emulated_full_flow_provider_calls_tool(self):
"""Full flow: provider calls file_search function → search → follow-up → OpenAI output."""
from litellm.responses.file_search.emulated_handler import (
aresponses_with_emulated_file_search,
)
first_resp = self._make_mock_responses_api_response(include_function_call=True)
final_resp = self._make_mock_responses_api_response(text="Deep research enables multi-step queries.")
search_result = MagicMock()
search_result.file_id = "file-xyz"
search_result.filename = "research.pdf"
search_result.score = 0.95
search_result.content = [{"type": "text", "text": "deep research context..."}]
mock_search_response = MagicMock()
mock_search_response.data = [search_result]
with patch(
"litellm.responses.file_search.emulated_handler._call_aresponses",
new=AsyncMock(side_effect=[first_resp, final_resp]),
), patch(
"litellm.vector_stores.main.asearch",
new=AsyncMock(return_value=mock_search_response),
):
result = await aresponses_with_emulated_file_search(
input="What is deep research?",
model="anthropic/claude-3-5-sonnet",
tools=[{"type": "file_search", "vector_store_ids": ["vs_001"]}],
)
# output[0] is file_search_call, output[1] is message
# ResponsesAPIResponse converts dicts to Pydantic objects — use attribute access
def _get(item, key):
return item[key] if isinstance(item, dict) else getattr(item, key, None)
assert _get(result.output[0], "type") == "file_search_call"
assert _get(result.output[0], "status") == "completed"
assert _get(result.output[1], "type") == "message"
content0 = _get(result.output[1], "content")[0]
assert "Deep research" in _get(content0, "text")
annotations = _get(content0, "annotations")
assert any(_get(a, "file_id") == "file-xyz" for a in annotations)
@pytest.mark.asyncio
async def test_H11b_emulated_full_flow_primary_queries_schema(self):
"""Primary path: provider returns queries (plural array) as defined in the tool schema."""
from litellm.responses.file_search.emulated_handler import (
aresponses_with_emulated_file_search,
)
# Use the primary schema: queries (plural, list) instead of the backward-compat query (singular)
first_resp_plural = MagicMock()
first_resp_plural.output = [
{
"type": "function_call",
"name": "litellm_file_search",
"call_id": "call_plural",
"arguments": '{"queries": ["what is deep research?", "multi-step reasoning"], "vector_store_id": "vs_001"}',
}
]
first_resp_plural.id = "resp_plural"
first_resp_plural.created_at = 1700000000
first_resp_plural.model = "claude-3-5-sonnet"
first_resp_plural.usage = None
final_resp = self._make_mock_responses_api_response(text="Deep research uses multiple queries.")
search_result = MagicMock()
search_result.file_id = "file-multi"
search_result.filename = "multi.pdf"
search_result.score = 0.9
search_result.content = [{"type": "text", "text": "multi-query context"}]
mock_search_response = MagicMock()
mock_search_response.data = [search_result]
with patch(
"litellm.responses.file_search.emulated_handler._call_aresponses",
new=AsyncMock(side_effect=[first_resp_plural, final_resp]),
), patch(
"litellm.vector_stores.main.asearch",
new=AsyncMock(return_value=mock_search_response),
):
result = await aresponses_with_emulated_file_search(
input="What is deep research?",
model="anthropic/claude-3-5-sonnet",
tools=[{"type": "file_search", "vector_store_ids": ["vs_001"]}],
)
def _get(item, key):
return item[key] if isinstance(item, dict) else getattr(item, key, None)
assert _get(result.output[0], "type") == "file_search_call"
# Two queries were issued, both should appear in the output
assert len(_get(result.output[0], "queries")) == 2
assert _get(result.output[1], "type") == "message"
@pytest.mark.asyncio
async def test_H12_emulated_flow_provider_answers_without_tool_call(self):
"""If provider answers directly (no tool call), still return OpenAI format."""
from litellm.responses.file_search.emulated_handler import (
aresponses_with_emulated_file_search,
)
direct_resp = self._make_mock_responses_api_response(text="I already know the answer.")
with patch(
"litellm.responses.file_search.emulated_handler._call_aresponses",
new=AsyncMock(return_value=direct_resp),
):
result = await aresponses_with_emulated_file_search(
input="What is 2+2?",
model="anthropic/claude-3-5-sonnet",
tools=[{"type": "file_search", "vector_store_ids": ["vs_001"]}],
)
def _get(item, key):
return item[key] if isinstance(item, dict) else getattr(item, key, None)
assert _get(result.output[0], "type") == "file_search_call"
assert _get(result.output[1], "type") == "message"
assert "I already know" in _get(_get(result.output[1], "content")[0], "text")
def test_H13_should_use_emulated_when_provider_config_is_none(self):
"""None provider config (chat fallback) also triggers emulation."""
from litellm.responses.file_search.emulated_handler import (
should_use_emulated_file_search,
)
tools = [{"type": "file_search", "vector_store_ids": ["vs_abc"]}]
assert should_use_emulated_file_search(tools, None) is True
@pytest.mark.asyncio
async def test_H15_sub_calls_carry_internal_call_flag(self):
"""Both internal aresponses sub-calls receive _is_litellm_internal_call=True.
This ensures wrapper_async skips success/failure callbacks for sub-calls so
billing fires exactly once (on the outer call) with the synthesized result.
"""
from litellm.responses.file_search.emulated_handler import (
aresponses_with_emulated_file_search,
)
first_resp = self._make_mock_responses_api_response(include_function_call=True)
final_resp = self._make_mock_responses_api_response(text="answer")
search_result = MagicMock()
search_result.file_id = "file-h15"
search_result.filename = "h15.pdf"
search_result.score = 0.9
search_result.content = [{"type": "text", "text": "context"}]
mock_search_response = MagicMock()
mock_search_response.data = [search_result]
captured_kwargs: list = []
async def _capture(*args, **kwargs):
captured_kwargs.append(dict(kwargs))
return captured_kwargs.__len__() == 1 and first_resp or final_resp
with patch(
"litellm.responses.file_search.emulated_handler._call_aresponses",
new=AsyncMock(side_effect=[first_resp, final_resp]),
) as mock_call, patch(
"litellm.vector_stores.main.asearch",
new=AsyncMock(return_value=mock_search_response),
):
# Intercept kwargs before the mock returns
original_side_effect = [first_resp, final_resp]
call_kwargs: list = []
async def _intercept(**kwargs): # type: ignore[misc]
call_kwargs.append(dict(kwargs))
return original_side_effect.pop(0)
mock_call.side_effect = _intercept
await aresponses_with_emulated_file_search(
input="What is H15?",
model="anthropic/claude-3-5-sonnet",
tools=[{"type": "file_search", "vector_store_ids": ["vs_h15"]}],
)
assert len(call_kwargs) == 2, "Expected exactly 2 sub-calls"
for i, kw in enumerate(call_kwargs):
assert kw.get("_is_litellm_internal_call") is True, (
f"Sub-call {i} must carry _is_litellm_internal_call=True to suppress "
"billing callbacks in wrapper_async"
)

View file

@ -1,3 +1,9 @@
from unittest.mock import MagicMock, patch
import pytest
import litellm
from litellm.llms.vertex_ai.batches.handler import VertexAIBatchPrediction
from litellm.llms.vertex_ai.batches.transformation import VertexAIBatchTransformation
@ -36,3 +42,124 @@ def test_output_file_id_falls_back_to_output_uri_prefix_with_predictions_jsonl()
output_file_id
== "gs://test-bucket/litellm-vertex-files/publishers/google/models/gemini-2.5-pro/prediction-model-456/predictions.jsonl"
)
def test_vertex_ai_cancel_batch():
"""Test that vertex_ai cancel_batch calls the correct API endpoint"""
handler = VertexAIBatchPrediction(gcs_bucket_name="test-bucket")
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {
"name": "projects/test-project/locations/us-central1/batchPredictionJobs/123456",
"state": "JOB_STATE_CANCELLING",
"createTime": "2024-03-17T10:00:00.000000Z",
"inputConfig": {
"gcsSource": {
"uris": ["gs://test-bucket/input.jsonl"]
}
},
"outputConfig": {
"gcsDestination": {
"outputUriPrefix": "gs://test-bucket/output"
}
}
}
with patch("litellm.llms.vertex_ai.batches.handler._get_httpx_client") as mock_client:
mock_client.return_value.post.return_value = mock_response
mock_client.return_value.get.return_value = mock_response
with patch.object(handler, "_ensure_access_token") as mock_auth:
mock_auth.return_value = ("fake-token", "test-project")
response = handler.cancel_batch(
_is_async=False,
batch_id="123456",
api_base=None,
vertex_credentials=None,
vertex_project="test-project",
vertex_location="us-central1",
timeout=600.0,
max_retries=None,
)
assert response.id == "123456"
assert response.status == "cancelling"
mock_client.return_value.post.assert_called_once()
mock_client.return_value.get.assert_called_once()
call_args = mock_client.return_value.post.call_args
assert ":cancel" in call_args.kwargs["url"]
def test_vertex_ai_cancel_batch_forwards_timeout():
"""Test that timeout is forwarded to the POST (cancel) HTTP call.
Note: the follow-up GET (retrieve) call does not accept a timeout
parameter in the underlying HTTP handler, so it is intentionally omitted.
"""
def test_vertex_ai_cancel_batch_custom_proxy_retrieve_url():
"""Retrieve URL should go through the custom proxy, not bypass it"""
handler = VertexAIBatchPrediction(gcs_bucket_name="test-bucket")
mock_response = MagicMock()
mock_response.status_code = 200
mock_response.json.return_value = {
"name": "projects/test-project/locations/us-central1/batchPredictionJobs/123456",
"state": "JOB_STATE_CANCELLING",
"createTime": "2024-03-17T10:00:00.000000Z",
"inputConfig": {"gcsSource": {"uris": ["gs://test-bucket/input.jsonl"]}},
"outputConfig": {"gcsDestination": {"outputUriPrefix": "gs://test-bucket/output"}},
}
with patch("litellm.llms.vertex_ai.batches.handler._get_httpx_client") as mock_client:
mock_client.return_value.post.return_value = mock_response
mock_client.return_value.get.return_value = mock_response
with patch.object(handler, "_ensure_access_token") as mock_auth:
mock_auth.return_value = ("fake-token", "test-project")
handler.cancel_batch(
_is_async=False,
batch_id="123456",
api_base="https://my-proxy.example.com",
vertex_credentials=None,
vertex_project="test-project",
vertex_location="us-central1",
timeout=600.0,
max_retries=None,
)
post_url = mock_client.return_value.post.call_args.kwargs["url"]
get_url = mock_client.return_value.get.call_args.kwargs["url"]
assert "my-proxy.example.com" in post_url
assert ":cancel" in post_url
assert "my-proxy.example.com" in get_url
assert ":cancel" not in get_url
assert "googleapis.com" not in get_url
@pytest.mark.asyncio
async def test_litellm_cancel_batch_vertex_ai():
"""Test that litellm.cancel_batch works with vertex_ai provider"""
mock_response = MagicMock()
mock_response.id = "batch_123"
mock_response.status = "cancelling"
with patch("litellm.batches.main.vertex_ai_batches_instance") as mock_instance:
mock_instance.cancel_batch.return_value = mock_response
response = litellm.cancel_batch(
batch_id="batch_123",
custom_llm_provider="vertex_ai",
vertex_project="test-project",
vertex_location="us-central1",
)
assert mock_instance.cancel_batch.called
assert response.id == "batch_123"
assert response.status == "cancelling"

View file

@ -1329,3 +1329,78 @@ def test_non_org_admin_with_organizations_list():
organization_memberships=[membership],
)
assert _user_is_org_admin({"organizations": ["org-1"]}, user_obj) is False
@pytest.mark.asyncio
async def test_initialize_pass_through_registers_wildcard_for_auth_subpath():
"""
Test that initialize_pass_through_endpoints registers both base path and
wildcard path in openai_routes when auth=true and include_subpath=true,
and that subpath requests pass is_llm_api_route.
Also verifies:
- Dedup: calling init twice does not duplicate entries
- Cleanup: removing the endpoint cleans up openai_routes
"""
from litellm.proxy._types import LiteLLMRoutes
from litellm.proxy.pass_through_endpoints.pass_through_endpoints import (
InitPassThroughEndpointHelpers,
initialize_pass_through_endpoints,
)
base_path = "/v1/ocr/nvidia/community/nemoretriever-ocr-v1"
wildcard_path = base_path + "/*"
endpoint_config = {
"path": base_path,
"target": "https://httpbin.org/post",
"include_subpath": True,
"auth": True,
"headers": {"content-type": "application/json"},
}
original_routes = LiteLLMRoutes.openai_routes.value[:]
try:
with patch(
"litellm.proxy.proxy_server.app",
MagicMock(),
), patch(
"litellm.proxy.proxy_server.premium_user",
True,
), patch(
"litellm.proxy.proxy_server.config_passthrough_endpoints",
None,
):
await initialize_pass_through_endpoints([endpoint_config])
# Both base and wildcard paths should be registered
assert base_path in LiteLLMRoutes.openai_routes.value
assert wildcard_path in LiteLLMRoutes.openai_routes.value
# Subpath requests should pass the auth route check
assert RouteChecks.is_llm_api_route(base_path) is True
assert RouteChecks.is_llm_api_route(base_path + "/v1/infer") is True
# Calling init again should not duplicate entries
await initialize_pass_through_endpoints([endpoint_config])
assert LiteLLMRoutes.openai_routes.value.count(base_path) == 1
assert LiteLLMRoutes.openai_routes.value.count(wildcard_path) == 1
# Removing the endpoint should clean up openai_routes
# remove_endpoint_routes takes endpoint_id (UUID portion of
# the route key "{id}:exact:{path}:{methods}")
registered = InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes()
endpoint_ids = {k.split(":")[0] for k in registered}
for eid in endpoint_ids:
InitPassThroughEndpointHelpers.remove_endpoint_routes(eid)
assert base_path not in LiteLLMRoutes.openai_routes.value
assert wildcard_path not in LiteLLMRoutes.openai_routes.value
finally:
LiteLLMRoutes.openai_routes.value[:] = original_routes
# Clean up any routes registered during this test to avoid
# polluting the module-level _registered_pass_through_routes
registered = InitPassThroughEndpointHelpers.get_all_registered_pass_through_routes()
for k in registered:
InitPassThroughEndpointHelpers.remove_endpoint_routes(
k.split(":")[0]
)

View file

@ -32,7 +32,6 @@ from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessin
from litellm.proxy.utils import ProxyLogging
from litellm.types.guardrails import GuardrailEventHooks
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
@ -86,10 +85,10 @@ class TestHasPostCallGuardrails:
with patch("litellm.callbacks", [PostCallGuardrail()]):
assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is True
def test_returns_true_for_event_hook_none(self):
"""event_hook=None means 'all events', including post_call."""
def test_returns_false_for_event_hook_none(self):
"""event_hook=None is not an explicit post_call registration for deferral."""
with patch("litellm.callbacks", [AllEventsGuardrail()]):
assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is True
assert ProxyBaseLLMRequestProcessing._has_post_call_guardrails() is False
def test_returns_false_for_pre_call_only(self):
with patch("litellm.callbacks", [PreCallGuardrail()]):
@ -112,7 +111,10 @@ class TestHasPostCallGuardrails:
super().__init__(
guardrail_name="list-post",
default_on=True,
event_hook=[GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call],
event_hook=[
GuardrailEventHooks.pre_call,
GuardrailEventHooks.post_call,
],
)
with patch("litellm.callbacks", [ListGuardrail()]):
@ -418,7 +420,9 @@ class TestDeferredStreamingClosure:
await asyncio.sleep(0)
assert guardrail_called is True, "Guardrail hook should be called"
assert logger_called is False, "Non-guardrail logger should NOT be called by closure"
assert (
logger_called is False
), "Non-guardrail logger should NOT be called by closure"
@pytest.mark.asyncio
async def test_closure_passes_guardrail_modified_response_to_logging(self):
@ -463,8 +467,9 @@ class TestDeferredStreamingClosure:
await asyncio.sleep(0)
await asyncio.sleep(0)
assert logged_response is modified_response, \
"Logging must receive the guardrail-modified response"
assert (
logged_response is modified_response
), "Logging must receive the guardrail-modified response"
@pytest.mark.asyncio
async def test_closure_logs_even_on_guardrail_exception(self):
@ -511,11 +516,13 @@ class TestDeferredStreamingClosure:
await asyncio.sleep(0)
await asyncio.sleep(0)
assert logging_called is True, \
"Logging must fire even when guardrail raises HTTPException"
assert mock_logging_obj.model_call_details["metadata"].get(
"guardrail_blocked"
) is True, "guardrail_blocked must be set for HTTPException"
assert (
logging_called is True
), "Logging must fire even when guardrail raises HTTPException"
assert (
mock_logging_obj.model_call_details["metadata"].get("guardrail_blocked")
is True
), "guardrail_blocked must be set for HTTPException"
@pytest.mark.asyncio
async def test_transient_error_does_not_set_guardrail_blocked(self):
@ -556,9 +563,10 @@ class TestDeferredStreamingClosure:
await asyncio.sleep(0)
assert mock_logging_obj.model_call_details["metadata"].get(
"guardrail_blocked"
) is not True, "guardrail_blocked must NOT be set for transient errors"
assert (
mock_logging_obj.model_call_details["metadata"].get("guardrail_blocked")
is not True
), "guardrail_blocked must NOT be set for transient errors"
@pytest.mark.asyncio
async def test_production_closure_integration(self):
@ -620,10 +628,10 @@ class TestDeferredStreamingClosure:
await asyncio.sleep(0)
await asyncio.sleep(0)
assert hook_called is True, \
"Production closure must call guardrail hook"
assert logged_response is modified_response, \
"Production closure must pass guardrail-modified response to logging"
assert hook_called is True, "Production closure must call guardrail hook"
assert (
logged_response is modified_response
), "Production closure must pass guardrail-modified response to logging"
@pytest.mark.asyncio
async def test_apply_guardrail_path_uses_unified_guardrail(self):
@ -686,10 +694,12 @@ class TestDeferredStreamingClosure:
await asyncio.sleep(0)
await asyncio.sleep(0)
assert unified_hook_called is True, \
"apply_guardrail guardrails must be dispatched through UnifiedLLMGuardrails"
assert logged_response is not None, \
"Logging must fire after unified guardrail path"
assert (
unified_hook_called is True
), "apply_guardrail guardrails must be dispatched through UnifiedLLMGuardrails"
assert (
logged_response is not None
), "Logging must fire after unified guardrail path"
@pytest.mark.asyncio
async def test_hooks_receive_merged_guardrail_data(self):
@ -742,11 +752,10 @@ class TestDeferredStreamingClosure:
merged["_merged_marker"] = True
return merged
with patch("litellm.callbacks", [guardrail]), \
patch(
"litellm.proxy.utils._check_and_merge_model_level_guardrails",
side_effect=mock_merge,
):
with patch("litellm.callbacks", [guardrail]), patch(
"litellm.proxy.utils._check_and_merge_model_level_guardrails",
side_effect=mock_merge,
):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data=captured_data,
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
@ -756,8 +765,9 @@ class TestDeferredStreamingClosure:
)
assert hook_received_data is not None, "Guardrail hook must be called"
assert hook_received_data.get("_merged_marker") is True, \
"Hook must receive guardrail_data (merged), not original captured_data"
assert (
hook_received_data.get("_merged_marker") is True
), "Hook must receive guardrail_data (merged), not original captured_data"
assert "model-guardrail" in hook_received_data.get("metadata", {}).get(
"guardrails", []
), "Hook data must contain model-level guardrails"
@ -773,6 +783,7 @@ class TestDeferredStreamingClosure:
be silently skipped at execution time if captured_data (unmerged) were
passed instead of guardrail_data (merged)."""
import copy
from litellm.types.utils import GenericGuardrailAPIInputs
unified_received_data = None
@ -812,6 +823,7 @@ class TestDeferredStreamingClosure:
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
UnifiedLLMGuardrails,
)
original_unified_hook = UnifiedLLMGuardrails.async_post_call_success_hook
async def tracking_unified_hook(self, user_api_key_dict, data, response):
@ -819,16 +831,14 @@ class TestDeferredStreamingClosure:
unified_received_data = data
return response
with patch("litellm.callbacks", [guardrail]), \
patch(
"litellm.proxy.utils._check_and_merge_model_level_guardrails",
side_effect=mock_merge,
), \
patch.object(
UnifiedLLMGuardrails,
"async_post_call_success_hook",
tracking_unified_hook,
):
with patch("litellm.callbacks", [guardrail]), patch(
"litellm.proxy.utils._check_and_merge_model_level_guardrails",
side_effect=mock_merge,
), patch.object(
UnifiedLLMGuardrails,
"async_post_call_success_hook",
tracking_unified_hook,
):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data=captured_data,
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
@ -837,14 +847,15 @@ class TestDeferredStreamingClosure:
cache_hit=False,
)
assert unified_received_data is not None, \
"UnifiedLLMGuardrails must be called for apply_guardrail guardrails"
assert unified_received_data.get("_merged_marker") is True, \
"UnifiedLLMGuardrails must receive guardrail_data (merged), not captured_data"
assert "model-apply-guardrail" in unified_received_data.get(
"metadata", {}
).get("guardrails", []), \
"UnifiedLLMGuardrails data must contain model-level guardrails"
assert (
unified_received_data is not None
), "UnifiedLLMGuardrails must be called for apply_guardrail guardrails"
assert (
unified_received_data.get("_merged_marker") is True
), "UnifiedLLMGuardrails must receive guardrail_data (merged), not captured_data"
assert "model-apply-guardrail" in unified_received_data.get("metadata", {}).get(
"guardrails", []
), "UnifiedLLMGuardrails data must contain model-level guardrails"
@pytest.mark.asyncio
async def test_multiple_guardrails_all_receive_merged_data(self):
@ -887,11 +898,10 @@ class TestDeferredStreamingClosure:
merged["_merged_marker"] = True
return merged
with patch("litellm.callbacks", [guardrail_a, guardrail_b]), \
patch(
"litellm.proxy.utils._check_and_merge_model_level_guardrails",
side_effect=mock_merge,
):
with patch("litellm.callbacks", [guardrail_a, guardrail_b]), patch(
"litellm.proxy.utils._check_and_merge_model_level_guardrails",
side_effect=mock_merge,
):
await ProxyBaseLLMRequestProcessing._run_deferred_stream_guardrails(
captured_data=captured_data,
captured_user_api_key_dict=UserAPIKeyAuth(api_key="test"),
@ -901,10 +911,10 @@ class TestDeferredStreamingClosure:
)
for name in ("guardrail-a", "guardrail-b"):
assert name in received_data_per_guardrail, \
f"{name} must be called"
assert received_data_per_guardrail[name].get("_merged_marker") is True, \
f"{name} must receive guardrail_data (merged), not captured_data"
assert name in received_data_per_guardrail, f"{name} must be called"
assert (
received_data_per_guardrail[name].get("_merged_marker") is True
), f"{name} must receive guardrail_data (merged), not captured_data"
@pytest.mark.asyncio
async def test_logging_fires_even_if_guardrail_init_raises(self):
@ -940,5 +950,6 @@ class TestDeferredStreamingClosure:
await asyncio.sleep(0)
await asyncio.sleep(0)
assert logging_called is True, \
"Logging must fire even when guardrail initialization raises"
assert (
logging_called is True
), "Logging must fire even when guardrail initialization raises"

View file

@ -0,0 +1,383 @@
"""
Unit tests for prompt management support in the Responses API.
Covers:
A) str input is coerced to a message list before merging with the template
B) list input is merged with the template
C) no prompt_id → hook is skipped, input is unchanged
D) model override from the prompt template is applied
E) prompt_template_optional_params flow into the request
F) non-message items in input are filtered out
G) model override re-resolves provider
H) async path calls async_get_chat_completion_prompt
I) async path propagates optional params to downstream handler
"""
import asyncio
from typing import List
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.types.llms.openai import AllMessageValues
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
def _make_logging_obj(
merged_model: str,
merged_messages: List[AllMessageValues],
should_run: bool = True,
merged_optional_params: dict = None,
) -> MagicMock:
"""Return a mock LiteLLMLoggingObj pre-configured for prompt management."""
if merged_optional_params is None:
merged_optional_params = {}
logging_obj = MagicMock()
logging_obj.__class__ = LiteLLMLoggingObj
logging_obj.should_run_prompt_management_hooks.return_value = should_run
prompt_return = (merged_model, merged_messages, merged_optional_params)
logging_obj.get_chat_completion_prompt.return_value = prompt_return
logging_obj.async_get_chat_completion_prompt = AsyncMock(
return_value=prompt_return
)
logging_obj.model_call_details = {}
return logging_obj
def _patch_responses_dispatch():
"""Patch everything after the prompt management block so tests stay unit-level."""
return [
patch(
"litellm.responses.main.litellm.get_llm_provider",
return_value=("gpt-4o", "openai", None, None),
),
patch(
"litellm.responses.mcp.litellm_proxy_mcp_handler."
"LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway",
return_value=False,
),
patch(
"litellm.responses.main.ProviderConfigManager"
".get_provider_responses_api_config",
return_value=None,
),
patch(
"litellm.responses.main.litellm_completion_transformation_handler"
".response_api_handler",
return_value=MagicMock(),
),
]
# ---------------------------------------------------------------------------
# Tests
# ---------------------------------------------------------------------------
class TestResponsesAPIPromptManagement:
def test_str_input_coerced_and_merged(self):
"""[A] str input is wrapped into a message list before being passed to the hook."""
template_messages: List[AllMessageValues] = [
{"role": "system", "content": "You are a summariser."}, # type: ignore[list-item]
]
client_message: List[AllMessageValues] = [
{"role": "user", "content": "Tell me about AI."}, # type: ignore[list-item]
]
expected_merged = template_messages + client_message
logging_obj = _make_logging_obj(
merged_model="openai/gpt-4o",
merged_messages=expected_merged,
)
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3]:
import litellm
litellm.responses(
input="Tell me about AI.",
model="gpt-4o",
prompt_id="summariser-prompt",
prompt_variables={},
litellm_logging_obj=logging_obj,
)
logging_obj.get_chat_completion_prompt.assert_called_once()
call_kwargs = logging_obj.get_chat_completion_prompt.call_args.kwargs
# str was coerced to a single user message before being passed to the hook
assert call_kwargs["messages"] == [
{"role": "user", "content": "Tell me about AI."}
]
assert call_kwargs["prompt_id"] == "summariser-prompt"
def test_list_input_merged_with_template(self):
"""[B] list input is passed directly to the hook and merged with the template."""
template_messages: List[AllMessageValues] = [
{"role": "system", "content": "You are helpful."}, # type: ignore[list-item]
]
client_messages = [
{"role": "user", "content": [{"type": "input_text", "text": "Hello"}]},
]
expected_merged = template_messages + client_messages # type: ignore[operator]
logging_obj = _make_logging_obj(
merged_model="openai/gpt-4o",
merged_messages=expected_merged, # type: ignore[arg-type]
)
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3]:
import litellm
litellm.responses(
input=client_messages, # type: ignore[arg-type]
model="gpt-4o",
prompt_id="helper-prompt",
litellm_logging_obj=logging_obj,
)
logging_obj.get_chat_completion_prompt.assert_called_once()
call_kwargs = logging_obj.get_chat_completion_prompt.call_args.kwargs
assert call_kwargs["messages"] == client_messages
def test_no_prompt_id_skips_hook(self):
"""[C] When prompt_id is absent, prompt management hooks are not called."""
logging_obj = _make_logging_obj(
merged_model="openai/gpt-4o",
merged_messages=[],
should_run=False,
)
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3]:
import litellm
litellm.responses(
input="Hello",
model="gpt-4o",
litellm_logging_obj=logging_obj,
)
logging_obj.get_chat_completion_prompt.assert_not_called()
def test_optional_params_from_template_applied(self):
"""[E] prompt_template_optional_params (e.g. temperature) flow into the request."""
template_messages: List[AllMessageValues] = [
{"role": "user", "content": "Hello"}, # type: ignore[list-item]
]
# Simulate get_chat_completion_prompt returning merged optional params
# that include a template-defined temperature
merged_kwargs = {"temperature": 0.2}
logging_obj = MagicMock()
logging_obj.__class__ = LiteLLMLoggingObj
logging_obj.should_run_prompt_management_hooks.return_value = True
logging_obj.get_chat_completion_prompt.return_value = (
"openai/gpt-4o",
template_messages,
merged_kwargs,
)
logging_obj.model_call_details = {}
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3] as mock_handler:
import litellm
litellm.responses(
input="Hello",
model="gpt-4o",
prompt_id="t",
litellm_logging_obj=logging_obj,
)
# temperature from the template should reach the downstream handler via local_vars
handler_call_kwargs = mock_handler.call_args.kwargs
request_params = handler_call_kwargs.get("responses_api_request", {})
assert request_params.get("temperature") == 0.2
def test_model_override_from_template(self):
"""[D] Model returned by the prompt hook overrides the original request model."""
template_messages: List[AllMessageValues] = [
{"role": "user", "content": "{{query}}"}, # type: ignore[list-item]
]
logging_obj = _make_logging_obj(
merged_model="openai/gpt-4o-mini", # overridden model from template
merged_messages=template_messages,
)
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3] as mock_handler:
import litellm
litellm.responses(
input="What is AI?",
model="gpt-4o",
prompt_id="query-prompt",
prompt_variables={"query": "What is AI?"},
litellm_logging_obj=logging_obj,
)
# The model passed to the downstream handler should be the overridden one
handler_call_kwargs = mock_handler.call_args.kwargs
assert handler_call_kwargs.get("model") == "openai/gpt-4o-mini"
def test_non_message_input_items_filtered(self):
"""[F] Non-message items in ResponseInputParam (e.g. function_call_output) are
filtered out before being passed to the prompt hook, avoiding malformed merges."""
template_messages: List[AllMessageValues] = [
{"role": "system", "content": "You are helpful."}, # type: ignore[list-item]
]
mixed_input = [
{"role": "user", "content": "Hello"},
{"type": "function_call_output", "call_id": "abc", "output": "42"},
]
logging_obj = _make_logging_obj(
merged_model="openai/gpt-4o",
merged_messages=template_messages + [{"role": "user", "content": "Hello"}], # type: ignore[operator]
)
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3]:
import litellm
litellm.responses(
input=mixed_input, # type: ignore[arg-type]
model="gpt-4o",
prompt_id="filter-test",
litellm_logging_obj=logging_obj,
)
call_kwargs = logging_obj.get_chat_completion_prompt.call_args.kwargs
passed_messages = call_kwargs["messages"]
assert all(isinstance(m, dict) and "role" in m for m in passed_messages)
assert len(passed_messages) == 1
def test_model_override_re_resolves_provider(self):
"""[G] When the prompt template overrides the model to a different provider,
custom_llm_provider is re-resolved so downstream routing uses the correct provider."""
template_messages: List[AllMessageValues] = [
{"role": "user", "content": "Hi"}, # type: ignore[list-item]
]
logging_obj = _make_logging_obj(
merged_model="anthropic/claude-3-5-sonnet",
merged_messages=template_messages,
)
patches = _patch_responses_dispatch()
with (
patch(
"litellm.responses.main.litellm.get_llm_provider",
side_effect=[
("gpt-4o", "openai", None, None),
("claude-3-5-sonnet", "anthropic", None, None),
],
),
patches[1],
patches[2],
patches[3] as mock_handler,
):
import litellm
litellm.responses(
input="Hi",
model="gpt-4o",
prompt_id="cross-provider",
litellm_logging_obj=logging_obj,
)
handler_call_kwargs = mock_handler.call_args.kwargs
assert handler_call_kwargs.get("custom_llm_provider") == "anthropic"
class TestAsyncResponsesAPIPromptManagement:
"""Tests for the async aresponses() prompt management path.
aresponses() calls async_get_chat_completion_prompt at the outer async
level, then pops prompt_id from kwargs and passes merged_optional_params
via an internal kwarg. The sync responses() path sees no prompt_id and
skips the sync hook entirely — preventing double-merge of template messages.
"""
@pytest.mark.asyncio
async def test_async_calls_async_hook_not_sync(self):
"""[H] aresponses() invokes async_get_chat_completion_prompt and the
sync get_chat_completion_prompt is NOT called (no double-merge)."""
template_messages: List[AllMessageValues] = [
{"role": "system", "content": "You are helpful."}, # type: ignore[list-item]
]
logging_obj = _make_logging_obj(
merged_model="openai/gpt-4o",
merged_messages=template_messages + [{"role": "user", "content": "Hi"}], # type: ignore[list-item]
)
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3]:
import litellm
await litellm.aresponses(
input="Hi",
model="gpt-4o",
prompt_id="async-test",
prompt_variables={},
litellm_logging_obj=logging_obj,
)
logging_obj.async_get_chat_completion_prompt.assert_called_once()
logging_obj.get_chat_completion_prompt.assert_not_called()
call_kwargs = logging_obj.async_get_chat_completion_prompt.call_args.kwargs
assert call_kwargs["prompt_id"] == "async-test"
@pytest.mark.asyncio
async def test_async_optional_params_propagated(self):
"""[I] Template-defined optional params (e.g. temperature) from the async
hook reach the downstream handler — they are NOT silently discarded."""
template_messages: List[AllMessageValues] = [
{"role": "user", "content": "Hello"}, # type: ignore[list-item]
]
logging_obj = _make_logging_obj(
merged_model="openai/gpt-4o",
merged_messages=template_messages,
merged_optional_params={"temperature": 0.7},
)
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3] as mock_handler:
import litellm
await litellm.aresponses(
input="Hello",
model="gpt-4o",
prompt_id="async-temp",
litellm_logging_obj=logging_obj,
)
logging_obj.get_chat_completion_prompt.assert_not_called()
handler_call_kwargs = mock_handler.call_args.kwargs
request_params = handler_call_kwargs.get("responses_api_request", {})
assert request_params.get("temperature") == 0.7
@pytest.mark.asyncio
async def test_async_non_message_items_filtered(self):
"""[J] Non-message items are filtered in the async path too."""
template_messages: List[AllMessageValues] = [
{"role": "system", "content": "Be helpful."}, # type: ignore[list-item]
]
mixed_input = [
{"role": "user", "content": "Hello"},
{"type": "function_call_output", "call_id": "abc", "output": "42"},
]
logging_obj = _make_logging_obj(
merged_model="openai/gpt-4o",
merged_messages=template_messages + [{"role": "user", "content": "Hello"}], # type: ignore[operator]
)
patches = _patch_responses_dispatch()
with patches[0], patches[1], patches[2], patches[3]:
import litellm
await litellm.aresponses(
input=mixed_input, # type: ignore[arg-type]
model="gpt-4o",
prompt_id="async-filter",
litellm_logging_obj=logging_obj,
)
logging_obj.async_get_chat_completion_prompt.assert_called_once()
logging_obj.get_chat_completion_prompt.assert_not_called()
call_kwargs = logging_obj.async_get_chat_completion_prompt.call_args.kwargs
passed_messages = call_kwargs["messages"]
assert all(isinstance(m, dict) and "role" in m for m in passed_messages)
assert len(passed_messages) == 1

View file

@ -657,3 +657,284 @@ def test_cache_key_does_not_double_hash_user_api_key_hash():
user_key=user_api_key_hash,
)
assert key.endswith(user_api_key_hash)
def test_get_effective_flags_returns_per_group_config():
"""
_get_effective_flags should return per-group flags when the model group has an entry
in model_group_affinity_config, and global flags otherwise.
"""
callback = DeploymentAffinityCheck(
cache=AsyncMock(),
ttl_seconds=60,
enable_user_key_affinity=True,
enable_responses_api_affinity=True,
enable_session_id_affinity=False,
model_group_affinity_config={
"gpt-4": ["deployment_affinity"],
"claude-3": ["session_affinity", "responses_api_deployment_check"],
},
)
# gpt-4: only deployment_affinity
user_key, responses_api, session_id = callback._get_effective_flags("gpt-4")
assert user_key is True
assert responses_api is False
assert session_id is False
# claude-3: session_affinity + responses_api_deployment_check
user_key, responses_api, session_id = callback._get_effective_flags("claude-3")
assert user_key is False
assert responses_api is True
assert session_id is True
# unconfigured-model: falls back to global flags
user_key, responses_api, session_id = callback._get_effective_flags(
"unconfigured-model"
)
assert user_key is True
assert responses_api is True
assert session_id is False
@pytest.mark.asyncio
async def test_model_group_affinity_config_only_applies_to_configured_group():
"""
When model_group_affinity_config is set without global optional_pre_call_checks,
only configured model groups should get affinity behavior.
"""
mock_response_data = {
"id": "resp_mock-resp-per-group",
"object": "response",
"created_at": 1741476542,
"status": "completed",
"model": "openai/gpt-4",
"output": [
{
"type": "message",
"id": "msg_pg",
"status": "completed",
"role": "assistant",
"content": [{"type": "output_text", "text": "Per-group response"}],
}
],
"parallel_tool_calls": True,
"usage": {"input_tokens": 5, "output_tokens": 5, "total_tokens": 10},
"text": {"format": {"type": "text"}},
"error": None,
"previous_response_id": None,
}
router = litellm.Router(
model_list=[
{
"model_name": "gpt-4",
"litellm_params": {
"model": "azure/gpt-4-deploy-1",
"api_key": "mock-key-1",
"api_base": "https://mock-gpt4-1.openai.azure.com",
"api_version": "2024-02-01",
},
"model_info": {"base_model": "gpt-4"},
},
{
"model_name": "gpt-4",
"litellm_params": {
"model": "azure/gpt-4-deploy-2",
"api_key": "mock-key-2",
"api_base": "https://mock-gpt4-2.openai.azure.com",
"api_version": "2024-02-01",
},
"model_info": {"base_model": "gpt-4"},
},
{
"model_name": "claude-3",
"litellm_params": {
"model": "azure/claude-3-deploy-1",
"api_key": "mock-key-3",
"api_base": "https://mock-claude-1.openai.azure.com",
"api_version": "2024-02-01",
},
"model_info": {"base_model": "claude-3"},
},
{
"model_name": "claude-3",
"litellm_params": {
"model": "azure/claude-3-deploy-2",
"api_key": "mock-key-4",
"api_base": "https://mock-claude-2.openai.azure.com",
"api_version": "2024-02-01",
},
"model_info": {"base_model": "claude-3"},
},
],
# No global optional_pre_call_checks — only per-group
model_group_affinity_config={
"gpt-4": ["deployment_affinity"],
},
)
user_api_key_hash = "test-per-group-key"
choice_calls = {"count": 0}
def deterministic_choice(seq):
choice_calls["count"] += 1
if choice_calls["count"] == 1:
return seq[0]
return seq[1] if len(seq) > 1 else seq[0]
with patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post, patch(
"litellm.router_strategy.simple_shuffle.random.choice",
side_effect=deterministic_choice,
):
mock_post.return_value = MockResponse(mock_response_data, 200)
# gpt-4: affinity should work — second request pinned to same deployment
first = await router.aresponses(
model="gpt-4",
input="Hello",
truncation="auto",
litellm_metadata={"user_api_key_hash": user_api_key_hash},
)
first_model_id = first._hidden_params["model_id"]
second = await router.aresponses(
model="gpt-4",
input="Follow-up",
truncation="auto",
litellm_metadata={"user_api_key_hash": user_api_key_hash},
)
assert second._hidden_params["model_id"] == first_model_id
# claude-3: no affinity configured — should NOT be pinned
choice_calls["count"] = 0
first_claude = await router.aresponses(
model="claude-3",
input="Hello",
truncation="auto",
litellm_metadata={"user_api_key_hash": user_api_key_hash},
)
first_claude_id = first_claude._hidden_params["model_id"]
second_claude = await router.aresponses(
model="claude-3",
input="Follow-up",
truncation="auto",
litellm_metadata={"user_api_key_hash": user_api_key_hash},
)
# With deterministic choice and len>1, second call picks seq[1]
assert second_claude._hidden_params["model_id"] != first_claude_id
@pytest.mark.asyncio
async def test_model_group_affinity_config_falls_back_to_global():
"""
When both global optional_pre_call_checks and model_group_affinity_config are set,
unconfigured model groups should use the global settings.
"""
callback = DeploymentAffinityCheck(
cache=DualCache(),
ttl_seconds=60,
enable_user_key_affinity=True,
enable_responses_api_affinity=False,
enable_session_id_affinity=False,
model_group_affinity_config={
"claude-3": ["session_affinity"],
},
)
stable_model_map_key = "gpt-4"
user_key = "test-fallback-key"
healthy_deployments = [
{
"model_name": stable_model_map_key,
"litellm_params": {"model": "openai/gpt-4"},
"model_info": {"id": "deployment-1"},
},
{
"model_name": stable_model_map_key,
"litellm_params": {"model": "openai/gpt-4"},
"model_info": {"id": "deployment-2"},
},
]
# Set up affinity cache for gpt-4 (should work since global has deployment_affinity)
await callback.async_pre_call_deployment_hook(
kwargs={
"model_info": {"id": "deployment-1"},
"metadata": {
"user_api_key_hash": user_key,
"deployment_model_name": stable_model_map_key,
},
},
call_type=None,
)
# gpt-4 not in model_group_affinity_config, so global flags apply (user_key affinity ON)
filtered = await callback.async_filter_deployments(
model="gpt-4",
healthy_deployments=healthy_deployments,
messages=None,
request_kwargs={"metadata": {"user_api_key_hash": user_key}},
parent_otel_span=None,
)
assert len(filtered) == 1
assert filtered[0]["model_info"]["id"] == "deployment-1"
@pytest.mark.asyncio
async def test_model_group_affinity_config_overrides_global():
"""
When model_group_affinity_config specifies session_affinity for a model group,
user-key affinity (from global config) should NOT apply to that group.
"""
callback = DeploymentAffinityCheck(
cache=DualCache(),
ttl_seconds=60,
enable_user_key_affinity=True,
enable_responses_api_affinity=False,
enable_session_id_affinity=False,
model_group_affinity_config={
"claude-3": ["session_affinity"],
},
)
stable_model_map_key = "claude-3"
user_key = "test-override-key"
healthy_deployments = [
{
"model_name": stable_model_map_key,
"litellm_params": {"model": "anthropic/claude-3-opus"},
"model_info": {"id": "deployment-1"},
},
{
"model_name": stable_model_map_key,
"litellm_params": {"model": "anthropic/claude-3-opus"},
"model_info": {"id": "deployment-2"},
},
]
# Set up user-key affinity cache for claude-3
cache_key = DeploymentAffinityCheck.get_affinity_cache_key(
model_group=stable_model_map_key, user_key=user_key
)
await callback.cache.async_set_cache(
cache_key, {"model_id": "deployment-1"}, ttl=60
)
# claude-3 has per-group config (session_affinity only), so user-key affinity
# should NOT apply even though it's globally enabled
filtered = await callback.async_filter_deployments(
model="claude-3",
healthy_deployments=healthy_deployments,
messages=None,
request_kwargs={"metadata": {"user_api_key_hash": user_key}},
parent_otel_span=None,
)
# All deployments returned (user-key affinity disabled for this group)
assert len(filtered) == 2

View file

@ -38,22 +38,28 @@ def test_check_provider_match_azure_ai_allows_openai_and_azure():
This is needed for Azure Model Router which can route to OpenAI models.
"""
# azure_ai should match openai models
assert _check_provider_match(
model_info={"litellm_provider": "openai"},
custom_llm_provider="azure_ai"
) is True
assert (
_check_provider_match(
model_info={"litellm_provider": "openai"}, custom_llm_provider="azure_ai"
)
is True
)
# azure_ai should match azure models
assert _check_provider_match(
model_info={"litellm_provider": "azure"},
custom_llm_provider="azure_ai"
) is True
assert (
_check_provider_match(
model_info={"litellm_provider": "azure"}, custom_llm_provider="azure_ai"
)
is True
)
# azure_ai should NOT match other providers
assert _check_provider_match(
model_info={"litellm_provider": "anthropic"},
custom_llm_provider="azure_ai"
) is False
assert (
_check_provider_match(
model_info={"litellm_provider": "anthropic"}, custom_llm_provider="azure_ai"
)
is False
)
def test_check_provider_match_github_allows_upstream_provider_metadata():
@ -61,20 +67,29 @@ def test_check_provider_match_github_allows_upstream_provider_metadata():
Test that github provider can match upstream provider metadata.
GitHub Models can provide models from multiple providers.
"""
assert _check_provider_match(
model_info={"litellm_provider": "openai"},
custom_llm_provider="github",
) is True
assert (
_check_provider_match(
model_info={"litellm_provider": "openai"},
custom_llm_provider="github",
)
is True
)
assert _check_provider_match(
model_info={"litellm_provider": "github"},
custom_llm_provider="github",
) is True
assert (
_check_provider_match(
model_info={"litellm_provider": "github"},
custom_llm_provider="github",
)
is True
)
assert _check_provider_match(
model_info={"litellm_provider": "anthropic"},
custom_llm_provider="github",
) is True
assert (
_check_provider_match(
model_info={"litellm_provider": "anthropic"},
custom_llm_provider="github",
)
is True
)
def test_supports_function_calling_github_openai_alias():
@ -604,7 +619,10 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"cache_read_input_token_cost": {"type": "number"},
"cache_read_input_token_cost_above_200k_tokens": {"type": "number"},
"cache_read_input_token_cost_above_272k_tokens": {"type": "number"},
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": {"type": "number"},
"cache_read_input_token_cost_batches": {"type": "number"},
"cache_creation_input_token_cost_above_1hr_above_200k_tokens": {
"type": "number"
},
"cache_read_input_audio_token_cost": {"type": "number"},
"cache_read_input_token_cost_per_audio_token": {"type": "number"},
"cache_read_input_image_token_cost": {"type": "number"},
@ -623,8 +641,12 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"input_cost_per_token_above_272k_tokens": {"type": "number"},
"cache_read_input_token_cost_flex": {"type": "number"},
"cache_read_input_token_cost_priority": {"type": "number"},
"cache_read_input_token_cost_above_200k_tokens_priority": {"type": "number"},
"cache_read_input_token_cost_above_272k_tokens_priority": {"type": "number"},
"cache_read_input_token_cost_above_200k_tokens_priority": {
"type": "number"
},
"cache_read_input_token_cost_above_272k_tokens_priority": {
"type": "number"
},
"input_cost_per_token_flex": {"type": "number"},
"input_cost_per_token_priority": {"type": "number"},
"input_cost_per_token_above_200k_tokens_priority": {"type": "number"},
@ -743,6 +765,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
"supports_multimodal": {"type": "boolean"},
"uses_embed_content": {"type": "boolean"},
"supports_reasoning": {"type": "boolean"},
"supports_minimal_reasoning_effort": {"type": "boolean"},
"supports_none_reasoning_effort": {"type": "boolean"},
"supports_xhigh_reasoning_effort": {"type": "boolean"},
"supports_service_tier": {"type": "boolean"},
@ -839,7 +862,9 @@ def test_aaamodel_prices_and_context_window_json_is_valid():
},
}
prod_json = os.path.join(os.path.dirname(__file__), "..", "..", "model_prices_and_context_window.json")
prod_json = os.path.join(
os.path.dirname(__file__), "..", "..", "model_prices_and_context_window.json"
)
with open(prod_json, "r") as model_prices_file:
actual_json = json.load(model_prices_file)
assert isinstance(actual_json, dict)
@ -880,8 +905,10 @@ def test_max_tokens_consistency():
from pathlib import Path
# Load the model configuration
config_path = Path(__file__).parent.parent.parent / "model_prices_and_context_window.json"
with open(config_path, 'r') as f:
config_path = (
Path(__file__).parent.parent.parent / "model_prices_and_context_window.json"
)
with open(config_path, "r") as f:
models = json.load(f)
inconsistencies = []
@ -893,17 +920,19 @@ def test_max_tokens_consistency():
# Check if both max_tokens and max_output_tokens exist
if isinstance(config, dict):
max_tokens = config.get('max_tokens')
max_output_tokens = config.get('max_output_tokens')
max_tokens = config.get("max_tokens")
max_output_tokens = config.get("max_output_tokens")
# Only validate if both exist
if max_tokens is not None and max_output_tokens is not None:
if max_tokens != max_output_tokens:
inconsistencies.append({
'model': model_name,
'max_tokens': max_tokens,
'max_output_tokens': max_output_tokens
})
inconsistencies.append(
{
"model": model_name,
"max_tokens": max_tokens,
"max_output_tokens": max_output_tokens,
}
)
if inconsistencies:
error_msg = f"\n\n❌ Found {len(inconsistencies)} models with max_tokens != max_output_tokens:\n\n"
@ -2381,13 +2410,14 @@ def test_register_model_with_scientific_notation():
# Use a truly unique model name with uuid to avoid conflicts when tests run in parallel
test_model_name = f"test-scientific-notation-model-{uuid.uuid4().hex[:12]}"
# Clear LRU caches that might have stale data
from litellm.utils import (
_invalidate_model_cost_lowercase_map,
)
_invalidate_model_cost_lowercase_map()
model_cost_dict = {
test_model_name: {
"max_tokens": 8192,
@ -2406,7 +2436,7 @@ def test_register_model_with_scientific_notation():
assert registered_model["output_cost_per_token"] == 6e-07
assert registered_model["litellm_provider"] == "openai"
assert registered_model["mode"] == "chat"
# Clean up after test
if test_model_name in litellm.model_cost:
del litellm.model_cost[test_model_name]
@ -2734,7 +2764,9 @@ def test_model_info_for_openrouter_kimi_k2_5():
model_cost = json.load(f)
model_info = model_cost.get("openrouter/moonshotai/kimi-k2.5")
assert model_info is not None, "Model not found in model_prices_and_context_window.json"
assert (
model_info is not None
), "Model not found in model_prices_and_context_window.json"
assert model_info["litellm_provider"] == "openrouter"
assert model_info["mode"] == "chat"
@ -2778,7 +2810,9 @@ def test_model_info_for_fireworks_short_form_models():
"fireworks_ai/accounts/fireworks/models/glm-4p7",
]:
info = model_cost.get(key)
assert info is not None, f"{key} not found in model_prices_and_context_window.json"
assert (
info is not None
), f"{key} not found in model_prices_and_context_window.json"
assert info["litellm_provider"] == "fireworks_ai"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == 6e-07
@ -2792,7 +2826,9 @@ def test_model_info_for_fireworks_short_form_models():
"fireworks_ai/accounts/fireworks/models/minimax-m2p1",
]:
info = model_cost.get(key)
assert info is not None, f"{key} not found in model_prices_and_context_window.json"
assert (
info is not None
), f"{key} not found in model_prices_and_context_window.json"
assert info["litellm_provider"] == "fireworks_ai"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == 3e-07
@ -2801,7 +2837,9 @@ def test_model_info_for_fireworks_short_form_models():
# kimi-k2p5: short-form only (long-form already existed)
info = model_cost.get("fireworks_ai/kimi-k2p5")
assert info is not None, "fireworks_ai/kimi-k2p5 not found in model_prices_and_context_window.json"
assert (
info is not None
), "fireworks_ai/kimi-k2p5 not found in model_prices_and_context_window.json"
assert info["litellm_provider"] == "fireworks_ai"
assert info["mode"] == "chat"
assert info["input_cost_per_token"] == 6e-07
@ -3047,7 +3085,9 @@ class TestProxyLoggingBudgetAlerts:
user_info = MagicMock()
# Should not raise an error
await proxy_logging.budget_alerts(type="organization_budget", user_info=user_info)
await proxy_logging.budget_alerts(
type="organization_budget", user_info=user_info
)
async def test_budget_alerts_with_both_slack_and_email(self):
"""Test that budget_alerts calls both slack and email instances when both are in alerting."""
@ -3103,11 +3143,13 @@ class TestProxyLoggingBudgetAlerts:
type=alert_type, user_info=user_info
)
async def test_budget_alerts_soft_budget_with_alert_emails_bypasses_alerting_none(self):
async def test_budget_alerts_soft_budget_with_alert_emails_bypasses_alerting_none(
self,
):
"""
Test that soft_budget alerts with alert_emails bypass the alerting=None check
and send emails even when alerting is None.
This tests the new logic that allows team-specific soft budget email alerts
via metadata.soft_budget_alerting_emails to work even when global alerting is disabled.
"""
@ -3143,7 +3185,9 @@ class TestProxyLoggingBudgetAlerts:
type="soft_budget", user_info=user_info
)
async def test_budget_alerts_soft_budget_without_alert_emails_respects_alerting_none(self):
async def test_budget_alerts_soft_budget_without_alert_emails_respects_alerting_none(
self,
):
"""
Test that soft_budget alerts WITHOUT alert_emails still respect alerting=None
and do not send emails when alerting is None.
@ -3176,7 +3220,9 @@ class TestProxyLoggingBudgetAlerts:
proxy_logging.slack_alerting_instance.budget_alerts.assert_not_called()
proxy_logging.email_logging_instance.budget_alerts.assert_not_called()
async def test_budget_alerts_soft_budget_with_empty_alert_emails_respects_alerting_none(self):
async def test_budget_alerts_soft_budget_with_empty_alert_emails_respects_alerting_none(
self,
):
"""
Test that soft_budget alerts with empty alert_emails list still respect alerting=None.
"""
@ -3317,7 +3363,10 @@ def test_last_assistant_with_tool_calls_has_no_thinking_blocks_issue_18926():
{"type": "thinking", "thinking": "Let me analyze the requirements..."}
],
"tool_calls": [
{"id": "toolu_1", "function": {"name": "file_editor", "arguments": "{}"}}
{
"id": "toolu_1",
"function": {"name": "file_editor", "arguments": "{}"},
}
],
},
{
@ -3330,7 +3379,10 @@ def test_last_assistant_with_tool_calls_has_no_thinking_blocks_issue_18926():
# NO thinking_blocks - Claude sometimes doesn't include them
"content": [{"type": "text", "text": "Let me explore more..."}],
"tool_calls": [
{"id": "toolu_2", "function": {"name": "file_editor", "arguments": "{}"}}
{
"id": "toolu_2",
"function": {"name": "file_editor", "arguments": "{}"},
}
],
},
]
@ -3343,10 +3395,9 @@ def test_last_assistant_with_tool_calls_has_no_thinking_blocks_issue_18926():
# So we should NOT drop thinking - the combination tells us thinking is in use
# The fix uses both checks: only drop if last has none AND no message has any
should_drop_thinking = (
last_assistant_with_tool_calls_has_no_thinking_blocks(messages)
and not any_assistant_message_has_thinking_blocks(messages)
)
should_drop_thinking = last_assistant_with_tool_calls_has_no_thinking_blocks(
messages
) and not any_assistant_message_has_thinking_blocks(messages)
assert should_drop_thinking is False
@ -3558,34 +3609,67 @@ class TestGetOptionalParamsDeepSeek:
class TestIsStreamingRequest:
def test_stream_true_in_kwargs(self):
assert _is_streaming_request(kwargs={"stream": True}, call_type="acompletion") is True
assert (
_is_streaming_request(kwargs={"stream": True}, call_type="acompletion")
is True
)
def test_stream_false_in_kwargs(self):
assert _is_streaming_request(kwargs={"stream": False}, call_type="acompletion") is False
assert (
_is_streaming_request(kwargs={"stream": False}, call_type="acompletion")
is False
)
def test_no_stream_in_kwargs(self):
assert _is_streaming_request(kwargs={}, call_type="acompletion") is False
def test_generate_content_stream_string(self):
assert _is_streaming_request(kwargs={}, call_type=CallTypes.generate_content_stream.value) is True
assert (
_is_streaming_request(
kwargs={}, call_type=CallTypes.generate_content_stream.value
)
is True
)
def test_agenerate_content_stream_string(self):
assert _is_streaming_request(kwargs={}, call_type=CallTypes.agenerate_content_stream.value) is True
assert (
_is_streaming_request(
kwargs={}, call_type=CallTypes.agenerate_content_stream.value
)
is True
)
def test_generate_content_stream_enum(self):
assert _is_streaming_request(kwargs={}, call_type=CallTypes.generate_content_stream) is True
assert (
_is_streaming_request(
kwargs={}, call_type=CallTypes.generate_content_stream
)
is True
)
def test_agenerate_content_stream_enum(self):
assert _is_streaming_request(kwargs={}, call_type=CallTypes.agenerate_content_stream) is True
assert (
_is_streaming_request(
kwargs={}, call_type=CallTypes.agenerate_content_stream
)
is True
)
def test_non_streaming_call_type_string(self):
assert _is_streaming_request(kwargs={}, call_type="acompletion") is False
def test_non_streaming_call_type_enum(self):
assert _is_streaming_request(kwargs={}, call_type=CallTypes.acompletion) is False
assert (
_is_streaming_request(kwargs={}, call_type=CallTypes.acompletion) is False
)
def test_stream_true_overrides_non_streaming_call_type(self):
assert _is_streaming_request(kwargs={"stream": True}, call_type=CallTypes.acompletion) is True
assert (
_is_streaming_request(
kwargs={"stream": True}, call_type=CallTypes.acompletion
)
is True
)
class TestCallbackAsyncSyncSeparation:
@ -3679,37 +3763,27 @@ class TestMetadataNoneHandling:
def test_metadata_none_get_previous_models(self):
"""kwargs.get("metadata") or {} should return {} when metadata is None."""
kwargs = {"metadata": None}
previous_models = (kwargs.get("metadata") or {}).get(
"previous_models", None
)
previous_models = (kwargs.get("metadata") or {}).get("previous_models", None)
assert previous_models is None
def test_metadata_none_model_group_check(self):
"""'model_group' in (kwargs.get("metadata") or {}) should not raise TypeError."""
kwargs = {"metadata": None}
_is_litellm_router_call = "model_group" in (
kwargs.get("metadata") or {}
)
_is_litellm_router_call = "model_group" in (kwargs.get("metadata") or {})
assert _is_litellm_router_call is False
def test_metadata_missing_key(self):
"""Should work when metadata key is completely absent."""
kwargs = {}
previous_models = (kwargs.get("metadata") or {}).get(
"previous_models", None
)
previous_models = (kwargs.get("metadata") or {}).get("previous_models", None)
assert previous_models is None
def test_metadata_present_with_values(self):
"""Should work when metadata has actual values."""
kwargs = {"metadata": {"previous_models": ["model1"], "model_group": "test"}}
previous_models = (kwargs.get("metadata") or {}).get(
"previous_models", None
)
previous_models = (kwargs.get("metadata") or {}).get("previous_models", None)
assert previous_models == ["model1"]
_is_litellm_router_call = "model_group" in (
kwargs.get("metadata") or {}
)
_is_litellm_router_call = "model_group" in (kwargs.get("metadata") or {})
assert _is_litellm_router_call is True
def test_metadata_none_causes_error_with_old_pattern(self):