mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge pull request #24211 from BerriAI/litellm_dev_sameer_16_march_week
Litellm dev sameer 16 march week
This commit is contained in:
commit
1986f1034e
44 changed files with 5762 additions and 529 deletions
106
docs/my-website/blog/gpt_5_4_mini_nano/index.md
Normal file
106
docs/my-website/blog/gpt_5_4_mini_nano/index.md
Normal 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.
|
||||
48
docs/my-website/docs/prompt_management.md
Normal file
48
docs/my-website/docs/prompt_management.md
Normal 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.
|
||||
|
|
@ -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) |
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
241
docs/my-website/docs/tutorials/file_search_responses_api.md
Normal file
241
docs/my-website/docs/tutorials/file_search_responses_api.md
Normal 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
|
||||
151
docs/my-website/docs/tutorials/vertex_ai_pay_go.md
Normal file
151
docs/my-website/docs/tutorials/vertex_ai_pay_go.md
Normal 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
|
||||
|
||||

|
||||
|
||||
**`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.
|
||||
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
63
docs/my-website/static/img/vertex_cost_tracking_flow.svg
Normal file
63
docs/my-website/static/img/vertex_cost_tracking_flow.svg
Normal 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", 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:"Anthropic Sans", -apple-system, "system-ui", "Segoe UI", sans-serif;font-size:12px;font-weight:400;text-anchor:end;dominant-baseline:central">⑤</text>
|
||||
|
||||
</svg>
|
||||
|
After Width: | Height: | Size: 12 KiB |
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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", [])
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
0
litellm/responses/file_search/__init__.py
Normal file
0
litellm/responses/file_search/__init__.py
Normal file
592
litellm/responses/file_search/emulated_handler.py
Normal file
592
litellm/responses/file_search/emulated_handler.py
Normal 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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
182
tests/proxy_unit_tests/test_response_polling_pre_call_checks.py
Normal file
182
tests/proxy_unit_tests/test_response_polling_pre_call_checks.py
Normal 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()
|
||||
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
892
tests/test_litellm/llms/test_file_search_responses.py
Normal file
892
tests/test_litellm/llms/test_file_search_responses.py
Normal 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"
|
||||
)
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
383
tests/test_litellm/responses/test_responses_prompt_management.py
Normal file
383
tests/test_litellm/responses/test_responses_prompt_management.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue