diff --git a/docs/my-website/blog/gpt_5_4_mini_nano/index.md b/docs/my-website/blog/gpt_5_4_mini_nano/index.md new file mode 100644 index 00000000000..6d7c2b33f72 --- /dev/null +++ b/docs/my-website/blog/gpt_5_4_mini_nano/index.md @@ -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 + + + + +**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?"}] + }' +``` + + + + +```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) +``` + + + + +## 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. diff --git a/docs/my-website/docs/prompt_management.md b/docs/my-website/docs/prompt_management.md new file mode 100644 index 00000000000..c4e606674b1 --- /dev/null +++ b/docs/my-website/docs/prompt_management.md @@ -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. diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index b1d52b506e7..91445863f76 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -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) | diff --git a/docs/my-website/docs/proxy/load_balancing.md b/docs/my-website/docs/proxy/load_balancing.md index 5bf39d179f6..74b3e8a5117 100644 --- a/docs/my-website/docs/proxy/load_balancing.md +++ b/docs/my-website/docs/proxy/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: diff --git a/docs/my-website/docs/proxy/prompt_management.md b/docs/my-website/docs/proxy/prompt_management.md index 08307ba99ec..5a3e411e984 100644 --- a/docs/my-website/docs/proxy/prompt_management.md +++ b/docs/my-website/docs/proxy/prompt_management.md @@ -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: diff --git a/docs/my-website/docs/response_api.md b/docs/my-website/docs/response_api.md index fb55ae9f9d0..0c428000c72 100644 --- a/docs/my-website/docs/response_api.md +++ b/docs/my-website/docs/response_api.md @@ -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. + + + + +```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"], + }, +) +``` + + + + +```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 +``` + + + + +**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. diff --git a/docs/my-website/docs/tutorials/file_search_responses_api.md b/docs/my-website/docs/tutorials/file_search_responses_api.md new file mode 100644 index 00000000000..d74bdf9cb83 --- /dev/null +++ b/docs/my-website/docs/tutorials/file_search_responses_api.md @@ -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 + + + + +### 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) +``` + + + + +### 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) +``` + + + + +### 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 diff --git a/docs/my-website/docs/tutorials/vertex_ai_pay_go.md b/docs/my-website/docs/tutorials/vertex_ai_pay_go.md new file mode 100644 index 00000000000..87197e5bad5 --- /dev/null +++ b/docs/my-website/docs/tutorials/vertex_ai_pay_go.md @@ -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` + + + + +```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) +``` + + + + +```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"}]}' +``` + + + + +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!"}]}]}' +``` + + + + +### How cost tracking works + +![Vertex AI Priority PayGo Cost Tracking Flow](/img/vertex_cost_tracking_flow.svg) + +**`trafficType` → `service_tier` mapping** + +| `usageMetadata.trafficType` | `service_tier` | Pricing keys used | +|---|---|---| +| `ON_DEMAND` | `None` | `input_cost_per_token` | +| `ON_DEMAND_PRIORITY` | `"priority"` | `input_cost_per_token_priority` | +| `FLEX` / `BATCH` | `"flex"` | `input_cost_per_token_flex` | + +If a tier-specific key is missing, LiteLLM falls back to standard pricing keys. + +--- + +## Standard PayGo vs Provisioned Throughput + +This is a different header from priority routing: + +| Header value | Behavior | +|---|---| +| `X-Vertex-AI-LLM-Request-Type: shared` | Force standard PayGo (bypass PT) | +| `X-Vertex-AI-LLM-Request-Type: dedicated` | Force Provisioned Throughput only (`429` if exhausted) | + +### Native route example + +```python +import litellm + +response = litellm.completion( + model="vertex_ai/gemini-2.0-flash", + messages=[{"role": "user", "content": "Hello!"}], + vertex_project="YOUR_PROJECT_ID", + vertex_location="us-central1", + extra_headers={"X-Vertex-AI-LLM-Request-Type": "shared"}, +) +``` + +### Pass-through example + +```bash +MODEL_ID="gemini-2.0-flash-001" +PROJECT_ID="YOUR_PROJECT_ID" + +curl -X POST \ + "${LITELLM_PROXY_BASE_URL}/vertex_ai/v1/projects/${PROJECT_ID}/locations/global/publishers/google/models/${MODEL_ID}:generateContent" \ + -H "Authorization: Bearer sk-your-litellm-key" \ + -H "Content-Type: application/json" \ + -H "x-pass-X-Vertex-AI-LLM-Request-Type: shared" \ + -d '{ + "contents": [{"role": "user", "parts": [{"text": "Hello!"}]}] + }' +``` + +--- + +## Troubleshooting + +**Q: What does `403 Permission denied` or `IAM_PERMISSION_DENIED` mean?** +A: The service account or Application Default Credentials (ADC) user does not have the `roles/aiplatform.user` role. To resolve this, re-run the `gcloud projects add-iam-policy-binding`. + +**Q: What should I do if I get a `429 Quota exceeded` error?** +A: This means you've hit the per-region QPM (queries per minute) or TPM (tokens per minute) quota. You can: +- Request a quota increase from the [GCP Quotas console](https://console.cloud.google.com/iam-admin/quotas) +- Add more regions to your LiteLLM configuration for load balancing +- Upgrade to [Provisioned Throughput](https://cloud.google.com/vertex-ai/generative-ai/docs/provisioned-throughput) for guaranteed capacity + +**Q: How do I fix the `VERTEXAI_PROJECT not set` error?** +A: Either pass the `vertex_project` parameter explicitly in your LiteLLM call, or set the `VERTEXAI_PROJECT` environment variable before running your code. + diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index 4c0471fb8f4..4fc41190692 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -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", diff --git a/docs/my-website/static/img/vertex_cost_tracking_flow.svg b/docs/my-website/static/img/vertex_cost_tracking_flow.svg new file mode 100644 index 00000000000..c3b2e33a073 --- /dev/null +++ b/docs/my-website/static/img/vertex_cost_tracking_flow.svg @@ -0,0 +1,63 @@ + + + + + + + + + + + HTTP request + X-Vertex-AI-LLM-Shared-Request-Type: priority + + + + + Vertex AI + + + + + Vertex response + usageMetadata.trafficType = ON_DEMAND_PRIORITY + + + + + + + + + LiteLLM stores it + _hidden_params.provider_specific_fields.traffic_type + + + + + + + + + completion_cost() + Maps traffic_type → service_tier = "priority" + + + + + + + + + Pricing lookup + input/output_cost_per_token_priority + + + + ① + ② + ③ + ④ + ⑤ + + \ No newline at end of file diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 5530054170c..dc14937d46b 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -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( diff --git a/litellm/batches/main.py b/litellm/batches/main.py index 17d73aae6ad..23f444b1cea 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -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", diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index a5d6bc936bb..cec61405ebb 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -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 diff --git a/litellm/llms/azure_ai/agents/handler.py b/litellm/llms/azure_ai/agents/handler.py index 9eeec7f4e36..c3cd06ab4de 100644 --- a/litellm/llms/azure_ai/agents/handler.py +++ b/litellm/llms/azure_ai/agents/handler.py @@ -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", []) diff --git a/litellm/llms/base_llm/responses/transformation.py b/litellm/llms/base_llm/responses/transformation.py index f429930e002..eea53fe06ec 100644 --- a/litellm/llms/base_llm/responses/transformation.py +++ b/litellm/llms/base_llm/responses/transformation.py @@ -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 diff --git a/litellm/llms/openai/chat/gpt_5_transformation.py b/litellm/llms/openai/chat/gpt_5_transformation.py index bb5783011a3..fc48704cd10 100644 --- a/litellm/llms/openai/chat/gpt_5_transformation.py +++ b/litellm/llms/openai/chat/gpt_5_transformation.py @@ -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, ) diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 9d909fd4017..cafb745862d 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -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 diff --git a/litellm/llms/vertex_ai/batches/handler.py b/litellm/llms/vertex_ai/batches/handler.py index f0b181c9a61..2cb02942061 100644 --- a/litellm/llms/vertex_ai/batches/handler.py +++ b/litellm/llms/vertex_ai/batches/handler.py @@ -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 diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b2fabb4936f..c53ee943c58 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -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, diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 1aa14fff574..815393467de 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -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, diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index 38e5229eee1..160e9c23f01 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -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 diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 8147d17c10e..d1aebe4dceb 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -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 diff --git a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py index 0aa99685209..e2f7646c0aa 100644 --- a/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/pass_through_endpoints.py @@ -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 diff --git a/litellm/proxy/response_api_endpoints/endpoints.py b/litellm/proxy/response_api_endpoints/endpoints.py index e9c7cce0d73..8023853e263 100644 --- a/litellm/proxy/response_api_endpoints/endpoints.py +++ b/litellm/proxy/response_api_endpoints/endpoints.py @@ -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, diff --git a/litellm/proxy/response_polling/background_streaming.py b/litellm/proxy/response_polling/background_streaming.py index 7583f30eb2d..bcc98175773 100644 --- a/litellm/proxy/response_polling/background_streaming.py +++ b/litellm/proxy/response_polling/background_streaming.py @@ -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 diff --git a/litellm/responses/file_search/__init__.py b/litellm/responses/file_search/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/responses/file_search/emulated_handler.py b/litellm/responses/file_search/emulated_handler.py new file mode 100644 index 00000000000..74a7d443c6c --- /dev/null +++ b/litellm/responses/file_search/emulated_handler.py @@ -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, + ) diff --git a/litellm/responses/main.py b/litellm/responses/main.py index 2a320517a4c..c82574278ba 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -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, diff --git a/litellm/router.py b/litellm/router.py index 36046ebf302..25e5c9cb5d9 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -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) diff --git a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py index 8044f71d904..148b7fce0ee 100644 --- a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py @@ -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( diff --git a/litellm/types/router.py b/litellm/types/router.py index 5d28349b5e4..4257628e7cb 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -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 diff --git a/litellm/utils.py b/litellm/utils.py index c674190ba8c..088ee07d630 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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( diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b2fabb4936f..c53ee943c58 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -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, diff --git a/tests/llm_translation/test_azure_agents.py b/tests/llm_translation/test_azure_agents.py index 66a46d53383..3e6b1e00a79 100644 --- a/tests/llm_translation/test_azure_agents.py +++ b/tests/llm_translation/test_azure_agents.py @@ -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 diff --git a/tests/llm_translation/test_prompt_factory.py b/tests/llm_translation/test_prompt_factory.py index a6dcabe25ef..64556c3f26d 100644 --- a/tests/llm_translation/test_prompt_factory.py +++ b/tests/llm_translation/test_prompt_factory.py @@ -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 diff --git a/tests/proxy_unit_tests/test_response_polling_pre_call_checks.py b/tests/proxy_unit_tests/test_response_polling_pre_call_checks.py new file mode 100644 index 00000000000..45e4e9e4d3e --- /dev/null +++ b/tests/proxy_unit_tests/test_response_polling_pre_call_checks.py @@ -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() + diff --git a/tests/test_litellm/llms/openai/test_gpt5_transformation.py b/tests/test_litellm/llms/openai/test_gpt5_transformation.py index 47ae3c44c9e..aebab33e808 100644 --- a/tests/test_litellm/llms/openai/test_gpt5_transformation.py +++ b/tests/test_litellm/llms/openai/test_gpt5_transformation.py @@ -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" \ No newline at end of file + assert params["reasoning_effort"] == "high" diff --git a/tests/test_litellm/llms/test_file_search_responses.py b/tests/test_litellm/llms/test_file_search_responses.py new file mode 100644 index 00000000000..1864a296eb6 --- /dev/null +++ b/tests/test_litellm/llms/test_file_search_responses.py @@ -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" + ) diff --git a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_batch_transformation.py b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_batch_transformation.py index 1aab74ddc26..7310c68b4e0 100644 --- a/tests/test_litellm/llms/vertex_ai/test_vertex_ai_batch_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/test_vertex_ai_batch_transformation.py @@ -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" diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index f20c14aa611..83703cd4edd 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -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] + ) diff --git a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py index c4d2dce5876..ca22c5aab56 100644 --- a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py @@ -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" diff --git a/tests/test_litellm/responses/test_responses_prompt_management.py b/tests/test_litellm/responses/test_responses_prompt_management.py new file mode 100644 index 00000000000..f49679fc400 --- /dev/null +++ b/tests/test_litellm/responses/test_responses_prompt_management.py @@ -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 diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_deployment_affinity_check.py b/tests/test_litellm/router_utils/pre_call_checks/test_deployment_affinity_check.py index e500ad3ca6e..28311a30c0d 100644 --- a/tests/test_litellm/router_utils/pre_call_checks/test_deployment_affinity_check.py +++ b/tests/test_litellm/router_utils/pre_call_checks/test_deployment_affinity_check.py @@ -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 diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 64488e2fb6a..38b7b576d4f 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -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):