diff --git a/.circleci/config.yml b/.circleci/config.yml index 3f61ed5fa91..09db37fb8c9 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -22,7 +22,7 @@ commands: name: "Install local version of litellm-enterprise" command: | cd enterprise - python -m pip install -e . + python -m pip install --force-reinstall --no-deps -e . cd .. setup_litellm_test_deps: steps: diff --git a/.github/workflows/test-litellm-matrix.yml b/.github/workflows/test-litellm-matrix.yml index 1672a193161..21e0f9d29f5 100644 --- a/.github/workflows/test-litellm-matrix.yml +++ b/.github/workflows/test-litellm-matrix.yml @@ -100,7 +100,11 @@ jobs: - name: Setup litellm-enterprise run: | - cd enterprise && poetry run pip install -e . && cd .. + cd enterprise && poetry run pip install --force-reinstall --no-deps -e . && cd .. + + - name: Generate Prisma client + run: | + poetry run prisma generate --schema litellm/proxy/schema.prisma - name: Run tests - ${{ matrix.test-group.name }} run: | diff --git a/.github/workflows/test-litellm.yml b/.github/workflows/test-litellm.yml index dc9b48c28f6..b3db62f0a9f 100644 --- a/.github/workflows/test-litellm.yml +++ b/.github/workflows/test-litellm.yml @@ -43,7 +43,7 @@ jobs: - name: Setup litellm-enterprise as local package run: | cd enterprise - poetry run pip install -e . + poetry run pip install --force-reinstall --no-deps -e . cd .. - name: Run tests run: | diff --git a/.github/workflows/test-mcp.yml b/.github/workflows/test-mcp.yml index e19e67c9c4f..1c1cc82cde6 100644 --- a/.github/workflows/test-mcp.yml +++ b/.github/workflows/test-mcp.yml @@ -41,7 +41,7 @@ jobs: - name: Setup litellm-enterprise as local package run: | cd enterprise - python -m pip install -e . + python -m pip install --force-reinstall --no-deps -e . cd .. - name: Run MCP tests diff --git a/.github/workflows/test_server_root_path.yml b/.github/workflows/test_server_root_path.yml index bc559817503..c359e38bff9 100644 --- a/.github/workflows/test_server_root_path.yml +++ b/.github/workflows/test_server_root_path.yml @@ -26,7 +26,7 @@ jobs: uses: docker/build-push-action@v5 with: context: . - file: ./docker/Dockerfile.database + file: ./docker/Dockerfile.non_root tags: litellm-test:${{ github.sha }} load: true cache-from: type=gha diff --git a/docs/my-website/blog/claude_code_beta_headers/index.md b/docs/my-website/blog/claude_code_beta_headers/index.md index b5ec14e209a..44567f616aa 100644 --- a/docs/my-website/blog/claude_code_beta_headers/index.md +++ b/docs/my-website/blog/claude_code_beta_headers/index.md @@ -24,6 +24,8 @@ hide_table_of_contents: false **Severity:** High **Status:** Resolved +> **Note:** This fix will be available starting from `v1.81.13-nightly` or higher of LiteLLM. + ## Summary Claude Code began sending unsupported Anthropic beta headers to non-Anthropic providers (Bedrock, Azure AI, Vertex AI), causing `invalid beta flag` errors. LiteLLM was forwarding all beta headers without provider-specific validation. Users experienced request failures when routing Claude Code requests through LiteLLM to these providers. diff --git a/docs/my-website/blog/vllm_embeddings_incident/index.md b/docs/my-website/blog/vllm_embeddings_incident/index.md new file mode 100644 index 00000000000..a1ce8152857 --- /dev/null +++ b/docs/my-website/blog/vllm_embeddings_incident/index.md @@ -0,0 +1,117 @@ +--- +slug: vllm-embeddings-incident +title: "Incident Report: vLLM Embeddings Broken by encoding_format Parameter" +date: 2026-02-18T10: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 +tags: [incident-report, embeddings, vllm] +hide_table_of_contents: false +--- + +**Date:** Feb 16, 2026 +**Duration:** ~3 hours +**Severity:** High (for vLLM embedding users) +**Status:** Resolved + +## Summary + +A commit ([`dbcae4a`](https://github.com/BerriAI/litellm/commit/dbcae4aca5836770d0e9cd43abab0333c3d61ab2)) intended to fix OpenAI SDK behavior broke vLLM embeddings by explicitly passing `encoding_format=None` in API requests. vLLM rejects this with error: `"unknown variant \`\`, expected float or base64"`. + +- **vLLM embedding calls:** Complete failure - all requests rejected +- **Other providers:** No impact - OpenAI and other providers functioned normally +- **Other vLLM functionality:** No impact - only embeddings were affected + +{/* truncate */} + +--- + +## Background + +The `encoding_format` parameter for embeddings specifies whether vectors should be returned as `float` arrays or `base64` encoded strings. Different providers have different expectations: + +- **OpenAI SDK:** If `encoding_format` is omitted, the SDK adds a default value of `"float"` +- **vLLM:** Strictly validates `encoding_format` - only accepts `"float"`, `"base64"`, or complete omission. Rejects `None` or empty string values. + +```mermaid +flowchart TD + A["1. User calls litellm.embedding() + litellm/main.py"] --> B["2. Transform request for provider + litellm/llms/openai_like/embedding/handler.py"] + B --> C["3. Send request to vLLM endpoint"] + C -->|"encoding_format omitted"| D["4a. ✅ vLLM processes request"] + C -->|"encoding_format='float' or 'base64'"| D + C -->|"encoding_format=None or ''"| E["4b. ❌ vLLM rejects with error: + 'unknown variant, expected float or base64'"] + + style D fill:#d4edda,stroke:#28a745 + style E fill:#f8d7da,stroke:#dc3545 + style B fill:#fff3cd,stroke:#ffc107 +``` + +--- + +## Root cause + +A well-intentioned fix for OpenAI SDK behavior inadvertently broke vLLM embeddings: + +**The Breaking Change ([`dbcae4a`](https://github.com/BerriAI/litellm/commit/dbcae4aca5836770d0e9cd43abab0333c3d61ab2)):** + +In `litellm/main.py`, the code was changed to explicitly set `encoding_format=None` instead of omitting it: + +```python +# Added in dbcae4a +if encoding_format is not None: + optional_params["encoding_format"] = encoding_format +else: + # Omitting causes openai sdk to add default value of "float" + optional_params["encoding_format"] = None +``` + +This fix worked correctly for OpenAI - explicitly passing `None` prevented the SDK from adding its default value. However, vLLM's strict parameter validation rejected `None` values, causing all embedding requests to fail. + +--- + +## The Fix + +Fix deployed ([`55348dd`](https://github.com/BerriAI/litellm/commit/55348dd9c51b5b028f676d25ad023b8f052fc071)). The solution filters out `None` and empty string values from `optional_params` before sending requests to OpenAI-like providers (including vLLM). + +**In `litellm/llms/openai_like/embedding/handler.py`:** + +```python +# Before (broken) +data = {"model": model, "input": input, **optional_params} + +# After (fixed) +filtered_optional_params = {k: v for k, v in optional_params.items() if v not in (None, '')} +data = {"model": model, "input": input, **filtered_optional_params} +``` + +This ensures: +- Valid values (`"float"`, `"base64"`) are preserved and sent +- `None` and empty string values are filtered out (parameter omitted entirely) +- OpenAI SDK no longer adds defaults because liteLLM handles the parameter upstream + +--- + +## Remediation + +| # | Action | Status | Code | +|---|---|---|---| +| 1 | Filter `None` and empty string values in OpenAI-like embedding handler | ✅ Done | [`handler.py#L108`](https://github.com/BerriAI/litellm/blob/main/litellm/llms/openai_like/embedding/handler.py#L108) | +| 2 | Unit tests for parameter filtering (None, empty string, valid values) | ✅ Done | [`test_openai_like_embedding.py`](https://github.com/BerriAI/litellm/blob/main/tests/test_litellm/llms/openai_like/embedding/test_openai_like_embedding.py) | +| 3 | Transformation tests for hosted_vllm embedding config | ✅ Done | [`test_hosted_vllm_embedding_transformation.py`](https://github.com/BerriAI/litellm/blob/main/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py) | +| 4 | E2E tests with actual vLLM endpoint | ✅ Done | [`test_hosted_vllm_embedding_e2e.py`](https://github.com/BerriAI/litellm/blob/main/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_e2e.py) | +| 5 | Validate JSON payload structure matches vLLM expectations | ✅ Done | Tests verify exact JSON sent to endpoint | + +--- diff --git a/docs/my-website/docs/adding_provider/generic_guardrail_api.md b/docs/my-website/docs/adding_provider/generic_guardrail_api.md index 0931c349e48..eb567a69fcb 100644 --- a/docs/my-website/docs/adding_provider/generic_guardrail_api.md +++ b/docs/my-website/docs/adding_provider/generic_guardrail_api.md @@ -237,6 +237,7 @@ litellm_settings: mode: pre_call # or post_call, during_call api_base: https://your-guardrail-api.com api_key: os.environ/YOUR_GUARDRAIL_API_KEY # optional + unreachable_fallback: fail_closed # default: fail_closed. Set to fail_open to proceed if the guardrail endpoint is unreachable (network errors, or HTTP 502/503/504 from an upstream proxy/LB). additional_provider_specific_params: # your custom parameters threshold: 0.8 diff --git a/docs/my-website/docs/completion/message_sanitization.md b/docs/my-website/docs/completion/message_sanitization.md new file mode 100644 index 00000000000..17482c59339 --- /dev/null +++ b/docs/my-website/docs/completion/message_sanitization.md @@ -0,0 +1,465 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Message Sanitization for Tool Calling for anthropic models + +**Automatically fix common message formatting issues when using tool calling with `modify_params=True`** + +LiteLLM can automatically sanitize messages to handle common issues that occur during tool calling workflows, especially when using OpenAI-compatible clients with providers that have strict message format requirements (like Anthropic Claude). + +## Overview + +When `litellm.modify_params = True` is enabled, LiteLLM automatically sanitizes messages to fix three common issues: + +1. **Orphaned Tool Calls** - Assistant messages with tool_calls but missing tool results +2. **Orphaned Tool Results** - Tool messages that reference non-existent tool_call_ids +3. **Empty Message Content** - Messages with empty or whitespace-only text content + +This ensures your tool calling workflows work seamlessly across different LLM providers without manual message validation. + +## Why Message Sanitization? + +Different LLM providers have varying requirements for message formats, especially during tool calling: + +- **Anthropic Claude** requires every tool_call to have a corresponding tool result +- Some providers reject messages with empty content +- OpenAI-compatible clients may not always maintain perfect message consistency + +Without sanitization, these issues cause API errors that interrupt your workflows. With `modify_params=True`, LiteLLM handles these edge cases automatically. + +## Quick Start + + + + +```python +import litellm + +# Enable automatic message sanitization +litellm.modify_params = True + +# This will work even if messages have formatting issues +response = litellm.completion( + model="anthropic/claude-3-5-sonnet-20241022", + messages=[ + {"role": "user", "content": "What's the weather in Boston?"}, + { + "role": "assistant", + "tool_calls": [ + { + "id": "call_123", + "type": "function", + "function": {"name": "get_weather", "arguments": '{"city": "Boston"}'} + } + ] + # Missing tool result - LiteLLM will add a dummy result automatically + }, + {"role": "user", "content": "Thanks!"} + ], + tools=[{ + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather for a city", + "parameters": { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"] + } + } + }] +) +``` + + + + +```yaml +litellm_settings: + modify_params: true # Enable automatic message sanitization + +model_list: + - model_name: claude-3-5-sonnet + litellm_params: + model: anthropic/claude-3-5-sonnet-20241022 +``` + + + + +## Sanitization Cases + +### Case A: Orphaned Tool Calls (Missing Tool Results) + +**Problem:** An assistant message contains `tool_calls`, but no corresponding tool result messages follow. + +**Solution:** LiteLLM automatically adds dummy tool result messages for any missing tool results. + +**Example:** + +```python +import litellm +litellm.modify_params = True + +# Messages with orphaned tool calls +messages = [ + {"role": "user", "content": "Search for Python tutorials"}, + { + "role": "assistant", + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": {"name": "web_search", "arguments": '{"query": "Python tutorials"}'} + } + ] + }, + # Missing tool result here! + {"role": "user", "content": "What about JavaScript?"} +] + +# LiteLLM automatically adds: +# { +# "role": "tool", +# "tool_call_id": "call_abc123", +# "content": "[System: Tool execution skipped/interrupted by user. No result provided for tool 'web_search'.]" +# } + +response = litellm.completion( + model="anthropic/claude-3-5-sonnet-20241022", + messages=messages, + tools=[...] +) +``` + +**When this happens:** +- User interrupts tool execution +- Client loses tool results due to network issues +- Conversation flow changes before tool completes +- Multi-turn conversations where tools are optional + +### Case B: Orphaned Tool Results (Invalid tool_call_id) + +**Problem:** A tool message references a `tool_call_id` that doesn't exist in any previous assistant message. + +**Solution:** LiteLLM automatically removes these orphaned tool result messages. + +**Example:** + +```python +import litellm +litellm.modify_params = True + +# Messages with orphaned tool result +messages = [ + {"role": "user", "content": "Hello"}, + {"role": "assistant", "content": "Hi! How can I help?"}, + { + "role": "tool", + "tool_call_id": "call_nonexistent", # This tool_call_id doesn't exist! + "content": "Some result" + } +] + +# LiteLLM automatically removes the orphaned tool message + +response = litellm.completion( + model="anthropic/claude-3-5-sonnet-20241022", + messages=messages +) +``` + +**When this happens:** +- Message history is manually edited +- Tool results are duplicated or mismatched +- Conversation state is restored incorrectly +- Messages are merged from different conversations + +### Case C: Empty Message Content + +**Problem:** User or assistant messages have empty or whitespace-only content. + +**Solution:** LiteLLM replaces empty content with a system placeholder message. + +**Example:** + +```python +import litellm +litellm.modify_params = True + +# Messages with empty content +messages = [ + {"role": "user", "content": ""}, # Empty content + {"role": "assistant", "content": " "}, # Whitespace only +] + +# LiteLLM automatically replaces with: +# {"role": "user", "content": "[System: Empty message content sanitised to satisfy protocol]"} +# {"role": "assistant", "content": "[System: Empty message content sanitised to satisfy protocol]"} + +response = litellm.completion( + model="anthropic/claude-3-5-sonnet-20241022", + messages=messages +) +``` + +**When this happens:** +- UI sends empty messages +- Content is stripped during preprocessing +- Placeholder messages in conversation history +- Edge cases in message construction + +## Configuration + +### Enable Globally + + + + +```python +import litellm + +# Enable for all completion calls +litellm.modify_params = True +``` + + + + +```yaml +litellm_settings: + modify_params: true +``` + + + + +```bash +export LITELLM_MODIFY_PARAMS=True +``` + + + + +### Enable Per-Request + +```python +import litellm + +# Enable only for specific requests +response = litellm.completion( + model="anthropic/claude-3-5-sonnet-20241022", + messages=messages, + modify_params=True # Override global setting +) +``` + +## Supported Providers + +Message sanitization currently works with: + +- ✅ Anthropic (Claude) + +**Note:** While the sanitization logic is provider-agnostic, it is currently only applied in the Anthropic message transformation pipeline. Support for additional providers may be added in future releases. + +## Implementation Details + +### How It Works + +The message sanitization process runs **before** messages are converted to provider-specific formats: + +1. **Input:** OpenAI-format messages with potential issues +2. **Sanitization:** Three helper functions process the messages: + - `_sanitize_empty_text_content()` - Fixes empty content + - `_add_missing_tool_results()` - Adds dummy tool results + - `_is_orphaned_tool_result()` - Identifies orphaned results +3. **Output:** Clean, provider-compatible messages + +### Code Reference + +The sanitization logic is implemented in: +- `litellm/litellm_core_utils/prompt_templates/factory.py` +- Function: `sanitize_messages_for_tool_calling()` + +### Logging + +When sanitization occurs, LiteLLM logs debug messages: + +```python +import litellm +litellm.set_verbose = True # Enable debug logging + +# You'll see logs like: +# "_add_missing_tool_results: Found 1 orphaned tool calls. Adding dummy tool results." +# "_is_orphaned_tool_result: Found orphaned tool result with tool_call_id=call_123" +# "_sanitize_empty_text_content: Replaced empty text content in user message" +``` + +## Best Practices + +### 1. Enable for Production Workflows + +```python +# Recommended for production +litellm.modify_params = True + +# Ensures robust handling of edge cases +response = litellm.completion( + model="anthropic/claude-3-5-sonnet-20241022", + messages=messages, + tools=tools +) +``` + +### 2. Preserve Tool Results When Possible + +While sanitization handles missing tool results, it's better to provide actual results: + +```python +# Good: Provide actual tool results +messages = [ + {"role": "user", "content": "Search for Python"}, + {"role": "assistant", "tool_calls": [...]}, + {"role": "tool", "tool_call_id": "call_123", "content": "Actual search results"} +] + +# Fallback: Sanitization adds dummy result if missing +messages = [ + {"role": "user", "content": "Search for Python"}, + {"role": "assistant", "tool_calls": [...]}, + # Missing tool result - sanitization adds dummy +] +``` + +### 3. Monitor Sanitization Events + +Use logging to track when sanitization occurs: + +```python +import litellm +import logging + +# Enable debug logging +litellm.set_verbose = True +logging.basicConfig(level=logging.DEBUG) + +# Track sanitization events in your application +response = litellm.completion( + model="anthropic/claude-3-5-sonnet-20241022", + messages=messages +) +``` + +### 4. Test Edge Cases + +Ensure your application handles sanitized messages correctly: + +```python +import litellm +litellm.modify_params = True + +# Test orphaned tool calls +test_messages = [ + {"role": "user", "content": "Test"}, + {"role": "assistant", "tool_calls": [{"id": "call_1", "type": "function", "function": {"name": "test", "arguments": "{}"}}]}, + {"role": "user", "content": "Continue"} # No tool result +] + +response = litellm.completion( + model="anthropic/claude-3-5-sonnet-20241022", + messages=test_messages, + tools=[...] +) + +# Verify the response handles the dummy tool result appropriately +``` + +## Related Features + +- **[Drop Params](./drop_params.md)** - Drop unsupported parameters for specific providers +- **[Message Trimming](./message_trimming.md)** - Trim messages to fit token limits +- **[Function Calling](./function_call.md)** - Complete guide to tool/function calling +- **[Reasoning Content](../reasoning_content.md)** - Extended thinking with tool calling + +## Troubleshooting + +### Sanitization Not Working + +**Issue:** Messages still cause errors despite `modify_params=True` + +**Solution:** +1. Verify `modify_params` is enabled: + ```python + import litellm + print(litellm.modify_params) # Should be True + ``` + +2. Check if the issue is provider-specific: + ```python + litellm.set_verbose = True # Enable debug logging + ``` + +3. Ensure you're using a recent version of LiteLLM: + ```bash + pip install --upgrade litellm + ``` + +### Unexpected Dummy Tool Results + +**Issue:** Dummy tool results appear when you expect actual results + +**Cause:** Tool result messages are missing or have incorrect `tool_call_id` + +**Solution:** +1. Verify tool result messages have correct `tool_call_id`: + ```python + # Correct + {"role": "tool", "tool_call_id": "call_123", "content": "result"} + + # Incorrect - will be treated as orphaned + {"role": "tool", "tool_call_id": "wrong_id", "content": "result"} + ``` + +2. Ensure tool results immediately follow assistant messages with tool_calls + +### Performance Impact + +**Issue:** Concerned about performance overhead + +**Details:** Message sanitization has minimal performance impact: +- Runs in O(n) time where n = number of messages +- Only processes messages when `modify_params=True` +- Typically adds < 1ms to request processing time + +## FAQ + +**Q: Does sanitization modify my original messages?** + +A: No, sanitization creates a new list of messages. Your original messages remain unchanged. + +**Q: Can I disable specific sanitization cases?** + +A: Currently, all three cases are handled together when `modify_params=True`. To disable sanitization entirely, set `modify_params=False`. + +**Q: What happens to the dummy tool results?** + +A: Dummy tool results are sent to the LLM provider along with other messages. The model sees them as regular tool results with informative error messages. + +**Q: Does this work with streaming?** + +A: Yes, message sanitization works with both streaming and non-streaming requests. + +**Q: Is this related to `drop_params`?** + +A: No, they're separate features: +- `modify_params` - Modifies/fixes message content and structure +- `drop_params` - Removes unsupported API parameters + +Both can be enabled simultaneously. + +## See Also + +- [Reasoning Content with Tool Calling](../reasoning_content.md) +- [Function Calling Guide](./function_call.md) +- [Bedrock Provider Documentation](../providers/bedrock.md) +- [Anthropic Provider Documentation](../providers/anthropic.md) diff --git a/docs/my-website/docs/observability/datadog.md b/docs/my-website/docs/observability/datadog.md index 6f785be1013..9385b0020cf 100644 --- a/docs/my-website/docs/observability/datadog.md +++ b/docs/my-website/docs/observability/datadog.md @@ -253,3 +253,12 @@ LiteLLM supports customizing the following Datadog environment variables \* **Required when using Direct API** (default): `DD_API_KEY` and `DD_SITE` are required \* **Optional when using DataDog Agent**: Set `LITELLM_DD_AGENT_HOST` to use agent mode; `DD_API_KEY` and `DD_SITE` are not required for **Datadog Logs**. (**Note: `DD_API_KEY` IS REQUIRED for Datadog LLM Observability**) +## Automatic Tags + +LiteLLM automatically adds the following tags to your Datadog logs and metrics if the information is available in the request: + +| Tag | Description | Source | +|-----|-------------|--------| +| `team` | The team alias or ID associated with the API Key | `user_api_key_team_alias`, `team_alias`, `user_api_key_team_id`, or `team_id` in metadata | +| `request_tag` | Custom tags passed in the request | `request_tags` in logging payload | + diff --git a/docs/my-website/docs/providers/watsonx/rerank.md b/docs/my-website/docs/providers/watsonx/rerank.md new file mode 100644 index 00000000000..0900ce96781 --- /dev/null +++ b/docs/my-website/docs/providers/watsonx/rerank.md @@ -0,0 +1,52 @@ +# watsonx.ai Rerank + +## Overview + +| Property | Details | +|----------|--------------------------------------------------------------------------| +| Description | watsonx.ai rerank integration | +| Provider Route on LiteLLM | `watsonx/` | +| Supported Operations | `/ml/v1/text/rerank` | +| Link to Provider Doc | [IBM WatsonX.ai ↗](https://cloud.ibm.com/apidocs/watsonx-ai#text-rerank) | + +## Quick Start + +### **LiteLLM SDK** + +```python +import os +from litellm import rerank + +os.environ["WATSONX_APIKEY"] = "YOUR_WATSONX_APIKEY" +os.environ["WATSONX_API_BASE"] = "YOUR_WATSONX_API_BASE" +os.environ["WATSONX_PROJECT_ID"] = "YOUR_WATSONX_PROJECT_ID" + +query="Best programming language for beginners?" +documents=[ + "Python is great for beginners due to simple syntax.", + "JavaScript runs in browsers and is versatile.", + "Rust has a steep learning curve but is very safe.", +] + +response = rerank( + model="watsonx/cross-encoder/ms-marco-minilm-l-12-v2", + query=query, + documents=documents, + top_n=2, + return_documents=True, +) + +print(response) +``` + +### **LiteLLM Proxy** + +```yaml +model_list: + - model_name: cross-encoder/ms-marco-minilm-l-12-v2 + litellm_params: + model: watsonx/cross-encoder/ms-marco-minilm-l-12-v2 + api_key: os.environ/WATSONX_APIKEY + api_base: os.environ/WATSONX_API_BASE + project_id: os.environ/WATSONX_PROJECT_ID +``` diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index a2371232302..55a99f71128 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -358,7 +358,8 @@ 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. Currently supported: 'router_budget_limiting', 'prompt_caching' | +| 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`, `deployment_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). | | 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.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) | @@ -540,7 +541,7 @@ router_settings: | DEFAULT_IMAGE_WIDTH | Default width for images. Default is 300 | DEFAULT_IN_MEMORY_TTL | Default time-to-live for in-memory cache in seconds. Default is 5 | DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL | Default time-to-live in seconds for management objects (User, Team, Key, Organization) in memory cache. Default is 60 seconds. -| DEFAULT_MAX_LRU_CACHE_SIZE | Default maximum size for LRU cache. Default is 16 +| DEFAULT_MAX_LRU_CACHE_SIZE | Default maximum size for LRU cache. Default is 64 | DEFAULT_MAX_RECURSE_DEPTH | Default maximum recursion depth. Default is 100 | DEFAULT_MAX_RECURSE_DEPTH_SENSITIVE_DATA_MASKER | Default maximum recursion depth for sensitive data masker. Default is 10 | DEFAULT_MAX_RETRIES | Default maximum retry attempts. Default is 2 diff --git a/docs/my-website/docs/rerank.md b/docs/my-website/docs/rerank.md index 90f685d2bbd..9c76883d7fd 100644 --- a/docs/my-website/docs/rerank.md +++ b/docs/my-website/docs/rerank.md @@ -8,15 +8,15 @@ LiteLLM Follows the [cohere api request / response for the rerank api](https://c ## Overview -| Feature | Supported | Notes | -|---------|-----------|-------| -| Cost Tracking | ✅ | Works with all supported models | -| Logging | ✅ | Works across all integrations | -| End-user Tracking | ✅ | | -| Fallbacks | ✅ | Works between supported models | -| Loadbalancing | ✅ | Works between supported models | -| Guardrails | ✅ | Applies to input query only (not documents) | -| Supported Providers | Cohere, Together AI, Azure AI, DeepInfra, Nvidia NIM, Infinity, Fireworks AI, Voyage AI | | +| Feature | Supported | Notes | +|---------|-----------------------------------------------------------------------------------------------------|-------| +| Cost Tracking | ✅ | Works with all supported models | +| Logging | ✅ | Works across all integrations | +| End-user Tracking | ✅ | | +| Fallbacks | ✅ | Works between supported models | +| Loadbalancing | ✅ | Works between supported models | +| Guardrails | ✅ | Applies to input query only (not documents) | +| Supported Providers | Cohere, Together AI, Azure AI, DeepInfra, Nvidia NIM, Infinity, Fireworks AI, Voyage AI, watsonx.ai | | ## **LiteLLM Python SDK Usage** ### Quick Start @@ -123,17 +123,18 @@ curl http://0.0.0.0:4000/rerank \ #### ⚡️See all supported models and providers at [models.litellm.ai](https://models.litellm.ai/) -| Provider | Link to Usage | -|-------------|--------------------| -| Cohere (v1 + v2 clients) | [Usage](#quick-start) | -| Together AI| [Usage](../docs/providers/togetherai) | -| Azure AI| [Usage](../docs/providers/azure_ai#rerank-endpoint) | -| Jina AI| [Usage](../docs/providers/jina_ai) | -| AWS Bedrock| [Usage](../docs/providers/bedrock#rerank-api) | -| HuggingFace| [Usage](../docs/providers/huggingface_rerank) | -| Infinity| [Usage](../docs/providers/infinity) | -| vLLM| [Usage](../docs/providers/vllm#rerank-endpoint) | -| DeepInfra| [Usage](../docs/providers/deepinfra#rerank-endpoint) | -| Vertex AI| [Usage](../docs/providers/vertex#rerank-api) | -| Fireworks AI| [Usage](../docs/providers/fireworks_ai#rerank-endpoint) | -| Voyage AI| [Usage](../docs/providers/voyage#rerank) | \ No newline at end of file +| Provider | Link to Usage | +|--------------------------|------------------------------------------------------| +| Cohere (v1 + v2 clients) | [Usage](#quick-start) | +| Together AI | [Usage](../docs/providers/togetherai) | +| Azure AI | [Usage](../docs/providers/azure_ai#rerank-endpoint) | +| Jina AI | [Usage](../docs/providers/jina_ai) | +| AWS Bedrock | [Usage](../docs/providers/bedrock#rerank-api) | +| HuggingFace | [Usage](../docs/providers/huggingface_rerank) | +| Infinity | [Usage](../docs/providers/infinity) | +| vLLM | [Usage](../docs/providers/vllm#rerank-endpoint) | +| DeepInfra | [Usage](../docs/providers/deepinfra#rerank-endpoint) | +| Vertex AI | [Usage](../docs/providers/vertex#rerank-api) | +| Fireworks AI | [Usage](../docs/providers/fireworks_ai#rerank-endpoint) | +| Voyage AI | [Usage](../docs/providers/voyage#rerank) | +| IBM watsonx.ai | [Usage](../docs/providers/watsonx/rerank) | \ No newline at end of file diff --git a/docs/my-website/docs/response_api.md b/docs/my-website/docs/response_api.md index dd2b77712c4..65b7ad7773a 100644 --- a/docs/my-website/docs/response_api.md +++ b/docs/my-website/docs/response_api.md @@ -884,7 +884,12 @@ router = litellm.Router( }, }, ], - optional_pre_call_checks=["responses_api_deployment_check"], + # `responses_api_deployment_check` ensures Requests with `previous_response_id` + # are routed to the same deployment. `deployment_affinity` adds sticky sessions + # for requests without `previous_response_id` (useful for implicit caching). + optional_pre_call_checks=["responses_api_deployment_check", "deployment_affinity"], + # Optional (default is 3600 seconds / 1 hour) + deployment_affinity_ttl_seconds=3600, ) # Initial request @@ -911,7 +916,16 @@ follow_up = await router.aresponses( #### 1. Setup session continuity on proxy config.yaml -To enable session continuity for Responses API in your LiteLLM proxy, set `optional_pre_call_checks: ["responses_api_deployment_check"]` in your proxy config.yaml. +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 +- `deployment_affinity`: sticky sessions based on user key (applies even without `previous_response_id`) + +Notes: +- User-key affinity is keyed on `metadata.user_api_key_hash` (the API key hash). The OpenAI `user` request parameter is an end-user identifier and is intentionally not used for deployment affinity. +- `user_api_key_hash` is already SHA-256, and is used as-is (no double hashing). +- Affinity is scoped by a stable model identifier (the model-map key, e.g. `model_map_information.model_map_key`) so model aliases map to the same stickiness bucket. +- The mapping TTL is controlled by `deployment_affinity_ttl_seconds` (configured on Router init / proxy startup). ```yaml showLineNumbers title="config.yaml with Session Continuity" model_list: @@ -929,7 +943,11 @@ model_list: api_base: https://endpoint2.openai.azure.com router_settings: - optional_pre_call_checks: ["responses_api_deployment_check"] + optional_pre_call_checks: + - responses_api_deployment_check + - deployment_affinity + # Optional (default is 3600 seconds / 1 hour) + deployment_affinity_ttl_seconds: 3600 ``` #### 2. Use the OpenAI Python SDK to make requests to LiteLLM Proxy @@ -1356,8 +1374,3 @@ Response: - - - - - diff --git a/docs/my-website/docs/search/index.md b/docs/my-website/docs/search/index.md index 551a495261a..8a71edead06 100644 --- a/docs/my-website/docs/search/index.md +++ b/docs/my-website/docs/search/index.md @@ -276,6 +276,7 @@ The response follows Perplexity's search format with the following structure: | Firecrawl | `FIRECRAWL_API_KEY` | `firecrawl` | | SearXNG | `SEARXNG_API_BASE` (required) | `searxng` | | Linkup | `LINKUP_API_KEY` | `linkup` | +| DuckDuckGo | `DUCKDUCKGO_API_BASE` | `duckduckgo` | See the individual provider documentation for detailed setup instructions and provider-specific parameters. diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index f2d7f6f423d..3acfa3937a9 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -944,6 +944,7 @@ const sidebars = { "providers/anthropic_tool_search", "guides/code_interpreter", "completion/message_trimming", + "completion/message_sanitization", "completion/model_alias", "completion/mock_requests", "completion/predict_outputs", diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index b1cbeecd1ec..bda20e2f744 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -1051,6 +1051,168 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): """Handled in files_endpoints.py""" return [] + def _is_batch_polling_enabled(self) -> bool: + """ + Check if batch cost tracking is actually enabled and running. + Returns: + bool: True if batch cost tracking is active, False otherwise + """ + try: + # Import here to avoid circular dependencies + import litellm.proxy.proxy_server as proxy_server_module + + # Check if the scheduler has the batch cost checking job registered + scheduler = getattr(proxy_server_module, 'scheduler', None) + if scheduler is None: + return False + + # Check if the check_batch_cost_job exists in the scheduler + try: + job = scheduler.get_job('check_batch_cost_job') + if job is not None: + return True + except Exception: + # Job not found or scheduler doesn't support get_job + pass + + return False + except Exception as e: + verbose_logger.warning( + f"Error checking batch polling configuration: {e}. Assuming disabled." + ) + return False + + async def _get_batches_referencing_file( + self, file_id: str + ) -> List[Dict[str, Any]]: + """ + Find batches in non-terminal states that reference this file. + + Non-terminal states: validating, in_progress, finalizing + Terminal states: completed, complete, failed, expired, cancelled + + Args: + file_id: The unified file ID to check + + Returns: + List of batch objects referencing this file in non-terminal state + (max 10 for error message display) + """ + # Prepare list of file IDs to check (both unified and provider IDs) + file_ids_to_check = [file_id] + + # Get model-specific file IDs for this unified file ID if it's a managed file + try: + model_file_id_mapping = await self.get_model_file_id_mapping( + [file_id], litellm_parent_otel_span=None + ) + + if model_file_id_mapping and file_id in model_file_id_mapping: + # Add all provider file IDs for this unified file + provider_file_ids = list(model_file_id_mapping[file_id].values()) + file_ids_to_check.extend(provider_file_ids) + except Exception as e: + verbose_logger.debug( + f"Could not get model file ID mapping for {file_id}: {e}. " + f"Will only check unified file ID." + ) + MAX_MATCHES_TO_RETURN = 10 + + batches = await self.prisma_client.db.litellm_managedobjecttable.find_many( + where={ + "file_purpose": "batch", + "status": {"in": ["validating", "in_progress", "finalizing"]}, + }, + take=MAX_MATCHES_TO_RETURN, + order={"created_at": "desc"}, + ) + + referencing_batches = [] + for batch in batches: + try: + # Parse the batch file_object to check for file references + batch_data = json.loads(batch.file_object) if isinstance(batch.file_object, str) else batch.file_object + + # Extract file IDs from batch + # Batches typically reference the unified file ID in input_file_id + # Output and error files are generated by the provider + input_file_id = batch_data.get("input_file_id") + output_file_id = batch_data.get("output_file_id") + error_file_id = batch_data.get("error_file_id") + + referenced_file_ids = [fid for fid in [input_file_id, output_file_id, error_file_id] if fid] + + # Check if any referenced file ID matches the file we're trying to delete + if any(ref_id in file_ids_to_check for ref_id in referenced_file_ids): + referencing_batches.append({ + "batch_id": batch.unified_object_id, + "status": batch.status, + "created_at": batch.created_at, + }) + except Exception as e: + verbose_logger.warning( + f"Error parsing batch object {batch.unified_object_id}: {e}" + ) + continue + + return referencing_batches + + async def _check_file_deletion_allowed(self, file_id: str) -> None: + """ + Check if file deletion should be blocked due to batch references. + + Blocks deletion if: + 1. File is referenced by any batch in non-terminal state, AND + 2. Batch polling is configured (user wants cost tracking) + + Args: + file_id: The unified file ID to check + + Raises: + HTTPException: If file deletion should be blocked + """ + # Check if batch polling is enabled + if not self._is_batch_polling_enabled(): + # Batch polling not configured, allow deletion + return + + # Check if file is referenced by any non-terminal batches + referencing_batches = await self._get_batches_referencing_file(file_id) + + if referencing_batches: + # File is referenced by non-terminal batches and polling is enabled + MAX_BATCHES_IN_ERROR = 5 # Limit batches shown in error message for readability + + # Show up to MAX_BATCHES_IN_ERROR in the error message + batches_to_show = referencing_batches[:MAX_BATCHES_IN_ERROR] + batch_statuses = [f"{b['batch_id']}: {b['status']}" for b in batches_to_show] + + # Determine the count message + count_message = f"{len(referencing_batches)}" + if len(referencing_batches) >= 10: # MAX_MATCHES_TO_RETURN from _get_batches_referencing_file + count_message = "10+" + + error_message = ( + f"Cannot delete file {file_id}. " + f"The file is referenced by {count_message} batch(es) in non-terminal state" + ) + + # Add specific batch details if not too many + if len(referencing_batches) <= MAX_BATCHES_IN_ERROR: + error_message += f": {', '.join(batch_statuses)}. " + else: + error_message += f" (showing {MAX_BATCHES_IN_ERROR} most recent): {', '.join(batch_statuses)}. " + + error_message += ( + f"To delete this file before complete cost tracking, please delete or cancel the referencing batch(es) first. " + f"Alternatively, wait for all batches to complete processing." + ) + + raise HTTPException( + status_code=400, + detail=error_message, + ) + async def afile_delete( self, file_id: str, @@ -1059,6 +1221,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): **data: Dict, ) -> OpenAIFileObject: + # Check if file deletion should be blocked due to batch references + await self._check_file_deletion_allowed(file_id) + # file_id = convert_b64_uid_to_unified_uid(file_id) model_file_id_mapping = await self.get_model_file_id_mapping( [file_id], litellm_parent_otel_span diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260131150814_add_team_user_to_vector_stores/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260131150814_add_team_user_to_vector_stores/migration.sql index 2032f76a5de..1f5dc311bd6 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260131150814_add_team_user_to_vector_stores/migration.sql +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260131150814_add_team_user_to_vector_stores/migration.sql @@ -1,10 +1,13 @@ -- AlterTable -ALTER TABLE "LiteLLM_ManagedVectorStoresTable" ADD COLUMN "team_id" TEXT, -ADD COLUMN "user_id" TEXT; +ALTER TABLE "LiteLLM_ManagedVectorStoresTable" + ADD COLUMN IF NOT EXISTS "team_id" TEXT, + ADD COLUMN IF NOT EXISTS "user_id" TEXT; -- CreateIndex -CREATE INDEX "LiteLLM_ManagedVectorStoresTable_team_id_idx" ON "LiteLLM_ManagedVectorStoresTable"("team_id"); +CREATE INDEX IF NOT EXISTS "LiteLLM_ManagedVectorStoresTable_team_id_idx" + ON "LiteLLM_ManagedVectorStoresTable"("team_id"); -- CreateIndex -CREATE INDEX "LiteLLM_ManagedVectorStoresTable_user_id_idx" ON "LiteLLM_ManagedVectorStoresTable"("user_id"); +CREATE INDEX IF NOT EXISTS "LiteLLM_ManagedVectorStoresTable_user_id_idx" + ON "LiteLLM_ManagedVectorStoresTable"("user_id"); diff --git a/litellm/__init__.py b/litellm/__init__.py index ea7bb6a78d8..a994db85b11 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1355,6 +1355,7 @@ if TYPE_CHECKING: from .llms.vertex_ai.rerank.transformation import VertexAIRerankConfig as VertexAIRerankConfig from .llms.fireworks_ai.rerank.transformation import FireworksAIRerankConfig as FireworksAIRerankConfig from .llms.voyage.rerank.transformation import VoyageRerankConfig as VoyageRerankConfig + from .llms.watsonx.rerank.transformation import IBMWatsonXRerankConfig as IBMWatsonXRerankConfig from .llms.clarifai.chat.transformation import ClarifaiConfig as ClarifaiConfig from .llms.ai21.chat.transformation import AI21ChatConfig as AI21ChatConfig from .llms.meta_llama.chat.transformation import LlamaAPIConfig as LlamaAPIConfig diff --git a/litellm/_lazy_imports_registry.py b/litellm/_lazy_imports_registry.py index cf446d15156..943acc6320f 100644 --- a/litellm/_lazy_imports_registry.py +++ b/litellm/_lazy_imports_registry.py @@ -155,6 +155,7 @@ LLM_CONFIG_NAMES = ( "VertexAIRerankConfig", "FireworksAIRerankConfig", "VoyageRerankConfig", + "IBMWatsonXRerankConfig", "ClarifaiConfig", "AI21ChatConfig", "LlamaAPIConfig", @@ -672,6 +673,7 @@ _LLM_CONFIGS_IMPORT_MAP = { "FireworksAIRerankConfig", ), "VoyageRerankConfig": (".llms.voyage.rerank.transformation", "VoyageRerankConfig"), + "IBMWatsonXRerankConfig": (".llms.watsonx.rerank.transformation", "IBMWatsonXRerankConfig"), "ClarifaiConfig": (".llms.clarifai.chat.transformation", "ClarifaiConfig"), "AI21ChatConfig": (".llms.ai21.chat.transformation", "AI21ChatConfig"), "LlamaAPIConfig": (".llms.meta_llama.chat.transformation", "LlamaAPIConfig"), diff --git a/litellm/anthropic_beta_headers_config.json b/litellm/anthropic_beta_headers_config.json index 5dd8536f4c0..a06c7173ea3 100644 --- a/litellm/anthropic_beta_headers_config.json +++ b/litellm/anthropic_beta_headers_config.json @@ -148,5 +148,35 @@ "tool-search-tool-2025-10-19": "tool-search-tool-2025-10-19", "web-fetch-2025-09-10": null, "web-search-2025-03-05": "web-search-2025-03-05" + }, + "databricks": { + "advanced-tool-use-2025-11-20": "advanced-tool-use-2025-11-20", + "bash_20241022": null, + "bash_20250124": null, + "code-execution-2025-08-25": "code-execution-2025-08-25", + "compact-2026-01-12": "compact-2026-01-12", + "computer-use-2025-01-24": "computer-use-2025-01-24", + "computer-use-2025-11-24": "computer-use-2025-11-24", + "context-1m-2025-08-07": "context-1m-2025-08-07", + "context-management-2025-06-27": "context-management-2025-06-27", + "effort-2025-11-24": "effort-2025-11-24", + "fast-mode-2026-02-01": "fast-mode-2026-02-01", + "files-api-2025-04-14": "files-api-2025-04-14", + "structured-output-2024-03-01": null, + "fine-grained-tool-streaming-2025-05-14": "fine-grained-tool-streaming-2025-05-14", + "interleaved-thinking-2025-05-14": "interleaved-thinking-2025-05-14", + "mcp-client-2025-11-20": "mcp-client-2025-11-20", + "mcp-client-2025-04-04": "mcp-client-2025-04-04", + "mcp-servers-2025-12-04": null, + "oauth-2025-04-20": "oauth-2025-04-20", + "output-128k-2025-02-19": "output-128k-2025-02-19", + "prompt-caching-scope-2026-01-05": "prompt-caching-scope-2026-01-05", + "skills-2025-10-02": "skills-2025-10-02", + "structured-outputs-2025-11-13": "structured-outputs-2025-11-13", + "text_editor_20241022": null, + "text_editor_20250124": null, + "token-efficient-tools-2025-02-19": "token-efficient-tools-2025-02-19", + "web-fetch-2025-09-10": "web-fetch-2025-09-10", + "web-search-2025-03-05": "web-search-2025-03-05" } } \ No newline at end of file diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index e546a0dbb02..35fc93bbeb0 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -62,9 +62,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): def __init__(self): pass - def _handle_raw_dict_response_item( - self, item: Dict[str, Any], index: int - ) -> Tuple[Optional[Any], int]: + def _handle_raw_dict_response_item(self, item: Dict[str, Any], index: int) -> Tuple[Optional[Any], int]: """ Handle raw dict response items from Responses API (e.g., GPT-5 Codex format). @@ -107,13 +105,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if item_type == "function_call": # Extract provider_specific_fields if present and pass through as-is provider_specific_fields = item.get("provider_specific_fields") - if provider_specific_fields and not isinstance( - provider_specific_fields, dict - ): + if provider_specific_fields and not isinstance(provider_specific_fields, dict): provider_specific_fields = ( - dict(provider_specific_fields) - if hasattr(provider_specific_fields, "__dict__") - else {} + dict(provider_specific_fields) if hasattr(provider_specific_fields, "__dict__") else {} ) tool_call_dict = { @@ -129,9 +123,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if provider_specific_fields: tool_call_dict["provider_specific_fields"] = provider_specific_fields # Also add to function's provider_specific_fields for consistency - tool_call_dict["function"][ - "provider_specific_fields" - ] = provider_specific_fields + tool_call_dict["function"]["provider_specific_fields"] = provider_specific_fields msg = Message( content=None, @@ -169,7 +161,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): "type": "message", "role": role, "content": self._convert_content_to_responses_format( - content, role # type: ignore + content, + role, # type: ignore ), } ) @@ -186,7 +179,8 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): elif isinstance(content, list): # Transform list content to Responses API format tool_output = self._convert_content_to_responses_format( - content, "user" # Use "user" role to get input_* types + content, + "user", # Use "user" role to get input_* types ) else: # Fallback: convert unexpected types to input_text @@ -219,9 +213,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): { "type": "message", "role": role, - "content": self._convert_content_to_responses_format( - content, cast(str, role) - ), + "content": self._convert_content_to_responses_format(content, cast(str, role)), } ) @@ -344,9 +336,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): previous_response_id = optional_params.get("previous_response_id") if previous_response_id: # Use the existing session handler for responses API - verbose_logger.debug( - f"Chat provider: Warning ignoring previous response ID: {previous_response_id}" - ) + verbose_logger.debug(f"Chat provider: Warning ignoring previous response ID: {previous_response_id}") # Convert back to responses API format for the actual request @@ -368,9 +358,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): "client": client, } - verbose_logger.debug( - f"Chat provider: Final request model={api_model}, input_items={len(input_items)}" - ) + verbose_logger.debug(f"Chat provider: Final request model={api_model}, input_items={len(input_items)}") self._merge_responses_api_request_into_request_data( request_data, responses_api_request, instructions @@ -450,9 +438,11 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): LiteLLMCompletionResponsesConfig, ) - tool_call_dict = LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call( - tool_call_item=item, - index=tool_call_index, + tool_call_dict = ( + LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call( + tool_call_item=item, + index=tool_call_index, + ) ) accumulated_tool_calls.append(tool_call_dict) tool_call_index += 1 @@ -472,9 +462,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): tool_calls=accumulated_tool_calls, reasoning_content=reasoning_content, ) - choices.append( - Choices(message=msg, finish_reason="tool_calls", index=index) - ) + choices.append(Choices(message=msg, finish_reason="tool_calls", index=index)) reasoning_content = None return choices @@ -510,17 +498,10 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ) if len(choices) == 0: - if ( - raw_response.incomplete_details is not None - and raw_response.incomplete_details.reason is not None - ): - raise ValueError( - f"{model} unable to complete request: {raw_response.incomplete_details.reason}" - ) + if raw_response.incomplete_details is not None and raw_response.incomplete_details.reason is not None: + raise ValueError(f"{model} unable to complete request: {raw_response.incomplete_details.reason}") else: - raise ValueError( - f"Unknown items in responses API response: {raw_response.output}" - ) + raise ValueError(f"Unknown items in responses API response: {raw_response.output}") setattr(model_response, "choices", choices) @@ -529,11 +510,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): setattr( model_response, "usage", - ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage( - raw_response.usage - ), + ResponseAPILoggingUtils._transform_response_api_usage_to_chat_usage(raw_response.usage), ) - + # Preserve hidden params from the ResponsesAPIResponse, especially the headers # which contain important provider information like x-request-id raw_response_hidden_params = getattr(raw_response, "_hidden_params", {}) @@ -550,24 +529,18 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): model_response._hidden_params[key] = merged_headers else: model_response._hidden_params[key] = value - + return model_response def get_model_response_iterator( self, - streaming_response: Union[ - Iterator[str], AsyncIterator[str], "ModelResponse", "BaseModel" - ], + streaming_response: Union[Iterator[str], AsyncIterator[str], "ModelResponse", "BaseModel"], sync_stream: bool, json_mode: Optional[bool] = False, ) -> BaseModelResponseIterator: - return OpenAiResponsesToChatCompletionStreamIterator( - streaming_response, sync_stream, json_mode - ) + return OpenAiResponsesToChatCompletionStreamIterator(streaming_response, sync_stream, json_mode) - def _convert_content_str_to_input_text( - self, content: str, role: str - ) -> Dict[str, Any]: + def _convert_content_str_to_input_text(self, content: str, role: str) -> Dict[str, Any]: if role == "user" or role == "system" or role == "tool": return {"type": "input_text", "text": content} else: @@ -594,9 +567,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if actual_image_url is None: raise ValueError(f"Invalid image URL: {content_image_url}") - image_param = ResponseInputImageParam( - image_url=actual_image_url, detail="auto", type="input_image" - ) + image_param = ResponseInputImageParam(image_url=actual_image_url, detail="auto", type="input_image") if detail: image_param["detail"] = detail @@ -605,31 +576,29 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): def _convert_content_to_responses_format( self, - content: Union[ - str, - Iterable[ - Union["OpenAIMessageContentListBlock", "ChatCompletionThinkingBlock"] - ], + content: Optional[ + Union[ + str, + Iterable[Union["OpenAIMessageContentListBlock", "ChatCompletionThinkingBlock"]], + ] ], role: str, ) -> List[Dict[str, Any]]: """Convert chat completion content to responses API format""" from litellm.types.llms.openai import ChatCompletionImageObject - verbose_logger.debug( - f"Chat provider: Converting content to responses format - input type: {type(content)}" - ) + verbose_logger.debug(f"Chat provider: Converting content to responses format - input type: {type(content)}") - if isinstance(content, str): + if content is None: + return [self._convert_content_str_to_input_text("", role)] + elif isinstance(content, str): result = [self._convert_content_str_to_input_text(content, role)] verbose_logger.debug(f"Chat provider: String content -> {result}") return result elif isinstance(content, list): result = [] for i, item in enumerate(content): - verbose_logger.debug( - f"Chat provider: Processing content item {i}: {type(item)} = {item}" - ) + verbose_logger.debug(f"Chat provider: Processing content item {i}: {type(item)} = {item}") if isinstance(item, str): converted = self._convert_content_str_to_input_text(item, role) result.append(converted) @@ -638,9 +607,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # Handle multimodal content original_type = item.get("type") if original_type == "text": - converted = self._convert_content_str_to_input_text( - item.get("text", ""), role - ) + converted = self._convert_content_str_to_input_text(item.get("text", ""), role) result.append(converted) verbose_logger.debug(f"Chat provider: text -> {converted}") elif original_type == "image_url": @@ -652,18 +619,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ), ) result.append(converted) - verbose_logger.debug( - f"Chat provider: image_url -> {converted}" - ) + verbose_logger.debug(f"Chat provider: image_url -> {converted}") else: # Try to map other types to responses API format item_type = original_type or "input_text" if item_type == "image": converted = {"type": "input_image", **item} result.append(converted) - verbose_logger.debug( - f"Chat provider: image -> {converted}" - ) + verbose_logger.debug(f"Chat provider: image -> {converted}") elif item_type in [ "input_text", "input_image", @@ -675,18 +638,12 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): ]: # Already in responses API format result.append(item) - verbose_logger.debug( - f"Chat provider: passthrough -> {item}" - ) + verbose_logger.debug(f"Chat provider: passthrough -> {item}") else: # Default to input_text for unknown types - converted = self._convert_content_str_to_input_text( - str(item.get("text", item)), role - ) + converted = self._convert_content_str_to_input_text(str(item.get("text", item)), role) result.append(converted) - verbose_logger.debug( - f"Chat provider: unknown({original_type}) -> {converted}" - ) + verbose_logger.debug(f"Chat provider: unknown({original_type}) -> {converted}") verbose_logger.debug(f"Chat provider: Final converted content: {result}") return result else: @@ -694,17 +651,13 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): verbose_logger.debug(f"Chat provider: Other content type -> {result}") return result - def _convert_tools_to_responses_format( - self, tools: List[Dict[str, Any]] - ) -> List["ALL_RESPONSES_API_TOOL_PARAMS"]: + def _convert_tools_to_responses_format(self, tools: List[Dict[str, Any]]) -> List["ALL_RESPONSES_API_TOOL_PARAMS"]: """Convert chat completion tools to responses API tools format""" responses_tools: List["ALL_RESPONSES_API_TOOL_PARAMS"] = [] for tool in tools: # convert function tool from chat completion to responses API format if tool.get("type") == "function": - function_tool = cast( - ChatCompletionToolParamFunctionChunk, tool.get("function") - ) + function_tool = cast(ChatCompletionToolParamFunctionChunk, tool.get("function")) responses_tools.append( FunctionToolParam( name=function_tool["name"], @@ -730,9 +683,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): if not extra_body: return optional_params - supported_responses_api_params = set( - ResponsesAPIOptionalRequestParams.__annotations__.keys() - ) + supported_responses_api_params = set(ResponsesAPIOptionalRequestParams.__annotations__.keys()) # Also include params we handle specially supported_responses_api_params.update( { @@ -750,9 +701,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return optional_params - def _map_reasoning_effort( - self, reasoning_effort: Union[str, Dict[str, Any]] - ) -> Optional[Reasoning]: + def _map_reasoning_effort(self, reasoning_effort: Union[str, Dict[str, Any]]) -> Optional[Reasoning]: # If dict is passed, convert it directly to Reasoning object if isinstance(reasoning_effort, dict): return Reasoning(**reasoning_effort) # type: ignore[typeddict-item] @@ -760,8 +709,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): # Check if auto-summary is enabled via flag or environment variable # Priority: litellm.reasoning_auto_summary flag > LITELLM_REASONING_AUTO_SUMMARY env var auto_summary_enabled = ( - litellm.reasoning_auto_summary - or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true" + litellm.reasoning_auto_summary or os.getenv("LITELLM_REASONING_AUTO_SUMMARY", "false").lower() == "true" ) # If string is passed, map with optional summary based on flag/env var @@ -772,11 +720,15 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): elif reasoning_effort == "xhigh": return Reasoning(effort="xhigh", summary="detailed") if auto_summary_enabled else Reasoning(effort="xhigh") # type: ignore[typeddict-item] elif reasoning_effort == "medium": - return Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium") + return ( + Reasoning(effort="medium", summary="detailed") if auto_summary_enabled else Reasoning(effort="medium") + ) elif reasoning_effort == "low": return Reasoning(effort="low", summary="detailed") if auto_summary_enabled else Reasoning(effort="low") elif reasoning_effort == "minimal": - return Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal") + return ( + Reasoning(effort="minimal", summary="detailed") if auto_summary_enabled else Reasoning(effort="minimal") + ) return None def _add_web_search_tool( @@ -855,7 +807,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): return {"format": {"type": "text"}} return None - + @staticmethod def _convert_annotations_to_chat_format( annotations: Optional[List[Any]], @@ -908,9 +860,7 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): - def __init__( - self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False - ): + def __init__(self, streaming_response, sync_stream: bool, json_mode: Optional[bool] = False): super().__init__(streaming_response, sync_stream, json_mode) def _handle_string_chunk( @@ -923,9 +873,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): if not str_line or str_line.startswith("event:"): # ignore. - return GenericStreamingChunk( - text="", tool_use=None, is_finished=False, finish_reason="", usage=None - ) + return GenericStreamingChunk(text="", tool_use=None, is_finished=False, finish_reason="", usage=None) index = str_line.find("data:") if index != -1: str_line = str_line[index + 5 :] @@ -988,13 +936,9 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): if output_item.get("type") == "function_call": # Extract provider_specific_fields if present provider_specific_fields = output_item.get("provider_specific_fields") - if provider_specific_fields and not isinstance( - provider_specific_fields, dict - ): + if provider_specific_fields and not isinstance(provider_specific_fields, dict): provider_specific_fields = ( - dict(provider_specific_fields) - if hasattr(provider_specific_fields, "__dict__") - else {} + dict(provider_specific_fields) if hasattr(provider_specific_fields, "__dict__") else {} ) function_chunk = ChatCompletionToolCallFunctionChunk( @@ -1003,9 +947,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): ) if provider_specific_fields: - function_chunk["provider_specific_fields"] = ( - provider_specific_fields - ) + function_chunk["provider_specific_fields"] = provider_specific_fields tool_call_chunk = ChatCompletionToolCallChunk( id=output_item.get("call_id"), @@ -1040,9 +982,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): id=None, index=0, type="function", - function=ChatCompletionToolCallFunctionChunk( - name=None, arguments=content_part - ), + function=ChatCompletionToolCallFunctionChunk(name=None, arguments=content_part), ) ] ), @@ -1051,22 +991,16 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): ] ) else: - raise ValueError( - f"Chat provider: Invalid function argument delta {parsed_chunk}" - ) + raise ValueError(f"Chat provider: Invalid function argument delta {parsed_chunk}") elif event_type == "response.output_item.done": # New output item added output_item = parsed_chunk.get("item", {}) if output_item.get("type") == "function_call": # Extract provider_specific_fields if present provider_specific_fields = output_item.get("provider_specific_fields") - if provider_specific_fields and not isinstance( - provider_specific_fields, dict - ): + if provider_specific_fields and not isinstance(provider_specific_fields, dict): provider_specific_fields = ( - dict(provider_specific_fields) - if hasattr(provider_specific_fields, "__dict__") - else {} + dict(provider_specific_fields) if hasattr(provider_specific_fields, "__dict__") else {} ) function_chunk = ChatCompletionToolCallFunctionChunk( @@ -1076,9 +1010,7 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): # Add provider_specific_fields to function if present if provider_specific_fields: - function_chunk["provider_specific_fields"] = ( - provider_specific_fields - ) + function_chunk["provider_specific_fields"] = provider_specific_fields tool_call_chunk = ChatCompletionToolCallChunk( id=output_item.get("call_id"), @@ -1142,21 +1074,31 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): elif event_type == "response.completed": # Response is fully complete - now we can signal is_finished=True # This ensures we don't prematurely end the stream before tool_calls arrive + + # Check if response contains function_call items in output + # to determine correct finish_reason + response_data = parsed_chunk.get("response", {}) + output_items = response_data.get("output", []) if response_data else [] + + has_function_calls = any( + item.get("type") == "function_call" for item in output_items if isinstance(item, dict) + ) + + finish_reason = "tool_calls" if has_function_calls else "stop" + return ModelResponseStream( choices=[ StreamingChoices( index=0, delta=Delta(content=""), - finish_reason="stop", + finish_reason=finish_reason, ) ] ) else: pass # For any unhandled event types, create a minimal valid chunk or skip - verbose_logger.debug( - f"Chat provider: Unhandled event type '{event_type}', creating empty chunk" - ) + verbose_logger.debug(f"Chat provider: Unhandled event type '{event_type}', creating empty chunk") # Return a minimal valid chunk for unknown events return ModelResponseStream( @@ -1179,9 +1121,5 @@ class OpenAiResponsesToChatCompletionStreamIterator(BaseModelResponseIterator): Returns: ModelResponseStream: OpenAI-formatted streaming chunk """ - verbose_logger.debug( - f"Chat provider: transform_streaming_response called with chunk: {chunk}" - ) - return OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream( - chunk - ) + verbose_logger.debug(f"Chat provider: transform_streaming_response called with chunk: {chunk}") + return OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(chunk) diff --git a/litellm/constants.py b/litellm/constants.py index 3c11eb701dc..17ad742e419 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -287,7 +287,9 @@ MIN_NON_ZERO_TEMPERATURE = float(os.getenv("MIN_NON_ZERO_TEMPERATURE", 0.0001)) REPEATED_STREAMING_CHUNK_LIMIT = int( os.getenv("REPEATED_STREAMING_CHUNK_LIMIT", 100) ) # catch if model starts looping the same chunk while streaming. Uses high default to prevent false positives. -DEFAULT_MAX_LRU_CACHE_SIZE = int(os.getenv("DEFAULT_MAX_LRU_CACHE_SIZE", 16)) +# Shared maxsize for functools.lru_cache usage across hot paths. +# Defaulted to 64 to avoid cache thrash in multi-model production workloads. +DEFAULT_MAX_LRU_CACHE_SIZE = int(os.getenv("DEFAULT_MAX_LRU_CACHE_SIZE", 64)) _REALTIME_BODY_CACHE_SIZE = 1000 # Keep realtime helper caches bounded; workloads rarely exceed 1k models/intents INITIAL_RETRY_DELAY = float(os.getenv("INITIAL_RETRY_DELAY", 0.5)) MAX_RETRY_DELAY = float(os.getenv("MAX_RETRY_DELAY", 8.0)) @@ -576,6 +578,11 @@ OPENAI_CHAT_COMPLETION_PARAMS = [ "thinking", "web_search_options", "service_tier", + "store", + "prompt_cache_key", + "prompt_cache_retention", + "safety_identifier", + "verbosity", ] OPENAI_TRANSCRIPTION_PARAMS = [ diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py index dae0bb1c2c0..02df747792d 100644 --- a/litellm/cost_calculator.py +++ b/litellm/cost_calculator.py @@ -448,7 +448,9 @@ def cost_per_token( # noqa: PLR0915 elif custom_llm_provider == "anthropic": return anthropic_cost_per_token(model=model, usage=usage_block) elif custom_llm_provider == "bedrock": - return bedrock_cost_per_token(model=model, usage=usage_block) + return bedrock_cost_per_token( + model=model, usage=usage_block, service_tier=service_tier + ) elif custom_llm_provider == "openai": return openai_cost_per_token( model=model, usage=usage_block, service_tier=service_tier @@ -2146,4 +2148,3 @@ def handle_realtime_stream_cost_calculation( return total_cost - diff --git a/litellm/integrations/datadog/datadog_cost_management.py b/litellm/integrations/datadog/datadog_cost_management.py index 2eb94b59dd8..a961d4f9244 100644 --- a/litellm/integrations/datadog/datadog_cost_management.py +++ b/litellm/integrations/datadog/datadog_cost_management.py @@ -93,7 +93,9 @@ class DatadogCostManagementLogger(CustomBatchLogger): Aggregates costs by Provider, Model, and Date. Returns a list of DatadogFOCUSCostEntry. """ - aggregator: Dict[Tuple[str, str, str, Tuple[Tuple[str, str], ...]], DatadogFOCUSCostEntry] = {} + aggregator: Dict[ + Tuple[str, str, str, Tuple[Tuple[str, str], ...]], DatadogFOCUSCostEntry + ] = {} for log in logs: try: @@ -167,10 +169,20 @@ class DatadogCostManagementLogger(CustomBatchLogger): metadata = log.get("metadata", {}) if metadata: # Add user info - if "user_api_key_alias" in metadata: + # Add user info + if metadata.get("user_api_key_alias"): tags["user"] = str(metadata["user_api_key_alias"]) - if "user_api_key_team_alias" in metadata: - tags["team"] = str(metadata["user_api_key_team_alias"]) + + # Add Team Tag + team_tag = ( + metadata.get("user_api_key_team_alias") + or metadata.get("team_alias") # type: ignore + or metadata.get("user_api_key_team_id") + or metadata.get("team_id") # type: ignore + ) + + if team_tag: + tags["team"] = str(team_tag) # model_group is not in StandardLoggingMetadata TypedDict, so we need to access it via dict.get() model_group = metadata.get("model_group") # type: ignore[misc] if model_group: diff --git a/litellm/integrations/datadog/datadog_handler.py b/litellm/integrations/datadog/datadog_handler.py index e2f30f2f614..0406f1e5d20 100644 --- a/litellm/integrations/datadog/datadog_handler.py +++ b/litellm/integrations/datadog/datadog_handler.py @@ -55,4 +55,15 @@ def get_datadog_tags( request_tags = standard_logging_object.get("request_tags", []) or [] tags.extend(f"request_tag:{tag}" for tag in request_tags) + # Add Team Tag + metadata = standard_logging_object.get("metadata", {}) or {} + team_tag = ( + metadata.get("user_api_key_team_alias") + or metadata.get("team_alias") + or metadata.get("user_api_key_team_id") + or metadata.get("team_id") + ) + if team_tag: + tags.append(f"team:{team_tag}") + return ",".join(tags) diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 1675201f1f1..4c7afd5a57c 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -22,6 +22,10 @@ from typing import ( import litellm from litellm._logging import print_verbose, verbose_logger from litellm.integrations.custom_logger import CustomLogger +from litellm.litellm_core_utils.core_helpers import ( + get_litellm_metadata_from_kwargs, + get_metadata_variable_name_from_kwargs, +) from litellm.proxy._types import ( LiteLLM_DeletedVerificationToken, LiteLLM_TeamTable, @@ -1055,16 +1059,16 @@ class PrometheusLogger(CustomLogger): enum_values=enum_values, ) - if ( - standard_logging_payload["stream"] is True - ): # log successful streaming requests from logging event hook. - _labels = prometheus_label_factory( - supported_enum_labels=self.get_labels_for_metric( - metric_name="litellm_proxy_total_requests_metric" - ), - enum_values=enum_values, - ) - self.litellm_proxy_total_requests_metric.labels(**_labels).inc() + # increment litellm_proxy_total_requests_metric for all successful requests + # (both streaming and non-streaming) in this single location to prevent + # double-counting that occurs when async_post_call_success_hook also increments + _labels = prometheus_label_factory( + supported_enum_labels=self.get_labels_for_metric( + metric_name="litellm_proxy_total_requests_metric" + ), + enum_values=enum_values, + ) + self.litellm_proxy_total_requests_metric.labels(**_labels).inc() def _increment_token_metrics( self, @@ -1086,13 +1090,6 @@ class PrometheusLogger(CustomLogger): ): _tags = standard_logging_payload["request_tags"] - _labels = prometheus_label_factory( - supported_enum_labels=self.get_labels_for_metric( - metric_name="litellm_proxy_total_requests_metric" - ), - enum_values=enum_values, - ) - _labels = prometheus_label_factory( supported_enum_labels=self.get_labels_for_metric( metric_name="litellm_total_tokens_metric" @@ -1655,49 +1652,12 @@ class PrometheusLogger(CustomLogger): ): """ Proxy level tracking - triggered when the proxy responds with a success response to the client + + Note: litellm_proxy_total_requests_metric is NOT incremented here to avoid + double-counting. It is incremented in async_log_success_event which fires + for all successful requests (both streaming and non-streaming). """ - try: - from litellm.litellm_core_utils.litellm_logging import ( - StandardLoggingPayloadSetup, - ) - - if self._should_skip_metrics_for_invalid_key( - user_api_key_dict=user_api_key_dict - ): - return - - _metadata = data.get("metadata", {}) or {} - enum_values = UserAPIKeyLabelValues( - end_user=user_api_key_dict.end_user_id, - hashed_api_key=user_api_key_dict.api_key, - api_key_alias=user_api_key_dict.key_alias, - requested_model=data.get("model", ""), - team=user_api_key_dict.team_id, - team_alias=user_api_key_dict.team_alias, - user=user_api_key_dict.user_id, - user_email=user_api_key_dict.user_email, - status_code="200", - route=user_api_key_dict.request_route, - tags=StandardLoggingPayloadSetup._get_request_tags( - litellm_params=data, - proxy_server_request=data.get("proxy_server_request", {}), - ), - client_ip=_metadata.get("requester_ip_address"), - user_agent=_metadata.get("user_agent"), - ) - _labels = prometheus_label_factory( - supported_enum_labels=self.get_labels_for_metric( - metric_name="litellm_proxy_total_requests_metric" - ), - enum_values=enum_values, - ) - self.litellm_proxy_total_requests_metric.labels(**_labels).inc() - - except Exception as e: - verbose_logger.exception( - "prometheus Layer Error(): Exception occured - {}".format(str(e)) - ) - pass + pass def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any: """Get value from dict or Pydantic model.""" @@ -2004,7 +1964,7 @@ class PrometheusLogger(CustomLogger): api_base = standard_logging_payload["api_base"] _litellm_params = request_kwargs.get("litellm_params", {}) or {} - _metadata = _litellm_params.get("metadata", {}) + _metadata = get_litellm_metadata_from_kwargs(request_kwargs) litellm_model_name = request_kwargs.get("model", None) llm_provider = _litellm_params.get("custom_llm_provider", None) _model_info = _metadata.get("model_info") or {} @@ -2220,7 +2180,8 @@ class PrometheusLogger(CustomLogger): original_model_group, kwargs, ) - _metadata = kwargs.get("metadata", {}) + _metadata_key = get_metadata_variable_name_from_kwargs(kwargs) + _metadata = kwargs.get(_metadata_key) or {} standard_metadata: StandardLoggingMetadata = ( StandardLoggingPayloadSetup.get_standard_logging_metadata( metadata=_metadata @@ -2265,7 +2226,8 @@ class PrometheusLogger(CustomLogger): kwargs, ) _new_model = kwargs.get("model") - _metadata = kwargs.get("metadata", {}) + _metadata_key = get_metadata_variable_name_from_kwargs(kwargs) + _metadata = kwargs.get(_metadata_key) or {} _tags = cast(List[str], kwargs.get("tags") or []) standard_metadata: StandardLoggingMetadata = ( StandardLoggingPayloadSetup.get_standard_logging_metadata( diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index bdbbc7579b7..6a14e42c485 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -1335,7 +1335,11 @@ class Logging(LiteLLMLoggingBaseClass): ) # Store additional costs if provided (free-form dict for extensibility) - if additional_costs and isinstance(additional_costs, dict) and len(additional_costs) > 0: + if ( + additional_costs + and isinstance(additional_costs, dict) + and len(additional_costs) > 0 + ): self.cost_breakdown["additional_costs"] = additional_costs # Store discount information if provided @@ -4519,13 +4523,19 @@ class StandardLoggingPayloadSetup: requester_custom_headers=None, cold_storage_object_key=None, user_api_key_auth_metadata=None, + team_alias=None, + team_id=None, ) if isinstance(metadata, dict): for key in metadata.keys() & _STANDARD_LOGGING_METADATA_KEYS: clean_metadata[key] = metadata[key] # type: ignore user_api_key = metadata.get("user_api_key") - if user_api_key and isinstance(user_api_key, str) and is_valid_sha256_hash(user_api_key): + if ( + user_api_key + and isinstance(user_api_key, str) + and is_valid_sha256_hash(user_api_key) + ): clean_metadata["user_api_key_hash"] = user_api_key _potential_requester_metadata = metadata.get( "metadata", None @@ -5279,6 +5289,8 @@ def get_standard_logging_metadata( user_api_key_request_route=None, cold_storage_object_key=None, user_api_key_auth_metadata=None, + team_alias=None, + team_id=None, ) if isinstance(metadata, dict): # Update the clean_metadata with values from input metadata that match StandardLoggingMetadata fields diff --git a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py index 25ad0a570cb..a6e502a32b3 100644 --- a/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py +++ b/litellm/litellm_core_utils/llm_response_utils/convert_dict_to_response.py @@ -546,7 +546,11 @@ def convert_to_model_response_object( # noqa: PLR0915 message = litellm.Message(content=json_mode_content_str) finish_reason = "stop" if message is None: - provider_specific_fields = {} + # Preserve provider_specific_fields if already present + # in the response (e.g. from proxy passthrough) + provider_specific_fields = dict( + choice["message"].get("provider_specific_fields", None) or {} + ) message_keys = Message.model_fields.keys() for field in choice["message"].keys(): if field not in message_keys: diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index c907ed32b95..7b485501f61 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -2018,6 +2018,235 @@ def anthropic_process_openai_file_message( ) +def _sanitize_empty_text_content( + message: AllMessageValues, +) -> AllMessageValues: + """ + Case C: Sanitize empty text content + - Replace empty or whitespace-only text content with a placeholder message. + + Returns: + The message with sanitized content if needed, otherwise the original message + """ + if message.get("role") in ["user", "assistant"]: + content = message.get("content") + if isinstance(content, str): + if not content or not content.strip(): + message = cast(AllMessageValues, dict(message)) # Make a copy + message["content"] = "[System: Empty message content sanitised to satisfy protocol]" + verbose_logger.debug( + f"_sanitize_empty_text_content: Replaced empty text content in {message.get('role')} message" + ) + return message + + +def _add_missing_tool_results( # noqa: PLR0915 + current_message: AllMessageValues, + messages: List[AllMessageValues], + current_index: int, +) -> Tuple[List[AllMessageValues], int]: + """ + Case A: Missing tool_result for tool_use (orphaned tool calls) + - If an assistant message has tool_calls but no corresponding tool result follows, + add a dummy tool result message indicating the user did not provide the result. + + Returns: + A tuple of: + - List containing the assistant message, followed by existing tool results, + followed by any dummy tool results needed + - Number of original messages consumed (to adjust iteration index) + """ + result_messages: List[AllMessageValues] = [] + tool_calls = current_message.get("tool_calls") + + if not tool_calls or len(cast(list, tool_calls)) == 0: + return ([current_message], 0) + + # Collect all tool_call_ids from this assistant message + expected_tool_call_ids = set() + for tool_call in cast(list, tool_calls): + tool_call_id = None + if isinstance(tool_call, dict): + tool_call_id = tool_call.get("id") + else: + tool_call_id = getattr(tool_call, "id", None) + if tool_call_id: + expected_tool_call_ids.add(tool_call_id) + + # Collect actual tool result messages that follow this assistant message + found_tool_call_ids = set() + actual_tool_results: List[AllMessageValues] = [] + j = current_index + 1 + + while j < len(messages): + next_msg = messages[j] + next_role = next_msg.get("role") + + if next_role == "assistant": + break + + if next_role in ["tool", "function"]: + tool_call_id = next_msg.get("tool_call_id") + if tool_call_id and tool_call_id in expected_tool_call_ids: + found_tool_call_ids.add(tool_call_id) + actual_tool_results.append(next_msg) + + j += 1 + + # Find missing tool results + missing_tool_call_ids = expected_tool_call_ids - found_tool_call_ids + + if missing_tool_call_ids: + verbose_logger.debug( + f"_add_missing_tool_results: Found {len(missing_tool_call_ids)} orphaned tool calls. Adding dummy tool results." + ) + + result_messages.append(current_message) + + # Add existing tool results FIRST + result_messages.extend(actual_tool_results) + + # Then add dummy tool results for missing ones + for tool_call_id in missing_tool_call_ids: + tool_name = "unknown_tool" + for tool_call in cast(list, tool_calls): + tc_id = None + if isinstance(tool_call, dict): + tc_id = tool_call.get("id") + else: + tc_id = getattr(tool_call, "id", None) + + if tc_id == tool_call_id: + if isinstance(tool_call, dict): + function = tool_call.get("function", {}) + if isinstance(function, dict): + tool_name = function.get("name", "unknown_tool") + else: + tool_name = getattr(function, "name", "unknown_tool") + else: + function = getattr(tool_call, "function", None) + if function: + tool_name = getattr(function, "name", "unknown_tool") + break + + dummy_tool_result: ChatCompletionToolMessage = { + "role": "tool", + "tool_call_id": tool_call_id, + "content": f"[System: Tool execution skipped/interrupted by user. No result provided for tool '{tool_name}'.]", + } + result_messages.append(dummy_tool_result) + + # Return the messages and the number of original messages to skip + return (result_messages, len(actual_tool_results)) + + return ([current_message], 0) + + +def _is_orphaned_tool_result( + current_message: AllMessageValues, + sanitized_messages: List[AllMessageValues], +) -> bool: + """ + Case B: Orphaned tool_result (unexpected result) + - Check if a tool message references a tool_call_id that doesn't exist in the previous + assistant message. + + Returns: + True if this is an orphaned tool result that should be removed, False otherwise + """ + if current_message.get("role") not in ["tool", "function"]: + return False + + tool_call_id = current_message.get("tool_call_id") + + if not tool_call_id: + return False + + # Look back to find the most recent assistant message with tool_calls + found_matching_tool_call = False + + for j in range(len(sanitized_messages) - 1, -1, -1): + prev_msg = sanitized_messages[j] + if prev_msg.get("role") == "assistant": + tool_calls = prev_msg.get("tool_calls") + if tool_calls: + for tool_call in cast(list, tool_calls): + tc_id = None + if isinstance(tool_call, dict): + tc_id = tool_call.get("id") + else: + tc_id = getattr(tool_call, "id", None) + + if tc_id == tool_call_id: + found_matching_tool_call = True + break + + break + + if not found_matching_tool_call: + verbose_logger.debug( + "_is_orphaned_tool_result: Found orphaned tool result with redacted tool_call_id" + ) + return True + + return False + + +def sanitize_messages_for_tool_calling( + messages: List[AllMessageValues], +) -> List[AllMessageValues]: + """ + Sanitize messages for tool calling to handle common issues when modify_params=True: + + Case A: Missing tool_result for tool_use (orphaned tool calls) + - If an assistant message has tool_calls but no corresponding tool result follows, + add a dummy tool result message indicating the user did not provide the result. + + Case B: Orphaned tool_result (unexpected result) + - If a tool message references a tool_call_id that doesn't exist in the previous + assistant message, remove that tool message. + + Case C: Empty text content + - Replace empty or whitespace-only text content with a placeholder message. + + This function operates on OpenAI format messages before they are converted to + provider-specific formats. + """ + if not litellm.modify_params: + return messages + + sanitized_messages: List[AllMessageValues] = [] + i = 0 + + while i < len(messages): + current_message = messages[i] + + # Case C: Sanitize empty text content + current_message = _sanitize_empty_text_content(current_message) + + # Case A: Check if assistant message has tool_calls without following tool results + if current_message.get("role") == "assistant": + result_messages, messages_consumed = _add_missing_tool_results(current_message, messages, i) + + # If dummy tool results were added, extend sanitized_messages and skip consumed messages + if len(result_messages) > 1: + sanitized_messages.extend(result_messages) + # Skip the assistant message and any actual tool results that were included + i += 1 + messages_consumed + continue + + # Case B: Check for orphaned tool results + if _is_orphaned_tool_result(current_message, sanitized_messages): + i += 1 + continue # Skip this orphaned tool result + + # Add the message to sanitized list + sanitized_messages.append(current_message) + i += 1 + + return sanitized_messages + + def anthropic_messages_pt( # noqa: PLR0915 messages: List[AllMessageValues], model: str, @@ -2037,6 +2266,9 @@ def anthropic_messages_pt( # noqa: PLR0915 5. System messages are a separate param to the Messages API 6. Ensure we only accept role, content. (message.name is not supported) """ + # Sanitize messages for tool calling issues when modify_params=True + messages = sanitize_messages_for_tool_calling(messages) + # add role=tool support to allow function call result/error submission user_message_types = {"user", "tool", "function"} # reformat messages to ensure user/assistant are alternating, if there's either 2 consecutive 'user' messages or 2 consecutive 'assistant' message, merge them. diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 85a4790a9b9..a5f8fe22a2c 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -172,8 +172,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): @staticmethod def _is_claude_opus_4_6(model: str) -> bool: - """Check if the model is Claude Opus 4.5.""" - return "opus-4-6" in model.lower() or "opus_4_6" in model.lower() + """Check if the model is Claude Opus 4.5 or Sonnet 4.6.""" + return "opus-4-6" in model.lower() or "opus_4_6" in model.lower() or "sonnet-4-6" in model.lower() or "sonnet_4_6" in model.lower() or "sonnet-4.6" in model.lower() def get_supported_openai_params(self, model: str): params = [ @@ -881,6 +881,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): "opus-4-5", "opus-4.6", "opus-4-6", + "sonnet-4.6", + "sonnet-4-6", + "sonnet_4.6", + "sonnet_4_6", } ): _output_format = ( diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index c665e084261..0cceddd9acf 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -22,6 +22,15 @@ from litellm.types.llms.anthropic import ( from litellm.types.llms.openai import AllMessageValues +def is_anthropic_oauth_key(value: Optional[str]) -> bool: + """Check if a value contains an Anthropic OAuth token (sk-ant-oat*).""" + if value is None: + return False + # Handle both raw token and "Bearer " format + if value.startswith("Bearer "): + value = value[7:] + return value.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX) + def optionally_handle_anthropic_oauth( headers: dict, api_key: Optional[str] ) -> tuple[dict, Optional[str]]: diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index efbac13735c..8b21569546e 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -299,6 +299,26 @@ class LiteLLMAnthropicMessagesAdapter: """ return ["messages", "metadata", "system", "tool_choice", "tools", "thinking", "output_format"] + def _is_web_search_tool(self, tool: Dict[str, Any]) -> bool: + """ + Check if a tool is an Anthropic web search tool. + + Anthropic web search tools have: + - type starting with "web_search" (e.g., "web_search_20260209") + - name = "web_search" + + Args: + tool: Tool definition dict + + Returns: + True if this is a web search tool + """ + tool_type = tool.get("type", "") + tool_name = tool.get("name", "") + return ( + isinstance(tool_type, str) and tool_type.startswith("web_search") + ) or tool_name == "web_search" + def translate_anthropic_messages_to_openai( # noqa: PLR0915 self, messages: List[ @@ -872,10 +892,25 @@ class LiteLLMAnthropicMessagesAdapter: if "tools" in anthropic_message_request: tools = anthropic_message_request["tools"] if tools: - new_kwargs["tools"], tool_name_mapping = self.translate_anthropic_tools_to_openai( - tools=cast(List[AllAnthropicToolsValues], tools), - model=new_kwargs.get("model"), - ) + # Separate web search tools from regular tools + web_search_tools = [] + regular_tools = [] + for tool in tools: + if self._is_web_search_tool(cast(Dict[str, Any], tool)): + web_search_tools.append(tool) + else: + regular_tools.append(tool) + + # If web search tools are present, add web_search_options parameter + if web_search_tools: + new_kwargs["web_search_options"] = {} # type: ignore + + # Only translate regular tools (non-web-search) + if regular_tools: + new_kwargs["tools"], tool_name_mapping = self.translate_anthropic_tools_to_openai( + tools=cast(List[AllAnthropicToolsValues], regular_tools), + model=new_kwargs.get("model"), + ) ## CONVERT THINKING if "thinking" in anthropic_message_request: diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 304c707fa0b..dfaddb3c2b1 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -384,6 +384,14 @@ class BaseAWSLLM: model_id = BaseAWSLLM._get_model_id_from_model_with_spec( model_id, spec="moonshot" ) + elif "nova-2/" in model_id: + model_id = BaseAWSLLM._get_model_id_from_model_with_spec( + model_id, spec="nova-2" + ) + elif "nova/" in model_id: + model_id = BaseAWSLLM._get_model_id_from_model_with_spec( + model_id, spec="nova" + ) return model_id @staticmethod diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 25af852e09c..60a93b169c8 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -272,7 +272,18 @@ class BedrockConverseLLM(BaseAWSLLM): if unencoded_model_id is not None: modelId = self.encode_model_id(model_id=unencoded_model_id) else: - modelId = self.encode_model_id(model_id=model) + # Strip nova spec prefixes before encoding model ID for API URL + _model_for_id = model + _stripped = _model_for_id + for rp in ["bedrock/converse/", "bedrock/", "converse/"]: + if _stripped.startswith(rp): + _stripped = _stripped[len(rp):] + break + for _nova_prefix in ["nova-2/", "nova/"]: + if _stripped.startswith(_nova_prefix): + _model_for_id = _model_for_id.replace(_nova_prefix, "", 1) + break + modelId = self.encode_model_id(model_id=_model_for_id) fake_stream = litellm.AmazonConverseConfig().should_fake_stream( fake_stream=fake_stream, diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 5faae07e2b9..daac3e6a008 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -3,6 +3,7 @@ Translating between OpenAI's `/chat/completion` format and Amazon's `/converse` """ import copy +import json import time import types from typing import List, Literal, Optional, Tuple, Union, cast, overload @@ -85,9 +86,37 @@ BEDROCK_COMPUTER_USE_TOOLS = [ UNSUPPORTED_BEDROCK_CONVERSE_BETA_PATTERNS = [ "advanced-tool-use", # Bedrock Converse doesn't support advanced-tool-use beta headers "prompt-caching", # Prompt caching not supported in Converse API - "compact-2026-01-12", # The compact beta feature is not currently supported on the Converse and ConverseStream APIs + "compact-2026-01-12", # The compact beta feature is not currently supported on the Converse and ConverseStream APIs ] +# Models that support Bedrock's native structured outputs API (outputConfig.textFormat) +# Uses substring matching against the Bedrock model ID +# Ref: https://docs.aws.amazon.com/bedrock/latest/userguide/structured-output.html +BEDROCK_NATIVE_STRUCTURED_OUTPUT_MODELS = { + # Anthropic Claude 4.5+ + "claude-haiku-4-5", + "claude-sonnet-4-5", + "claude-opus-4-5", + "claude-opus-4-6", + # Qwen3 + "qwen3", + # DeepSeek + "deepseek-v3.1", + # Gemma 3 + "gemma-3", + # MiniMax + "minimax-m2", + # Mistral (magistral-small excluded: broken constrained decoding on Bedrock) + "ministral", + "mistral-large-3", + "voxtral", + # Moonshot + "kimi-k2", + # NVIDIA + "nemotron-nano", + # OpenAI (gpt-oss excluded: broken constrained decoding, works via tool-call fallback) +} + class AmazonConverseConfig(BaseConfig): """ @@ -270,45 +299,56 @@ class AmazonConverseConfig(BaseConfig): llm_provider="bedrock", ) - def _is_nova_lite_2_model(self, model: str) -> bool: + def _is_nova_2_model(self, model: str) -> bool: """ - Check if the model is a Nova Lite 2 model that supports reasoningConfig. + Check if the model is a Nova 2 model that supports reasoningConfig. - Nova Lite 2 models use a different reasoning configuration structure compared to + Nova 2 models use a different reasoning configuration structure compared to Anthropic's thinking parameter and GPT-OSS's reasoning_effort parameter. Supported models: - amazon.nova-2-lite-v1:0 + - amazon.nova-2-pro-preview-20251202-v1:0 - us.amazon.nova-2-lite-v1:0 - eu.amazon.nova-2-lite-v1:0 - apac.amazon.nova-2-lite-v1:0 + - (and other regional variants) Args: model: The model identifier Returns: - True if the model is a Nova Lite 2 model, False otherwise + True if the model is a Nova 2 model, False otherwise Examples: >>> config = AmazonConverseConfig() - >>> config._is_nova_lite_2_model("amazon.nova-2-lite-v1:0") + >>> config._is_nova_2_model("amazon.nova-2-lite-v1:0") True - >>> config._is_nova_lite_2_model("us.amazon.nova-2-lite-v1:0") + >>> config._is_nova_2_model("us.amazon.nova-2-lite-v1:0") True - >>> config._is_nova_lite_2_model("amazon.nova-pro-1-5-v1:0") + >>> config._is_nova_2_model("us.amazon.nova-2-pro-preview-20251202-v1:0") + True + >>> config._is_nova_2_model("amazon.nova-pro-1-5-v1:0") False - >>> config._is_nova_lite_2_model("amazon.nova-pro-v1:0") + >>> config._is_nova_2_model("amazon.nova-pro-v1:0") False """ - # Remove regional prefix if present (us., eu., apac.) + # Remove provider routing prefix if present (bedrock/converse/, bedrock/, converse/) model_without_region = model - for prefix in ["us.", "eu.", "apac."]: - if model.startswith(prefix): - model_without_region = model[len(prefix) :] + for routing_prefix in ["bedrock/converse/", "bedrock/", "converse/"]: + if model_without_region.startswith(routing_prefix): + model_without_region = model_without_region[len(routing_prefix) :] break - # Check if the model is specifically Nova Lite 2 - return "nova-2-lite" in model_without_region + # Remove regional prefix if present (us., eu., apac.) + for prefix in ["us.", "eu.", "apac."]: + if model_without_region.startswith(prefix): + model_without_region = model_without_region[len(prefix) :] + break + + # Check if the model is a Nova 2 model (matches nova-2-lite, nova-2-pro, etc.) + # Also check for nova-2/ spec prefix for imported models + return model_without_region.startswith("amazon.nova-2-") or model_without_region.startswith("nova-2/") def _map_web_search_options( self, web_search_options: dict, model: str @@ -396,7 +436,7 @@ class AmazonConverseConfig(BaseConfig): Different model families handle reasoning effort differently: - GPT-OSS models: Keep reasoning_effort as-is (passed to additionalModelRequestFields) - - Nova Lite 2 models: Transform to reasoningConfig structure + - Nova 2 models: Transform to reasoningConfig structure - Other models (Anthropic, etc.): Convert to thinking parameter Args: @@ -425,8 +465,8 @@ class AmazonConverseConfig(BaseConfig): # GPT-OSS models: keep reasoning_effort as-is # It will be passed through to additionalModelRequestFields optional_params["reasoning_effort"] = reasoning_effort - elif self._is_nova_lite_2_model(model): - # Nova Lite 2 models: transform to reasoningConfig + elif self._is_nova_2_model(model): + # Nova 2 models: transform to reasoningConfig reasoning_config = self._transform_reasoning_effort_to_reasoning_config( reasoning_effort ) @@ -480,6 +520,9 @@ class AmazonConverseConfig(BaseConfig): supported_params.append("tool_choice") supported_params.append("thinking") supported_params.append("reasoning_effort") + # For nova imported models, also add web_search_options + if "nova" in model.lower(): + supported_params.append("web_search_options") return supported_params ## Filter out 'cross-region' from model name @@ -514,8 +557,8 @@ class AmazonConverseConfig(BaseConfig): if "gpt-oss" in model: supported_params.append("reasoning_effort") - elif self._is_nova_lite_2_model(model): - # Nova Lite 2 models support reasoning_effort (transformed to reasoningConfig) + elif self._is_nova_2_model(model): + # Nova 2 models support reasoning_effort (transformed to reasoningConfig) # These models use a different reasoning structure than Anthropic's thinking parameter supported_params.append("reasoning_effort") elif ( @@ -714,6 +757,100 @@ class AmazonConverseConfig(BaseConfig): ) return _tool + @staticmethod + def _supports_native_structured_outputs(model: str) -> bool: + """Check if the Bedrock model supports native structured outputs (outputConfig.textFormat).""" + return any( + substring in model + for substring in BEDROCK_NATIVE_STRUCTURED_OUTPUT_MODELS + ) + + @staticmethod + def _add_additional_properties_to_schema(schema: dict) -> dict: + """ + Recursively ensure all object types in a JSON schema have + ``"additionalProperties": false``. + + Bedrock's native structured-outputs API requires this field to be + explicitly set on every object node, otherwise it returns a + validation error. + """ + if not isinstance(schema, dict): + return schema + + result = dict(schema) + + if result.get("type") == "object" and "additionalProperties" not in result: + result["additionalProperties"] = False + + # Recurse into nested schemas + if "properties" in result and isinstance(result["properties"], dict): + result["properties"] = { + k: AmazonConverseConfig._add_additional_properties_to_schema(v) + for k, v in result["properties"].items() + } + if "items" in result and isinstance(result["items"], dict): + result["items"] = AmazonConverseConfig._add_additional_properties_to_schema( + result["items"] + ) + for defs_key in ("$defs", "definitions"): + if defs_key in result and isinstance(result[defs_key], dict): + result[defs_key] = { + k: AmazonConverseConfig._add_additional_properties_to_schema(v) + for k, v in result[defs_key].items() + } + for key in ("anyOf", "allOf", "oneOf"): + if key in result and isinstance(result[key], list): + result[key] = [ + AmazonConverseConfig._add_additional_properties_to_schema(item) + for item in result[key] + ] + + return result + + @staticmethod + def _create_output_config_for_response_format( + json_schema: Optional[dict] = None, + name: Optional[str] = None, + description: Optional[str] = None, + ) -> "OutputConfigBlock": + """ + Build an outputConfig block for Bedrock's native structured outputs API. + + The Converse API expects: + { + "outputConfig": { + "textFormat": { + "type": "json_schema", + "structure": { + "jsonSchema": { + "schema": "", + "name": "optional", + "description": "optional" + } + } + } + } + } + """ + if json_schema is not None: + json_schema = AmazonConverseConfig._add_additional_properties_to_schema( + json_schema + ) + schema_str = json.dumps(json_schema) if json_schema is not None else "{}" + json_schema_def: JsonSchemaDefinition = {"schema": schema_str} + if name is not None: + json_schema_def["name"] = name + if description is not None: + json_schema_def["description"] = description + + return OutputConfigBlock( + textFormat=OutputFormat( + type="json_schema", + structure=OutputFormatStructure(jsonSchema=json_schema_def), + ) + ) + def _apply_tool_call_transformation( self, tools: List[OpenAIChatCompletionToolParam], @@ -806,8 +943,8 @@ class AmazonConverseConfig(BaseConfig): ) # Only update thinking tokens for non-GPT-OSS models and non-Nova-Lite-2 models - # Nova Lite 2 handles token budgeting differently through reasoningConfig - if "gpt-oss" not in model and not self._is_nova_lite_2_model(model): + # Nova 2 handles token budgeting differently through reasoningConfig + if "gpt-oss" not in model and not self._is_nova_2_model(model): self.update_optional_params_with_thinking_tokens( non_default_params=non_default_params, optional_params=optional_params ) @@ -843,45 +980,53 @@ class AmazonConverseConfig(BaseConfig): return optional_params json_schema: Optional[dict] = None + name: Optional[str] = None description: Optional[str] = None if "response_schema" in value: json_schema = value["response_schema"] elif "json_schema" in value: json_schema = value["json_schema"]["schema"] + name = value["json_schema"].get("name") description = value["json_schema"].get("description") if "type" in value and value["type"] == "text": return optional_params - """ - Follow similar approach to anthropic - translate to a single tool call. - - When using tools in this way: - https://docs.anthropic.com/en/docs/build-with-claude/tool-use#json-mode - - You usually want to provide a single tool - - You should set tool_choice (see Forcing tool use) to instruct the model to explicitly use that tool - - Remember that the model will pass the input to the tool, so the name of the tool and description should be from the model’s perspective. - """ - _tool = self._create_json_tool_call_for_response_format( - json_schema=json_schema, - description=description, - ) - optional_params = self._add_tools_to_optional_params( - optional_params=optional_params, tools=[_tool] - ) - - if ( - litellm.utils.supports_tool_choice( - model=model, custom_llm_provider=self.custom_llm_provider + if self._supports_native_structured_outputs(model) and json_schema is not None: + # Use Bedrock's native structured outputs API (outputConfig.textFormat) + # No synthetic tool injection, no fake_stream needed. + # Requires an explicit schema — json_object with no schema falls through + # to the tool-call path below. + output_config = self._create_output_config_for_response_format( + json_schema=json_schema, + name=name, + description=description, ) - and not is_thinking_enabled - ): - optional_params["tool_choice"] = ToolChoiceValuesBlock( - tool=SpecificToolChoiceBlock(name=RESPONSE_FORMAT_TOOL_NAME) + optional_params["outputConfig"] = output_config + else: + # Fallback: translate to a synthetic tool call + # https://docs.anthropic.com/en/docs/build-with-claude/tool-use#json-mode + _tool = self._create_json_tool_call_for_response_format( + json_schema=json_schema, + description=description, ) + optional_params = self._add_tools_to_optional_params( + optional_params=optional_params, tools=[_tool] + ) + + if ( + litellm.utils.supports_tool_choice( + model=model, custom_llm_provider=self.custom_llm_provider + ) + and not is_thinking_enabled + ): + optional_params["tool_choice"] = ToolChoiceValuesBlock( + tool=SpecificToolChoiceBlock(name=RESPONSE_FORMAT_TOOL_NAME) + ) + if non_default_params.get("stream", False) is True: + optional_params["fake_stream"] = True + optional_params["json_mode"] = True - if non_default_params.get("stream", False) is True: - optional_params["fake_stream"] = True - return optional_params def update_optional_params_with_thinking_tokens( @@ -1024,7 +1169,7 @@ class AmazonConverseConfig(BaseConfig): def _prepare_request_params( self, optional_params: dict, model: str - ) -> Tuple[dict, dict, dict]: + ) -> Tuple[dict, dict, dict, Optional[OutputConfigBlock]]: """Prepare and separate request parameters.""" # Filter out exception objects before deepcopy to prevent deepcopy failures # Exceptions should not be stored in optional_params (this is a defensive fix) @@ -1047,6 +1192,8 @@ class AmazonConverseConfig(BaseConfig): if request_metadata is not None: self._validate_request_metadata(request_metadata) + output_config: Optional[OutputConfigBlock] = inference_params.pop("outputConfig", None) + # keep supported params in 'inference_params', and set all model-specific params in 'additional_request_params' additional_request_params = { k: v for k, v in inference_params.items() if k not in total_supported_params @@ -1071,7 +1218,12 @@ class AmazonConverseConfig(BaseConfig): additional_request_params ) - return inference_params, additional_request_params, request_metadata + return ( + inference_params, + additional_request_params, + request_metadata, + output_config, + ) def _process_tools_and_beta( self, @@ -1125,22 +1277,44 @@ class AmazonConverseConfig(BaseConfig): # "computer-use-2025-01-24" for Claude Sonnet 4.5, Haiku 4.5, Opus 4.1, Sonnet 4, Opus 4, and Sonnet 3.7 # "computer-use-2024-10-22" for older models model_lower = model.lower() - if "opus-4.6" in model_lower or "opus_4.6" in model_lower or "opus-4-6" in model_lower or "opus_4_6" in model_lower: + if "opus-4.6" in model_lower or "opus_4.6" in model_lower or "opus-4-6" in model_lower or "opus_4_6" in model_lower or "sonnet-4.6" in model_lower or "sonnet_4.6" in model_lower or "sonnet-4-6" in model_lower or "sonnet_4_6" in model_lower: computer_use_header = "computer-use-2025-11-24" - elif "opus-4.5" in model_lower or "opus_4.5" in model_lower or "opus-4-5" in model_lower or "opus_4_5" in model_lower: + elif ( + "opus-4.5" in model_lower + or "opus_4.5" in model_lower + or "opus-4-5" in model_lower + or "opus_4_5" in model_lower + ): computer_use_header = "computer-use-2025-11-24" - elif any(pattern in model_lower for pattern in [ - "sonnet-4.5", "sonnet_4.5", "sonnet-4-5", "sonnet_4_5", - "haiku-4.5", "haiku_4.5", "haiku-4-5", "haiku_4_5", - "opus-4.1", "opus_4.1", "opus-4-1", "opus_4_1", - "sonnet-4", "sonnet_4", - "opus-4", "opus_4", - "sonnet-3.7", "sonnet_3.7", "sonnet-3-7", "sonnet_3_7" - ]): + elif any( + pattern in model_lower + for pattern in [ + "sonnet-4.5", + "sonnet_4.5", + "sonnet-4-5", + "sonnet_4_5", + "haiku-4.5", + "haiku_4.5", + "haiku-4-5", + "haiku_4_5", + "opus-4.1", + "opus_4.1", + "opus-4-1", + "opus_4_1", + "sonnet-4", + "sonnet_4", + "opus-4", + "opus_4", + "sonnet-3.7", + "sonnet_3.7", + "sonnet-3-7", + "sonnet_3_7", + ] + ): computer_use_header = "computer-use-2025-01-24" else: computer_use_header = "computer-use-2024-10-22" - + anthropic_beta_list.append(computer_use_header) # Transform computer use tools to proper Bedrock format transformed_computer_tools = self._transform_computer_use_tools( @@ -1214,6 +1388,7 @@ class AmazonConverseConfig(BaseConfig): inference_params, additional_request_params, request_metadata, + output_config, ) = self._prepare_request_params(optional_params, model) original_tools = inference_params.pop("tools", []) @@ -1256,6 +1431,9 @@ class AmazonConverseConfig(BaseConfig): if request_metadata is not None: data["requestMetadata"] = request_metadata + if output_config is not None: + data["outputConfig"] = output_config + return data async def _async_transform_request( @@ -1504,9 +1682,7 @@ class AmazonConverseConfig(BaseConfig): return message, returned_finish_reason - def _translate_message_content( - self, content_blocks: List[ContentBlock] - ) -> Tuple[ + def _translate_message_content(self, content_blocks: List[ContentBlock]) -> Tuple[ str, List[ChatCompletionToolCallChunk], Optional[List[BedrockConverseReasoningContentBlock]], @@ -1523,9 +1699,9 @@ class AmazonConverseConfig(BaseConfig): """ content_str = "" tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[ - List[BedrockConverseReasoningContentBlock] - ] = None + reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( + None + ) citationsContentBlocks: Optional[List[CitationsContentBlock]] = None for idx, content in enumerate(content_blocks): """ @@ -1652,9 +1828,9 @@ class AmazonConverseConfig(BaseConfig): chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"} content_str = "" tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[ - List[BedrockConverseReasoningContentBlock] - ] = None + reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( + None + ) citationsContentBlocks: Optional[List[CitationsContentBlock]] = None if message is not None: @@ -1673,17 +1849,17 @@ class AmazonConverseConfig(BaseConfig): provider_specific_fields["citationsContent"] = citationsContentBlocks if provider_specific_fields: - chat_completion_message[ - "provider_specific_fields" - ] = provider_specific_fields + chat_completion_message["provider_specific_fields"] = ( + provider_specific_fields + ) if reasoningContentBlocks is not None: - chat_completion_message[ - "reasoning_content" - ] = self._transform_reasoning_content(reasoningContentBlocks) - chat_completion_message[ - "thinking_blocks" - ] = self._transform_thinking_blocks(reasoningContentBlocks) + chat_completion_message["reasoning_content"] = ( + self._transform_reasoning_content(reasoningContentBlocks) + ) + chat_completion_message["thinking_blocks"] = ( + self._transform_thinking_blocks(reasoningContentBlocks) + ) chat_completion_message["content"] = content_str if ( json_mode is True @@ -1696,8 +1872,6 @@ class AmazonConverseConfig(BaseConfig): ) json_mode_content_str: Optional[str] = tools[0]["function"].get("arguments") if json_mode_content_str is not None: - import json - # Bedrock returns the response wrapped in a "properties" object # We need to extract the actual content from this wrapper try: @@ -1716,7 +1890,7 @@ class AmazonConverseConfig(BaseConfig): pass chat_completion_message["content"] = json_mode_content_str - else: + elif tools: chat_completion_message["tool_calls"] = tools ## CALCULATING USAGE - bedrock returns usage in the headers diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 4c87f6fa994..b779c892c67 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -404,7 +404,7 @@ def extract_model_name_from_bedrock_arn(model: str) -> str: def strip_bedrock_routing_prefix(model: str) -> str: """Strip LiteLLM routing prefixes from model name.""" - for prefix in ["bedrock/", "converse/", "invoke/", "openai/"]: + for prefix in ["bedrock/", "converse/", "invoke/", "openai/", "nova-2/", "nova/"]: if model.startswith(prefix): model = model.split("/", 1)[1] return model @@ -427,7 +427,20 @@ def get_bedrock_base_model(model: str) -> str: - "us.meta.llama3-2-11b-instruct-v1:0" -> "meta.llama3-2-11b-instruct-v1" - "bedrock/converse/model" -> "model" - "anthropic.claude-3-5-sonnet-20241022-v2:0:51k" -> "anthropic.claude-3-5-sonnet-20241022-v2:0" + - "bedrock/nova-2/arn:aws:..." -> "amazon.nova-2-custom" + - "bedrock/nova/arn:aws:..." -> "amazon.nova-custom" """ + # Detect nova spec prefixes before stripping them + stripped = model + for rp in ["bedrock/converse/", "bedrock/", "converse/"]: + if stripped.startswith(rp): + stripped = stripped[len(rp):] + break + if stripped.startswith("nova-2/"): + return "amazon.nova-2-custom" + elif stripped.startswith("nova/"): + return "amazon.nova-custom" + model = strip_bedrock_routing_prefix(model) model = extract_model_name_from_bedrock_arn(model) model = strip_bedrock_throughput_suffix(model) @@ -465,6 +478,14 @@ def is_claude_4_5_on_bedrock(model: str) -> bool: "opus_4.5", "opus-4-5", "opus_4_5", + "sonnet-4.6", + "sonnet_4.6", + "sonnet-4-6", + "sonnet_4_6", + "opus-4.6", + "opus_4.6", + "opus-4-6", + "opus_4_6", ] return any(pattern in model_lower for pattern in claude_4_5_patterns) @@ -594,6 +615,11 @@ class BedrockModelInfo(BaseLLMModelInfo): if prefix in model: return route_type + # Check for nova spec prefixes (nova/ and nova-2/) + _model_after_bedrock = model.replace("bedrock/", "", 1) + if _model_after_bedrock.startswith("nova-2/") or _model_after_bedrock.startswith("nova/"): + return "converse" + base_model = BedrockModelInfo.get_base_model(model) alt_model = BedrockModelInfo.get_non_litellm_routing_model_name(model=model) if ( diff --git a/litellm/llms/bedrock/cost_calculation.py b/litellm/llms/bedrock/cost_calculation.py index b20350d7325..ac99d4e36e7 100644 --- a/litellm/llms/bedrock/cost_calculation.py +++ b/litellm/llms/bedrock/cost_calculation.py @@ -3,7 +3,7 @@ Helper util for handling bedrock-specific cost calculation - e.g.: prompt caching """ -from typing import TYPE_CHECKING, Tuple +from typing import TYPE_CHECKING, Optional, Tuple from litellm.litellm_core_utils.llm_cost_calc.utils import generic_cost_per_token @@ -11,12 +11,17 @@ if TYPE_CHECKING: from litellm.types.utils import Usage -def cost_per_token(model: str, usage: "Usage") -> Tuple[float, float]: +def cost_per_token( + model: str, usage: "Usage", service_tier: Optional[str] = None +) -> Tuple[float, float]: """ Calculates the cost per token for a given model, prompt tokens, and completion tokens. Follows the same logic as Anthropic's cost per token calculation. """ return generic_cost_per_token( - model=model, usage=usage, custom_llm_provider="bedrock" - ) \ No newline at end of file + model=model, + usage=usage, + custom_llm_provider="bedrock", + service_tier=service_tier, + ) diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index 477fa3316d1..03885ff2080 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -180,6 +180,14 @@ class AmazonAnthropicClaudeMessagesConfig( "opus_4", # Opus 4 "sonnet-4", "sonnet_4", # Sonnet 4 + "sonnet-4.6", + "sonnet_4.6", + "sonnet-4-6", + "sonnet_4_6", + "opus-4.6", + "opus_4.6", + "opus-4-6", + "opus_4_6", ] return any(pattern in model_lower for pattern in supported_patterns) @@ -251,6 +259,11 @@ class AmazonAnthropicClaudeMessagesConfig( "opus_4.6", "opus-4-6", "opus_4_6", + #sonnet 4.6 + "sonnet-4.6", + "sonnet_4.6", + "sonnet-4-6", + "sonnet_4_6", ] return any(pattern in model_lower for pattern in supported_patterns) @@ -285,7 +298,7 @@ class AmazonAnthropicClaudeMessagesConfig( programmatic_tool_calling_used or input_examples_used ): beta_set.discard(ANTHROPIC_TOOL_SEARCH_BETA_HEADER) - if "opus-4" in model.lower() or "opus_4" in model.lower(): + if self._supports_tool_search_on_bedrock(model): beta_set.add("tool-search-tool-2025-10-19") def _convert_output_format_to_inline_schema( @@ -420,10 +433,8 @@ class AmazonAnthropicClaudeMessagesConfig( beta_set=beta_set, ) - # --- Custom logic: if tool-search-tool-2025-10-19 is present, add tool-examples-2025-10-29 --- if "tool-search-tool-2025-10-19" in beta_set: beta_set.add("tool-examples-2025-10-29") - # ------------------------------------------------------------------------------ if beta_set: anthropic_messages_request["anthropic_beta"] = list(beta_set) diff --git a/litellm/llms/duckduckgo/search/__init__.py b/litellm/llms/duckduckgo/search/__init__.py new file mode 100644 index 00000000000..c0019637838 --- /dev/null +++ b/litellm/llms/duckduckgo/search/__init__.py @@ -0,0 +1,6 @@ +""" +DuckDuckGo Search API module. +""" +from litellm.llms.duckduckgo.search.transformation import DuckDuckGoSearchConfig + +__all__ = ["DuckDuckGoSearchConfig"] diff --git a/litellm/llms/duckduckgo/search/transformation.py b/litellm/llms/duckduckgo/search/transformation.py new file mode 100644 index 00000000000..509d69041fb --- /dev/null +++ b/litellm/llms/duckduckgo/search/transformation.py @@ -0,0 +1,252 @@ +""" +Calls DuckDuckGo's Instant Answer API to search the web. + +DuckDuckGo API Reference: https://duckduckgo.com/api +""" +from typing import Dict, List, Literal, Optional, TypedDict, Union +from urllib.parse import urlencode + +import httpx + +from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +from litellm.llms.base_llm.search.transformation import ( + BaseSearchConfig, + SearchResponse, + SearchResult, +) +from litellm.secret_managers.main import get_secret_str + + +class _DuckDuckGoSearchRequestRequired(TypedDict): + """Required fields for DuckDuckGo Search API request.""" + q: str # Required - search query + + +class DuckDuckGoSearchRequest(_DuckDuckGoSearchRequestRequired, total=False): + """ + DuckDuckGo Instant Answer API request format. + Based on: https://duckduckgo.com/api + """ + format: str # Optional - output format ('json', 'xml'), default 'json' + pretty: int # Optional - pretty print (0 or 1), default 1 + no_redirect: int # Optional - skip HTTP redirects (0 or 1), default 0 + no_html: int # Optional - remove HTML from text (0 or 1), default 0 + skip_disambig: int # Optional - skip disambiguation results (0 or 1), default 0 + + +class DuckDuckGoSearchConfig(BaseSearchConfig): + DUCKDUCKGO_API_BASE = "https://api.duckduckgo.com" + + @staticmethod + def ui_friendly_name() -> str: + return "DuckDuckGo" + + def get_http_method(self) -> Literal["GET", "POST"]: + """ + Get HTTP method for search requests. + DuckDuckGo Instant Answer API uses GET requests. + + Returns: + HTTP method 'GET' + """ + return "GET" + + def validate_environment( + self, + headers: Dict, + api_key: Optional[str] = None, + api_base: Optional[str] = None, + **kwargs, + ) -> Dict: + """ + Validate environment and return headers. + DuckDuckGo Instant Answer API does not require authentication. + """ + # DuckDuckGo API is free and doesn't require API key + headers["Content-Type"] = "application/json" + return headers + + def get_complete_url( + self, + api_base: Optional[str], + optional_params: dict, + data: Optional[Union[Dict, List[Dict]]] = None, + **kwargs, + ) -> str: + """ + Get complete URL for Search endpoint. + DuckDuckGo uses query parameters, so we construct the URL with the query. + """ + api_base = api_base or get_secret_str("DUCKDUCKGO_API_BASE") or self.DUCKDUCKGO_API_BASE + + # Build query parameters from the transformed request body + if data and isinstance(data, dict) and "_duckduckgo_params" in data: + params = data["_duckduckgo_params"] + query_string = urlencode(params, doseq=True) + return f"{api_base}/?{query_string}" + + return api_base + + + def transform_search_request( + self, + query: Union[str, List[str]], + optional_params: dict, + **kwargs, + ) -> Dict: + """ + Transform Search request to DuckDuckGo API format. + + Args: + query: Search query (string or list of strings). DuckDuckGo only supports single string queries. + optional_params: Optional parameters for the request + - max_results: Maximum number of search results (DuckDuckGo API doesn't directly support this, used for filtering) + - format: Output format ('json', 'xml') + - pretty: Pretty print (0 or 1) + - no_redirect: Skip HTTP redirects (0 or 1) + - no_html: Remove HTML from text (0 or 1) + - skip_disambig: Skip disambiguation results (0 or 1) + + Returns: + Dict with typed request data following DuckDuckGoSearchRequest spec + """ + if isinstance(query, list): + # DuckDuckGo only supports single string queries + query = " ".join(query) + + request_data: DuckDuckGoSearchRequest = { + "q": query, + "format": "json", # Always use JSON format + } + + # Convert to dict before dynamic key assignments + result_data = dict(request_data) + + if "max_results" in optional_params: + result_data["_max_results"] = optional_params["max_results"] + + # Pass through DuckDuckGo-specific parameters + ddg_params = ["pretty", "no_redirect", "no_html", "skip_disambig"] + for param in ddg_params: + if param in optional_params: + result_data[param] = optional_params[param] + + return { + "_duckduckgo_params": result_data, + } + + def transform_search_response( + self, + raw_response: httpx.Response, + logging_obj: LiteLLMLoggingObj, + **kwargs, + ) -> SearchResponse: + """ + Transform DuckDuckGo API response to LiteLLM unified SearchResponse format. + + DuckDuckGo → LiteLLM mappings: + - RelatedTopics[].Text → SearchResult.title + snippet + - RelatedTopics[].FirstURL → SearchResult.url + - RelatedTopics[].Text → SearchResult.snippet + - No date/last_updated fields in DuckDuckGo response (set to None) + + Args: + raw_response: Raw httpx response from DuckDuckGo API + logging_obj: Logging object for tracking + + Returns: + SearchResponse with standardized format + """ + response_json = raw_response.json() + + # Extract max_results from the request URL params + query_params = raw_response.request.url.params if raw_response.request else {} + max_results = None + if "_max_results" in query_params: + try: + max_results = int(query_params["_max_results"]) + except (ValueError, TypeError): + pass + + # Transform results to SearchResult objects + results = [] + + # DuckDuckGo can return results in different fields + # Priority: Abstract > Answer > RelatedTopics + + # Check if there's an Abstract with URL + if response_json.get("AbstractURL") and response_json.get("AbstractText"): + abstract_result = SearchResult( + title=response_json.get("Heading", ""), + url=response_json.get("AbstractURL", ""), + snippet=response_json.get("AbstractText", ""), + date=None, + last_updated=None, + ) + results.append(abstract_result) + + # Process RelatedTopics + related_topics = response_json.get("RelatedTopics", []) + for topic in related_topics: + # Stop if we've reached max_results + if max_results is not None and len(results) >= max_results: + break + + if isinstance(topic, dict): + # Check if it's a direct result + if "FirstURL" in topic and "Text" in topic: + text = topic.get("Text", "") + url = topic.get("FirstURL", "") + + # Try to split title and snippet + if " - " in text: + parts = text.split(" - ", 1) + title = parts[0] + snippet = parts[1] if len(parts) > 1 else text + else: + title = text[:50] + "..." if len(text) > 50 else text + snippet = text + + search_result = SearchResult( + title=title, + url=url, + snippet=snippet, + date=None, + last_updated=None, + ) + results.append(search_result) + + # Check if it contains nested topics + elif "Topics" in topic: + nested_topics = topic.get("Topics", []) + for nested_topic in nested_topics: + # Stop if we've reached max_results + if max_results is not None and len(results) >= max_results: + break + + if "FirstURL" in nested_topic and "Text" in nested_topic: + text = nested_topic.get("Text", "") + url = nested_topic.get("FirstURL", "") + + # Try to split title and snippet + if " - " in text: + parts = text.split(" - ", 1) + title = parts[0] + snippet = parts[1] if len(parts) > 1 else text + else: + title = text[:50] + "..." if len(text) > 50 else text + snippet = text + + search_result = SearchResult( + title=title, + url=url, + snippet=snippet, + date=None, + last_updated=None, + ) + results.append(search_result) + + return SearchResponse( + results=results, + object="search", + ) diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 5b9840d95b0..59f52e2b81c 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -770,14 +770,36 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): class OpenAIChatCompletionStreamingHandler(BaseModelResponseIterator): + def _map_reasoning_to_reasoning_content(self, choices: list) -> list: + """ + Map 'reasoning' field to 'reasoning_content' field in delta. + + Some OpenAI-compatible providers (e.g., GLM-5, hosted_vllm) return + delta.reasoning, but LiteLLM expects delta.reasoning_content. + + Args: + choices: List of choice objects from the streaming chunk + + Returns: + List of choices with reasoning field mapped to reasoning_content + """ + for choice in choices: + delta = choice.get("delta", {}) + if "reasoning" in delta: + delta["reasoning_content"] = delta.pop("reasoning") + return choices + def chunk_parser(self, chunk: dict) -> ModelResponseStream: try: + choices = chunk.get("choices", []) + choices = self._map_reasoning_to_reasoning_content(choices) + kwargs = { "id": chunk["id"], "object": "chat.completion.chunk", "created": chunk.get("created"), "model": chunk.get("model"), - "choices": chunk.get("choices", []), + "choices": choices, } if "usage" in chunk and chunk["usage"] is not None: kwargs["usage"] = chunk["usage"] diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index bef83b6d35e..daa82a46bdc 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -1072,7 +1072,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): elif param == "modalities" and isinstance(value, list): response_modalities = self.map_response_modalities(value) optional_params["responseModalities"] = response_modalities - elif param == "web_search_options" and value and isinstance(value, dict): + elif param == "web_search_options" and isinstance(value, dict): _tools = self._map_web_search_options(value) optional_params = self._add_tools_to_optional_params( optional_params, [_tools] diff --git a/litellm/llms/watsonx/__init__.py b/litellm/llms/watsonx/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/watsonx/chat/__init__.py b/litellm/llms/watsonx/chat/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/watsonx/completion/__init__.py b/litellm/llms/watsonx/completion/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/watsonx/embed/__init__.py b/litellm/llms/watsonx/embed/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/watsonx/rerank/__init__.py b/litellm/llms/watsonx/rerank/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/litellm/llms/watsonx/rerank/transformation.py b/litellm/llms/watsonx/rerank/transformation.py new file mode 100644 index 00000000000..7b4c2a07c3c --- /dev/null +++ b/litellm/llms/watsonx/rerank/transformation.py @@ -0,0 +1,204 @@ +""" +Transformation logic for IBM watsonx.ai's /ml/v1/text/rerank endpoint. + +Docs - https://cloud.ibm.com/apidocs/watsonx-ai#text-rerank +""" + +import uuid +from typing import Any, Dict, List, Optional, Union, cast + +import httpx + +from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj +from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig +from litellm.secret_managers.main import get_secret_str +from litellm.types.llms.watsonx import ( + WatsonXAIEndpoint, +) +from litellm.types.rerank import ( + RerankResponse, + RerankResponseMeta, + RerankTokens, +) + +from ..common_utils import IBMWatsonXMixin, _generate_watsonx_token, _get_api_params + + +class IBMWatsonXRerankConfig(IBMWatsonXMixin, BaseRerankConfig): + """ + IBM watsonx.ai Rerank API configuration + """ + + def get_complete_url( + self, + api_base: Optional[str], + model: str, + optional_params: Optional[dict] = None, + ) -> str: + base_url = self._get_base_url(api_base=api_base) + endpoint = WatsonXAIEndpoint.RERANK.value + + url = base_url.rstrip("/") + endpoint + + params = optional_params or {} + + complete_url = self._add_api_version_to_url(url=url, api_version=(params.get("api_version", None))) + return complete_url + + def get_supported_cohere_rerank_params(self, model: str) -> list: + return [ + "query", + "documents", + "top_n", + "return_documents", + "max_tokens_per_doc", + ] + + def validate_environment( # type: ignore[override] + self, + headers: dict, + model: str, + api_key: Optional[str] = None, + optional_params: Optional[dict] = None, + ) -> Dict: + optional_params = optional_params or {} + + default_headers = { + "Content-Type": "application/json", + "Accept": "application/json", + } + + if "Authorization" in headers: + return {**default_headers, **headers} + token = cast( + Optional[str], + optional_params.pop("token", None) or get_secret_str("WATSONX_TOKEN"), + ) + zen_api_key = cast( + Optional[str], + optional_params.pop("zen_api_key", None) or get_secret_str("WATSONX_ZENAPIKEY"), + ) + if token: + headers["Authorization"] = f"Bearer {token}" + elif zen_api_key: + headers["Authorization"] = f"ZenApiKey {zen_api_key}" + else: + token = _generate_watsonx_token(api_key=api_key, token=token) + # build auth headers + headers["Authorization"] = f"Bearer {token}" + return {**default_headers, **headers} + + def map_cohere_rerank_params( + self, + non_default_params: Optional[dict], + model: str, + drop_params: bool, + query: str, + documents: List[Union[str, Dict[str, Any]]], + custom_llm_provider: Optional[str] = None, + top_n: Optional[int] = None, + rank_fields: Optional[List[str]] = None, + return_documents: Optional[bool] = True, + max_chunks_per_doc: Optional[int] = None, + max_tokens_per_doc: Optional[int] = None, + ) -> Dict: + """ + Map Cohere rerank params to IBM watsonx.ai rerank params + """ + optional_rerank_params = {} + if non_default_params is not None: + for k, v in non_default_params.items(): + if k == "query" and v is not None: + optional_rerank_params["query"] = v + elif k == "documents" and v is not None: + optional_rerank_params["inputs"] = [ + {"text": el} if isinstance(el, str) else el for el in v + ] + elif k == "top_n" and v is not None: + optional_rerank_params.setdefault("parameters", {}).setdefault("return_options", {})["top_n"] = v + elif k == "return_documents" and v is not None and isinstance(v, bool): + optional_rerank_params.setdefault("parameters", {}).setdefault("return_options", {})["inputs"] = v + elif k == "max_tokens_per_doc" and v is not None: + optional_rerank_params.setdefault("parameters", {})["truncate_input_tokens"] = v + + # IBM watsonx.ai require one of below parameters + elif k == "project_id" and v is not None: + optional_rerank_params["project_id"] = v + elif k == "space_id" and v is not None: + optional_rerank_params["space_id"] = v + + return dict(optional_rerank_params) + + def transform_rerank_request( + self, + model: str, + optional_rerank_params: Dict, + headers: dict, + ) -> dict: + """ + Transform request to IBM watsonx.ai rerank format + """ + watsonx_api_params = _get_api_params(params=optional_rerank_params, model=model) + watsonx_auth_payload = self._prepare_payload( + model=model, + api_params=watsonx_api_params, + ) + + return optional_rerank_params | watsonx_auth_payload + + def transform_rerank_response( + self, + model: str, + raw_response: httpx.Response, + model_response: RerankResponse, + logging_obj: LiteLLMLoggingObj, + api_key: Optional[str] = None, + request_data: dict = {}, + optional_params: dict = {}, + litellm_params: dict = {}, + ) -> RerankResponse: + """ + Transform IBM watsonx.ai rerank response to LiteLLM RerankResponse format + """ + try: + raw_response_json = raw_response.json() + except Exception as e: + raise self.get_error_class( + error_message=f"Failed to parse response: {str(e)}", + status_code=raw_response.status_code, + headers=raw_response.headers, + ) + + _results: Optional[List[dict]] = raw_response_json.get("results") + if _results is None: + raise ValueError(f"No results found in the response={raw_response_json}") + + transformed_results = [] + + for result in _results: + transformed_result: Dict[str, Any] = { + "index": result["index"], + "relevance_score": result["score"], + } + + if "input" in result: + if isinstance(result["input"], str): + transformed_result["document"] = {"text": result["input"]} + else: + transformed_result["document"] = result["input"] + + transformed_results.append(transformed_result) + + response_id = raw_response_json.get("id") or raw_response_json.get("model_id") or str(uuid.uuid4()) + + # Extract usage information + _tokens = RerankTokens( + input_tokens=raw_response_json.get("input_token_count", 0), + ) + rerank_meta = RerankResponseMeta(tokens=_tokens) + + return RerankResponse( + id=response_id, + results=transformed_results, # type: ignore + meta=rerank_meta, + ) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 9a9acb91986..2c89e53ada9 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -8294,6 +8294,37 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 346 }, + "us/claude-sonnet-4-6": { + "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06, + "cache_read_input_token_cost": 3.3e-07, + "cache_read_input_token_cost_above_200k_tokens": 6.6e-07, + "input_cost_per_token": 3.3e-06, + "input_cost_per_token_above_200k_tokens": 6.6e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 1.65e-05, + "output_cost_per_token_above_200k_tokens": 2.475e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 346, + "inference_geo": "us" + }, "claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, @@ -22465,6 +22496,20 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "mistral/devstral-small-latest": { + "input_cost_per_token": 1e-07, + "litellm_provider": "mistral", + "max_input_tokens": 256000, + "max_output_tokens": 256000, + "max_tokens": 256000, + "mode": "chat", + "output_cost_per_token": 3e-07, + "source": "https://docs.mistral.ai/models/devstral-small-2-25-12", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "mistral/labs-devstral-small-2512": { "input_cost_per_token": 1e-07, "litellm_provider": "mistral", @@ -22479,6 +22524,34 @@ "supports_response_schema": true, "supports_tool_choice": true }, + "mistral/devstral-latest": { + "input_cost_per_token": 4e-07, + "litellm_provider": "mistral", + "max_input_tokens": 256000, + "max_output_tokens": 256000, + "max_tokens": 256000, + "mode": "chat", + "output_cost_per_token": 2e-06, + "source": "https://mistral.ai/news/devstral-2-vibe-cli", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, + "mistral/devstral-medium-latest": { + "input_cost_per_token": 4e-07, + "litellm_provider": "mistral", + "max_input_tokens": 256000, + "max_output_tokens": 256000, + "max_tokens": 256000, + "mode": "chat", + "output_cost_per_token": 2e-06, + "source": "https://mistral.ai/news/devstral-2-vibe-cli", + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true + }, "mistral/devstral-2512": { "input_cost_per_token": 4e-07, "litellm_provider": "mistral", @@ -37270,5 +37343,13 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 } + }, + "duckduckgo/search": { + "litellm_provider": "duckduckgo", + "mode": "search", + "input_cost_per_query": 0.0, + "metadata": { + "notes": "DuckDuckGo Instant Answer API is free and does not require an API key." + } } } \ No newline at end of file diff --git a/litellm/proxy/auth/model_checks.py b/litellm/proxy/auth/model_checks.py index 71ae1348f39..32f209a763e 100644 --- a/litellm/proxy/auth/model_checks.py +++ b/litellm/proxy/auth/model_checks.py @@ -223,12 +223,14 @@ def get_known_models_from_wildcard( except ValueError: # safely fail return [] - if litellm_params is None: # need litellm params to extract litellm model name - return [] - - try: - provider = litellm_params.model.split("/", 1)[0] - except ValueError: + # Use provider from litellm_params when available, otherwise from wildcard prefix + # (e.g., "openai" from "openai/*" - needed for BYOK where wildcard isn't in router) + if litellm_params is not None: + try: + provider = litellm_params.model.split("/", 1)[0] + except ValueError: + provider = wildcard_provider_prefix + else: provider = wildcard_provider_prefix # get all known provider models @@ -282,7 +284,7 @@ def _get_wildcard_models( ## get litellm params from model if llm_router is not None: model_list = llm_router.get_model_list(model_name=model) - if model_list is not None: + if model_list: for router_model in model_list: wildcard_models = get_known_models_from_wildcard( wildcard_model=model, @@ -291,11 +293,22 @@ def _get_wildcard_models( ), ) all_wildcard_models.extend(wildcard_models) + else: + # Router has no deployment for this wildcard (e.g., BYOK team models) + # Fall back to expanding from known provider models + wildcard_models = get_known_models_from_wildcard( + wildcard_model=model, litellm_params=None + ) + if wildcard_models: + models_to_remove.add(model) + all_wildcard_models.extend(wildcard_models) else: # get all known provider models - wildcard_models = get_known_models_from_wildcard(wildcard_model=model) + wildcard_models = get_known_models_from_wildcard( + wildcard_model=model, litellm_params=None + ) - if wildcard_models is not None: + if wildcard_models: models_to_remove.add(model) all_wildcard_models.extend(wildcard_models) diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py index 8d32d95f0ac..a0c2113b7ab 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -18,6 +18,9 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail" additional_provider_specific_params=getattr( litellm_params, "additional_provider_specific_params", {} ), + unreachable_fallback=getattr( + litellm_params, "unreachable_fallback", "fail_closed" + ), guardrail_name=guardrail.get("guardrail_name", ""), event_hook=litellm_params.mode, default_on=litellm_params.default_on, diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/example_config.yaml b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/example_config.yaml index 7ad33b24608..a4dae103626 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/example_config.yaml +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/example_config.yaml @@ -14,6 +14,7 @@ litellm_settings: mode: pre_call # Options: pre_call, post_call, during_call, [pre_call, post_call] api_key: os.environ/GENERIC_GUARDRAIL_API_KEY # Optional if using Bearer auth api_base: http://localhost:8080 # Required. Endpoint /beta/litellm_basic_guardrail_api is automatically appended + unreachable_fallback: fail_closed # Options: fail_closed (default, raise), fail_open (proceed if endpoint unreachable or upstream returns 502/503/504) default_on: false # Set to true to apply to all requests by default additional_provider_specific_params: # Any additional parameters your guardrail API needs diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py index 9cded6f0ac2..1892424e86d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/generic_guardrail_api.py @@ -9,9 +9,11 @@ import fnmatch import os from typing import TYPE_CHECKING, Any, Dict, Literal, Optional +import httpx + from litellm._logging import verbose_proxy_logger from litellm._version import version as litellm_version -from litellm.exceptions import GuardrailRaisedException +from litellm.exceptions import GuardrailRaisedException, Timeout from litellm.integrations.custom_guardrail import ( CustomGuardrail, log_guardrail_information, @@ -34,17 +36,19 @@ if TYPE_CHECKING: GUARDRAIL_NAME = "generic_guardrail_api" # Headers whose values are forwarded as-is (case-insensitive). Glob patterns supported (e.g. x-stainless-*, x-litellm*). -_HEADER_VALUE_ALLOWLIST = frozenset({ - "host", - "accept-encoding", - "connection", - "accept", - "content-type", - "user-agent", - "x-stainless-*", - "x-litellm-*", - "content-length", -}) +_HEADER_VALUE_ALLOWLIST = frozenset( + { + "host", + "accept-encoding", + "connection", + "accept", + "content-type", + "user-agent", + "x-stainless-*", + "x-litellm-*", + "content-length", + } +) # Placeholder for headers that exist but are not on the allowlist (we don't expose their value). _HEADER_PRESENT_PLACEHOLDER = "[present]" @@ -166,6 +170,7 @@ class GenericGuardrailAPI(CustomGuardrail): api_base: Optional[str] = None, api_key: Optional[str] = None, additional_provider_specific_params: Optional[Dict[str, Any]] = None, + unreachable_fallback: Literal["fail_closed", "fail_open"] = "fail_closed", **kwargs, ): self.async_handler = get_async_httpx_client( @@ -196,6 +201,10 @@ class GenericGuardrailAPI(CustomGuardrail): additional_provider_specific_params or {} ) + self.unreachable_fallback: Literal["fail_closed", "fail_open"] = ( + unreachable_fallback + ) + # Set supported event hooks if "supported_event_hooks" not in kwargs: kwargs["supported_event_hooks"] = [ @@ -259,6 +268,54 @@ class GenericGuardrailAPI(CustomGuardrail): return result_metadata + def _fail_open_passthrough( + self, + *, + inputs: GenericGuardrailAPIInputs, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"], + error: Exception, + http_status_code: Optional[int] = None, + ) -> GenericGuardrailAPIInputs: + status_suffix = f" http_status_code={http_status_code}" if http_status_code else "" + verbose_proxy_logger.critical( + "Generic Guardrail API unreachable (fail-open). Proceeding without guardrail.%s " + "guardrail_name=%s api_base=%s input_type=%s litellm_call_id=%s litellm_trace_id=%s", + status_suffix, + getattr(self, "guardrail_name", None), + getattr(self, "api_base", None), + input_type, + getattr(logging_obj, "litellm_call_id", None) if logging_obj else None, + getattr(logging_obj, "litellm_trace_id", None) if logging_obj else None, + exc_info=error, + ) + # Keep flow going - treat as action=NONE (no modifications) + return_inputs: GenericGuardrailAPIInputs = {} + return_inputs.update(inputs) + return return_inputs + + def _build_guardrail_return_inputs( + self, + *, + texts: list, + images: Any, + tools: Any, + guardrail_response: GenericGuardrailAPIResponse, + ) -> GenericGuardrailAPIInputs: + # Action is NONE or no modifications needed + return_inputs = GenericGuardrailAPIInputs(texts=texts) + if guardrail_response.texts: + return_inputs["texts"] = guardrail_response.texts + if guardrail_response.images: + return_inputs["images"] = guardrail_response.images + elif images: + return_inputs["images"] = images + if guardrail_response.tools: + return_inputs["tools"] = guardrail_response.tools + elif tools: + return_inputs["tools"] = tools + return return_inputs + @log_guardrail_information async def apply_guardrail( self, @@ -313,7 +370,9 @@ class GenericGuardrailAPI(CustomGuardrail): # Extract user API key metadata user_metadata = self._extract_user_api_key_metadata(request_data) - inbound_headers = _extract_inbound_headers(request_data=request_data, logging_obj=logging_obj) + inbound_headers = _extract_inbound_headers( + request_data=request_data, logging_obj=logging_obj + ) # Create request payload guardrail_request = GenericGuardrailAPIRequest( @@ -370,23 +429,64 @@ class GenericGuardrailAPI(CustomGuardrail): should_wrap_with_default_message=False, ) - # Action is NONE or no modifications needed - return_inputs = GenericGuardrailAPIInputs(texts=texts) - if guardrail_response.texts: - return_inputs["texts"] = guardrail_response.texts - if guardrail_response.images: - return_inputs["images"] = guardrail_response.images - elif images: - return_inputs["images"] = images - if guardrail_response.tools: - return_inputs["tools"] = guardrail_response.tools - elif tools: - return_inputs["tools"] = tools - return return_inputs + return self._build_guardrail_return_inputs( + texts=texts, + images=images, + tools=tools, + guardrail_response=guardrail_response, + ) except GuardrailRaisedException: # Re-raise guardrail exceptions as-is raise + except Timeout as e: + # AsyncHTTPHandler wraps httpx.TimeoutException into litellm.Timeout + if self.unreachable_fallback == "fail_open": + return self._fail_open_passthrough( + inputs=inputs, + input_type=input_type, + logging_obj=logging_obj, + error=e, + ) + + verbose_proxy_logger.error( + "Generic Guardrail API: failed to make request: %s", str(e) + ) + raise Exception(f"Generic Guardrail API failed: {str(e)}") + except httpx.HTTPStatusError as e: + # Common reverse-proxy/LB failures can present as HTTP errors even when the backend is unreachable. + status_code = getattr(getattr(e, "response", None), "status_code", None) + if self.unreachable_fallback == "fail_open" and status_code in ( + 502, + 503, + 504, + ): + return self._fail_open_passthrough( + inputs=inputs, + input_type=input_type, + logging_obj=logging_obj, + error=e, + http_status_code=status_code, + ) + + verbose_proxy_logger.error( + "Generic Guardrail API: failed to make request: %s", str(e) + ) + raise Exception(f"Generic Guardrail API failed: {str(e)}") + except httpx.RequestError as e: + # Guardrail endpoint is unreachable (DNS/connect/timeout/etc) + if self.unreachable_fallback == "fail_open": + return self._fail_open_passthrough( + inputs=inputs, + input_type=input_type, + logging_obj=logging_obj, + error=e, + ) + + verbose_proxy_logger.error( + "Generic Guardrail API: failed to make request: %s", str(e) + ) + raise Exception(f"Generic Guardrail API failed: {str(e)}") except Exception as e: verbose_proxy_logger.error( "Generic Guardrail API: failed to make request: %s", str(e) diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py index c4ade2f1a85..7058e7644cb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py @@ -31,11 +31,15 @@ from litellm import Router from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import CustomGuardrail from litellm.proxy._types import UserAPIKeyAuth -from litellm.types.utils import GuardrailTracingDetail, ModelResponseStream +from litellm.types.utils import ( + GenericGuardrailAPIInputs, + GuardrailStatus, + GuardrailTracingDetail, + ModelResponseStream, +) if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj - from litellm.types.utils import GenericGuardrailAPIInputs, GuardrailStatus from litellm.types.guardrails import ( BlockedWord, @@ -1546,8 +1550,6 @@ class ContentFilterGuardrail(CustomGuardrail): Raises: HTTPException: If sensitive content is detected and action is BLOCK """ - from litellm.types.utils import GuardrailStatus - start_time = datetime.now() detections: List[ContentFilterDetection] = [] masked_entity_count: Dict[str, int] = {} @@ -1693,4 +1695,4 @@ class ContentFilterGuardrail(CustomGuardrail): LitellmContentFilterGuardrailConfigModel, ) - return LitellmContentFilterGuardrailConfigModel + return LitellmContentFilterGuardrailConfigModel \ No newline at end of file diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index 4d77af513a8..3e8cc46521d 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -239,6 +239,8 @@ def clean_headers( """ Removes litellm api key from headers """ + from litellm.llms.anthropic.common_utils import is_anthropic_oauth_key + clean_headers = {} litellm_key_lower = ( litellm_key_header_name.lower() if litellm_key_header_name is not None else None @@ -246,8 +248,13 @@ def clean_headers( for header, value in headers.items(): header_lower = header.lower() + # Preserve Authorization header if it contains Anthropic OAuth token (sk-ant-oat*) + # This allows OAuth tokens to be forwarded to Anthropic-compatible providers + # via add_provider_specific_headers_to_request() + if header_lower == "authorization" and is_anthropic_oauth_key(value): + clean_headers[header] = value # Check if header should be excluded: either in special headers cache or matches custom litellm key - if header_lower not in _SPECIAL_HEADERS_CACHE and ( + elif header_lower not in _SPECIAL_HEADERS_CACHE and ( litellm_key_lower is None or header_lower != litellm_key_lower ): clean_headers[header] = value @@ -1717,6 +1724,8 @@ def add_provider_specific_headers_to_request( data: dict, headers: dict, ): + from litellm.llms.anthropic.common_utils import is_anthropic_oauth_key + anthropic_headers = {} # boolean to indicate if a header was added added_header = False @@ -1726,6 +1735,14 @@ def add_provider_specific_headers_to_request( anthropic_headers[header] = header_value added_header = True + # Check for Authorization header with Anthropic OAuth token (sk-ant-oat*) + # This needs to be handled via provider-specific headers to ensure it only + # goes to Anthropic-compatible providers, not all providers in the router + for header, value in headers.items(): + if header.lower() == "authorization" and is_anthropic_oauth_key(value): + anthropic_headers[header] = value + added_header = True + break if added_header is True: # Anthropic headers work across multiple providers # Store as comma-separated list so retrieval can match any of them diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index acdb4efc21b..007e558cffc 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -20,6 +20,7 @@ from datetime import datetime, timedelta, timezone from typing import Any, Dict, List, Literal, Optional, Tuple, cast import fastapi +import prisma import yaml from fastapi import APIRouter, Depends, Header, HTTPException, Query, Request, status @@ -74,7 +75,6 @@ from litellm.proxy.utils import ( _hash_token_if_needed, handle_exception_on_proxy, is_valid_api_key, - jsonify_object, ) from litellm.router import Router from litellm.secret_managers.main import get_secret @@ -3052,7 +3052,7 @@ async def delete_key_aliases( ) -async def _rotate_master_key( +async def _rotate_master_key( # noqa: PLR0915 prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth, current_master_key: str, @@ -3095,13 +3095,17 @@ async def _rotate_master_key( should_create_model_in_db=False, ) if new_model: - new_models.append(jsonify_object(new_model.model_dump())) + _dumped = new_model.model_dump(exclude_none=True) + _dumped["litellm_params"] = prisma.Json(_dumped["litellm_params"]) # type: ignore[attr-defined] + _dumped["model_info"] = prisma.Json(_dumped["model_info"]) # type: ignore[attr-defined] + new_models.append(_dumped) verbose_proxy_logger.debug("Resetting proxy model table") - await prisma_client.db.litellm_proxymodeltable.delete_many() - verbose_proxy_logger.debug("Creating %s models", len(new_models)) - await prisma_client.db.litellm_proxymodeltable.create_many( - data=new_models, - ) + async with prisma_client.db.tx() as tx: + await tx.litellm_proxymodeltable.delete_many() + verbose_proxy_logger.debug("Creating %s models", len(new_models)) + await tx.litellm_proxymodeltable.create_many( + data=new_models, + ) # 3. process config table try: config = await prisma_client.db.litellm_config.find_many() @@ -3127,15 +3131,20 @@ async def _rotate_master_key( if encrypted_env_vars: await prisma_client.db.litellm_config.update( where={"param_name": "environment_variables"}, - data={"param_value": jsonify_object(encrypted_env_vars)}, + data={"param_value": prisma.Json(encrypted_env_vars)}, # type: ignore[attr-defined] ) # 4. process MCP server table - await rotate_mcp_server_credentials_master_key( - prisma_client=prisma_client, - touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, - new_master_key=new_master_key, - ) + try: + await rotate_mcp_server_credentials_master_key( + prisma_client=prisma_client, + touched_by=user_api_key_dict.user_id or LITELLM_PROXY_ADMIN_NAME, + new_master_key=new_master_key, + ) + except Exception as e: + verbose_proxy_logger.warning( + "Failed to rotate MCP server credentials: %s", str(e) + ) # 5. process credentials table try: @@ -3153,13 +3162,19 @@ async def _rotate_master_key( updated_patch=decrypted_cred, new_encryption_key=new_master_key, ) - credential_object_jsonified = jsonify_object( - encrypted_cred.model_dump() - ) + _cred_data = encrypted_cred.model_dump(exclude_none=True) + if "credential_values" in _cred_data: + _cred_data["credential_values"] = prisma.Json( # type: ignore[attr-defined] + _cred_data["credential_values"] + ) + if "credential_info" in _cred_data: + _cred_data["credential_info"] = prisma.Json( # type: ignore[attr-defined] + _cred_data["credential_info"] + ) await prisma_client.db.litellm_credentialstable.update( where={"credential_name": cred.credential_name}, data={ - **credential_object_jsonified, + **_cred_data, "updated_by": user_api_key_dict.user_id, }, ) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 1c2a6d378fc..7cfcef8155f 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -456,18 +456,26 @@ class ProxyLogging: def _init_litellm_callbacks(self, llm_router: Optional[Router] = None): self._add_proxy_hooks(llm_router) litellm.logging_callback_manager.add_litellm_callback(self.service_logging_obj) # type: ignore - for callback in litellm.callbacks: + + # Track string callbacks and their initialized instances so we can + # replace them in-place, preventing duplicates (string + instance) in + # litellm.callbacks which caused double-counting of metrics. + string_callbacks_to_replace: Dict[int, CustomLogger] = {} + + for idx, callback in enumerate(litellm.callbacks): if isinstance(callback, str): - callback = litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class( # type: ignore + initialized_callback = litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class( cast(_custom_logger_compatible_callbacks_literal, callback), internal_usage_cache=self.internal_usage_cache.dual_cache, llm_router=llm_router, ) - if callback is None: - continue + if initialized_callback is not None: + string_callbacks_to_replace[idx] = initialized_callback - litellm.logging_callback_manager.add_litellm_callback(callback) + # Replace string entries in litellm.callbacks with initialized instances + for idx, initialized_callback in string_callbacks_to_replace.items(): + litellm.callbacks[idx] = initialized_callback async def update_request_status( self, litellm_call_id: str, status: Literal["success", "fail"] diff --git a/litellm/rerank_api/main.py b/litellm/rerank_api/main.py index 8910d37fbe7..f47fd6323f0 100644 --- a/litellm/rerank_api/main.py +++ b/litellm/rerank_api/main.py @@ -10,6 +10,7 @@ from litellm.llms.base_llm.rerank.transformation import BaseRerankConfig from litellm.llms.bedrock.rerank.handler import BedrockRerankHandler from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.llms.together_ai.rerank.handler import TogetherAIRerank +from litellm.llms.watsonx.common_utils import IBMWatsonXMixin from litellm.rerank_api.rerank_utils import get_optional_rerank_params from litellm.secret_managers.main import get_secret, get_secret_str from litellm.types.rerank import RerankResponse @@ -29,7 +30,7 @@ async def arerank( model: str, query: str, documents: List[Union[str, Dict[str, Any]]], - custom_llm_provider: Optional[Literal["cohere", "together_ai", "deepinfra", "fireworks_ai", "voyage"]] = None, + custom_llm_provider: Optional[Literal["cohere", "together_ai", "deepinfra", "fireworks_ai", "voyage", "watsonx"]] = None, top_n: Optional[int] = None, rank_fields: Optional[List[str]] = None, return_documents: Optional[bool] = None, @@ -85,6 +86,7 @@ def rerank( # noqa: PLR0915 "deepinfra", "fireworks_ai", "voyage", + "watsonx", ] ] = None, top_n: Optional[int] = None, @@ -478,6 +480,31 @@ def rerank( # noqa: PLR0915 or get_secret_str("VOYAGE_API_BASE") ) + response = base_llm_http_handler.rerank( + model=model, + custom_llm_provider=_custom_llm_provider, + provider_config=rerank_provider_config, + optional_rerank_params=optional_rerank_params, + logging_obj=litellm_logging_obj, + timeout=optional_params.timeout, + api_key=api_key, + api_base=api_base, + _is_async=_is_async, + headers=headers or litellm.headers or {}, + client=client, + model_response=model_response, + ) + elif _custom_llm_provider == litellm.LlmProviders.WATSONX: + credentials = IBMWatsonXMixin.get_watsonx_credentials( + optional_params=dict(optional_params), api_key=dynamic_api_key, api_base=dynamic_api_base + ) + + api_key = credentials["api_key"] + api_base = credentials["api_base"] + + if credentials.get("token") is not None: + optional_rerank_params["token"] = credentials["token"] + response = base_llm_http_handler.rerank( model=model, custom_llm_provider=_custom_llm_provider, diff --git a/litellm/responses/main.py b/litellm/responses/main.py index e943789a1cd..2a1749b5626 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -600,8 +600,12 @@ def responses( # 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 - + 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( @@ -611,7 +615,7 @@ def responses( ), ) local_vars["input"] = input - + # Update tools with provider-specific file IDs if needed if tools: tools = cast( @@ -696,7 +700,10 @@ def responses( ) ) - # Pre Call logging + # Pre Call logging - preserve metadata for custom callbacks + # When called from completion bridge (codex models), metadata is in litellm_metadata + metadata_for_callbacks = metadata or kwargs.get("litellm_metadata") or {} + litellm_logging_obj.update_environment_variables( model=model, user=user, @@ -705,7 +712,7 @@ def responses( **responses_api_request_params, "aresponses": _is_async, "litellm_call_id": litellm_call_id, - "metadata": metadata, + "metadata": metadata_for_callbacks, }, custom_llm_provider=custom_llm_provider, ) diff --git a/litellm/router.py b/litellm/router.py index 888c97ca0b1..4e268a94a2d 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -113,12 +113,12 @@ from litellm.router_utils.handle_error import ( from litellm.router_utils.pre_call_checks.model_rate_limit_check import ( ModelRateLimitingCheck, ) +from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( + DeploymentAffinityCheck, +) from litellm.router_utils.pre_call_checks.prompt_caching_deployment_check import ( PromptCachingDeploymentCheck, ) -from litellm.router_utils.pre_call_checks.responses_api_deployment_check import ( - ResponsesApiDeploymentCheck, -) from litellm.router_utils.router_callbacks.track_deployment_metrics import ( increment_deployment_failures_for_current_minute, increment_deployment_successes_for_current_minute, @@ -293,6 +293,7 @@ class Router: router_general_settings: Optional[ RouterGeneralSettings ] = RouterGeneralSettings(), + deployment_affinity_ttl_seconds: int = 3600, ignore_invalid_deployments: bool = False, ) -> None: """ @@ -326,6 +327,7 @@ class Router: routing_strategy_args (dict): Additional args for latency-based routing. Defaults to {}. alerting_config (AlertingConfig): Slack alerting configuration. Defaults to None. provider_budget_config (ProviderBudgetConfig): Provider budget configuration. Use this to set llm_provider budget limits. example $100/day to OpenAI, $100/day to Azure, etc. Defaults to None. + deployment_affinity_ttl_seconds (int): TTL for user-key -> deployment affinity mapping. Defaults to 3600. ignore_invalid_deployments (bool): Ignores invalid deployments, and continues with other deployments. Default is to raise an error. Returns: Router: An instance of the litellm.Router class. @@ -604,6 +606,7 @@ class Router: litellm.failure_callback = [self.deployment_callback_on_failure] self.routing_strategy_args = routing_strategy_args self.provider_budget_config = provider_budget_config + self.deployment_affinity_ttl_seconds = deployment_affinity_ttl_seconds self.router_budget_logger: Optional[RouterBudgetLimiting] = None if RouterBudgetLimiting.should_init_router_budget_limiter( model_list=model_list, provider_budget_config=self.provider_budget_config @@ -1184,26 +1187,78 @@ class Router: def add_optional_pre_call_checks( self, optional_pre_call_checks: Optional[OptionalPreCallChecks] ): - if optional_pre_call_checks is not None: - for pre_call_check in optional_pre_call_checks: - _callback: Optional[CustomLogger] = None - if pre_call_check == "prompt_caching": - _callback = PromptCachingDeploymentCheck(cache=self.cache) - elif pre_call_check == "router_budget_limiting": - _callback = RouterBudgetLimiting( - dual_cache=self.cache, - provider_budget_config=self.provider_budget_config, - model_list=self.model_list, - ) - elif pre_call_check == "responses_api_deployment_check": - _callback = ResponsesApiDeploymentCheck() - elif pre_call_check == "enforce_model_rate_limits": - _callback = ModelRateLimitingCheck(dual_cache=self.cache) - if _callback is not None: - if self.optional_callbacks is None: - self.optional_callbacks = [] - self.optional_callbacks.append(_callback) - litellm.logging_callback_manager.add_litellm_callback(_callback) + if optional_pre_call_checks is None: + return + + # --------------------------------------------------------------------- + # Unified deployment affinity (session stickiness) + # --------------------------------------------------------------------- + enable_user_key_affinity = "deployment_affinity" in optional_pre_call_checks + enable_responses_api_affinity = ( + "responses_api_deployment_check" in optional_pre_call_checks + ) + if enable_user_key_affinity or enable_responses_api_affinity: + if self.optional_callbacks is None: + self.optional_callbacks = [] + + existing_affinity_callback: Optional[DeploymentAffinityCheck] = None + for cb in self.optional_callbacks: + if isinstance(cb, DeploymentAffinityCheck): + existing_affinity_callback = cb + break + + if existing_affinity_callback is not None: + existing_affinity_callback.enable_user_key_affinity = ( + existing_affinity_callback.enable_user_key_affinity + or enable_user_key_affinity + ) + existing_affinity_callback.enable_responses_api_affinity = ( + existing_affinity_callback.enable_responses_api_affinity + or enable_responses_api_affinity + ) + existing_affinity_callback.ttl_seconds = ( + self.deployment_affinity_ttl_seconds + ) + else: + affinity_callback = DeploymentAffinityCheck( + cache=self.cache, + ttl_seconds=self.deployment_affinity_ttl_seconds, + enable_user_key_affinity=enable_user_key_affinity, + enable_responses_api_affinity=enable_responses_api_affinity, + ) + self.optional_callbacks.append(affinity_callback) + litellm.logging_callback_manager.add_litellm_callback( + affinity_callback + ) + + # --------------------------------------------------------------------- + # Remaining optional pre-call checks + # --------------------------------------------------------------------- + for pre_call_check in optional_pre_call_checks: + _callback: Optional[CustomLogger] = None + if pre_call_check in ( + "deployment_affinity", + "responses_api_deployment_check", + ): + continue + if pre_call_check == "prompt_caching": + _callback = PromptCachingDeploymentCheck(cache=self.cache) + elif pre_call_check == "router_budget_limiting": + _callback = RouterBudgetLimiting( + dual_cache=self.cache, + provider_budget_config=self.provider_budget_config, + model_list=self.model_list, + ) + elif pre_call_check == "enforce_model_rate_limits": + _callback = ModelRateLimitingCheck(dual_cache=self.cache) + + if _callback is None: + continue + + if self.optional_callbacks is None: + self.optional_callbacks = [] + self.optional_callbacks.append(_callback) + litellm.logging_callback_manager.add_litellm_callback(_callback) def print_deployment(self, deployment: dict): """ @@ -7600,9 +7655,16 @@ class Router: Used by `.get_model_list` to get model list from model alias. """ returned_models: List[DeploymentTypedDict] = [] - for model_alias, model_value in self.model_group_alias.items(): - if model_name is not None and model_alias != model_name: - continue + + if model_name is not None: + # Fast path: direct dict lookup avoids scanning all aliases for non-alias model names. + if model_name not in self.model_group_alias: + return returned_models + alias_items = [(model_name, self.model_group_alias[model_name])] + else: + alias_items = list(self.model_group_alias.items()) + + for model_alias, model_value in alias_items: if isinstance(model_value, str): _router_model_name: str = model_value elif isinstance(model_value, dict): @@ -9099,4 +9161,3 @@ class Router: litellm._async_failure_callback = [] self.retry_policy = None self.flush_cache() - diff --git a/litellm/router_strategy/lowest_tpm_rpm_v2.py b/litellm/router_strategy/lowest_tpm_rpm_v2.py index bf3035fcc9f..70d4c6751db 100644 --- a/litellm/router_strategy/lowest_tpm_rpm_v2.py +++ b/litellm/router_strategy/lowest_tpm_rpm_v2.py @@ -335,13 +335,14 @@ class LowestTPMLoggingHandler_v2(BaseRoutingStrategy, CustomLogger): ): lowest_tpm = float("inf") potential_deployments = [] # if multiple deployments have the same low value + deployment_lookup = { + deployment.get("model_info", {}).get("id"): deployment + for deployment in healthy_deployments + } for item, item_tpm in all_deployments.items(): ## get the item from model list - _deployment = None item = item.split(":")[0] - for m in healthy_deployments: - if item == m["model_info"]["id"]: - _deployment = m + _deployment = deployment_lookup.get(item) if _deployment is None: continue # skip to next one elif item_tpm is None: diff --git a/litellm/router_utils/common_utils.py b/litellm/router_utils/common_utils.py index 10acc343abd..3b0273f4c5d 100644 --- a/litellm/router_utils/common_utils.py +++ b/litellm/router_utils/common_utils.py @@ -58,7 +58,7 @@ def filter_team_based_models( request_team_id = metadata.get("user_api_key_team_id") or litellm_metadata.get( "user_api_key_team_id" ) - ids_to_remove = [] + ids_to_remove = set() if isinstance(healthy_deployments, dict): return healthy_deployments for deployment in healthy_deployments: @@ -67,7 +67,7 @@ def filter_team_based_models( if model_team_id is None: continue if model_team_id != request_team_id: - ids_to_remove.append(deployment.get("model_info", {}).get("id")) + ids_to_remove.add(_model_info.get("id")) return [ deployment @@ -125,4 +125,3 @@ def filter_web_search_deployments( if len(healthy_deployments) > 0 and len(final_deployments) == 0: verbose_logger.warning("No deployments support web search for request") return final_deployments - diff --git a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py new file mode 100644 index 00000000000..d34607732b8 --- /dev/null +++ b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py @@ -0,0 +1,396 @@ +""" +Unified deployment affinity (session stickiness) for the Router. + +Features (independently enable-able): +1. Responses API continuity: when a `previous_response_id` is provided, route to the + deployment that generated the original response (highest priority). +2. API-key affinity: map an API key hash -> deployment id for a TTL and re-use that + deployment for subsequent requests to the same router deployment model name + (alias-safe, aligns to `model_map_information.model_map_key`). + +This is designed to support "implicit prompt caching" scenarios (no explicit cache_control), +where routing to a consistent deployment is still beneficial. +""" + +import hashlib +from typing import Any, Dict, List, Optional, cast + +from typing_extensions import TypedDict + +from litellm._logging import verbose_router_logger +from litellm.caching.dual_cache import DualCache +from litellm.integrations.custom_logger import CustomLogger, Span +from litellm.responses.utils import ResponsesAPIRequestUtils +from litellm.types.llms.openai import AllMessageValues +from litellm.types.utils import CallTypes + + +class DeploymentAffinityCacheValue(TypedDict): + model_id: str + + +class DeploymentAffinityCheck(CustomLogger): + """ + Router deployment affinity callback. + + NOTE: This is a Router-only callback intended to be wired through + `Router(optional_pre_call_checks=[...])`. + """ + + CACHE_KEY_PREFIX = "deployment_affinity:v1" + + def __init__( + self, + cache: DualCache, + ttl_seconds: int, + enable_user_key_affinity: bool, + enable_responses_api_affinity: bool, + ): + super().__init__() + self.cache = cache + self.ttl_seconds = ttl_seconds + self.enable_user_key_affinity = enable_user_key_affinity + self.enable_responses_api_affinity = enable_responses_api_affinity + + @staticmethod + def _looks_like_sha256_hex(value: str) -> bool: + if len(value) != 64: + return False + try: + int(value, 16) + except ValueError: + return False + return True + + @staticmethod + def _hash_user_key(user_key: str) -> str: + """ + Hash user identifiers before storing them in cache keys. + + This avoids putting raw API keys / user identifiers into Redis keys (and therefore + into logs/metrics), while keeping the cache key stable and a fixed length. + """ + # If the proxy already provides a stable SHA-256 (e.g. `metadata.user_api_key_hash`), + # keep it as-is to avoid double-hashing and to make correlation/debugging possible. + if DeploymentAffinityCheck._looks_like_sha256_hex(user_key): + return user_key.lower() + + return hashlib.sha256(user_key.encode("utf-8")).hexdigest() + + @staticmethod + def _get_model_map_key_from_litellm_model_name(litellm_model_name: str) -> Optional[str]: + """ + Best-effort derivation of a stable "model map key" for affinity scoping. + + The intent is to align with `standard_logging_payload.model_map_information.model_map_key`, + which is typically the base model identifier (stable across deployments/endpoints). + + Notes: + - When the model name is in "provider/model" format, the provider prefix is stripped. + - For Azure, the string after "azure/" is commonly an *Azure deployment name*, which may + differ across instances. If `base_model` is not explicitly set, we skip deriving a + model-map key from the model string to avoid generating unstable keys. + """ + if not litellm_model_name: + return None + + if "/" not in litellm_model_name: + return litellm_model_name + + provider_prefix, remainder = litellm_model_name.split("/", 1) + if provider_prefix == "azure": + return None + + return remainder + + @staticmethod + def _get_model_map_key_from_deployment(deployment: dict) -> Optional[str]: + """ + Derive a stable model-map key from a router deployment dict. + + Primary source: `deployment.model_name` (Router's canonical group name after + alias resolution). This is stable across provider-specific deployments (e.g., + Azure/Vertex/Bedrock for the same logical model) and aligns with + `model_map_information.model_map_key` in standard logging. + + Prefer `base_model` when available (important for Azure), otherwise fall back to + parsing `litellm_params.model`. + """ + model_name = deployment.get("model_name") + if isinstance(model_name, str) and model_name: + return model_name + + model_info = deployment.get("model_info") + if isinstance(model_info, dict): + base_model = model_info.get("base_model") + if isinstance(base_model, str) and base_model: + return base_model + + litellm_params = deployment.get("litellm_params") + if isinstance(litellm_params, dict): + base_model = litellm_params.get("base_model") + if isinstance(base_model, str) and base_model: + return base_model + litellm_model_name = litellm_params.get("model") + if isinstance(litellm_model_name, str) and litellm_model_name: + return DeploymentAffinityCheck._get_model_map_key_from_litellm_model_name( + litellm_model_name + ) + + return None + + @staticmethod + def _get_stable_model_map_key_from_deployments( + healthy_deployments: List[dict], + ) -> Optional[str]: + """ + Only use model-map key scoping when it is stable across the deployment set. + + This prevents accidentally keying on per-deployment identifiers like Azure deployment + names (when `base_model` is not configured). + """ + if not healthy_deployments: + return None + + keys: List[str] = [] + for deployment in healthy_deployments: + key = DeploymentAffinityCheck._get_model_map_key_from_deployment(deployment) + if key is None: + return None + keys.append(key) + + unique_keys = set(keys) + if len(unique_keys) != 1: + return None + return keys[0] + + @staticmethod + def _shorten_for_logs(value: str, keep: int = 8) -> str: + if len(value) <= keep: + return value + return f"{value[:keep]}..." + + @classmethod + def get_affinity_cache_key(cls, model_group: str, user_key: str) -> str: + hashed_user_key = cls._hash_user_key(user_key=user_key) + return f"{cls.CACHE_KEY_PREFIX}:{model_group}:{hashed_user_key}" + + @staticmethod + def _get_user_key_from_metadata_dict(metadata: dict) -> Optional[str]: + # NOTE: affinity is keyed on the *API key hash* provided by the proxy (not the + # OpenAI `user` parameter, which is an end-user identifier). + user_key = metadata.get("user_api_key_hash") + if user_key is None: + return None + return str(user_key) + + @staticmethod + def _iter_metadata_dicts(request_kwargs: dict) -> List[dict]: + """ + Return all metadata dicts available on the request. + + Depending on the endpoint, Router may populate `metadata` or `litellm_metadata`. + Users may also send one or both, so we check both (rather than using `or`). + """ + metadata_dicts: List[dict] = [] + for key in ("litellm_metadata", "metadata"): + md = request_kwargs.get(key) + if isinstance(md, dict): + metadata_dicts.append(md) + return metadata_dicts + + @staticmethod + def _get_user_key_from_request_kwargs(request_kwargs: dict) -> Optional[str]: + """ + Extract a stable affinity key from request kwargs. + + Source (proxy): `metadata.user_api_key_hash` + + Note: the OpenAI `user` parameter is an end-user identifier and is intentionally + not used for deployment affinity. + """ + # Check metadata dicts (Proxy usage) + for metadata in DeploymentAffinityCheck._iter_metadata_dicts(request_kwargs): + user_key = DeploymentAffinityCheck._get_user_key_from_metadata_dict( + metadata=metadata + ) + if user_key is not None: + return user_key + + return None + + @staticmethod + def _find_deployment_by_model_id(healthy_deployments: List[dict], model_id: str) -> Optional[dict]: + for deployment in healthy_deployments: + model_info = deployment.get("model_info") + if not isinstance(model_info, dict): + continue + deployment_model_id = model_info.get("id") + if deployment_model_id is not None and str(deployment_model_id) == str(model_id): + return deployment + return None + + async def async_filter_deployments( + self, + model: str, + healthy_deployments: List, + messages: Optional[List[AllMessageValues]], + request_kwargs: Optional[dict] = None, + parent_otel_span: Optional[Span] = None, + ) -> List[dict]: + """ + Optionally filter healthy deployments based on: + 1. `previous_response_id` (Responses API continuity) [highest priority] + 2. cached API-key deployment affinity + """ + request_kwargs = request_kwargs or {} + typed_healthy_deployments = cast(List[dict], healthy_deployments) + + # 1) Responses API continuity (high priority) + if self.enable_responses_api_affinity: + previous_response_id = request_kwargs.get("previous_response_id") + if previous_response_id is not None: + responses_model_id = ResponsesAPIRequestUtils.get_model_id_from_response_id(str(previous_response_id)) + if responses_model_id is not None: + deployment = self._find_deployment_by_model_id( + healthy_deployments=typed_healthy_deployments, + model_id=responses_model_id, + ) + if deployment is not None: + verbose_router_logger.debug( + "DeploymentAffinityCheck: previous_response_id pinning -> deployment=%s", + responses_model_id, + ) + return [deployment] + + # 2) User key -> deployment affinity + if not self.enable_user_key_affinity: + return typed_healthy_deployments + + user_key = self._get_user_key_from_request_kwargs(request_kwargs=request_kwargs) + if user_key is None: + return typed_healthy_deployments + + stable_model_map_key = self._get_stable_model_map_key_from_deployments( + healthy_deployments=typed_healthy_deployments + ) + if stable_model_map_key is None: + return typed_healthy_deployments + + cache_key = self.get_affinity_cache_key( + model_group=stable_model_map_key, user_key=user_key + ) + cache_result = await self.cache.async_get_cache(key=cache_key) + + model_id: Optional[str] = None + if isinstance(cache_result, dict): + model_id = cast(Optional[str], cache_result.get("model_id")) + elif isinstance(cache_result, str): + # Backwards / safety: allow raw string values. + model_id = cache_result + + if not model_id: + return typed_healthy_deployments + + deployment = self._find_deployment_by_model_id( + healthy_deployments=typed_healthy_deployments, + model_id=model_id, + ) + if deployment is None: + verbose_router_logger.debug( + "DeploymentAffinityCheck: pinned deployment=%s not found in healthy_deployments", + model_id, + ) + return typed_healthy_deployments + + verbose_router_logger.debug( + "DeploymentAffinityCheck: api-key affinity hit -> deployment=%s user_key=%s", + model_id, + self._shorten_for_logs(user_key), + ) + return [deployment] + + async def async_pre_call_deployment_hook( + self, kwargs: Dict[str, Any], call_type: Optional[CallTypes] + ) -> Optional[dict]: + """ + Persist/update the API-key -> deployment mapping for this request. + + Why pre-call? + - 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: + return None + + user_key = self._get_user_key_from_request_kwargs(request_kwargs=kwargs) + if user_key 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 + + if model_info is None: + for metadata in metadata_dicts: + maybe_model_info = metadata.get("model_info") + if isinstance(maybe_model_info, dict): + model_info = maybe_model_info + break + + if model_info is None: + # Router sets `model_info` after selecting a deployment. If it's missing, this is + # likely a non-router call or a call path that doesn't support affinity. + return None + + model_id = model_info.get("id") + if not model_id: + verbose_router_logger.warning( + "DeploymentAffinityCheck: model_id missing; skipping affinity cache update." + ) + 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 + + try: + cache_key = self.get_affinity_cache_key( + model_group=deployment_model_name, user_key=user_key + ) + await self.cache.async_set_cache( + cache_key, + DeploymentAffinityCacheValue(model_id=str(model_id)), + ttl=self.ttl_seconds, + ) + + verbose_router_logger.debug( + "DeploymentAffinityCheck: set affinity mapping model_map_key=%s deployment=%s ttl=%s user_key=%s", + deployment_model_name, + model_id, + self.ttl_seconds, + self._shorten_for_logs(user_key), + ) + except Exception as e: + # Non-blocking: affinity is a best-effort optimization. + verbose_router_logger.debug( + "DeploymentAffinityCheck: failed to set affinity cache. model_map_key=%s error=%s", + deployment_model_name, + e, + ) + + return None diff --git a/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py b/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py index b030fc28c84..5ae3c20baf3 100644 --- a/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py +++ b/litellm/router_utils/pre_call_checks/responses_api_deployment_check.py @@ -10,6 +10,7 @@ This is different from the normal behavior of the router, which does not have ro If previous_response_id is provided, route to the deployment that returned the previous response """ +import warnings from typing import List, Optional from litellm.integrations.custom_logger import CustomLogger, Span @@ -18,6 +19,17 @@ from litellm.types.llms.openai import AllMessageValues class ResponsesApiDeploymentCheck(CustomLogger): + def __init__(self) -> None: + super().__init__() + warnings.warn( + ( + "ResponsesApiDeploymentCheck is deprecated. " + "Use DeploymentAffinityCheck(enable_responses_api_affinity=True) instead." + ), + DeprecationWarning, + stacklevel=2, + ) + async def async_filter_deployments( self, model: str, diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 3e12ccc6a48..bbbd60dab7c 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -652,6 +652,15 @@ class BaseLitellmParams( description="Additional provider-specific parameters for generic guardrail APIs", ) + unreachable_fallback: Literal["fail_closed", "fail_open"] = Field( + default="fail_closed", + description=( + "Behavior when a guardrail endpoint is unreachable due to network errors. " + "NOTE: This is currently only implemented by guardrail='generic_guardrail_api'. " + "'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed." + ), + ) + # Custom code guardrail params custom_code: Optional[str] = Field( default=None, @@ -693,6 +702,7 @@ class LitellmParams( "mode", "default_action", "on_disallowed_action", + "unreachable_fallback", mode="before", check_fields=False, ) diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index 998c60ab60d..54237dfb37a 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -302,6 +302,33 @@ class PerformanceConfigBlock(TypedDict): latency: Literal["optimized", "throughput"] +class JsonSchemaDefinition(TypedDict, total=False): + """JSON schema structured output format options for Bedrock Converse API.""" + + schema: Required[str] # JSON string, not dict + name: str + description: str + + +class OutputFormatStructure(TypedDict, total=False): + """The structure that the model's output must adhere to (union type).""" + + jsonSchema: Required[JsonSchemaDefinition] + + +class OutputFormat(TypedDict): + """Structured output parameters to control the model's response.""" + + type: Literal["json_schema"] + structure: OutputFormatStructure + + +class OutputConfigBlock(TypedDict, total=False): + """Output configuration for a model response in Converse/ConverseStream.""" + + textFormat: OutputFormat + + class CommonRequestObject( TypedDict, total=False ): # common request object across sync + async flows @@ -314,6 +341,7 @@ class CommonRequestObject( performanceConfig: Optional[PerformanceConfigBlock] serviceTier: Optional[ServiceTierBlock] requestMetadata: Optional[Dict[str, str]] + outputConfig: Optional[OutputConfigBlock] class RequestObject(CommonRequestObject, total=False): diff --git a/litellm/types/llms/watsonx.py b/litellm/types/llms/watsonx.py index 137090b032e..21e58500c6f 100644 --- a/litellm/types/llms/watsonx.py +++ b/litellm/types/llms/watsonx.py @@ -63,6 +63,7 @@ class WatsonXAIEndpoint(str, Enum): EMBEDDINGS = "/ml/v1/text/embeddings" PROMPTS = "/ml/v1/prompts" AVAILABLE_MODELS = "/ml/v1/foundation_model_specs" + RERANK = "/ml/v1/text/rerank" class WatsonXModelPattern(str, Enum): diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py index 21f6f5b3b4e..94f219a5fc6 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/generic_guardrail_api.py @@ -31,6 +31,14 @@ class GenericGuardrailAPIOptionalParams(BaseModel): description="Additional provider-specific parameters to send with the guardrail request", ) + unreachable_fallback: Optional[Literal["fail_closed", "fail_open"]] = Field( + default="fail_closed", + description=( + "Behavior when the guardrail endpoint is unreachable due to network errors. " + "'fail_closed' raises an error (default). 'fail_open' logs a critical error and allows the request to proceed." + ), + ) + class GenericGuardrailAPIConfigModel( GuardrailConfigModel[GenericGuardrailAPIOptionalParams], @@ -52,9 +60,9 @@ class GenericGuardrailAPIRequest(BaseModel): input_type: Literal["request", "response"] litellm_call_id: Optional[str] = None # the call id of the individual LLM call - litellm_trace_id: Optional[ - str - ] = None # the trace id of the LLM call - useful if there are multiple LLM calls for the same conversation + litellm_trace_id: Optional[str] = ( + None # the trace id of the LLM call - useful if there are multiple LLM calls for the same conversation + ) structured_messages: Optional[List[AllMessageValues]] = None images: Optional[List[str]] = None tools: Optional[List[ChatCompletionToolParam]] = None diff --git a/litellm/types/router.py b/litellm/types/router.py index f78789c9772..3abe9f202aa 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -802,6 +802,7 @@ OptionalPreCallChecks = List[ "prompt_caching", "router_budget_limiting", "responses_api_deployment_check", + "deployment_affinity", "forward_client_headers_by_model_group", "enforce_model_rate_limits", ] diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 5f8798c7712..1fc0e55dab3 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2532,6 +2532,8 @@ class StandardLoggingMetadata(StandardLoggingUserAPIKeyMetadata): cold_storage_object_key: Optional[ str ] # S3/GCS object key for cold storage retrieval + team_alias: Optional[str] + team_id: Optional[str] class StandardLoggingAdditionalHeaders(TypedDict, total=False): @@ -3197,7 +3199,7 @@ class SearchProviders(str, Enum): FIRECRAWL = "firecrawl" SEARXNG = "searxng" LINKUP = "linkup" - + DUCKDUCKGO = "duckduckgo" # Create a set of all search provider values for quick lookup SearchProvidersSet = {provider.value for provider in SearchProviders} diff --git a/litellm/utils.py b/litellm/utils.py index 33f60f8d0a7..6a18fcc9e35 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -8151,6 +8151,8 @@ class ProviderConfigManager: return litellm.FireworksAIRerankConfig() elif litellm.LlmProviders.VOYAGE == provider: return litellm.VoyageRerankConfig() + elif litellm.LlmProviders.WATSONX == provider: + return litellm.IBMWatsonXRerankConfig() return litellm.CohereRerankConfig() @staticmethod @@ -8776,6 +8778,7 @@ class ProviderConfigManager: """ from litellm.llms.brave.search.transformation import BraveSearchConfig from litellm.llms.dataforseo.search.transformation import DataForSEOSearchConfig + from litellm.llms.duckduckgo.search.transformation import DuckDuckGoSearchConfig from litellm.llms.exa_ai.search.transformation import ExaAISearchConfig from litellm.llms.firecrawl.search.transformation import FirecrawlSearchConfig from litellm.llms.google_pse.search.transformation import GooglePSESearchConfig @@ -8798,6 +8801,7 @@ class ProviderConfigManager: SearchProviders.FIRECRAWL: FirecrawlSearchConfig, SearchProviders.SEARXNG: SearXNGSearchConfig, SearchProviders.LINKUP: LinkupSearchConfig, + SearchProviders.DUCKDUCKGO: DuckDuckGoSearchConfig, } config_class = PROVIDER_TO_CONFIG_MAP.get(provider, None) if config_class is None: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 471d4b9ab7a..2c89e53ada9 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -8294,6 +8294,37 @@ "supports_vision": true, "tool_use_system_prompt_tokens": 346 }, + "us/claude-sonnet-4-6": { + "cache_creation_input_token_cost": 4.125e-06, + "cache_creation_input_token_cost_above_200k_tokens": 8.25e-06, + "cache_read_input_token_cost": 3.3e-07, + "cache_read_input_token_cost_above_200k_tokens": 6.6e-07, + "input_cost_per_token": 3.3e-06, + "input_cost_per_token_above_200k_tokens": 6.6e-06, + "litellm_provider": "anthropic", + "max_input_tokens": 200000, + "max_output_tokens": 64000, + "max_tokens": 64000, + "mode": "chat", + "output_cost_per_token": 1.65e-05, + "output_cost_per_token_above_200k_tokens": 2.475e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "supports_assistant_prefill": true, + "supports_computer_use": true, + "supports_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true, + "tool_use_system_prompt_tokens": 346, + "inference_geo": "us" + }, "claude-sonnet-4-5-20250929-v1:0": { "cache_creation_input_token_cost": 3.75e-06, "cache_read_input_token_cost": 3e-07, @@ -37312,5 +37343,13 @@ "search_context_size_low": 0.01, "search_context_size_medium": 0.01 } + }, + "duckduckgo/search": { + "litellm_provider": "duckduckgo", + "mode": "search", + "input_cost_per_query": 0.0, + "metadata": { + "notes": "DuckDuckGo Instant Answer API is free and does not require an API key." + } } } \ No newline at end of file diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 366da0c0b46..328398a296a 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -761,6 +761,23 @@ "interactions": true } }, + "duckduckgo": { + "display_name": "DuckDuckGo (`duckduckgo`)", + "url": "https://docs.litellm.ai/docs/search/duckduckgo", + "endpoints": { + "chat_completions": false, + "messages": false, + "responses": false, + "embeddings": false, + "image_generations": false, + "audio_transcriptions": false, + "audio_speech": false, + "moderations": false, + "batches": false, + "rerank": false, + "search": true + } + }, "elevenlabs": { "display_name": "ElevenLabs (`elevenlabs`)", "url": "https://docs.litellm.ai/docs/providers/elevenlabs", diff --git a/tests/code_coverage_tests/enforce_llms_folder_style.py b/tests/code_coverage_tests/enforce_llms_folder_style.py index f684d884a6b..7e6fd8e6fd6 100644 --- a/tests/code_coverage_tests/enforce_llms_folder_style.py +++ b/tests/code_coverage_tests/enforce_llms_folder_style.py @@ -16,6 +16,7 @@ SEARCH_PROVIDERS = [ "firecrawl", "searxng", "linkup", + "duckduckgo", ] ALLOWED_FILES_IN_LLMS_FOLDER = [ @@ -73,8 +74,8 @@ def run_lint_check(unique_names): def main(): - llms_dir = "./litellm/llms/" # Update this path if needed - # llms_dir = "../../litellm/llms/" # LOCAL TESTING + # llms_dir = "./litellm/llms/" # Update this path if needed + llms_dir = "litellm/litellm/llms" # LOCAL TESTING unique_names = get_unique_names_from_llms_dir(llms_dir) print("Unique names in llms directory:", sorted(list(unique_names))) diff --git a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index c39454728a8..08b9351f9a3 100644 --- a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -792,7 +792,8 @@ async def test_async_post_call_success_hook(prometheus_logger): """ Test for the async_post_call_success_hook method - it should increment the litellm_proxy_total_requests_metric + litellm_proxy_total_requests_metric is NOT incremented here to avoid double-counting. + It is incremented in async_log_success_event instead. """ # Mock the prometheus metric prometheus_logger.litellm_proxy_total_requests_metric = MagicMock() @@ -817,23 +818,8 @@ async def test_async_post_call_success_hook(prometheus_logger): data=data, user_api_key_dict=user_api_key_dict, response=response ) - # Assert total requests metric was incremented with correct labels - prometheus_logger.litellm_proxy_total_requests_metric.labels.assert_called_once_with( - end_user=None, - hashed_api_key="test_key", - api_key_alias="test_alias", - requested_model="gpt-3.5-turbo", - team="test_team", - team_alias="test_team_alias", - user="test_user", - status_code="200", - user_email=None, - route=user_api_key_dict.request_route, - model_id=None, - client_ip=None, - user_agent=None, - ) - prometheus_logger.litellm_proxy_total_requests_metric.labels().inc.assert_called_once() + # Assert total requests metric was NOT incremented (moved to async_log_success_event) + prometheus_logger.litellm_proxy_total_requests_metric.labels.assert_not_called() def test_set_llm_deployment_success_metrics(prometheus_logger): diff --git a/tests/enterprise/litellm_enterprise/integrations/test_prometheus.py b/tests/enterprise/litellm_enterprise/integrations/test_prometheus.py index 4a08c208fda..212c5d4a322 100644 --- a/tests/enterprise/litellm_enterprise/integrations/test_prometheus.py +++ b/tests/enterprise/litellm_enterprise/integrations/test_prometheus.py @@ -545,6 +545,9 @@ async def test_request_counter_semantic_validation(mock_prometheus_logger): CRITICAL TEST: Validates that request counters are incremented by 1, not by token count. This test specifically catches the bug where litellm_proxy_total_requests_metric is incorrectly incremented by total_tokens instead of 1. + + The metric is now ONLY incremented in async_log_success_event (for both streaming + and non-streaming) to prevent double-counting. """ from datetime import datetime, timedelta from unittest.mock import MagicMock @@ -583,18 +586,18 @@ async def test_request_counter_semantic_validation(mock_prometheus_logger): }, } - # Call the success event + # Call the success event - should increment for both streaming and non-streaming await mock_prometheus_logger.async_log_success_event( kwargs, None, kwargs["start_time"], kwargs["end_time"] ) - # CRITICAL ASSERTION: Request counter should not be incremented + # CRITICAL ASSERTION: Request counter should be incremented by 1 total_requests_metric = mock_prometheus_logger.litellm_proxy_total_requests_metric assert ( - len(total_requests_metric.inc_calls) == 0 - ), "Request metric should not be incremented" + len(total_requests_metric.inc_calls) == 1 + ), "Request metric should be incremented once in async_log_success_event" - # Call the post-call logging hook + # Call the post-call logging hook - should NOT increment (to prevent double-counting) await mock_prometheus_logger.async_post_call_success_hook( data={}, user_api_key_dict=UserAPIKeyAuth( @@ -607,11 +610,11 @@ async def test_request_counter_semantic_validation(mock_prometheus_logger): response=MagicMock(), ) - # CRITICAL ASSERTION: Request counter be incremented by 1 + # CRITICAL ASSERTION: Request counter should still be 1 (not incremented again) total_requests_metric = mock_prometheus_logger.litellm_proxy_total_requests_metric assert ( len(total_requests_metric.inc_calls) == 1 - ), "Request metric should not be incremented" + ), "Request metric should not be incremented again in async_post_call_success_hook" # Check that ALL request counter increments are by 1 (not by token count) for inc_value in total_requests_metric.inc_calls: @@ -684,8 +687,8 @@ async def test_multiple_requests_counter_semantics(mock_prometheus_logger): expected_total_tokens = num_requests * tokens_per_request # 3 * 500 = 1500 # With the bug, total_request_increments would be 1500 instead of 3 - assert total_request_increments == 0, ( - f"SEMANTIC BUG: Request counter total increments = 0, " + assert total_request_increments == num_requests, ( + f"SEMANTIC BUG: Request counter total increments = {total_request_increments}, " f"expected {num_requests}. This suggests request counters are being incremented " f"by token counts instead of request counts." ) diff --git a/tests/litellm/llms/bedrock/test_nova_imported_models.py b/tests/litellm/llms/bedrock/test_nova_imported_models.py new file mode 100644 index 00000000000..e3677aaf9e6 --- /dev/null +++ b/tests/litellm/llms/bedrock/test_nova_imported_models.py @@ -0,0 +1,92 @@ +""" +Tests for Nova imported/custom model support via spec prefixes (nova/, nova-2/). +""" + +import pytest + +from litellm.llms.bedrock.common_utils import ( + BedrockModelInfo, + get_bedrock_base_model, + strip_bedrock_routing_prefix, +) +from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig + + +NOVA_ARN = "arn:aws:bedrock:us-east-1:123456789012:custom-model-deployment/a1b2c3d4e5f6" +NOVA_MODEL = f"bedrock/nova/{NOVA_ARN}" +NOVA2_MODEL = f"bedrock/nova-2/{NOVA_ARN}" + + +class TestGetBedrockRoute: + def test_nova_prefix_routes_to_converse(self): + assert BedrockModelInfo.get_bedrock_route(NOVA_MODEL) == "converse" + + def test_nova2_prefix_routes_to_converse(self): + assert BedrockModelInfo.get_bedrock_route(NOVA2_MODEL) == "converse" + + def test_plain_arn_routes_to_invoke(self): + # Without spec prefix, ARN doesn't match converse models + result = BedrockModelInfo.get_bedrock_route(f"bedrock/{NOVA_ARN}") + assert result == "invoke" + + +class TestGetBedrockBaseModel: + def test_nova_prefix_returns_sentinel(self): + assert get_bedrock_base_model(f"nova/{NOVA_ARN}") == "amazon.nova-custom" + + def test_nova2_prefix_returns_sentinel(self): + assert get_bedrock_base_model(f"nova-2/{NOVA_ARN}") == "amazon.nova-2-custom" + + def test_bedrock_nova_prefix_returns_sentinel(self): + assert get_bedrock_base_model(NOVA_MODEL) == "amazon.nova-custom" + + def test_bedrock_nova2_prefix_returns_sentinel(self): + assert get_bedrock_base_model(NOVA2_MODEL) == "amazon.nova-2-custom" + + +class TestStripBedrockRoutingPrefix: + def test_strips_nova_prefix(self): + result = strip_bedrock_routing_prefix(f"nova/{NOVA_ARN}") + assert result == NOVA_ARN + + def test_strips_nova2_prefix(self): + result = strip_bedrock_routing_prefix(f"nova-2/{NOVA_ARN}") + assert result == NOVA_ARN + + +class TestIsNova2Model: + def setup_method(self): + self.config = AmazonConverseConfig() + + def test_standard_nova2_model(self): + assert self.config._is_nova_2_model("amazon.nova-2-lite-v1:0") is True + + def test_nova2_imported_model(self): + assert self.config._is_nova_2_model(NOVA2_MODEL) is True + + def test_nova_imported_model_is_not_nova2(self): + assert self.config._is_nova_2_model(NOVA_MODEL) is False + + def test_plain_nova_model(self): + assert self.config._is_nova_2_model("amazon.nova-pro-v1:0") is False + + +class TestGetSupportedOpenaiParams: + def setup_method(self): + self.config = AmazonConverseConfig() + + def test_nova_imported_has_tools_and_web_search(self): + params = self.config.get_supported_openai_params(NOVA_MODEL) + assert "tools" in params + assert "tool_choice" in params + assert "web_search_options" in params + + def test_nova2_imported_has_reasoning_effort(self): + params = self.config.get_supported_openai_params(NOVA2_MODEL) + assert "reasoning_effort" in params + assert "web_search_options" in params + + def test_nova2_imported_has_tools(self): + params = self.config.get_supported_openai_params(NOVA2_MODEL) + assert "tools" in params + assert "tool_choice" in params diff --git a/tests/litellm/proxy/test_init_litellm_callbacks.py b/tests/litellm/proxy/test_init_litellm_callbacks.py new file mode 100644 index 00000000000..a3cd84faa90 --- /dev/null +++ b/tests/litellm/proxy/test_init_litellm_callbacks.py @@ -0,0 +1,175 @@ +""" +Unit tests for ProxyLogging._init_litellm_callbacks. + +Validates that string callbacks in litellm.callbacks are replaced in-place +with their initialized instances, preventing duplicate entries (string + instance) +that caused double-counting of metrics like litellm_proxy_total_requests_metric. +""" + +from typing import List, Union +from unittest.mock import MagicMock, patch + +import pytest + +import litellm +from litellm.integrations.custom_logger import CustomLogger + + +class FakeCustomLogger(CustomLogger): + """A minimal CustomLogger subclass for testing.""" + + pass + + +class TestInitLitellmCallbacks: + """Tests for ProxyLogging._init_litellm_callbacks.""" + + def _make_proxy_logging(self): + """Create a ProxyLogging instance with mocked dependencies.""" + from litellm.proxy.utils import ProxyLogging + + mock_cache = MagicMock() + proxy_logging = ProxyLogging(user_api_key_cache=mock_cache) + return proxy_logging + + @patch( + "litellm.proxy.utils.ProxyLogging._add_proxy_hooks", + new_callable=lambda: lambda self, *a, **kw: None, + ) + def test_should_replace_string_callback_with_instance(self, _mock_hooks): + """ + When litellm.callbacks contains a string callback (e.g. "lago"), + _init_litellm_callbacks should replace the string with the initialized + CustomLogger instance, not leave both the string and instance in the list. + """ + fake_logger = FakeCustomLogger() + + # Start with a string callback in litellm.callbacks + litellm.callbacks = ["lago"] # type: ignore + + proxy_logging = self._make_proxy_logging() + + with patch( + "litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class", + return_value=fake_logger, + ): + proxy_logging._init_litellm_callbacks(llm_router=None) + + # The string "lago" should be replaced by the instance, not appended + string_entries = [c for c in litellm.callbacks if isinstance(c, str)] + instance_entries = [ + c for c in litellm.callbacks if isinstance(c, FakeCustomLogger) + ] + + assert len(string_entries) == 0, ( + f"String callbacks should have been replaced, but found: {string_entries}" + ) + assert len(instance_entries) == 1, ( + f"Expected exactly one FakeCustomLogger instance, found {len(instance_entries)}" + ) + assert instance_entries[0] is fake_logger + + # Clean up + litellm.callbacks = [] # type: ignore + + @patch( + "litellm.proxy.utils.ProxyLogging._add_proxy_hooks", + new_callable=lambda: lambda self, *a, **kw: None, + ) + def test_should_not_duplicate_existing_instance_callbacks(self, _mock_hooks): + """ + When litellm.callbacks already contains a CustomLogger instance (not a string), + _init_litellm_callbacks should not create a duplicate. + """ + existing_logger = FakeCustomLogger() + + litellm.callbacks = [existing_logger] # type: ignore + + proxy_logging = self._make_proxy_logging() + + proxy_logging._init_litellm_callbacks(llm_router=None) + + # Count how many FakeCustomLogger instances are in litellm.callbacks + instance_count = sum( + 1 for c in litellm.callbacks if isinstance(c, FakeCustomLogger) + ) + assert instance_count == 1, ( + f"Expected exactly 1 FakeCustomLogger instance, found {instance_count}. " + f"litellm.callbacks = {litellm.callbacks}" + ) + + # Clean up + litellm.callbacks = [] # type: ignore + + @patch( + "litellm.proxy.utils.ProxyLogging._add_proxy_hooks", + new_callable=lambda: lambda self, *a, **kw: None, + ) + def test_should_handle_unrecognized_string_callback(self, _mock_hooks): + """ + When _init_custom_logger_compatible_class returns None for a string callback, + the string should remain in litellm.callbacks (not crash). + """ + litellm.callbacks = ["unknown_callback"] # type: ignore + + proxy_logging = self._make_proxy_logging() + + with patch( + "litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class", + return_value=None, + ): + proxy_logging._init_litellm_callbacks(llm_router=None) + + # The unknown string callback should still be there (not replaced, not crashed) + assert "unknown_callback" in litellm.callbacks + + # Clean up + litellm.callbacks = [] # type: ignore + + @patch( + "litellm.proxy.utils.ProxyLogging._add_proxy_hooks", + new_callable=lambda: lambda self, *a, **kw: None, + ) + def test_should_replace_multiple_string_callbacks(self, _mock_hooks): + """ + When litellm.callbacks contains multiple string callbacks, + each should be replaced with its corresponding initialized instance. + """ + fake_logger_a = FakeCustomLogger() + fake_logger_b = FakeCustomLogger() + + litellm.callbacks = ["callback_a", "callback_b"] # type: ignore + + proxy_logging = self._make_proxy_logging() + + call_count = 0 + + def mock_init_class(callback_name, **kwargs): + nonlocal call_count + call_count += 1 + if call_count == 1: + return fake_logger_a + return fake_logger_b + + with patch( + "litellm.litellm_core_utils.litellm_logging._init_custom_logger_compatible_class", + side_effect=mock_init_class, + ): + proxy_logging._init_litellm_callbacks(llm_router=None) + + string_entries = [c for c in litellm.callbacks if isinstance(c, str)] + instance_entries = [ + c for c in litellm.callbacks if isinstance(c, FakeCustomLogger) + ] + + assert len(string_entries) == 0, ( + f"All string callbacks should have been replaced: {string_entries}" + ) + assert len(instance_entries) == 2, ( + f"Expected 2 FakeCustomLogger instances, found {len(instance_entries)}" + ) + assert instance_entries[0] is fake_logger_a + assert instance_entries[1] is fake_logger_b + + # Clean up + litellm.callbacks = [] # type: ignore diff --git a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py index c151150f634..3b2087d25e9 100644 --- a/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py +++ b/tests/llm_translation/test_llm_response_utils/test_convert_dict_to_chat_completion.py @@ -1037,6 +1037,193 @@ def test_convert_to_model_response_object_with_empty_dict_error(): assert result.choices[0].message.content == "Hello!" +def test_convert_to_model_response_object_preserves_provider_specific_fields_from_proxy(): + """ + Test that provider_specific_fields (e.g. Anthropic citations) are preserved + when the response already contains them (e.g. from a proxy passthrough). + + Regression test for https://github.com/BerriAI/litellm/issues/21153 + """ + citations = [ + [ + { + "type": "web_search_result_location", + "cited_text": "The Sony WH-1000XM5 remains one of the best...", + "url": "https://example.com/headphones-review", + "title": "Best Headphones 2025", + "supported_text": "Based on current reviews...", + } + ], + ] + web_search_results = [ + { + "url": "https://example.com/headphones-review", + "title": "Best Headphones 2025", + "snippet": "The Sony WH-1000XM5 remains one of the best...", + } + ] + + response_object = { + "id": "chatcmpl-proxy-123", + "object": "chat.completion", + "created": 1728933352, + "model": "anthropic/claude-opus-4-5-20251101", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "Based on current reviews, the Sony WH-1000XM5 remains one of the best headphones.", + "tool_calls": [ + { + "id": "call_ws_123", + "type": "function", + "function": { + "name": "web_search", + "arguments": '{"query": "best headphones 2025"}', + }, + } + ], + "provider_specific_fields": { + "citations": citations, + "web_search_results": web_search_results, + }, + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 50, + "completion_tokens": 20, + "total_tokens": 70, + }, + } + + result = convert_to_model_response_object( + model_response_object=ModelResponse(), + response_object=response_object, + stream=False, + start_time=datetime.now(), + end_time=datetime.now(), + hidden_params=None, + _response_headers=None, + convert_tool_call_to_json_mode=False, + ) + + assert isinstance(result, ModelResponse) + assert result.id == "chatcmpl-proxy-123" + + choice = result.choices[0] + assert choice.message.content == "Based on current reviews, the Sony WH-1000XM5 remains one of the best headphones." + assert choice.message.provider_specific_fields is not None + assert "citations" in choice.message.provider_specific_fields + assert choice.message.provider_specific_fields["citations"] == citations + assert "web_search_results" in choice.message.provider_specific_fields + assert choice.message.provider_specific_fields["web_search_results"] == web_search_results + + +def test_convert_to_model_response_object_provider_specific_fields_merges_extra_keys(): + """ + Test that provider_specific_fields from the response are merged with + any extra non-standard keys present in the message dict. + + Regression test for https://github.com/BerriAI/litellm/issues/21153 + """ + response_object = { + "id": "chatcmpl-merge-123", + "object": "chat.completion", + "created": 1728933352, + "model": "some-model", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "Hello!", + "provider_specific_fields": { + "citations": [{"url": "https://example.com"}], + }, + "custom_extra_field": "extra_value", + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15, + }, + } + + result = convert_to_model_response_object( + model_response_object=ModelResponse(), + response_object=response_object, + stream=False, + start_time=datetime.now(), + end_time=datetime.now(), + hidden_params=None, + _response_headers=None, + convert_tool_call_to_json_mode=False, + ) + + assert isinstance(result, ModelResponse) + psf = result.choices[0].message.provider_specific_fields + assert psf is not None + # Both the existing provider_specific_fields and the extra key should be present + assert "citations" in psf + assert psf["citations"] == [{"url": "https://example.com"}] + assert "custom_extra_field" in psf + assert psf["custom_extra_field"] == "extra_value" + + +def test_convert_to_model_response_object_no_provider_specific_fields_still_works(): + """ + Test that responses without provider_specific_fields continue to work as before. + + Ensures the fix for https://github.com/BerriAI/litellm/issues/21153 + doesn't break normal responses. + """ + response_object = { + "id": "chatcmpl-normal-123", + "object": "chat.completion", + "created": 1728933352, + "model": "gpt-4o", + "choices": [ + { + "index": 0, + "message": { + "role": "assistant", + "content": "Hello!", + "refusal": None, + }, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 10, + "completion_tokens": 5, + "total_tokens": 15, + }, + } + + result = convert_to_model_response_object( + model_response_object=ModelResponse(), + response_object=response_object, + stream=False, + start_time=datetime.now(), + end_time=datetime.now(), + hidden_params=None, + _response_headers=None, + convert_tool_call_to_json_mode=False, + ) + + assert isinstance(result, ModelResponse) + psf = result.choices[0].message.provider_specific_fields + # refusal is not a Message model field, so it should be in provider_specific_fields + assert psf is not None + assert "refusal" in psf + + def test_convert_to_model_response_object_with_error_code_only(): """ Test that errors with only a code (no message) are still treated as real errors. diff --git a/tests/llm_translation/test_optional_params.py b/tests/llm_translation/test_optional_params.py index 4699c31c378..6ecac7b36a2 100644 --- a/tests/llm_translation/test_optional_params.py +++ b/tests/llm_translation/test_optional_params.py @@ -1894,28 +1894,28 @@ def test_validate_openai_optional_params_stop_truncation(): result = validate_openai_optional_params(stop=stop_sequences) assert result == ["stop1", "stop2", "stop3", "stop4"] assert len(result) == 4 - + # Test with exactly 4 stop sequences - should not truncate stop_sequences_4 = ["stop1", "stop2", "stop3", "stop4"] result = validate_openai_optional_params(stop=stop_sequences_4) assert result == ["stop1", "stop2", "stop3", "stop4"] assert len(result) == 4 - + # Test with less than 4 stop sequences - should not truncate stop_sequences_2 = ["stop1", "stop2"] result = validate_openai_optional_params(stop=stop_sequences_2) assert result == ["stop1", "stop2"] assert len(result) == 2 - + # Test with single stop sequence as string - should return as is stop_string = "stop1" result = validate_openai_optional_params(stop=stop_string) assert result == "stop1" - + # Test with None - should return None result = validate_openai_optional_params(stop=None) assert result is None - + # Test with empty list - should return empty list result = validate_openai_optional_params(stop=[]) assert result == [] @@ -1928,7 +1928,7 @@ def test_validate_openai_optional_params_disable_stop_sequence_limit(): """ # Save original value original_value = litellm.disable_stop_sequence_limit - + try: # Test with disable_stop_sequence_limit = True - should NOT truncate litellm.disable_stop_sequence_limit = True @@ -1936,7 +1936,7 @@ def test_validate_openai_optional_params_disable_stop_sequence_limit(): result = validate_openai_optional_params(stop=stop_sequences) assert result == ["stop1", "stop2", "stop3", "stop4", "stop5", "stop6"] assert len(result) == 6 - + # Test with disable_stop_sequence_limit = False - should truncate to 4 litellm.disable_stop_sequence_limit = False stop_sequences = ["stop1", "stop2", "stop3", "stop4", "stop5", "stop6"] @@ -1965,19 +1965,83 @@ def test_validate_openai_optional_params_integration(): mock_response.usage.prompt_tokens = 10 mock_response.usage.completion_tokens = 5 mock_response.usage.total_tokens = 15 - - mock_client.return_value.chat.completions.create.return_value = mock_response - + + mock_client.return_value.chat.completions.create.return_value = ( + mock_response + ) + # Call completion with more than 4 stop sequences response = litellm.completion( model="gpt-3.5-turbo", messages=[{"role": "user", "content": "Hello"}], stop=["stop1", "stop2", "stop3", "stop4", "stop5", "stop6"], - mock_response="Test response" # This will use mock + mock_response="Test response", # This will use mock ) - + # Verify the call was made (stop sequences should be truncated internally) assert response is not None except Exception as e: # Should not raise an exception pytest.fail(f"validate_openai_optional_params integration failed: {e}") + + +def test_drop_store_param_for_anthropic(): + """ + Test that the OpenAI-specific `store` parameter is correctly dropped + when calling Anthropic with drop_params=True. + + `store` is an OpenAI Chat Completion parameter (for storing completions + for distillation/evals) that Anthropic does not support. Without proper + handling, it leaks through to the Anthropic API and causes a + "store: Extra inputs are not permitted" error. + + Ref: https://github.com/BerriAI/litellm/issues/19700 + """ + optional_params = get_optional_params( + model="claude-sonnet-4-20250514", + custom_llm_provider="anthropic", + drop_params=True, + store=True, + ) + assert "store" not in optional_params + + +def test_additional_drop_params_store_for_anthropic(): + """ + Test that `additional_drop_params=["store"]` correctly strips the `store` + parameter for non-OpenAI providers like Anthropic. + + Ref: https://github.com/BerriAI/litellm/issues/19700 + """ + optional_params = get_optional_params( + model="claude-sonnet-4-20250514", + custom_llm_provider="anthropic", + additional_drop_params=["store"], + store=True, + ) + assert "store" not in optional_params + + +def test_store_in_openai_chat_completion_params(): + """ + Test that `store` is recognized as a standard OpenAI Chat Completion + parameter. This ensures it is correctly handled by helper functions + like `get_standard_openai_params()` and provider configs that rely on + `OPENAI_CHAT_COMPLETION_PARAMS`. + + Without `store` in this list, functions that filter by known OpenAI + params will silently drop it for OpenAI calls or incorrectly treat + it as a provider-specific param for non-OpenAI providers. + + Ref: https://github.com/BerriAI/litellm/issues/19700 + """ + from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS + + assert "store" in OPENAI_CHAT_COMPLETION_PARAMS + + # Verify get_standard_openai_params recognizes store + from litellm.utils import get_standard_openai_params + + result = get_standard_openai_params({"store": True, "temperature": 0.7}) + assert "store" in result + assert result["store"] is True diff --git a/tests/proxy_e2e_anthropic_messages_tests/test_all_beta_headers.py b/tests/proxy_e2e_anthropic_messages_tests/test_all_beta_headers.py index cf688d7deae..14a301add57 100644 --- a/tests/proxy_e2e_anthropic_messages_tests/test_all_beta_headers.py +++ b/tests/proxy_e2e_anthropic_messages_tests/test_all_beta_headers.py @@ -32,6 +32,8 @@ def get_all_supported_anthropic_beta_headers(provider: str): "model_name,provider_name", [ ("claude-sonnet-4-5-20250929", "anthropic"), + ("azure-ai-claude-opus-4.5", "azure_ai"), + ("vertex-ai-claude-opus-4-6", "vertex_ai"), ], ) async def test_anthropic_messages_with_all_beta_headers(model_name, provider_name): diff --git a/tests/proxy_e2e_anthropic_messages_tests/test_config.yaml b/tests/proxy_e2e_anthropic_messages_tests/test_config.yaml index ae1514e4e42..fbbb6d4114c 100644 --- a/tests/proxy_e2e_anthropic_messages_tests/test_config.yaml +++ b/tests/proxy_e2e_anthropic_messages_tests/test_config.yaml @@ -34,6 +34,13 @@ model_list: litellm_params: model: "bedrock/converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0" aws_region_name: "us-east-1" + + # Azure AI models + - model_name: azure-ai-claude-opus-4.5 + litellm_params: + model: "azure_ai/claude-opus-4.5" + api_key: os.environ/AZURE_AI_API_KEY + api_base: os.environ/AZURE_AI_API_BASE # Vertex AI models - model_name: vertex-ai-claude-opus-4-6 diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index e34cad66ba6..1d9bb15b217 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -1928,6 +1928,24 @@ def test_get_known_models_from_wildcard( assert all(model in wildcard_models for model in expected_models) +def test_get_known_models_from_wildcard_without_litellm_params(): + """ + Test wildcard expansion without litellm_params (BYOK case - team has openai/* + but no deployment in router config). + """ + from litellm.proxy.auth.model_checks import get_known_models_from_wildcard + + wildcard_models = get_known_models_from_wildcard( + wildcard_model="openai/*", litellm_params=None + ) + # Should return expanded OpenAI models (gpt-4o, gpt-4o-mini, etc.) + assert len(wildcard_models) > 0 + assert all(m.startswith("openai/") for m in wildcard_models) + # Check for common OpenAI models + model_ids = [m.split("/", 1)[1] for m in wildcard_models] + assert "gpt-4o" in model_ids or "gpt-3.5-turbo" in model_ids + + @pytest.mark.parametrize( "data, user_api_key_dict, expected_model", [ diff --git a/tests/router_unit_tests/test_get_model_list_alias_optimization.py b/tests/router_unit_tests/test_get_model_list_alias_optimization.py new file mode 100644 index 00000000000..31d992b6646 --- /dev/null +++ b/tests/router_unit_tests/test_get_model_list_alias_optimization.py @@ -0,0 +1,50 @@ +from litellm import Router + + +class NoItemsAliasDict(dict): + def items(self): + raise AssertionError("Unexpected full alias iteration via items()") + + +def test_get_model_list_from_model_alias_should_not_iterate_for_non_alias_lookup(): + router = Router( + model_list=[ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": {"model": "gpt-3.5-turbo"}, + } + ], + model_group_alias={"alias-1": "gpt-4"}, + ) + router.model_group_alias = NoItemsAliasDict( + {f"alias-{idx}": "gpt-4" for idx in range(200)} + ) + + model_alias_list = router.get_model_list_from_model_alias( + model_name="gpt-3.5-turbo" + ) + assert model_alias_list == [] + + +def test_map_team_model_should_not_iterate_aliases_for_non_alias_team_model_name(): + router = Router( + model_list=[ + { + "model_name": "gpt-3.5-turbo", + "litellm_params": {"model": "gpt-3.5-turbo"}, + "model_info": { + "team_id": "team-1", + "team_public_model_name": "team-model", + }, + } + ], + model_group_alias={"alias-1": "gpt-4"}, + ) + router.model_group_alias = NoItemsAliasDict( + {f"alias-{idx}": "gpt-4" for idx in range(200)} + ) + + assert ( + router.map_team_model(team_model_name="team-model", team_id="team-1") + == "gpt-3.5-turbo" + ) diff --git a/tests/search_tests/test_duckduckgo_search.py b/tests/search_tests/test_duckduckgo_search.py new file mode 100644 index 00000000000..13df1cff9d4 --- /dev/null +++ b/tests/search_tests/test_duckduckgo_search.py @@ -0,0 +1,356 @@ +""" +Tests for DuckDuckGo Search API integration. +""" +import os +import sys +import pytest +from unittest.mock import AsyncMock, patch, MagicMock + +sys.path.insert( + 0, os.path.abspath("../..") +) + +import litellm +from tests.search_tests.base_search_unit_tests import BaseSearchTest + + +class TestDuckDuckGoSearch(BaseSearchTest): + """ + Tests for DuckDuckGo Search functionality. + """ + + def get_search_provider(self) -> str: + """ + Return search_provider for DuckDuckGo Search. + """ + return "duckduckgo" + + @pytest.mark.asyncio + async def test_basic_search(self): + """ + Test basic search functionality with a simple query. + """ + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + litellm._turn_on_debug() + search_provider = self.get_search_provider() + print("Search Provider=", search_provider) + + try: + response = await litellm.asearch( + query="india", + search_provider=search_provider, + ) + print("Search response=", response.model_dump_json(indent=4)) + + print(f"\n{'='*80}") + print(f"Response type: {type(response)}") + print(f"Response object: {response.object if hasattr(response, 'object') else 'N/A'}") + + # Check if response has expected Search format + assert hasattr(response, "results"), "Response should have 'results' attribute" + assert hasattr(response, "object"), "Response should have 'object' attribute" + assert response.object == "search", f"Expected object='search', got '{response.object}'" + + # Validate results structure + assert isinstance(response.results, list), "results should be a list" + assert len(response.results) > 0, "Should have at least one result" + + # Check first result structure + first_result = response.results[0] + assert hasattr(first_result, "title"), "Result should have 'title' attribute" + assert hasattr(first_result, "url"), "Result should have 'url' attribute" + assert hasattr(first_result, "snippet"), "Result should have 'snippet' attribute" + + print(f"Total results: {len(response.results)}") + print(f"First result title: {first_result.title}") + print(f"First result URL: {first_result.url}") + print(f"First result snippet: {first_result.snippet[:100]}...") + print(f"{'='*80}\n") + + assert len(first_result.title) > 0, "Title should not be empty" + assert len(first_result.url) > 0, "URL should not be empty" + assert len(first_result.snippet) > 0, "Snippet should not be empty" + + # Validate cost tracking in _hidden_params + assert hasattr(response, "_hidden_params"), "Response should have '_hidden_params' attribute" + hidden_params = response._hidden_params + assert "response_cost" in hidden_params, "_hidden_params should contain 'response_cost'" + + response_cost = hidden_params["response_cost"] + assert response_cost is not None, "response_cost should not be None" + assert isinstance(response_cost, (int, float)), "response_cost should be a number" + assert response_cost == 0, "response_cost should be 0" + + print(f"Cost tracking: ${response_cost:.6f}") + + except Exception as e: + pytest.fail(f"Search call failed: {str(e)}") + + + def test_search_response_structure(self): + """ + Test that the Search response has the correct structure. + """ + litellm.set_verbose = True + search_provider = self.get_search_provider() + + response = litellm.search( + query="india", + search_provider=search_provider, + ) + + # Validate response structure + assert hasattr(response, "results"), "Response should have 'results' attribute" + assert hasattr(response, "object"), "Response should have 'object' attribute" + + assert isinstance(response.results, list), "results should be a list" + assert len(response.results) > 0, "Should have at least one result" + assert response.object == "search", "object should be 'search'" + + # Validate first result structure + first_result = response.results[0] + assert hasattr(first_result, "title"), "Result should have 'title' attribute" + assert hasattr(first_result, "url"), "Result should have 'url' attribute" + assert hasattr(first_result, "snippet"), "Result should have 'snippet' attribute" + assert isinstance(first_result.title, str), "title should be a string" + assert isinstance(first_result.url, str), "url should be a string" + assert isinstance(first_result.snippet, str), "snippet should be a string" + + print(f"\nResponse structure validated:") + print(f" - object: {response.object}") + print(f" - results: {len(response.results)}") + print(f" - first result has all required fields") + +class TestDuckDuckGoSearchMocked: + """ + Tests for DuckDuckGo Search functionality with mocked network responses. + """ + + @pytest.mark.asyncio + async def test_duckduckgo_search_request_payload(self): + """ + Test that validates the DuckDuckGo search request payload structure without making real API calls. + """ + # Create a mock response matching DuckDuckGo API format + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "Abstract": "", + "AbstractSource": "Wikipedia", + "AbstractText": "Python is a high-level programming language.", + "AbstractURL": "https://en.wikipedia.org/wiki/Python_(programming_language)", + "Answer": "", + "AnswerType": "", + "Definition": "", + "DefinitionSource": "", + "DefinitionURL": "", + "Entity": "", + "Heading": "Python (programming language)", + "Image": "", + "ImageHeight": 0, + "ImageIsLogo": 0, + "ImageWidth": 0, + "Infobox": "", + "Redirect": "", + "RelatedTopics": [ + { + "FirstURL": "https://duckduckgo.com/Python_programming", + "Icon": { + "Height": "", + "URL": "/i/python.png", + "Width": "" + }, + "Result": "Python Programming A general-purpose programming language.", + "Text": "Python Programming - A general-purpose programming language." + }, + { + "FirstURL": "https://duckduckgo.com/Python_packages", + "Icon": { + "Height": "", + "URL": "", + "Width": "" + }, + "Result": "Python Packages Package management in Python.", + "Text": "Python Packages - Package management in Python." + } + ], + "Results": [], + "Type": "A", + "meta": { + "attribution": None, + "blockgroup": None, + "created_date": None, + "description": "Wikipedia", + "designer": None, + "dev_date": None, + "dev_milestone": "live", + "developer": [ + { + "name": "DDG Team", + "type": "ddg", + "url": "http://www.duckduckhack.com" + } + ], + "example_query": "python programming", + "id": "wikipedia_fathead", + "is_stackexchange": None, + "js_callback_name": "wikipedia", + "live_date": None, + "maintainer": { + "github": "duckduckgo" + }, + "name": "Wikipedia", + "perl_module": "DDG::Fathead::Wikipedia", + "producer": None, + "production_state": "online", + "repo": "fathead", + "signal_from": "wikipedia_fathead", + "src_domain": "en.wikipedia.org", + "src_id": 1, + "src_name": "Wikipedia", + "src_options": { + "directory": "", + "is_fanon": 0, + "is_mediawiki": 1, + "is_wikipedia": 1, + "language": "en", + "min_abstract_length": "20", + "skip_abstract": 0, + "skip_abstract_paren": 0, + "skip_end": "0", + "skip_icon": 0, + "skip_image_name": 0, + "skip_qr": "", + "source_skip": "", + "src_info": "" + }, + "src_url": None, + "status": "live", + "tab": "About", + "topic": [ + "productivity" + ], + "unsafe": 0 + } + } + + # Mock the httpx AsyncClient get method (DuckDuckGo uses GET) + with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get", new_callable=AsyncMock) as mock_get: + mock_get.return_value = mock_response + + # Make the search call + response = await litellm.asearch( + query="python programming", + search_provider="duckduckgo", + max_results=5 + ) + + # Verify the get method was called once + assert mock_get.call_count == 1 + + # Get the actual call arguments + call_args = mock_get.call_args + + # Verify URL contains the query with proper URL encoding + url = call_args.kwargs["url"] + assert "api.duckduckgo.com" in url + # URL should be properly encoded with %20 for spaces + assert ("q=python+programming" in url or "q=python%20programming" in url) + assert "format=json" in url + + # Verify response structure + assert hasattr(response, "results") + assert hasattr(response, "object") + assert response.object == "search" + assert len(response.results) > 0 + + # Verify first result (Abstract) + first_result = response.results[0] + assert first_result.title == "Python (programming language)" + assert first_result.url == "https://en.wikipedia.org/wiki/Python_(programming_language)" + assert "Python is a high-level programming language" in first_result.snippet + + # Verify related topics are included + assert len(response.results) >= 2 # Abstract + at least one related topic + + @pytest.mark.asyncio + async def test_duckduckgo_search_disambiguation(self): + """ + Test handling of disambiguation results from DuckDuckGo. + """ + # Create a mock response with disambiguation type + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "Abstract": "", + "AbstractSource": "Wikipedia", + "AbstractText": "", + "AbstractURL": "https://en.wikipedia.org/wiki/India_(disambiguation)", + "Answer": "", + "AnswerType": "", + "Definition": "", + "DefinitionSource": "", + "DefinitionURL": "", + "Entity": "", + "Heading": "India", + "Image": "", + "ImageHeight": 0, + "ImageIsLogo": 0, + "ImageWidth": 0, + "Infobox": "", + "Redirect": "", + "RelatedTopics": [ + { + "FirstURL": "https://duckduckgo.com/India", + "Icon": { + "Height": "", + "URL": "/i/cef47a13.png", + "Width": "" + }, + "Result": "India A country in South Asia.", + "Text": "India - A country in South Asia." + }, + { + "Name": "Related Topics", + "Topics": [ + { + "FirstURL": "https://duckduckgo.com/d/Indus", + "Icon": { + "Height": "", + "URL": "", + "Width": "" + }, + "Result": "Indus See related meanings for the word 'Indus'.", + "Text": "Indus - See related meanings for the word 'Indus'." + } + ] + } + ], + "Results": [], + "Type": "D", + "meta": {} + } + + # Mock the httpx AsyncClient get method + with patch("litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.get", new_callable=AsyncMock) as mock_get: + mock_get.return_value = mock_response + + # Make the search call + response = await litellm.asearch( + query="India", + search_provider="duckduckgo" + ) + + # Verify response structure + assert hasattr(response, "results") + assert hasattr(response, "object") + assert response.object == "search" + + # Should have results from both direct topics and nested topics + assert len(response.results) >= 2 + + # Verify nested topics are processed + urls = [result.url for result in response.results] + assert any("India" in url for url in urls) + assert any("Indus" in url for url in urls) diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py index f8a082ee30c..3021fff9a22 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_transformation.py @@ -9,9 +9,7 @@ from unittest.mock import ANY, MagicMock, Mock, patch import httpx import pytest -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system-path +sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system-path import litellm @@ -119,9 +117,7 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_imag function_call_output = item break - assert ( - function_call_output is not None - ), "function_call_output not found in response" + assert function_call_output is not None, "function_call_output not found in response" assert function_call_output["call_id"] == "call_abc123" # Check that the output is correctly transformed @@ -131,12 +127,8 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_imag image_item = output[0] # Should be transformed to Responses API format - assert ( - image_item["type"] == "input_image" - ), f"Expected type 'input_image', got '{image_item.get('type')}'" - assert ( - image_item["image_url"] == test_image_base64 - ), "image_url should be a flat string, not a nested object" + assert image_item["type"] == "input_image", f"Expected type 'input_image', got '{image_item.get('type')}'" + assert image_item["image_url"] == test_image_base64, "image_url should be a flat string, not a nested object" assert "detail" in image_item, "detail field should be present" print("✓ Tool result with image correctly transformed to Responses API format") @@ -198,9 +190,7 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_text function_call_output = item break - assert ( - function_call_output is not None - ), "function_call_output not found in response" + assert function_call_output is not None, "function_call_output not found in response" assert function_call_output["call_id"] == "call_abc123" # Check that the output is correctly transformed to use input_text, not output_text @@ -210,12 +200,10 @@ def test_convert_chat_completion_messages_to_responses_api_tool_result_with_text text_item = output[0] # Should be transformed to use input_text for tool results in Responses API format - assert ( - text_item["type"] == "input_text" - ), f"Expected type 'input_text' for tool result, got '{text_item.get('type')}'" - assert ( - text_item["text"] == "15 degrees" - ), f"Expected text '15 degrees', got '{text_item.get('text')}'" + assert text_item["type"] == "input_text", ( + f"Expected type 'input_text' for tool result, got '{text_item.get('type')}'" + ) + assert text_item["text"] == "15 degrees", f"Expected text '15 degrees', got '{text_item.get('text')}'" print("✓ Tool result with text correctly transformed to use input_text for Responses API format") @@ -226,9 +214,7 @@ def test_openai_responses_chunk_parser_reasoning_summary(): ) from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "delta": "**Compar", @@ -260,9 +246,7 @@ def test_chunk_parser_string_output_text_delta_produces_text(): ) from litellm.types.utils import ModelResponseStream - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = {"type": "response.output_text.delta", "delta": "literal text"} @@ -283,9 +267,7 @@ def test_chunk_parser_enum_output_text_delta_produces_text(): from litellm.types.llms.openai import ResponsesAPIStreamEvents from litellm.types.utils import ModelResponseStream - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = {"type": ResponsesAPIStreamEvents.OUTPUT_TEXT_DELTA, "delta": "enum text"} @@ -306,9 +288,7 @@ def test_chunk_parser_function_call_added_produces_tool_use(): from litellm.types.llms.openai import ResponsesAPIStreamEvents from litellm.types.utils import ModelResponseStream - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "type": ResponsesAPIStreamEvents.OUTPUT_ITEM_ADDED, @@ -393,9 +373,7 @@ Tomorrow will bring its petitions and promises, but for now the city breathes slow and wide, and I learn to carry this small calm home.""" - output_text = ResponseOutputText( - annotations=[], text=poem_text, type="output_text", logprobs=[] - ) + output_text = ResponseOutputText(annotations=[], text=poem_text, type="output_text", logprobs=[]) output_message = ResponseOutputMessage( id="msg_04c8021b8b3188a00068e9ae0b92f4819dac64d85b4abb67ec", content=[output_text], @@ -407,9 +385,7 @@ and I learn to carry this small calm home.""" # Create usage information usage = ResponseAPIUsage( input_tokens=16, - input_tokens_details=InputTokensDetails( - audio_tokens=None, cached_tokens=0, text_tokens=None - ), + input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None), output_tokens=195, output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None), total_tokens=211, @@ -621,9 +597,7 @@ def test_transform_request_single_char_keys_not_matched(): assert result_correct.get("metadata") == {"user_id": "123"} assert result_correct.get("previous_response_id") == "resp_abc" - print( - "✓ Single-character keys are not incorrectly matched to metadata/previous_response_id" - ) + print("✓ Single-character keys are not incorrectly matched to metadata/previous_response_id") # ============================================================================= @@ -643,9 +617,7 @@ def test_message_done_does_not_emit_is_finished(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "type": "response.output_item.done", @@ -657,9 +629,9 @@ def test_message_done_does_not_emit_is_finished(): # After the fix, message completion should NOT set finish_reason # ModelResponseStream doesn't have is_finished - check finish_reason instead assert len(result.choices) > 0, "result should have choices" - assert ( - result.choices[0].finish_reason is None or result.choices[0].finish_reason == "" - ), "message completion should not emit finish_reason" + assert result.choices[0].finish_reason is None or result.choices[0].finish_reason == "", ( + "message completion should not emit finish_reason" + ) def test_response_completed_emits_is_finished(): @@ -671,9 +643,7 @@ def test_response_completed_emits_is_finished(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = {"type": "response.completed"} @@ -681,9 +651,91 @@ def test_response_completed_emits_is_finished(): # response.completed should emit finish_reason='stop' assert len(result.choices) > 0, "result should have choices" - assert ( - result.choices[0].finish_reason == "stop" - ), "response.completed should emit finish_reason='stop'" + assert result.choices[0].finish_reason == "stop", "response.completed should emit finish_reason='stop'" + + +def test_response_completed_with_function_calls_emits_tool_calls_finish_reason(): + """ + Test that response.completed with function_call items in output emits finish_reason='tool_calls'. + + This is a regression test for an issue where response.completed always returned + finish_reason='stop' even when the response contained tool calls, causing agents + like OpenCode to incorrectly conclude the stream ended without tools to execute. + + When the response.completed event includes function_call items in its output, + the finish_reason should be 'tool_calls' to signal the client that tools need + to be executed. + """ + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + OpenAiResponsesToChatCompletionStreamIterator, + ) + + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + + # Simulate a response.completed event with function_call in output + # This matches what Azure/OpenAI sends for gpt-5.1-codex-mini and similar models + chunk = { + "type": "response.completed", + "response": { + "id": "resp_123", + "status": "completed", + "output": [ + { + "type": "function_call", + "id": "call_abc123", + "call_id": "call_abc123", + "name": "read_file", + "arguments": '{"path": "/tmp/test.py"}', + "status": "completed", + } + ], + }, + } + + result = iterator.chunk_parser(chunk) + + # response.completed with function_call should emit finish_reason='tool_calls' + assert len(result.choices) > 0, "result should have choices" + assert result.choices[0].finish_reason == "tool_calls", ( + "response.completed with function_call output should emit finish_reason='tool_calls'" + ) + + +def test_response_completed_with_message_only_emits_stop_finish_reason(): + """ + Test that response.completed with only message output (no function_call) emits finish_reason='stop'. + """ + from litellm.completion_extras.litellm_responses_transformation.transformation import ( + OpenAiResponsesToChatCompletionStreamIterator, + ) + + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) + + # Simulate a response.completed event with only message output + chunk = { + "type": "response.completed", + "response": { + "id": "resp_456", + "status": "completed", + "output": [ + { + "type": "message", + "id": "msg_xyz", + "role": "assistant", + "content": [{"type": "output_text", "text": "Hello, world!"}], + "status": "completed", + } + ], + }, + } + + result = iterator.chunk_parser(chunk) + + # response.completed with only message should emit finish_reason='stop' + assert len(result.choices) > 0, "result should have choices" + assert result.choices[0].finish_reason == "stop", ( + "response.completed with only message output should emit finish_reason='stop'" + ) def test_function_call_done_emits_is_finished(): @@ -695,9 +747,7 @@ def test_function_call_done_emits_is_finished(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) chunk = { "type": "response.output_item.done", @@ -713,13 +763,10 @@ def test_function_call_done_emits_is_finished(): # function_call completion should emit finish_reason='tool_calls' assert len(result.choices) > 0, "result should have choices" - assert ( - result.choices[0].finish_reason == "tool_calls" - ), "function_call should emit finish_reason='tool_calls'" - assert ( - result.choices[0].delta.tool_calls is not None - and len(result.choices[0].delta.tool_calls) > 0 - ), "function_call should include tool_calls" + assert result.choices[0].finish_reason == "tool_calls", "function_call should emit finish_reason='tool_calls'" + assert result.choices[0].delta.tool_calls is not None and len(result.choices[0].delta.tool_calls) > 0, ( + "function_call should include tool_calls" + ) def test_text_plus_tool_calls_sequence(): @@ -734,9 +781,7 @@ def test_text_plus_tool_calls_sequence(): OpenAiResponsesToChatCompletionStreamIterator, ) - iterator = OpenAiResponsesToChatCompletionStreamIterator( - streaming_response=None, sync_stream=True - ) + iterator = OpenAiResponsesToChatCompletionStreamIterator(streaming_response=None, sync_stream=True) # Simulate the sequence from OpenAI Responses API chunks = [ @@ -775,26 +820,21 @@ def test_text_plus_tool_calls_sequence(): # Check message done (index 2) does NOT have finish_reason set message_done_result = results[2] assert len(message_done_result.choices) > 0, "message done should have choices" - assert ( - message_done_result.choices[0].finish_reason is None - or message_done_result.choices[0].finish_reason == "" - ), "message done should not have finish_reason" + assert message_done_result.choices[0].finish_reason is None or message_done_result.choices[0].finish_reason == "", ( + "message done should not have finish_reason" + ) # Check function_call done (index 5) DOES have finish_reason='tool_calls' function_done_result = results[5] - assert ( - len(function_done_result.choices) > 0 - ), "function_call done should have choices" - assert ( - function_done_result.choices[0].finish_reason == "tool_calls" - ), "function_call done should have finish_reason='tool_calls'" + assert len(function_done_result.choices) > 0, "function_call done should have choices" + assert function_done_result.choices[0].finish_reason == "tool_calls", ( + "function_call done should have finish_reason='tool_calls'" + ) # Check response.completed (index 6) has finish_reason='stop' completed_result = results[6] assert len(completed_result.choices) > 0, "response.completed should have choices" - assert ( - completed_result.choices[0].finish_reason == "stop" - ), "response.completed should have finish_reason='stop'" + assert completed_result.choices[0].finish_reason == "stop", "response.completed should have finish_reason='stop'" # ============================================================================= @@ -1012,11 +1052,11 @@ def test_multiple_tool_calls_in_single_choice(): def test_map_reasoning_effort_adds_summary_detailed(): """ Test that _map_reasoning_effort behavior with reasoning_auto_summary flag. - + By default (flag=False), summary should NOT be added to avoid: 1. Breaking for users without verified OpenAI orgs (400 errors) 2. Making requests more expensive by including summary reasoning tokens - + When flag is enabled (flag=True or env var), summary="detailed" is added. """ import os @@ -1030,64 +1070,68 @@ def test_map_reasoning_effort_adds_summary_detailed(): # Test all string effort levels - DEFAULT BEHAVIOR (no summary) effort_levels = ["none", "low", "medium", "high", "xhigh", "minimal"] - + # Save original flag value original_flag = litellm.reasoning_auto_summary original_env = os.environ.get("LITELLM_REASONING_AUTO_SUMMARY") - + try: # Test 1: Default behavior (flag=False, no env var) - NO summary litellm.reasoning_auto_summary = False if "LITELLM_REASONING_AUTO_SUMMARY" in os.environ: del os.environ["LITELLM_REASONING_AUTO_SUMMARY"] - + for effort in effort_levels: result = handler._map_reasoning_effort(effort) - + assert result is not None, f"Result should not be None for effort={effort}" assert result["effort"] == effort, f"Effort should be {effort}" assert "summary" not in result, f"Summary should NOT be present by default for effort={effort}" - + print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}' (no summary by default)") - + # Test 2: With flag enabled - summary IS added litellm.reasoning_auto_summary = True - + for effort in effort_levels: result = handler._map_reasoning_effort(effort) - + assert result is not None, f"Result should not be None for effort={effort}" assert result["effort"] == effort, f"Effort should be {effort}" - assert result["summary"] == "detailed", f"Summary should be 'detailed' when flag is enabled for effort={effort}" - - print(f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed' (flag enabled)") - + assert result["summary"] == "detailed", ( + f"Summary should be 'detailed' when flag is enabled for effort={effort}" + ) + + print( + f"✓ reasoning_effort='{effort}' correctly maps to effort='{effort}', summary='detailed' (flag enabled)" + ) + # Test 3: With env var enabled (flag disabled) - summary IS added litellm.reasoning_auto_summary = False os.environ["LITELLM_REASONING_AUTO_SUMMARY"] = "true" - + result = handler._map_reasoning_effort("high") assert result["summary"] == "detailed", "Summary should be 'detailed' when env var is enabled" print("✓ LITELLM_REASONING_AUTO_SUMMARY env var works correctly") - + # Test 4: Dict input is passed through as-is (no modification) litellm.reasoning_auto_summary = False if "LITELLM_REASONING_AUTO_SUMMARY" in os.environ: del os.environ["LITELLM_REASONING_AUTO_SUMMARY"] - + dict_input = {"effort": "high", "summary": "custom_summary"} result_dict = handler._map_reasoning_effort(dict_input) assert result_dict["effort"] == "high" assert result_dict["summary"] == "custom_summary" print("✓ Dict input is passed through without modification") - + # Test 5: None/unknown values return None result_unknown = handler._map_reasoning_effort("unknown_value") assert result_unknown is None print("✓ Unknown reasoning_effort values return None") - + print("✓ All reasoning_effort behaviors work correctly with flag/env var control") - + finally: # Restore original values litellm.reasoning_auto_summary = original_flag @@ -1100,10 +1144,10 @@ def test_map_reasoning_effort_adds_summary_detailed(): def test_transform_response_preserves_annotations(): """ Test that annotations from Responses API are preserved when transforming to Chat Completions format. - + This is a regression test for the bug where annotations (like url_citation) were being dropped during the transformation from ResponsesAPIResponse to ModelResponse. - + The fix ensures annotations are extracted from ResponseOutputText content items and passed through to the Message object in the Chat Completions response. """ @@ -1162,13 +1206,9 @@ def test_transform_response_preserves_annotations(): # Create usage information usage = ResponseAPIUsage( input_tokens=10, - input_tokens_details=InputTokensDetails( - audio_tokens=None, cached_tokens=0, text_tokens=None - ), + input_tokens_details=InputTokensDetails(audio_tokens=None, cached_tokens=0, text_tokens=None), output_tokens=20, - output_tokens_details=OutputTokensDetails( - reasoning_tokens=0, text_tokens=None - ), + output_tokens_details=OutputTokensDetails(reasoning_tokens=0, text_tokens=None), total_tokens=30, cost=None, ) diff --git a/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py b/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py new file mode 100644 index 00000000000..852077dcf0c --- /dev/null +++ b/tests/test_litellm/enterprise/proxy/test_file_deletion_blocking.py @@ -0,0 +1,525 @@ +""" +Tests for file deletion blocking when referenced by non-terminal batches. + +This tests the feature where file deletion is blocked when: +1. File is referenced by a batch in non-terminal state (validating, in_progress, finalizing) +2. Batch polling is configured (proxy_batch_polling_interval > 0) + +This ensures cost tracking is not disrupted by premature file deletion. +""" + +import base64 +import json +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import HTTPException + +from litellm.proxy._types import UserAPIKeyAuth + + +def _make_unified_file_id(file_id: str = "file-abc123") -> str: + """Create a base64-encoded unified file ID.""" + raw = f"litellm_proxy:application/json;unified_id,test-{file_id};target_model_names,azure-gpt-4;llm_output_file_id,{file_id};llm_output_file_model_id,model-123" + return base64.urlsafe_b64encode(raw.encode()).decode().rstrip("=") + + +def _make_unified_batch_id(batch_id: str = "batch-123") -> str: + """Create a base64-encoded unified batch ID.""" + raw = f"litellm_proxy;model_id:model-deploy-xyz;llm_batch_id:{batch_id};llm_output_file_id:file-output" + return base64.urlsafe_b64encode(raw.encode()).decode().rstrip("=") + + +def _make_user_api_key_dict(user_id: str = "user-A") -> UserAPIKeyAuth: + return UserAPIKeyAuth( + api_key="sk-test", + user_id=user_id, + parent_otel_span=None, + ) + + +def _make_batch_db_record( + unified_object_id: str, + status: str, + file_object: dict, + created_by: str = "user-A", +): + """Create a mock batch database record.""" + mock_batch = MagicMock() + mock_batch.unified_object_id = unified_object_id + mock_batch.status = status + mock_batch.file_object = json.dumps(file_object) + mock_batch.created_by = created_by + mock_batch.created_at = 1700000000 + return mock_batch + + +def _make_managed_files_instance_with_batches( + file_id: str, + batches: list, + file_created_by: str = "user-A", +): + """ + Create a _PROXY_LiteLLMManagedFiles instance with mocked DB and batches. + + Args: + file_id: The unified file ID + batches: List of batch records to return from DB + file_created_by: The user who created the file + """ + from litellm_enterprise.proxy.hooks.managed_files import ( + _PROXY_LiteLLMManagedFiles, + ) + + # Mock file record + mock_file_record = MagicMock() + mock_file_record.unified_file_id = file_id + mock_file_record.created_by = file_created_by + mock_file_record.model_mappings = {"model-123": "provider-file-abc"} + + # Mock prisma + mock_prisma = MagicMock() + + # Mock file table queries + mock_prisma.db.litellm_managedfiletable.find_first = AsyncMock( + return_value=mock_file_record + ) + mock_prisma.db.litellm_managedfiletable.delete = AsyncMock( + return_value=mock_file_record + ) + + # Mock batch/object table queries + mock_prisma.db.litellm_managedobjecttable.find_many = AsyncMock( + return_value=batches + ) + + # Mock cache + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value={ + "unified_file_id": file_id, + "model_mappings": {"model-123": "provider-file-abc"}, + "flat_model_file_ids": ["provider-file-abc"], + }) + mock_cache.async_set_cache = AsyncMock() + + instance = _PROXY_LiteLLMManagedFiles( + internal_usage_cache=mock_cache, + prisma_client=mock_prisma, + ) + return instance + + +# --- Test: Batch polling configuration check --- + + +def test_is_batch_polling_enabled_when_job_registered(): + """Test that batch polling is detected as enabled when scheduler job is registered.""" + from litellm_enterprise.proxy.hooks.managed_files import ( + _PROXY_LiteLLMManagedFiles, + ) + + instance = _PROXY_LiteLLMManagedFiles( + internal_usage_cache=MagicMock(), + prisma_client=MagicMock(), + ) + + # Mock scheduler with registered job + mock_scheduler = MagicMock() + mock_job = MagicMock() + mock_scheduler.get_job.return_value = mock_job + + with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler): + assert instance._is_batch_polling_enabled() is True + + +def test_is_batch_polling_disabled_when_job_not_registered(): + """Test that batch polling is detected as disabled when scheduler job is not registered.""" + from litellm_enterprise.proxy.hooks.managed_files import ( + _PROXY_LiteLLMManagedFiles, + ) + + instance = _PROXY_LiteLLMManagedFiles( + internal_usage_cache=MagicMock(), + prisma_client=MagicMock(), + ) + + # Mock scheduler without registered job + mock_scheduler = MagicMock() + mock_scheduler.get_job.return_value = None + + with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler): + assert instance._is_batch_polling_enabled() is False + + +def test_is_batch_polling_disabled_when_no_scheduler(): + """Test that batch polling is detected as disabled when scheduler is not available.""" + from litellm_enterprise.proxy.hooks.managed_files import ( + _PROXY_LiteLLMManagedFiles, + ) + + instance = _PROXY_LiteLLMManagedFiles( + internal_usage_cache=MagicMock(), + prisma_client=MagicMock(), + ) + + with patch("litellm.proxy.proxy_server.scheduler", None): + assert instance._is_batch_polling_enabled() is False + + +# --- Test: Finding batches referencing files --- + + +@pytest.mark.asyncio +async def test_get_batches_referencing_file_finds_batch_with_input_file(): + """Test finding a batch that references the file as input_file_id.""" + unified_file_id = _make_unified_file_id("file-input-123") + unified_batch_id = _make_unified_batch_id("batch-123") + + batch_file_object = { + "id": "batch-123", + "input_file_id": unified_file_id, # Batch references this file + "status": "validating", + } + + batch_record = _make_batch_db_record( + unified_object_id=unified_batch_id, + status="validating", + file_object=batch_file_object, + ) + + managed_files = _make_managed_files_instance_with_batches( + file_id=unified_file_id, + batches=[batch_record], + ) + + referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id) + + assert len(referencing_batches) == 1 + assert referencing_batches[0]["batch_id"] == unified_batch_id + assert referencing_batches[0]["status"] == "validating" + + +@pytest.mark.asyncio +async def test_get_batches_referencing_file_finds_batch_with_output_file(): + """Test finding a batch that references the file as output_file_id.""" + unified_file_id = _make_unified_file_id("file-output-456") + unified_batch_id = _make_unified_batch_id("batch-456") + + batch_file_object = { + "id": "batch-456", + "input_file_id": "file-input-different", + "output_file_id": unified_file_id, # Batch references this file + "status": "in_progress", + } + + batch_record = _make_batch_db_record( + unified_object_id=unified_batch_id, + status="in_progress", + file_object=batch_file_object, + ) + + managed_files = _make_managed_files_instance_with_batches( + file_id=unified_file_id, + batches=[batch_record], + ) + + referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id) + + assert len(referencing_batches) == 1 + assert referencing_batches[0]["status"] == "in_progress" + + +@pytest.mark.asyncio +async def test_get_batches_referencing_file_ignores_terminal_batches(): + """Test that batches in terminal states are not returned.""" + unified_file_id = _make_unified_file_id("file-123") + unified_batch_id = _make_unified_batch_id("batch-completed") + + batch_file_object = { + "id": "batch-completed", + "input_file_id": unified_file_id, + "status": "completed", + } + + # Batch is in terminal state in DB + batch_record = _make_batch_db_record( + unified_object_id=unified_batch_id, + status="completed", # Terminal state + file_object=batch_file_object, + ) + + managed_files = _make_managed_files_instance_with_batches( + file_id=unified_file_id, + batches=[], # Query returns no batches (terminal states filtered out) + ) + + referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id) + + assert len(referencing_batches) == 0 + + +@pytest.mark.asyncio +async def test_get_batches_referencing_file_finds_multiple_batches(): + """Test finding multiple batches referencing the same file.""" + unified_file_id = _make_unified_file_id("file-shared") + + batch1 = _make_batch_db_record( + unified_object_id=_make_unified_batch_id("batch-1"), + status="validating", + file_object={"id": "batch-1", "input_file_id": unified_file_id, "status": "validating"}, + ) + + batch2 = _make_batch_db_record( + unified_object_id=_make_unified_batch_id("batch-2"), + status="in_progress", + file_object={"id": "batch-2", "input_file_id": unified_file_id, "status": "in_progress"}, + ) + + managed_files = _make_managed_files_instance_with_batches( + file_id=unified_file_id, + batches=[batch1, batch2], + ) + + referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id) + + assert len(referencing_batches) == 2 + statuses = [b["status"] for b in referencing_batches] + assert "validating" in statuses + assert "in_progress" in statuses + + +# --- Test: File deletion blocking logic --- + + +@pytest.mark.asyncio +async def test_file_deletion_blocked_when_batch_polling_enabled_and_batch_references_file(): + """ + Test that file deletion is blocked when: + 1. Batch cost tracking job is registered (polling enabled) + 2. File is referenced by a non-terminal batch + """ + unified_file_id = _make_unified_file_id("file-to-delete") + unified_batch_id = _make_unified_batch_id("batch-active") + + batch_file_object = { + "id": "batch-active", + "input_file_id": unified_file_id, + "status": "validating", + } + + batch_record = _make_batch_db_record( + unified_object_id=unified_batch_id, + status="validating", + file_object=batch_file_object, + ) + + managed_files = _make_managed_files_instance_with_batches( + file_id=unified_file_id, + batches=[batch_record], + ) + + # Mock scheduler with registered batch cost job + mock_scheduler = MagicMock() + mock_scheduler.get_job.return_value = MagicMock() # Job exists + + with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler): + with pytest.raises(HTTPException) as exc_info: + await managed_files._check_file_deletion_allowed(unified_file_id) + + assert exc_info.value.status_code == 400 + error_detail = exc_info.value.detail + assert "Cannot delete file" in error_detail + assert unified_file_id in error_detail + assert "validating" in error_detail + assert "delete or cancel the referencing batch" in error_detail.lower() + + +@pytest.mark.asyncio +async def test_file_deletion_allowed_when_batch_polling_disabled(): + """ + Test that file deletion is allowed when batch cost tracking job is not registered, + even if there are non-terminal batches referencing the file. + """ + unified_file_id = _make_unified_file_id("file-to-delete") + unified_batch_id = _make_unified_batch_id("batch-active") + + batch_file_object = { + "id": "batch-active", + "input_file_id": unified_file_id, + "status": "validating", + } + + batch_record = _make_batch_db_record( + unified_object_id=unified_batch_id, + status="validating", + file_object=batch_file_object, + ) + + managed_files = _make_managed_files_instance_with_batches( + file_id=unified_file_id, + batches=[batch_record], + ) + + # Mock scheduler without registered job (batch cost tracking disabled) + mock_scheduler = MagicMock() + mock_scheduler.get_job.return_value = None + + with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler): + # Should not raise an exception + await managed_files._check_file_deletion_allowed(unified_file_id) + + +@pytest.mark.asyncio +async def test_file_deletion_allowed_when_no_batches_reference_file(): + """ + Test that file deletion is allowed when no batches reference the file, + even when batch cost tracking is enabled. + """ + unified_file_id = _make_unified_file_id("file-to-delete") + + managed_files = _make_managed_files_instance_with_batches( + file_id=unified_file_id, + batches=[], # No batches reference this file + ) + + # Mock scheduler with registered job (batch cost tracking enabled) + mock_scheduler = MagicMock() + mock_scheduler.get_job.return_value = MagicMock() + + with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler): + # Should not raise an exception + await managed_files._check_file_deletion_allowed(unified_file_id) + + +@pytest.mark.asyncio +async def test_afile_delete_calls_check_deletion_allowed(): + """ + Test that afile_delete calls _check_file_deletion_allowed before deleting. + """ + unified_file_id = _make_unified_file_id("file-to-delete") + unified_batch_id = _make_unified_batch_id("batch-active") + + batch_file_object = { + "id": "batch-active", + "input_file_id": unified_file_id, + "status": "in_progress", + } + + batch_record = _make_batch_db_record( + unified_object_id=unified_batch_id, + status="in_progress", + file_object=batch_file_object, + ) + + managed_files = _make_managed_files_instance_with_batches( + file_id=unified_file_id, + batches=[batch_record], + ) + + # Mock llm_router + mock_router = MagicMock() + mock_router.afile_delete = AsyncMock() + + # Mock scheduler with registered job + mock_scheduler = MagicMock() + mock_scheduler.get_job.return_value = MagicMock() + + with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler): + with pytest.raises(HTTPException) as exc_info: + await managed_files.afile_delete( + file_id=unified_file_id, + litellm_parent_otel_span=None, + llm_router=mock_router, + ) + + # Should raise error before calling router delete + assert exc_info.value.status_code == 400 + mock_router.afile_delete.assert_not_called() + + +@pytest.mark.asyncio +async def test_database_limit_respected(): + """ + Test that we only fetch 10 batches from DB (not 500). + This is a performance optimization - we only fetch what we need. + """ + unified_file_id = _make_unified_file_id("file-shared") + + # Create exactly 10 batches (what DB will return with take=10) + ten_batches = [] + for i in range(10): + batch = _make_batch_db_record( + unified_object_id=_make_unified_batch_id(f"batch-{i}"), + status="validating", + file_object={ + "id": f"batch-{i}", + "input_file_id": unified_file_id, + "status": "validating" + }, + ) + ten_batches.append(batch) + + # Mock will return only 10 batches (as DB would with take=10) + managed_files = _make_managed_files_instance_with_batches( + file_id=unified_file_id, + batches=ten_batches, + ) + + referencing_batches = await managed_files._get_batches_referencing_file(unified_file_id) + + # Should return all 10 that reference the file + assert len(referencing_batches) == 10 + + # Verify error message handles "10+" case (since we got exactly 10, might be more in DB) + mock_scheduler = MagicMock() + mock_scheduler.get_job.return_value = MagicMock() + + with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler): + with pytest.raises(HTTPException) as exc_info: + await managed_files._check_file_deletion_allowed(unified_file_id) + + error_detail = exc_info.value.detail + # When we get exactly 10 matches, show "10+" to indicate there might be more + assert "10+ batch(es)" in error_detail + + +@pytest.mark.asyncio +async def test_error_message_includes_batch_details(): + """ + Test that the error message includes helpful information about the blocking batches. + """ + unified_file_id = _make_unified_file_id("file-to-delete") + batch1_id = _make_unified_batch_id("batch-1") + batch2_id = _make_unified_batch_id("batch-2") + + batch1 = _make_batch_db_record( + unified_object_id=batch1_id, + status="validating", + file_object={"id": "batch-1", "input_file_id": unified_file_id, "status": "validating"}, + ) + + batch2 = _make_batch_db_record( + unified_object_id=batch2_id, + status="in_progress", + file_object={"id": "batch-2", "output_file_id": unified_file_id, "status": "in_progress"}, + ) + + managed_files = _make_managed_files_instance_with_batches( + file_id=unified_file_id, + batches=[batch1, batch2], + ) + + # Mock scheduler with registered job + mock_scheduler = MagicMock() + mock_scheduler.get_job.return_value = MagicMock() + + with patch("litellm.proxy.proxy_server.scheduler", mock_scheduler): + with pytest.raises(HTTPException) as exc_info: + await managed_files._check_file_deletion_allowed(unified_file_id) + + error_detail = exc_info.value.detail + assert "2 batch(es)" in error_detail + assert "validating" in error_detail + assert "in_progress" in error_detail + assert "complete cost tracking" in error_detail.lower() + assert "delete or cancel the referencing batch" in error_detail.lower() diff --git a/tests/test_litellm/integrations/datadog/test_datadog_tags_regression.py b/tests/test_litellm/integrations/datadog/test_datadog_tags_regression.py new file mode 100644 index 00000000000..3f1d2be4137 --- /dev/null +++ b/tests/test_litellm/integrations/datadog/test_datadog_tags_regression.py @@ -0,0 +1,91 @@ +import os +import sys +from unittest.mock import patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../")) + +from litellm.integrations.datadog.datadog_handler import get_datadog_tags +from litellm.integrations.datadog.datadog_cost_management import ( + DatadogCostManagementLogger, +) +from litellm.types.utils import StandardLoggingPayload, StandardLoggingMetadata + + +class TestDatadogTagsRegression: + @pytest.fixture + def mock_env_vars(self): + """Mock environment variables to isolate environment.""" + with patch.dict( + os.environ, + { + "DD_ENV": "test-env", + "DD_SERVICE": "test-service", + "DD_VERSION": "1.0.0", + "HOSTNAME": "test-host", + "POD_NAME": "test-pod", + "DD_API_KEY": "mock-api-key", + "DD_APP_KEY": "mock-app-key", + }, + ): + yield + + def test_get_datadog_tags_regression(self, mock_env_vars): + """ + Regression Test: Ensure that get_datadog_tags still produces basic tags correctly + AND now includes the new team tag when provided. + """ + # Case 1: Legacy behavior (no team info) + payload_legacy = StandardLoggingPayload(metadata={}) + tags_legacy = get_datadog_tags(payload_legacy) + + # Verify base tags exist (legacy requirement) + assert "env:test-env" in tags_legacy + assert "service:test-service" in tags_legacy + # Verify NO team tag (should not invent one) + assert "team:" not in tags_legacy + + # Case 2: New feature (team info provided) + payload_with_team = StandardLoggingPayload( + metadata=StandardLoggingMetadata(user_api_key_team_alias="regression-team") + ) + tags_with_team = get_datadog_tags(payload_with_team) + + # Verify base tags STILL exist + assert "env:test-env" in tags_with_team + assert "service:test-service" in tags_with_team + # Verify NEW team tag is added + assert "team:regression-team" in tags_with_team + + @pytest.mark.asyncio + async def test_datadog_cost_management_tags_regression(self, mock_env_vars): + """ + Regression Test: Ensure DatadogCostManagementLogger extracts tags correctly, + preserving existing behavior while adding the team tag capability. + """ + logger = DatadogCostManagementLogger() + + # Case 1: Legacy metadata (user alias only) + payload_legacy = StandardLoggingPayload( + metadata=StandardLoggingMetadata(user_api_key_alias="legacy-user") + ) + + tags_legacy = logger._extract_tags(payload_legacy) + + assert tags_legacy["env"] == "test-env" + assert tags_legacy["user"] == "legacy-user" + assert "team" not in tags_legacy # Should not exist + + # Case 2: New metadata (team alias) + payload_new = StandardLoggingPayload( + metadata=StandardLoggingMetadata( + user_api_key_alias="new-user", user_api_key_team_alias="new-team-alias" + ) + ) + + tags_new = logger._extract_tags(payload_new) + + assert tags_new["env"] == "test-env" + assert tags_new["user"] == "new-user" + assert tags_new["team"] == "new-team-alias" # New feature verified diff --git a/tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py b/tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py index 4a9fa3de5fd..4bfa3a581e3 100644 --- a/tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py +++ b/tests/test_litellm/integrations/test_prometheus_client_ip_user_agent.py @@ -1,10 +1,12 @@ +from unittest.mock import AsyncMock, MagicMock, patch + import pytest -from unittest.mock import MagicMock, patch + from litellm.integrations.prometheus import PrometheusLogger +from litellm.proxy._types import UserAPIKeyAuth from litellm.types.integrations.prometheus import ( UserAPIKeyLabelValues, ) -from litellm.proxy._types import UserAPIKeyAuth @pytest.mark.asyncio @@ -72,10 +74,12 @@ async def test_async_post_call_failure_hook_includes_client_ip_user_agent(): @pytest.mark.asyncio async def test_async_post_call_success_hook_includes_client_ip_user_agent(): """ - Test that async_post_call_success_hook includes client_ip and user_agent in UserAPIKeyLabelValues + Test that async_log_success_event includes client_ip and user_agent in UserAPIKeyLabelValues. + + Note: After PR #21159, the metric increment was moved from async_post_call_success_hook + to async_log_success_event to prevent double-counting. """ # Mocking - # Mocking with patch( "litellm.integrations.prometheus.PrometheusLogger.__init__", return_value=None ): @@ -84,16 +88,43 @@ async def test_async_post_call_success_hook_includes_client_ip_user_agent(): logger.get_labels_for_metric = MagicMock( return_value=["client_ip", "user_agent"] ) + logger._should_skip_metrics_for_invalid_key = MagicMock(return_value=False) + logger._increment_top_level_request_and_spend_metrics = MagicMock() + logger._increment_token_metrics = MagicMock() + logger._increment_remaining_budget_metrics = AsyncMock() + logger._set_virtual_key_rate_limit_metrics = MagicMock() + logger._set_latency_metrics = MagicMock() + logger.set_llm_deployment_success_metrics = MagicMock() + logger._increment_cache_metrics = MagicMock() - data = { + kwargs = { "model": "gpt-4", - "metadata": { - "requester_ip_address": "192.168.1.1", - "user_agent": "success-agent", + "litellm_params": { + "metadata": {} + }, + "start_time": None, + "standard_logging_object": { + "model_group": "gpt-4", + "model_id": "model_1", + "api_base": "http://api.base", + "custom_llm_provider": "openai", + "completion_tokens": 10, + "total_tokens": 20, + "response_cost": 0.01, + "request_tags": [], + "metadata": { + "user_api_key_user_id": "user_1", + "user_api_key_hash": "hash_1", + "user_api_key_alias": "alias_1", + "user_api_key_team_id": "team_1", + "user_api_key_team_alias": "team_alias_1", + "user_api_key_user_email": "test@example.com", + "user_api_key_request_route": "/chat/completions", + "requester_ip_address": "192.168.1.1", + "user_agent": "success-agent", + }, }, } - user_api_key_dict = UserAPIKeyAuth(token="test_token") - response = MagicMock() # Mock prometheus_label_factory to inspect arguments with patch( @@ -101,10 +132,11 @@ async def test_async_post_call_success_hook_includes_client_ip_user_agent(): ) as mock_label_factory: mock_label_factory.return_value = {} - await logger.async_post_call_success_hook( - data=data, - user_api_key_dict=user_api_key_dict, - response=response, + await logger.async_log_success_event( + kwargs=kwargs, + response_obj=None, + start_time=None, + end_time=None, ) # Verification @@ -114,8 +146,8 @@ async def test_async_post_call_success_hook_includes_client_ip_user_agent(): calls = mock_label_factory.call_args_list found = False for call in calls: - kwargs = call.kwargs - enum_values = kwargs.get("enum_values") + kwargs_args = call.kwargs + enum_values = kwargs_args.get("enum_values") if isinstance(enum_values, UserAPIKeyLabelValues): if ( enum_values.client_ip == "192.168.1.1" diff --git a/tests/test_litellm/integrations/test_prometheus_none_metadata.py b/tests/test_litellm/integrations/test_prometheus_none_metadata.py new file mode 100644 index 00000000000..fff2e48bf5a --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_none_metadata.py @@ -0,0 +1,176 @@ +""" +Unit tests for Prometheus handling of None metadata in litellm_params. + +When the Responses API sends streaming requests, litellm_params.metadata +can be None, causing AttributeError: 'NoneType' object has no attribute 'get' +in set_llm_deployment_success_metrics. +""" + +import os +import sys +from datetime import datetime + +import pytest +from prometheus_client import REGISTRY + +sys.path.insert(0, os.path.abspath("../../..")) + +from litellm.integrations.prometheus import PrometheusLogger +from litellm.types.integrations.prometheus import UserAPIKeyLabelValues + + +@pytest.fixture(scope="function") +def prometheus_logger(): + """Create a PrometheusLogger instance for testing.""" + collectors = list(REGISTRY._collector_to_names.keys()) + for collector in collectors: + REGISTRY.unregister(collector) + return PrometheusLogger() + + +class TestNoneMetadataHandling: + """ + Test that Prometheus metrics don't crash when metadata is None. + + This targets the bug where Responses API streaming sets + litellm_params["metadata"] = None, causing: + _metadata.get("model_info") -> AttributeError + """ + + def test_set_llm_deployment_success_metrics_with_none_metadata( + self, prometheus_logger + ): + """ + set_llm_deployment_success_metrics should not raise when + litellm_params.metadata is None. + """ + request_kwargs = { + "litellm_params": { + "metadata": None, # Bug trigger + "custom_llm_provider": "openai", + }, + "model": "gpt-4o", + "standard_logging_object": { + "api_base": "https://api.openai.com", + "hidden_params": { + "additional_headers": None, + "litellm_overhead_time_ms": None, + }, + "metadata": { + "user_api_key_hash": "test-key", + "user_api_key_alias": None, + "user_api_key_team_id": None, + "user_api_key_team_alias": None, + }, + "model": "gpt-4o", + "response_cost": 0.001, + }, + } + enum_values = UserAPIKeyLabelValues( + end_user=None, + hashed_api_key="test-key", + api_key_alias=None, + team=None, + team_alias=None, + requested_model="gpt-4o", + ) + + # Should not raise AttributeError + prometheus_logger.set_llm_deployment_success_metrics( + request_kwargs=request_kwargs, + start_time=datetime.now(), + end_time=datetime.now(), + enum_values=enum_values, + output_tokens=10.0, + ) + + def test_set_llm_deployment_success_metrics_with_missing_litellm_params( + self, prometheus_logger + ): + """ + set_llm_deployment_success_metrics should not raise when + litellm_params is missing entirely. + """ + request_kwargs = { + "model": "gpt-4o", + "standard_logging_object": { + "api_base": "https://api.openai.com", + "hidden_params": { + "additional_headers": None, + "litellm_overhead_time_ms": None, + }, + "metadata": { + "user_api_key_hash": "test-key", + "user_api_key_alias": None, + "user_api_key_team_id": None, + "user_api_key_team_alias": None, + }, + "model": "gpt-4o", + "response_cost": 0.001, + }, + } + enum_values = UserAPIKeyLabelValues( + end_user=None, + hashed_api_key="test-key", + api_key_alias=None, + team=None, + team_alias=None, + requested_model="gpt-4o", + ) + + # Should not raise + prometheus_logger.set_llm_deployment_success_metrics( + request_kwargs=request_kwargs, + start_time=datetime.now(), + end_time=datetime.now(), + enum_values=enum_values, + output_tokens=10.0, + ) + + def test_set_llm_deployment_success_metrics_with_litellm_metadata_key( + self, prometheus_logger + ): + """ + set_llm_deployment_success_metrics should pick up litellm_metadata + when metadata is None, using get_litellm_metadata_from_kwargs. + """ + request_kwargs = { + "litellm_params": { + "metadata": None, + "litellm_metadata": {"model_info": {"id": "test-model-id"}}, + "custom_llm_provider": "openai", + }, + "model": "gpt-4o", + "standard_logging_object": { + "api_base": "https://api.openai.com", + "hidden_params": { + "additional_headers": None, + "litellm_overhead_time_ms": None, + }, + "metadata": { + "user_api_key_hash": "test-key", + "user_api_key_alias": None, + "user_api_key_team_id": None, + "user_api_key_team_alias": None, + }, + "model": "gpt-4o", + "response_cost": 0.001, + }, + } + enum_values = UserAPIKeyLabelValues( + end_user=None, + hashed_api_key="test-key", + api_key_alias=None, + team=None, + team_alias=None, + requested_model="gpt-4o", + ) + + # Should not raise, and should pick up litellm_metadata + prometheus_logger.set_llm_deployment_success_metrics( + request_kwargs=request_kwargs, + start_time=datetime.now(), + end_time=datetime.now(), + enum_values=enum_values, + output_tokens=10.0, + ) diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index f9e5c6d0252..1ea1374cfb3 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -1811,3 +1811,131 @@ def test_translate_openai_response_to_anthropic_input_tokens_no_cache(): # Validate: input_tokens should equal prompt_tokens when no caching assert anthropic_response["usage"]["input_tokens"] == 100 assert anthropic_response["usage"]["output_tokens"] == 50 + + +# ===================================================================== +# Web Search Tool Transformation Tests +# ===================================================================== + + +def test_is_web_search_tool(): + """Test detection of Anthropic web search tools.""" + adapter = LiteLLMAnthropicMessagesAdapter() + + # Tool with type starting with "web_search" should be detected + web_search_tool_with_type = { + "type": "web_search_20260209", + "name": "web_search", + } + assert adapter._is_web_search_tool(web_search_tool_with_type) is True + + # Tool with name "web_search" should be detected + web_search_tool_with_name = { + "name": "web_search", + } + assert adapter._is_web_search_tool(web_search_tool_with_name) is True + + # Regular function tool should not be detected + regular_tool = { + "name": "get_weather", + "description": "Get weather info", + "input_schema": {"type": "object"}, + } + assert adapter._is_web_search_tool(regular_tool) is False + + +def test_translate_anthropic_to_openai_with_web_search_tool(): + """ + Test that Anthropic web search tools are converted to web_search_options parameter. + + When a user sends an Anthropic /v1/messages request with {"type": "web_search_20260209"} + tool, it should be transformed to OpenAI format with web_search_options: {} parameter. + """ + from litellm.types.llms.anthropic import AnthropicMessagesRequest + + anthropic_request = AnthropicMessagesRequest( + model="gemini-2.5-flash-lite", + max_tokens=4096, + messages=[ + { + "role": "user", + "content": "Search for the current prices of AAPL and GOOGL", + } + ], + tools=[ + { + "type": "web_search_20260209", + "name": "web_search", + } + ], + ) + + adapter = LiteLLMAnthropicMessagesAdapter() + openai_request, tool_name_mapping = adapter.translate_anthropic_to_openai( + anthropic_message_request=anthropic_request + ) + + # web_search_options should be added + assert "web_search_options" in openai_request + assert openai_request["web_search_options"] == {} + + # web search tool should NOT be in the tools array + assert "tools" not in openai_request or openai_request.get("tools") == [] + + # tool_name_mapping should be empty since no regular tools were present + assert tool_name_mapping == {} + + +def test_translate_anthropic_to_openai_with_mixed_tools(): + """ + Test that web search tools are separated from regular tools. + + When a request has both web search tools and regular function tools, + only the regular tools should be in the tools array, and web_search_options + should be added. + """ + from litellm.types.llms.anthropic import AnthropicMessagesRequest + + anthropic_request = AnthropicMessagesRequest( + model="gemini-2.5-flash-lite", + max_tokens=4096, + messages=[ + { + "role": "user", + "content": "Get weather and search the web", + } + ], + tools=[ + { + "type": "web_search_20260209", + "name": "web_search", + }, + { + "name": "get_weather", + "description": "Get weather information", + "input_schema": { + "type": "object", + "properties": { + "location": {"type": "string"} + }, + }, + }, + ], + ) + + adapter = LiteLLMAnthropicMessagesAdapter() + openai_request, tool_name_mapping = adapter.translate_anthropic_to_openai( + anthropic_message_request=anthropic_request + ) + + # web_search_options should be added + assert "web_search_options" in openai_request + assert openai_request["web_search_options"] == {} + + # Only get_weather tool should be in the tools array + assert "tools" in openai_request + assert len(openai_request["tools"]) == 1 + assert openai_request["tools"][0]["function"]["name"] == "get_weather" + + # tool_name_mapping should be empty for short tool names + assert tool_name_mapping == {} diff --git a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py index a321a24540f..ebffb56446e 100644 --- a/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/test_litellm/llms/anthropic/test_anthropic_common_utils.py @@ -283,3 +283,142 @@ class TestPassthroughOAuth: assert updated_headers["x-api-key"] == FAKE_REGULAR_KEY assert "authorization" not in updated_headers + + +class TestIsAnthropicOAuthKey: + """Tests for is_anthropic_oauth_key helper function.""" + + def test_oauth_token_raw(self): + """Raw OAuth token should be detected.""" + from litellm.llms.anthropic.common_utils import is_anthropic_oauth_key + + assert is_anthropic_oauth_key("sk-ant-oat01-abc123") is True + assert is_anthropic_oauth_key("sk-ant-oat02-xyz789") is True + + def test_oauth_token_bearer_format(self): + """Bearer-prefixed OAuth token should be detected.""" + from litellm.llms.anthropic.common_utils import is_anthropic_oauth_key + + assert is_anthropic_oauth_key("Bearer sk-ant-oat01-abc123") is True + assert is_anthropic_oauth_key("Bearer sk-ant-oat02-xyz789") is True + + def test_non_oauth_tokens(self): + """Non-OAuth values should return False.""" + from litellm.llms.anthropic.common_utils import is_anthropic_oauth_key + + assert is_anthropic_oauth_key(None) is False + assert is_anthropic_oauth_key("") is False + assert is_anthropic_oauth_key("sk-ant-api01-abc123") is False + assert is_anthropic_oauth_key("Bearer sk-ant-api01-abc123") is False + + def test_case_sensitivity(self): + """OAuth prefix matching should be case-sensitive.""" + from litellm.llms.anthropic.common_utils import is_anthropic_oauth_key + + assert is_anthropic_oauth_key("sk-ant-OAT01-abc123") is False + assert is_anthropic_oauth_key("SK-ANT-OAT01-abc123") is False + + def test_just_prefix(self): + """Just the prefix with no suffix should still match.""" + from litellm.llms.anthropic.common_utils import is_anthropic_oauth_key + + assert is_anthropic_oauth_key("sk-ant-oat") is True + + +class TestProxyOAuthHeaderForwarding: + """Tests for proxy-layer OAuth header preservation and forwarding.""" + + def test_clean_headers_preserves_oauth_authorization(self): + """clean_headers should preserve Authorization header with OAuth tokens.""" + from starlette.datastructures import Headers + + from litellm.proxy.litellm_pre_call_utils import clean_headers + + raw_headers = Headers( + raw=[ + (b"authorization", f"Bearer {FAKE_OAUTH_TOKEN}".encode()), + (b"content-type", b"application/json"), + ] + ) + cleaned = clean_headers(raw_headers) + + assert "authorization" in cleaned + assert cleaned["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" + assert cleaned["content-type"] == "application/json" + + def test_clean_headers_strips_non_oauth_authorization(self): + """clean_headers should strip Authorization header with regular API keys.""" + from starlette.datastructures import Headers + + from litellm.proxy.litellm_pre_call_utils import clean_headers + + raw_headers = Headers( + raw=[ + (b"authorization", b"Bearer sk-regular-key-123"), + (b"content-type", b"application/json"), + ] + ) + cleaned = clean_headers(raw_headers) + + assert "authorization" not in cleaned + assert cleaned["content-type"] == "application/json" + + def test_add_provider_specific_headers_forwards_oauth(self): + """add_provider_specific_headers_to_request should forward OAuth Authorization + as a ProviderSpecificHeader scoped to Anthropic-compatible providers.""" + from litellm.proxy.litellm_pre_call_utils import ( + add_provider_specific_headers_to_request, + ) + + data: dict = {} + headers = { + "authorization": f"Bearer {FAKE_OAUTH_TOKEN}", + "content-type": "application/json", + } + + add_provider_specific_headers_to_request(data=data, headers=headers) + + assert "provider_specific_header" in data + psh = data["provider_specific_header"] + assert "anthropic" in psh["custom_llm_provider"] + assert "bedrock" in psh["custom_llm_provider"] + assert "vertex_ai" in psh["custom_llm_provider"] + assert psh["extra_headers"]["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" + + def test_add_provider_specific_headers_ignores_non_oauth(self): + """add_provider_specific_headers_to_request should not create a + ProviderSpecificHeader for non-OAuth Authorization headers.""" + from litellm.proxy.litellm_pre_call_utils import ( + add_provider_specific_headers_to_request, + ) + + data: dict = {} + headers = { + "authorization": "Bearer sk-regular-key-123", + "content-type": "application/json", + } + + add_provider_specific_headers_to_request(data=data, headers=headers) + + assert "provider_specific_header" not in data + + def test_add_provider_specific_headers_combines_anthropic_and_oauth(self): + """When both anthropic-beta and OAuth Authorization are present, both + should be included in the ProviderSpecificHeader.""" + from litellm.proxy.litellm_pre_call_utils import ( + add_provider_specific_headers_to_request, + ) + + data: dict = {} + headers = { + "authorization": f"Bearer {FAKE_OAUTH_TOKEN}", + "anthropic-beta": "oauth-2025-04-20", + "content-type": "application/json", + } + + add_provider_specific_headers_to_request(data=data, headers=headers) + + assert "provider_specific_header" in data + psh = data["provider_specific_header"] + assert psh["extra_headers"]["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" + assert psh["extra_headers"]["anthropic-beta"] == "oauth-2025-04-20" diff --git a/tests/test_litellm/llms/anthropic/test_message_sanitization.py b/tests/test_litellm/llms/anthropic/test_message_sanitization.py new file mode 100644 index 00000000000..973f2897884 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/test_message_sanitization.py @@ -0,0 +1,380 @@ +""" +Test message sanitization for Anthropic API when modify_params=True + +Tests three cases: +A. Missing tool_result for tool_use (orphaned tool calls) +B. Orphaned tool_result without matching tool_use +C. Empty text content +""" + +import pytest +import sys +import os + +# Add the parent directory to the path so we can import litellm +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))) + +import litellm +from litellm.litellm_core_utils.prompt_templates.factory import ( + sanitize_messages_for_tool_calling, + anthropic_messages_pt, +) + + +class TestMessageSanitization: + """Test message sanitization for tool calling scenarios""" + + def setup_method(self): + """Setup for each test""" + # Save original modify_params value + self.original_modify_params = litellm.modify_params + litellm.modify_params = True + + def teardown_method(self): + """Cleanup after each test""" + # Restore original modify_params value + litellm.modify_params = self.original_modify_params + + def test_case_a_orphaned_tool_call_single(self): + """ + Test Case A: Assistant message with tool_calls but no tool result + Should add a dummy tool result message + """ + messages = [ + { + "role": "user", + "content": "What is the weather in Nashik?" + }, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "toolu_01Kus2cC3ydjBW7UK4GJqBP4", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "Nashik, India"}' + } + } + ] + } + ] + + sanitized = sanitize_messages_for_tool_calling(messages) + + # Should have 3 messages: user, assistant, and dummy tool result + assert len(sanitized) == 3 + assert sanitized[0]["role"] == "user" + assert sanitized[1]["role"] == "assistant" + assert sanitized[2]["role"] == "tool" + assert sanitized[2]["tool_call_id"] == "toolu_01Kus2cC3ydjBW7UK4GJqBP4" + assert "skipped" in sanitized[2]["content"].lower() or "interrupted" in sanitized[2]["content"].lower() + assert "get_weather" in sanitized[2]["content"] + + def test_case_a_orphaned_tool_call_multiple(self): + """ + Test Case A: Assistant message with multiple tool_calls, some missing results + """ + messages = [ + { + "role": "user", + "content": "Get weather for Nashik and Mumbai" + }, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "Nashik"}' + } + }, + { + "id": "call_2", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "Mumbai"}' + } + } + ] + }, + { + "role": "tool", + "tool_call_id": "call_1", + "content": "Weather in Nashik: 25°C" + } + ] + + sanitized = sanitize_messages_for_tool_calling(messages) + + # Should have 4 messages: user, assistant, tool result for call_1, dummy for call_2 + assert len(sanitized) == 4 + assert sanitized[0]["role"] == "user" + assert sanitized[1]["role"] == "assistant" + assert sanitized[2]["tool_call_id"] == "call_1" # Original tool result (first in tool_calls) + assert sanitized[3]["tool_call_id"] == "call_2" # Dummy added for missing call_2 + + def test_case_b_orphaned_tool_result(self): + """ + Test Case B: Tool result without matching tool_call in previous assistant message + Should remove the orphaned tool result + """ + messages = [ + { + "role": "user", + "content": "Hello" + }, + { + "role": "assistant", + "content": "Hi there!" + }, + { + "role": "tool", + "tool_call_id": "nonexistent_id", + "content": "Some result" + } + ] + + sanitized = sanitize_messages_for_tool_calling(messages) + + # Should have only 2 messages, orphaned tool result removed + assert len(sanitized) == 2 + assert sanitized[0]["role"] == "user" + assert sanitized[1]["role"] == "assistant" + + def test_case_b_valid_tool_result_preserved(self): + """ + Test Case B: Valid tool result with matching tool_call should be preserved + """ + 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": "Boston"}' + } + } + ] + }, + { + "role": "tool", + "tool_call_id": "call_123", + "content": "Weather: 20°C" + } + ] + + sanitized = sanitize_messages_for_tool_calling(messages) + + # All messages should be preserved + assert len(sanitized) == 3 + assert sanitized[2]["role"] == "tool" + assert sanitized[2]["tool_call_id"] == "call_123" + + def test_case_c_empty_text_content_user(self): + """ + Test Case C: Empty text content in user message + Should replace with placeholder + """ + messages = [ + { + "role": "user", + "content": "" + }, + { + "role": "assistant", + "content": "Hello!" + } + ] + + sanitized = sanitize_messages_for_tool_calling(messages) + + assert len(sanitized) == 2 + assert sanitized[0]["role"] == "user" + assert sanitized[0]["content"] == "[System: Empty message content sanitised to satisfy protocol]" + + def test_case_c_whitespace_only_content(self): + """ + Test Case C: Whitespace-only content + Should replace with placeholder + """ + messages = [ + { + "role": "user", + "content": " \n \t " + }, + { + "role": "assistant", + "content": " " + } + ] + + sanitized = sanitize_messages_for_tool_calling(messages) + + assert len(sanitized) == 2 + assert sanitized[0]["content"] == "[System: Empty message content sanitised to satisfy protocol]" + assert sanitized[1]["content"] == "[System: Empty message content sanitised to satisfy protocol]" + + def test_case_c_valid_content_preserved(self): + """ + Test Case C: Valid non-empty content should be preserved + """ + messages = [ + { + "role": "user", + "content": "Hello" + }, + { + "role": "assistant", + "content": "Hi there!" + } + ] + + sanitized = sanitize_messages_for_tool_calling(messages) + + assert len(sanitized) == 2 + assert sanitized[0]["content"] == "Hello" + assert sanitized[1]["content"] == "Hi there!" + + def test_combined_cases(self): + """ + Test combination of multiple cases + """ + messages = [ + { + "role": "user", + "content": "Get weather" + }, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "NYC"}' + } + } + ] + }, + # Missing tool result for call_1 + { + "role": "user", + "content": "" # Empty content + }, + { + "role": "assistant", + "content": "Response" + }, + { + "role": "tool", + "tool_call_id": "orphaned_id", # Orphaned tool result + "content": "Some data" + } + ] + + sanitized = sanitize_messages_for_tool_calling(messages) + + # Should have: user, assistant, dummy tool result, user (sanitized), assistant + # Orphaned tool result should be removed + assert len(sanitized) == 5 + assert sanitized[0]["role"] == "user" + assert sanitized[1]["role"] == "assistant" + assert sanitized[2]["role"] == "tool" + assert sanitized[2]["tool_call_id"] == "call_1" # Dummy added + assert sanitized[3]["role"] == "user" + assert sanitized[3]["content"] == "[System: Empty message content sanitised to satisfy protocol]" + assert sanitized[4]["role"] == "assistant" + + def test_modify_params_false_no_sanitization(self): + """ + Test that sanitization is skipped when modify_params=False + """ + litellm.modify_params = False + + messages = [ + { + "role": "user", + "content": "" + }, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "call_1", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{}' + } + } + ] + } + ] + + sanitized = sanitize_messages_for_tool_calling(messages) + + # Messages should be unchanged + assert len(sanitized) == 2 + assert sanitized[0]["content"] == "" + assert len(sanitized[1].get("tool_calls", [])) == 1 + + def test_anthropic_messages_pt_integration(self): + """ + Test that sanitization is integrated into anthropic_messages_pt + """ + litellm.modify_params = True + + messages = [ + { + "role": "user", + "content": "What is the weather in Nashik?" + }, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "toolu_01Kus2cC3ydjBW7UK4GJqBP4", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "Nashik, India"}' + } + } + ] + } + ] + + # This should not raise an error and should add dummy tool result + result = anthropic_messages_pt( + messages=messages, + model="claude-sonnet-4-5", + llm_provider="anthropic" + ) + + # Should have at least 2 messages (user and assistant) + # The tool result will be merged into user content + assert len(result) >= 2 + assert result[0]["role"] == "user" + assert result[1]["role"] == "assistant" + + +if __name__ == "__main__": + pytest.main([__file__, "-v"]) diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index ddbb0454cac..c773db21074 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -2701,37 +2701,37 @@ def test_empty_assistant_message_handling(): assert result[1]["content"][0]["text"] == "I'm doing well, thank you!" -def test_is_nova_lite_2_model(): - """Test the _is_nova_lite_2_model() method for detecting Nova 2 models.""" +def test_is_nova_2_model(): + """Test the _is_nova_2_model() method for detecting Nova 2 models.""" config = AmazonConverseConfig() # Test with amazon.nova-2-lite-v1:0 - assert config._is_nova_lite_2_model("amazon.nova-2-lite-v1:0") is True + assert config._is_nova_2_model("amazon.nova-2-lite-v1:0") is True # Test with regional variants - assert config._is_nova_lite_2_model("us.amazon.nova-2-lite-v1:0") is True - assert config._is_nova_lite_2_model("eu.amazon.nova-2-lite-v1:0") is True - assert config._is_nova_lite_2_model("apac.amazon.nova-2-lite-v1:0") is True + assert config._is_nova_2_model("us.amazon.nova-2-lite-v1:0") is True + assert config._is_nova_2_model("eu.amazon.nova-2-lite-v1:0") is True + assert config._is_nova_2_model("apac.amazon.nova-2-lite-v1:0") is True # Test with other Nova 2 variants (pro, micro) - assert config._is_nova_lite_2_model("amazon.nova-pro-1-5-v1:0") is False - assert config._is_nova_lite_2_model("amazon.nova-micro-1-5-v1:0") is False - assert config._is_nova_lite_2_model("us.amazon.nova-pro-1-5-v1:0") is False - assert config._is_nova_lite_2_model("eu.amazon.nova-micro-1-5-v1:0") is False + assert config._is_nova_2_model("amazon.nova-pro-1-5-v1:0") is False + assert config._is_nova_2_model("amazon.nova-micro-1-5-v1:0") is False + assert config._is_nova_2_model("us.amazon.nova-pro-1-5-v1:0") is False + assert config._is_nova_2_model("eu.amazon.nova-micro-1-5-v1:0") is False # Test with non-Nova-1.5 lite models (should return False) - assert config._is_nova_lite_2_model("amazon.nova-lite-v1:0") is False - assert config._is_nova_lite_2_model("amazon.nova-pro-v1:0") is False - assert config._is_nova_lite_2_model("amazon.nova-micro-v1:0") is False + assert config._is_nova_2_model("amazon.nova-lite-v1:0") is False + assert config._is_nova_2_model("amazon.nova-pro-v1:0") is False + assert config._is_nova_2_model("amazon.nova-micro-v1:0") is False # Test with Nova v1:0 models (should return False) - assert config._is_nova_lite_2_model("us.amazon.nova-lite-v1:0") is False - assert config._is_nova_lite_2_model("eu.amazon.nova-pro-v1:0") is False + assert config._is_nova_2_model("us.amazon.nova-lite-v1:0") is False + assert config._is_nova_2_model("eu.amazon.nova-pro-v1:0") is False # Test with completely different models (should return False) - assert config._is_nova_lite_2_model("anthropic.claude-3-5-sonnet-20240620-v1:0") is False - assert config._is_nova_lite_2_model("meta.llama3-70b-instruct-v1:0") is False - assert config._is_nova_lite_2_model("mistral.mistral-7b-instruct-v0:2") is False + assert config._is_nova_2_model("anthropic.claude-3-5-sonnet-20240620-v1:0") is False + assert config._is_nova_2_model("meta.llama3-70b-instruct-v1:0") is False + assert config._is_nova_2_model("mistral.mistral-7b-instruct-v0:2") is False def test_thinking_with_max_completion_tokens(): @@ -2936,6 +2936,447 @@ def test_drop_thinking_param_when_thinking_blocks_missing(): litellm.modify_params = original_modify_params +def test_supports_native_structured_outputs(): + """Test model detection for native structured outputs support.""" + config = AmazonConverseConfig() + + # Supported models + assert config._supports_native_structured_outputs( + "anthropic.claude-sonnet-4-5-20250929-v1:0" + ) + assert config._supports_native_structured_outputs( + "anthropic.claude-haiku-4-5-20251001-v1:0" + ) + assert config._supports_native_structured_outputs( + "anthropic.claude-opus-4-6-v1:0" + ) + assert config._supports_native_structured_outputs( + "eu.anthropic.claude-opus-4-5-20260101-v1:0" + ) + assert config._supports_native_structured_outputs("qwen.qwen3-235b-instruct-v1:0") + assert config._supports_native_structured_outputs("mistral.mistral-large-3-v1:0") + assert config._supports_native_structured_outputs("deepseek.deepseek-v3.1-v1:0") + + # Unsupported models — should fall back to tool-call approach + assert not config._supports_native_structured_outputs( + "anthropic.claude-3-5-sonnet-20241022-v2:0" + ) + assert not config._supports_native_structured_outputs( + "anthropic.claude-sonnet-4-20250514-v1:0" + ) + assert not config._supports_native_structured_outputs( + "meta.llama3-3-70b-instruct-v1:0" + ) + assert not config._supports_native_structured_outputs( + "amazon.nova-pro-v1:0" + ) + # Excluded despite AWS listing them: broken constrained decoding on Bedrock + assert not config._supports_native_structured_outputs( + "openai.gpt-oss-120b-1:0" + ) + assert not config._supports_native_structured_outputs( + "mistral.magistral-small-2509" + ) + + +def test_create_output_config_for_response_format(): + """Test outputConfig dict creation from JSON schema.""" + config = AmazonConverseConfig() + + schema = { + "type": "object", + "properties": { + "name": {"type": "string"}, + "age": {"type": "integer"}, + }, + "required": ["name", "age"], + } + + output_config = config._create_output_config_for_response_format( + json_schema=schema, + name="PersonInfo", + description="A person's info", + ) + + assert "textFormat" in output_config + text_format = output_config["textFormat"] + assert text_format["type"] == "json_schema" + assert "structure" in text_format + + json_schema_def = text_format["structure"]["jsonSchema"] + assert json_schema_def["name"] == "PersonInfo" + assert json_schema_def["description"] == "A person's info" + # schema field must be a JSON string, not a dict + assert isinstance(json_schema_def["schema"], str) + parsed_schema = json.loads(json_schema_def["schema"]) + # additionalProperties: false is injected by normalization + expected = {**schema, "additionalProperties": False} + assert parsed_schema == expected + + +def test_translate_response_format_native_output_config(): + """For supported models, _translate_response_format_param should produce outputConfig.""" + config = AmazonConverseConfig() + + response_format = { + "type": "json_schema", + "json_schema": { + "name": "WeatherResult", + "description": "Weather info", + "schema": { + "type": "object", + "properties": { + "temp": {"type": "number"}, + }, + "required": ["temp"], + }, + }, + } + + optional_params: dict = {} + result = config._translate_response_format_param( + value=response_format, + model="anthropic.claude-sonnet-4-5-20250929-v1:0", + optional_params=optional_params, + non_default_params={"response_format": response_format}, + is_thinking_enabled=False, + ) + + # Should have outputConfig, NOT tools + assert "outputConfig" in result + assert "tools" not in result + assert "tool_choice" not in result + assert result["json_mode"] is True + # No fake_stream for native approach + assert "fake_stream" not in result + + # Verify the schema content (additionalProperties: false is added by normalization) + schema_str = result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["schema"] + parsed_schema = json.loads(schema_str) + expected_schema = {**response_format["json_schema"]["schema"], "additionalProperties": False} + assert parsed_schema == expected_schema + assert ( + result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["name"] + == "WeatherResult" + ) + + +def test_translate_response_format_fallback_tool_call(): + """For unsupported models, should fall back to tool-call approach.""" + config = AmazonConverseConfig() + + response_format = { + "type": "json_schema", + "json_schema": { + "name": "WeatherResult", + "schema": { + "type": "object", + "properties": { + "temp": {"type": "number"}, + }, + }, + }, + } + + optional_params: dict = {} + result = config._translate_response_format_param( + value=response_format, + model="anthropic.claude-3-5-sonnet-20241022-v2:0", + optional_params=optional_params, + non_default_params={"response_format": response_format}, + is_thinking_enabled=False, + ) + + # Should use tool-call approach, NOT outputConfig + assert "outputConfig" not in result + assert "tools" in result + assert result["json_mode"] is True + + +def test_native_structured_output_no_fake_stream(): + """When using native structured outputs with streaming, fake_stream should NOT be set.""" + config = AmazonConverseConfig() + + response_format = { + "type": "json_schema", + "json_schema": { + "name": "Result", + "schema": { + "type": "object", + "properties": { + "answer": {"type": "string"}, + }, + }, + }, + } + + optional_params: dict = {} + result = config._translate_response_format_param( + value=response_format, + model="anthropic.claude-sonnet-4-5-20250929-v1:0", + optional_params=optional_params, + non_default_params={"response_format": response_format, "stream": True}, + is_thinking_enabled=False, + ) + + assert "outputConfig" in result + assert result["json_mode"] is True + # No fake_stream for native approach + assert "fake_stream" not in result + + # Verify the schema content + schema_str = result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["schema"] + assert json.loads(schema_str) == { + "type": "object", + "properties": {"answer": {"type": "string"}}, + "additionalProperties": False, + } + + +def test_transform_request_with_output_config(): + """Test that outputConfig flows through _transform_request_helper into the final request.""" + from litellm.types.llms.bedrock import OutputConfigBlock, OutputFormat, OutputFormatStructure, JsonSchemaDefinition + + config = AmazonConverseConfig() + + output_config = OutputConfigBlock( + textFormat=OutputFormat( + type="json_schema", + structure=OutputFormatStructure( + jsonSchema=JsonSchemaDefinition( + schema='{"type": "object", "properties": {"x": {"type": "string"}}, "additionalProperties": false}', + name="TestSchema", + ) + ), + ) + ) + + messages = [{"role": "user", "content": "test"}] + optional_params = { + "outputConfig": output_config, + "json_mode": True, + } + + result = config._transform_request( + model="anthropic.claude-sonnet-4-5-20250929-v1:0", + messages=messages, + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + assert "outputConfig" in result + assert result["outputConfig"]["textFormat"]["type"] == "json_schema" + assert result["outputConfig"]["textFormat"]["structure"]["jsonSchema"]["name"] == "TestSchema" + + +def test_transform_response_native_structured_output(): + """Test response handling when model returns JSON as text content (native structured output).""" + response_json = { + "output": { + "message": { + "role": "assistant", + "content": [ + { + "text": '{"temp": 62, "description": "Mild and foggy"}' + } + ], + } + }, + "stopReason": "end_turn", + "usage": { + "inputTokens": 10, + "outputTokens": 20, + "totalTokens": 30, + }, + } + + class MockResponse: + def json(self): + return response_json + + @property + def text(self): + return json.dumps(response_json) + + config = AmazonConverseConfig() + model_response = ModelResponse() + # json_mode=True but no tool_call in response — native structured output path + optional_params = {"json_mode": True} + + result = config._transform_response( + model="anthropic.claude-sonnet-4-5-20250929-v1:0", + response=MockResponse(), + model_response=model_response, + stream=False, + logging_obj=None, + optional_params=optional_params, + api_key=None, + data={}, + messages=[], + encoding=None, + ) + + # Content should be the JSON text directly + assert result.choices[0].message.content == '{"temp": 62, "description": "Mild and foggy"}' + # Should NOT have tool_calls + assert result.choices[0].message.tool_calls is None + assert result.choices[0].finish_reason == "stop" + + +def test_add_additional_properties_simple_object(): + """Object schemas without additionalProperties get it set to false.""" + schema = { + "type": "object", + "properties": { + "city": {"type": "string"}, + "country": {"type": "string"}, + }, + "required": ["city", "country"], + } + result = AmazonConverseConfig._add_additional_properties_to_schema(schema) + assert result["additionalProperties"] is False + # Original should not be mutated + assert "additionalProperties" not in schema + + +def test_add_additional_properties_already_set(): + """If additionalProperties is already set, don't overwrite it.""" + schema = { + "type": "object", + "properties": {"x": {"type": "string"}}, + "additionalProperties": True, + } + result = AmazonConverseConfig._add_additional_properties_to_schema(schema) + assert result["additionalProperties"] is True + + +def test_add_additional_properties_nested(): + """Recursively processes nested object types in properties, items, $defs, anyOf.""" + schema = { + "type": "object", + "properties": { + "address": { + "type": "object", + "properties": { + "street": {"type": "string"}, + "zip": {"type": "string"}, + }, + }, + "tags": { + "type": "array", + "items": { + "type": "object", + "properties": {"name": {"type": "string"}}, + }, + }, + }, + "$defs": { + "Metadata": { + "type": "object", + "properties": {"key": {"type": "string"}}, + } + }, + "anyOf": [ + { + "type": "object", + "properties": {"variant": {"type": "string"}}, + } + ], + } + result = AmazonConverseConfig._add_additional_properties_to_schema(schema) + # Top-level + assert result["additionalProperties"] is False + # Nested property object + assert result["properties"]["address"]["additionalProperties"] is False + # Array items object + assert result["properties"]["tags"]["items"]["additionalProperties"] is False + # $defs object + assert result["$defs"]["Metadata"]["additionalProperties"] is False + # anyOf object + assert result["anyOf"][0]["additionalProperties"] is False + + +def test_add_additional_properties_non_object(): + """Non-object schemas are returned unchanged.""" + schema = {"type": "string"} + result = AmazonConverseConfig._add_additional_properties_to_schema(schema) + assert "additionalProperties" not in result + assert result == {"type": "string"} + + +def test_add_additional_properties_definitions(): + """Recursively processes object types inside 'definitions' (not just '$defs').""" + schema = { + "type": "object", + "properties": { + "item": {"$ref": "#/definitions/Item"}, + }, + "definitions": { + "Item": { + "type": "object", + "properties": { + "name": {"type": "string"}, + "details": { + "type": "object", + "properties": {"weight": {"type": "number"}}, + }, + }, + } + }, + } + result = AmazonConverseConfig._add_additional_properties_to_schema(schema) + # Top-level + assert result["additionalProperties"] is False + # definitions object + assert result["definitions"]["Item"]["additionalProperties"] is False + # Nested object inside definitions + assert result["definitions"]["Item"]["properties"]["details"]["additionalProperties"] is False + + +def test_json_object_no_schema_falls_back_to_tool_call(): + """response_format: {type: json_object} with no schema should use tool-call fallback, + even for models that support native structured outputs.""" + config = AmazonConverseConfig() + optional_params: dict = {} + non_default_params = {"response_format": {"type": "json_object"}} + + result = config._translate_response_format_param( + value=non_default_params["response_format"], + model="anthropic.claude-sonnet-4-5-20250929-v1:0", + optional_params=optional_params, + non_default_params=non_default_params, + is_thinking_enabled=False, + ) + + # Should NOT use native outputConfig (no schema provided) + assert "outputConfig" not in result + # Should use tool-call fallback + assert "tools" in result + assert result["json_mode"] is True + + +def test_output_config_applies_additional_properties(): + """_create_output_config_for_response_format normalizes the schema.""" + schema = { + "type": "object", + "properties": { + "name": {"type": "string"}, + "nested": { + "type": "object", + "properties": {"val": {"type": "integer"}}, + }, + }, + } + output_config = AmazonConverseConfig._create_output_config_for_response_format( + json_schema=schema, name="test_schema" + ) + parsed = json.loads(output_config["textFormat"]["structure"]["jsonSchema"]["schema"]) + assert parsed["additionalProperties"] is False + assert parsed["properties"]["nested"]["additionalProperties"] is False + + + class TestBedrockMinThinkingBudgetTokens: """Test that thinking.budget_tokens is clamped to the Bedrock minimum (1024).""" diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation_nova_2.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation_nova_2.py index 23243dac201..bac7aa08a04 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation_nova_2.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation_nova_2.py @@ -1,7 +1,10 @@ """ Unit tests for Amazon Nova 2 reasoning configuration transformation. -Tests the _transform_reasoning_effort_to_reasoning_config method in AmazonConverseConfig. +Tests request transformation, response parsing, multi-turn message translation, +and model detection for Nova 2 Lite and Nova 2 Pro via the Bedrock Converse API. + +Reference: https://docs.aws.amazon.com/nova/latest/nova2-userguide/using-converse-api.html """ import pytest @@ -12,6 +15,7 @@ sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path +import httpx import litellm from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConfig @@ -323,248 +327,52 @@ class TestNova15SupportedParameters: assert "response_format" in supported_params -class TestNova15ResponseParsing: - """Test suite for Nova 2 response parsing.""" +class TestNova2ResponseParsing: + """Test that reasoningContent blocks are parsed into reasoning_content strings.""" - def test_transform_reasoning_content_single_block(self): - """Test that reasoning content is extracted correctly from a single block.""" + def test_should_extract_single_reasoning_block(self): config = AmazonConverseConfig() - - reasoning_blocks = [ - {"reasoningText": {"text": "Let me think through this step by step..."}} - ] - - result = config._transform_reasoning_content(reasoning_blocks) - + result = config._transform_reasoning_content( + [{"reasoningText": {"text": "Let me think through this step by step..."}}] + ) assert result == "Let me think through this step by step..." - def test_transform_reasoning_content_multiple_blocks(self): - """Test that reasoning content is concatenated from multiple blocks.""" + def test_should_concatenate_multiple_reasoning_blocks(self): config = AmazonConverseConfig() - - reasoning_blocks = [ - {"reasoningText": {"text": "First, I need to analyze the problem. "}}, - {"reasoningText": {"text": "Then, I'll consider the solution."}}, - ] - - result = config._transform_reasoning_content(reasoning_blocks) - + result = config._transform_reasoning_content( + [ + {"reasoningText": {"text": "First, I need to analyze the problem. "}}, + {"reasoningText": {"text": "Then, I'll consider the solution."}}, + ] + ) assert ( result == "First, I need to analyze the problem. Then, I'll consider the solution." ) - def test_transform_reasoning_content_empty_blocks(self): - """Test that empty reasoning blocks return empty string.""" + def test_should_return_empty_string_for_empty_blocks(self): config = AmazonConverseConfig() - - reasoning_blocks = [] - - result = config._transform_reasoning_content(reasoning_blocks) - - assert result == "" - - def test_transform_thinking_blocks_with_text(self): - """Test that thinking blocks are populated correctly with text.""" - config = AmazonConverseConfig() - - reasoning_blocks = [{"reasoningText": {"text": "My reasoning process..."}}] - - result = config._transform_thinking_blocks(reasoning_blocks) - - assert len(result) == 1 - assert result[0]["type"] == "thinking" - assert result[0]["thinking"] == "My reasoning process..." - assert "signature" not in result[0] - - def test_transform_thinking_blocks_with_signature(self): - """Test that signature field is preserved when present.""" - config = AmazonConverseConfig() - - reasoning_blocks = [ - { - "reasoningText": { - "text": "My reasoning...", - "signature": "signature-hash-12345", - } - } - ] - - result = config._transform_thinking_blocks(reasoning_blocks) - - assert len(result) == 1 - assert result[0]["type"] == "thinking" - assert result[0]["thinking"] == "My reasoning..." - assert result[0]["signature"] == "signature-hash-12345" - - def test_transform_thinking_blocks_with_redacted_content(self): - """Test that redacted content blocks are handled correctly.""" - config = AmazonConverseConfig() - - reasoning_blocks = [ - {"reasoningText": {"text": "First part of reasoning..."}}, - {"redactedContent": {}}, - {"reasoningText": {"text": "Second part after redaction..."}}, - ] - - result = config._transform_thinking_blocks(reasoning_blocks) - - assert len(result) == 3 - assert result[0]["type"] == "thinking" - assert result[0]["thinking"] == "First part of reasoning..." - assert result[1]["type"] == "redacted_thinking" - assert result[2]["type"] == "thinking" - assert result[2]["thinking"] == "Second part after redaction..." - - def test_transform_thinking_blocks_multiple_blocks(self): - """Test that multiple thinking blocks are all transformed.""" - config = AmazonConverseConfig() - - reasoning_blocks = [ - {"reasoningText": {"text": "Step 1: Analyze the problem"}}, - { - "reasoningText": { - "text": "Step 2: Consider solutions", - "signature": "sig-abc", - } - }, - {"reasoningText": {"text": "Step 3: Choose best approach"}}, - ] - - result = config._transform_thinking_blocks(reasoning_blocks) - - assert len(result) == 3 - assert all(block["type"] == "thinking" for block in result) - assert result[0]["thinking"] == "Step 1: Analyze the problem" - assert result[1]["thinking"] == "Step 2: Consider solutions" - assert result[1]["signature"] == "sig-abc" - assert result[2]["thinking"] == "Step 3: Choose best approach" - - def test_transform_thinking_blocks_empty_list(self): - """Test that empty thinking blocks list returns empty list.""" - config = AmazonConverseConfig() - - reasoning_blocks = [] - - result = config._transform_thinking_blocks(reasoning_blocks) - - assert result == [] - - def test_response_parsing_integration(self): - """Test that response parsing works end-to-end with Nova 2 structure.""" - config = AmazonConverseConfig() - - # Simulate a Nova 2 response with reasoning content - reasoning_blocks = [ - { - "reasoningText": { - "text": "Let me analyze this carefully. ", - "signature": "test-signature", - } - }, - {"reasoningText": {"text": "Based on my analysis, the answer is clear."}}, - ] - - # Test reasoning content extraction - reasoning_content = config._transform_reasoning_content(reasoning_blocks) - assert ( - reasoning_content - == "Let me analyze this carefully. Based on my analysis, the answer is clear." - ) - - # Test thinking blocks transformation - thinking_blocks = config._transform_thinking_blocks(reasoning_blocks) - assert len(thinking_blocks) == 2 - assert thinking_blocks[0]["thinking"] == "Let me analyze this carefully. " - assert thinking_blocks[0]["signature"] == "test-signature" - assert ( - thinking_blocks[1]["thinking"] - == "Based on my analysis, the answer is clear." - ) + assert config._transform_reasoning_content([]) == "" -class TestNova15StreamingResponseParsing: - """Test suite for Nova 2 streaming response parsing.""" +class TestNova2StreamingResponseParsing: + """Test that streaming reasoningContent deltas produce reasoning_content on the delta.""" - def test_streaming_reasoning_content_start_event(self): - """Test that streaming start event with reasoningContent is handled correctly.""" + def test_should_extract_reasoning_content_from_delta(self): from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") - - # Simulate a start event with redacted reasoning content - chunk_data = { - "start": {"reasoningContent": {"redactedContent": {}}}, - "contentBlockIndex": 0, - } - - result = handler.converse_chunk_parser(chunk_data) - - # Verify thinking blocks are populated - assert result.choices[0].delta.thinking_blocks is not None - assert len(result.choices[0].delta.thinking_blocks) == 1 - assert result.choices[0].delta.thinking_blocks[0]["type"] == "redacted_thinking" - - def test_streaming_reasoning_content_delta_text(self): - """Test that streaming delta event with reasoning text is handled correctly.""" - from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder - - handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") - - # Simulate a delta event with reasoning text chunk_data = { "delta": {"reasoningContent": {"text": "Let me think about this..."}}, "contentBlockIndex": 0, } - result = handler.converse_chunk_parser(chunk_data) - - # Verify reasoning content is extracted assert result.choices[0].delta.reasoning_content == "Let me think about this..." - # Verify thinking blocks are populated - assert result.choices[0].delta.thinking_blocks is not None - assert len(result.choices[0].delta.thinking_blocks) == 1 - assert result.choices[0].delta.thinking_blocks[0]["type"] == "thinking" - assert ( - result.choices[0].delta.thinking_blocks[0]["thinking"] - == "Let me think about this..." - ) - - def test_streaming_reasoning_content_delta_signature(self): - """Test that streaming delta event with signature is handled correctly.""" + def test_should_accumulate_multiple_reasoning_deltas(self): from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") - - # Simulate a delta event with signature - chunk_data = { - "delta": {"reasoningContent": {"signature": "signature-hash-xyz"}}, - "contentBlockIndex": 0, - } - - result = handler.converse_chunk_parser(chunk_data) - - # Verify reasoning content is set to empty string for consistency - assert result.choices[0].delta.reasoning_content == "" - - # Verify thinking blocks are populated with signature - assert result.choices[0].delta.thinking_blocks is not None - assert len(result.choices[0].delta.thinking_blocks) == 1 - assert result.choices[0].delta.thinking_blocks[0]["type"] == "thinking" - assert ( - result.choices[0].delta.thinking_blocks[0]["signature"] - == "signature-hash-xyz" - ) - assert result.choices[0].delta.thinking_blocks[0]["thinking"] == "" - - def test_streaming_reasoning_content_multiple_deltas(self): - """Test that multiple reasoning content deltas are accumulated correctly.""" - from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder - - handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") - - # Simulate multiple delta events chunks = [ { "delta": {"reasoningContent": {"text": "First, "}}, @@ -579,30 +387,15 @@ class TestNova15StreamingResponseParsing: "contentBlockIndex": 0, }, ] - - results = [] - for chunk_data in chunks: - result = handler.converse_chunk_parser(chunk_data) - results.append(result) - - # Verify each delta has the correct reasoning content + results = [handler.converse_chunk_parser(c) for c in chunks] assert results[0].choices[0].delta.reasoning_content == "First, " assert results[1].choices[0].delta.reasoning_content == "I need to analyze " assert results[2].choices[0].delta.reasoning_content == "the problem." - # Verify thinking blocks are populated for each delta - for result in results: - assert result.choices[0].delta.thinking_blocks is not None - assert len(result.choices[0].delta.thinking_blocks) == 1 - assert result.choices[0].delta.thinking_blocks[0]["type"] == "thinking" - - def test_streaming_reasoning_then_text_content(self): - """Test that reasoning content followed by text content is handled correctly.""" + def test_should_stream_reasoning_then_text(self): from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") - - # Simulate reasoning content followed by text content chunks = [ { "delta": {"reasoningContent": {"text": "Let me think..."}}, @@ -611,184 +404,286 @@ class TestNova15StreamingResponseParsing: {"delta": {"text": "Based on my reasoning, "}, "contentBlockIndex": 1}, {"delta": {"text": "the answer is 42."}, "contentBlockIndex": 1}, ] - - results = [] - for chunk_data in chunks: - result = handler.converse_chunk_parser(chunk_data) - results.append(result) - - # Verify first chunk has reasoning content + results = [handler.converse_chunk_parser(c) for c in chunks] assert results[0].choices[0].delta.reasoning_content == "Let me think..." - assert results[0].choices[0].delta.thinking_blocks is not None - - # Verify subsequent chunks have text content assert results[1].choices[0].delta.content == "Based on my reasoning, " assert results[2].choices[0].delta.content == "the answer is 42." - def test_streaming_redacted_content_delta(self): - """Test that streaming delta with redacted content is handled correctly.""" + def test_should_populate_provider_specific_fields(self): from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") - - # Simulate a delta event with redacted content - chunk_data = { - "delta": {"reasoningContent": {"redactedContent": {}}}, - "contentBlockIndex": 0, - } - - result = handler.converse_chunk_parser(chunk_data) - - # Verify reasoning content is set to empty string for consistency - assert result.choices[0].delta.reasoning_content == "" - - # Verify thinking blocks contain redacted block - assert result.choices[0].delta.thinking_blocks is not None - assert len(result.choices[0].delta.thinking_blocks) == 1 - assert result.choices[0].delta.thinking_blocks[0]["type"] == "redacted_thinking" - - def test_streaming_provider_specific_fields(self): - """Test that provider_specific_fields are populated in streaming responses.""" - from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder - - handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") - - # Simulate a delta event with reasoning content chunk_data = { "delta": {"reasoningContent": {"text": "Reasoning text"}}, "contentBlockIndex": 0, } - result = handler.converse_chunk_parser(chunk_data) + psf = result.choices[0].delta.provider_specific_fields + assert psf is not None + assert psf["reasoningContent"]["text"] == "Reasoning text" - # Verify provider_specific_fields are populated - assert result.choices[0].delta.provider_specific_fields is not None - assert "reasoningContent" in result.choices[0].delta.provider_specific_fields - assert ( - result.choices[0].delta.provider_specific_fields["reasoningContent"]["text"] - == "Reasoning text" - ) - - def test_streaming_mixed_content_blocks(self): - """Test streaming with mixed content blocks (reasoning, text, tool calls).""" + def test_should_stream_reasoning_with_tool_calls(self): from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") - - # Simulate a complex streaming scenario chunks = [ - # Start with reasoning { - "delta": { - "reasoningContent": { - "text": "I need to call a tool to get information." - } - }, + "delta": {"reasoningContent": {"text": "I need to call a tool."}}, "contentBlockIndex": 0, }, - # Tool use start { "start": {"toolUse": {"toolUseId": "tool-123", "name": "get_weather"}}, "contentBlockIndex": 1, }, - # Tool use delta { "delta": {"toolUse": {"input": '{"location": "NYC"}'}}, "contentBlockIndex": 1, }, - # Text response {"delta": {"text": "The weather is sunny."}, "contentBlockIndex": 2}, ] - - results = [] - for chunk_data in chunks: - result = handler.converse_chunk_parser(chunk_data) - results.append(result) - - # Verify reasoning content in first chunk - assert ( - results[0].choices[0].delta.reasoning_content - == "I need to call a tool to get information." - ) - - # Verify tool call in second and third chunks - assert results[1].choices[0].delta.tool_calls is not None + results = [handler.converse_chunk_parser(c) for c in chunks] + assert results[0].choices[0].delta.reasoning_content == "I need to call a tool." assert ( results[1].choices[0].delta.tool_calls[0]["function"]["name"] == "get_weather" ) - assert results[2].choices[0].delta.tool_calls is not None - - # Verify text content in fourth chunk assert results[3].choices[0].delta.content == "The weather is sunny." - def test_extract_reasoning_content_str_with_text(self): - """Test extract_reasoning_content_str method with text.""" - from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder - handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") +# --------------------------------------------------------------------------- +# Model detection — _is_nova_2_model covers both Lite and Pro +# --------------------------------------------------------------------------- - reasoning_block = {"text": "This is reasoning text"} +NOVA_2_LITE = "amazon.nova-2-lite-v1:0" +NOVA_2_PRO = "us.amazon.nova-2-pro-preview-20251202-v1:0" - result = handler.extract_reasoning_content_str(reasoning_block) - assert result == "This is reasoning text" +class TestNova2ModelDetection: + """Verify _is_nova_2_model identifies all Nova 2 variants (lite, pro, regional, routed).""" - def test_extract_reasoning_content_str_without_text(self): - """Test extract_reasoning_content_str method without text (e.g., signature only).""" - from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder + @pytest.mark.parametrize( + "model", + [ + "amazon.nova-2-lite-v1:0", + "amazon.nova-2-pro-preview-20251202-v1:0", + "us.amazon.nova-2-lite-v1:0", + "us.amazon.nova-2-pro-preview-20251202-v1:0", + "eu.amazon.nova-2-lite-v1:0", + "apac.amazon.nova-2-pro-preview-20251202-v1:0", + "bedrock/converse/amazon.nova-2-lite-v1:0", + "bedrock/converse/us.amazon.nova-2-pro-preview-20251202-v1:0", + "bedrock/amazon.nova-2-lite-v1:0", + "converse/us.amazon.nova-2-lite-v1:0", + "converse/amazon.nova-2-pro-preview-20251202-v1:0", + ], + ) + def test_should_recognize_nova_2_models(self, model): + assert AmazonConverseConfig()._is_nova_2_model(model) is True - handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") + @pytest.mark.parametrize( + "model", + [ + "amazon.nova-pro-v1:0", + "amazon.nova-lite-v1:0", + "amazon.nova-pro-1-5-v1:0", + "anthropic.claude-3-sonnet-20240229-v1:0", + "us.amazon.nova-pro-v1:0", + ], + ) + def test_should_not_match_non_nova_2_models(self, model): + assert AmazonConverseConfig()._is_nova_2_model(model) is False - reasoning_block = {"signature": "sig-123"} - result = handler.extract_reasoning_content_str(reasoning_block) +# --------------------------------------------------------------------------- +# End-to-end request body — reasoningConfig in additionalModelRequestFields +# --------------------------------------------------------------------------- - assert result is None - def test_translate_thinking_blocks_streaming_text(self): - """Test translate_thinking_blocks method with text.""" - from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder +class TestNova2EndToEndRequest: + """Verify transform_request places reasoningConfig correctly for both model variants.""" - handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") + def _build_request(self, model, effort, **extra): + config = AmazonConverseConfig() + optional_params = config.map_openai_params( + non_default_params={"reasoning_effort": effort, **extra}, + optional_params={}, + model=model, + drop_params=False, + ) + return config.transform_request( + model=model, + messages=[{"role": "user", "content": "What is 2+2?"}], + optional_params=optional_params, + litellm_params={}, + headers={}, + ) - thinking_block = {"text": "Thinking content"} + @pytest.mark.parametrize("model", [NOVA_2_LITE, NOVA_2_PRO]) + def test_should_place_reasoning_config_in_additional_model_request_fields( + self, model + ): + body = self._build_request(model, "high") + additional = body.get("additionalModelRequestFields", {}) + assert additional["reasoningConfig"] == { + "type": "enabled", + "maxReasoningEffort": "high", + } + assert "reasoningConfig" not in body # not top-level + assert "thinking" not in body # not Anthropic-style - result = handler.translate_thinking_blocks(thinking_block) - - assert result is not None - assert len(result) == 1 - assert result[0]["type"] == "thinking" - assert result[0]["thinking"] == "Thinking content" - - def test_translate_thinking_blocks_streaming_signature(self): - """Test translate_thinking_blocks method with signature.""" - from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder - - handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") - - thinking_block = {"signature": "sig-abc"} - - result = handler.translate_thinking_blocks(thinking_block) - - assert result is not None - assert len(result) == 1 - assert result[0]["type"] == "thinking" - assert result[0]["signature"] == "sig-abc" + @pytest.mark.parametrize("model", [NOVA_2_LITE, NOVA_2_PRO]) + def test_should_coexist_with_inference_params(self, model): + body = self._build_request(model, "high", temperature=0.5, max_tokens=512) assert ( - result[0]["thinking"] == "" - ) # Empty string for consistency with Anthropic + body["additionalModelRequestFields"]["reasoningConfig"]["type"] == "enabled" + ) + inf = body.get("inferenceConfig", {}) + assert inf.get("temperature") == 0.5 + assert inf.get("maxTokens") == 512 - def test_translate_thinking_blocks_streaming_redacted(self): - """Test translate_thinking_blocks method with redacted content.""" - from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder - handler = AWSEventStreamDecoder(model="amazon.nova-2-lite-v1:0") +# --------------------------------------------------------------------------- +# End-to-end response — reasoningContent parsed to reasoning_content string +# --------------------------------------------------------------------------- - thinking_block = {"redactedContent": {}} - result = handler.translate_thinking_blocks(thinking_block) +class TestNova2EndToEndResponse: + """Verify transform_response produces reasoning_content from reasoningContent blocks.""" - assert result is not None - assert len(result) == 1 - assert result[0]["type"] == "redacted_thinking" + def _transform(self, content_blocks, model=NOVA_2_LITE): + config = AmazonConverseConfig() + body = { + "output": {"message": {"role": "assistant", "content": content_blocks}}, + "usage": {"inputTokens": 10, "outputTokens": 50, "totalTokens": 60}, + "stopReason": "end_turn", + "metrics": {"latencyMs": 100}, + } + resp = httpx.Response( + 200, json=body, request=httpx.Request("POST", "https://bedrock") + ) + return config.transform_response( + model=model, + raw_response=resp, + model_response=litellm.ModelResponse(), + logging_obj=None, + request_data={}, + messages=[], + optional_params={}, + litellm_params={}, + encoding=None, + api_key=None, + json_mode=None, + ) + + def test_should_extract_reasoning_content_as_string(self): + result = self._transform( + [ + {"reasoningContent": {"reasoningText": {"text": "Step 1. "}}}, + {"reasoningContent": {"reasoningText": {"text": "Step 2."}}}, + {"text": "The answer is 4."}, + ] + ) + msg = result.choices[0].message + assert msg.content == "The answer is 4." + assert msg.reasoning_content == "Step 1. Step 2." + + def test_should_include_raw_blocks_in_provider_specific_fields(self): + result = self._transform( + [ + {"reasoningContent": {"reasoningText": {"text": "thinking..."}}}, + {"text": "done"}, + ] + ) + psf = result.choices[0].message.get("provider_specific_fields", {}) + assert "reasoningContentBlocks" in psf + + def test_should_omit_reasoning_content_when_absent(self): + result = self._transform([{"text": "Plain answer."}]) + assert not getattr(result.choices[0].message, "reasoning_content", None) + + +# --------------------------------------------------------------------------- +# Multi-turn — reasoning_content round-trips back to Bedrock format +# --------------------------------------------------------------------------- + + +class TestNova2MultiTurnMessageTranslation: + """Verify that assistant messages carrying reasoning from a previous turn are + correctly translated to Bedrock content blocks via _bedrock_converse_messages_pt.""" + + def _to_bedrock(self, messages, model=NOVA_2_LITE): + from litellm.litellm_core_utils.prompt_templates.factory import ( + _bedrock_converse_messages_pt, + ) + + return _bedrock_converse_messages_pt( + messages=messages, + model=model, + llm_provider="bedrock_converse", + ) + + def test_should_inline_unsigned_thinking_blocks_as_text(self): + """Without a signature, reasoning text becomes a plain text block.""" + bedrock_msgs = self._to_bedrock( + [ + {"role": "user", "content": "What is 2+2?"}, + { + "role": "assistant", + "content": "4.", + "thinking_blocks": [ + {"type": "thinking", "thinking": "Simple addition"}, + ], + }, + {"role": "user", "content": "Sure?"}, + ] + ) + assistant = next(m for m in bedrock_msgs if m["role"] == "assistant") + texts = [b["text"] for b in assistant["content"] if "text" in b] + assert "Simple addition" in texts + assert "4." in texts + + def test_should_keep_signed_thinking_blocks_as_reasoning_content(self): + """With a signature, reasoning is preserved as a reasoningContent block.""" + bedrock_msgs = self._to_bedrock( + [ + {"role": "user", "content": "What is 2+2?"}, + { + "role": "assistant", + "content": "4.", + "thinking_blocks": [ + {"type": "thinking", "thinking": "math", "signature": "sig-1"}, + ], + }, + {"role": "user", "content": "Sure?"}, + ] + ) + assistant = next(m for m in bedrock_msgs if m["role"] == "assistant") + rc_blocks = [b for b in assistant["content"] if "reasoningContent" in b] + assert len(rc_blocks) >= 1 + assert rc_blocks[0]["reasoningContent"]["reasoningText"]["text"] == "math" + assert rc_blocks[0]["reasoningContent"]["reasoningText"]["signature"] == "sig-1" + + def test_should_translate_inline_content_list_thinking_type(self): + """content=[{type:'thinking',...},{type:'text',...}] should also round-trip.""" + bedrock_msgs = self._to_bedrock( + [ + {"role": "user", "content": "Hi"}, + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": "hmm", "signature": "sig-2"}, + {"type": "text", "text": "Hello!"}, + ], + }, + {"role": "user", "content": "Bye"}, + ] + ) + assistant = next(m for m in bedrock_msgs if m["role"] == "assistant") + rc_blocks = [b for b in assistant["content"] if "reasoningContent" in b] + text_blocks = [ + b + for b in assistant["content"] + if "text" in b and "reasoningContent" not in b + ] + assert len(rc_blocks) >= 1 + assert any("Hello!" in b["text"] for b in text_blocks) diff --git a/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py b/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py index f3842214e4b..4cb20154570 100644 --- a/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py +++ b/tests/test_litellm/llms/hosted_vllm/embedding/test_hosted_vllm_embedding_transformation.py @@ -232,23 +232,6 @@ class TestHostedVLLMEmbeddingTransformation: assert result["Authorization"] == "Bearer test-api-key" assert result["Content-Type"] == "application/json" - def test_validate_environment_without_api_key(self): - """Test environment validation without API key (uses fake-api-key).""" - headers = {} - - result = self.config.validate_environment( - headers=headers, - model=self.model, - messages=[], - optional_params={}, - litellm_params={}, - api_key=None, - ) - - # Should not include Authorization header with fake-api-key - assert "Authorization" not in result - assert result["Content-Type"] == "application/json" - def test_encoding_format_not_sent_in_actual_request(self): """ E2E test that encoding_format is not sent when not provided. @@ -306,61 +289,5 @@ class TestHostedVLLMEmbeddingTransformation: assert sent_data["model"] == "BAAI/bge-small-en-v1.5" assert sent_data["input"] == ["Hello world"] - def test_encoding_format_float_sent_in_actual_request(self): - """ - Test that encoding_format='float' is sent when explicitly provided. - """ - from litellm.llms.custom_httpx.http_handler import HTTPHandler - - client = HTTPHandler() - - with patch.object(client, "post") as mock_post: - # Mock response - mock_response = Mock() - mock_response.status_code = 200 - mock_response.headers = {"content-type": "application/json"} - mock_response.json.return_value = { - "object": "list", - "data": [ - { - "object": "embedding", - "index": 0, - "embedding": [0.1, 0.2, 0.3, 0.4, 0.5], - } - ], - "model": "BAAI/bge-small-en-v1.5", - "usage": { - "prompt_tokens": 5, - "total_tokens": 5, - }, - } - mock_response.text = json.dumps(mock_response.json.return_value) - mock_post.return_value = mock_response - - try: - litellm.embedding( - model=self.model, - input=["Hello world"], - api_base="https://test-vllm.example.com/v1", - encoding_format="float", - client=client, - ) - except Exception: - pass - - # Verify the request was made - mock_post.assert_called_once() - - # Get the data that was sent - call_kwargs = mock_post.call_args[1] - sent_data = json.loads(call_kwargs["data"]) - - # Assert that encoding_format IS in the sent data - assert "encoding_format" in sent_data, ( - "encoding_format='float' should be in request when provided" - ) - assert sent_data["encoding_format"] == "float" - - if __name__ == "__main__": pytest.main([__file__, "-v", "-s"]) diff --git a/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py b/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py index e6ab199168d..37959f74086 100644 --- a/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py +++ b/tests/test_litellm/llms/openai/chat/test_openai_gpt_transformation.py @@ -10,8 +10,8 @@ import pytest sys.path.insert(0, os.path.abspath("../../../../..")) from litellm.llms.openai.chat.gpt_transformation import ( - OpenAIGPTConfig, OpenAIChatCompletionStreamingHandler, + OpenAIGPTConfig, ) @@ -204,6 +204,81 @@ class TestOpenAIChatCompletionStreamingHandler: assert result.choices[0].delta.content == "Hello" assert not hasattr(result, "usage") or result.usage is None + def test_chunk_parser_maps_reasoning_to_reasoning_content(self): + """ + Test that chunk_parser maps 'reasoning' field to 'reasoning_content'. + + Some OpenAI-compatible providers (e.g., GLM-5, hosted_vllm) return + delta.reasoning, but LiteLLM expects delta.reasoning_content. + + Regression test for: Streaming responses with delta.reasoning field + coming back empty when using openai/ or hosted_vllm/ providers. + """ + handler = OpenAIChatCompletionStreamingHandler( + streaming_response=None, sync_stream=True + ) + + # Simulate a chunk with reasoning field (as returned by GLM-5) + chunk = { + "id": "chatcmpl-8e3d624de9b12528", + "object": "chat.completion.chunk", + "created": 1771411455, + "model": "glm-5", + "choices": [ + { + "index": 0, + "delta": { + "reasoning": "The capital of France", + "role": None, + }, + "finish_reason": None, + } + ], + } + + # Parse the chunk + parsed_chunk = handler.chunk_parser(chunk) + + # Verify that reasoning was mapped to reasoning_content + assert parsed_chunk.choices[0].delta.reasoning_content == "The capital of France" + # Verify that the original 'reasoning' field was removed + assert not hasattr(parsed_chunk.choices[0].delta, "reasoning") + + def test_chunk_parser_reasoning_field_not_present(self): + """ + Test that chunks without reasoning field still work correctly. + """ + handler = OpenAIChatCompletionStreamingHandler( + streaming_response=None, sync_stream=True + ) + + # Simulate a chunk without reasoning field + chunk = { + "id": "chatcmpl-test", + "object": "chat.completion.chunk", + "created": 1769511767, + "model": "gpt-4o", + "choices": [ + { + "delta": { + "content": "Regular content", + "role": "assistant", + }, + "finish_reason": None, + "index": 0, + } + ], + } + + # Parse the chunk + parsed_chunk = handler.chunk_parser(chunk) + + # Verify that content is present + assert parsed_chunk.choices[0].delta.content == "Regular content" + assert parsed_chunk.choices[0].delta.role == "assistant" + # Verify that reasoning_content is not set (it should be deleted by Delta.__init__) + assert not hasattr(parsed_chunk.choices[0].delta, "reasoning_content") + class TestPromptCacheKeyIntegration: """Tests for prompt_cache_key support""" diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index 581d1e603dd..6047da66b6d 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -3224,6 +3224,7 @@ def test_video_metadata_only_for_gemini_3(): def test_chunk_parser_handles_prompt_feedback_block(): """Test chunk_parser correctly handles promptFeedback.blockReason""" from unittest.mock import Mock + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( ModelResponseIterator, ) @@ -3260,6 +3261,7 @@ def test_chunk_parser_handles_prompt_feedback_block(): def test_chunk_parser_handles_prompt_feedback_safety_block(): """Test chunk_parser handles different blockReason types (SAFETY)""" from unittest.mock import Mock + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( ModelResponseIterator, ) @@ -3294,6 +3296,7 @@ def test_chunk_parser_handles_prompt_feedback_safety_block(): def test_chunk_parser_handles_prompt_feedback_block_with_usage(): """Test chunk_parser correctly extracts usageMetadata when promptFeedback.blockReason is present""" from unittest.mock import Mock + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( ModelResponseIterator, ) @@ -3429,3 +3432,80 @@ def test_vertex_ai_traffic_type_surfaced_in_responses_api(): assert responses_api_response.provider_specific_fields["traffic_type"] == "ON_DEMAND" + +def test_vertex_ai_web_search_options_parameter(): + """ + Test that web_search_options parameter is transformed to googleSearch tool. + + When a user provides web_search_options as a parameter (not as a tool in the tools array), + it should be transformed to Gemini's googleSearch tool. + + This is important for the /v1/messages -> chat/completions -> Gemini flow: + - Anthropic web search tool -> web_search_options parameter -> Gemini googleSearch tool + + Input (optional_params): + {"web_search_options": {}} + + Expected Output: + tools=[{"googleSearch": {}}] + """ + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + + v = VertexGeminiConfig() + + # Simulate the map_openai_params flow + optional_params = {} + + # When web_search_options is present, it should be mapped to a tool + web_search_options = {} + _tools = v._map_web_search_options(web_search_options) + + # Verify the tool is a googleSearch tool + assert "googleSearch" in _tools, f"Expected googleSearch in tool, got {_tools.keys()}" + assert _tools["googleSearch"] == {}, f"Expected empty googleSearch config, got {_tools['googleSearch']}" + + +def test_vertex_ai_web_search_options_in_map_openai_params(): + """ + Test that web_search_options is properly handled in map_openai_params. + + This tests the full flow where web_search_options parameter is converted + to a googleSearch tool and added to optional_params. + + Input: + optional_params with web_search_options: {} + + Expected: + optional_params should have tools with googleSearch + """ + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + VertexGeminiConfig, + ) + + v = VertexGeminiConfig() + + # Simulate optional_params passed to map_openai_params + optional_params = { + "web_search_options": {} + } + + # Call the transformation that happens in map_openai_params + # Lines 1075-1079 in vertex_and_google_ai_studio_gemini.py (after fix) + web_search_value = optional_params.get("web_search_options") + if isinstance(web_search_value, dict): # Fixed: removed 'value and' check to support empty dicts + _tools = v._map_web_search_options(web_search_value) + # Simulate _add_tools_to_optional_params + optional_params = v._add_tools_to_optional_params(optional_params, [_tools]) + + # Remove web_search_options as it's been transformed + optional_params.pop("web_search_options", None) + + # Verify the transformation + assert "tools" in optional_params, "tools should be added to optional_params" + assert len(optional_params["tools"]) == 1, "Should have exactly one tool" + assert "googleSearch" in optional_params["tools"][0], "Tool should be googleSearch" + assert optional_params["tools"][0]["googleSearch"] == {}, "googleSearch should be empty config" + assert "web_search_options" not in optional_params, "web_search_options should be removed after transformation" + diff --git a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py index dd0a3e36e46..2f9a0b63921 100644 --- a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py +++ b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_integration.py @@ -3,13 +3,9 @@ Integration tests for Vertex AI rerank functionality. These tests demonstrate end-to-end usage of the Vertex AI rerank feature. """ import importlib -import os -from unittest.mock import MagicMock, patch +from unittest.mock import MagicMock import httpx -import pytest - -from litellm.llms.vertex_ai.rerank.transformation import VertexAIRerankConfig class TestVertexAIRerankIntegration: @@ -20,16 +16,25 @@ class TestVertexAIRerankIntegration: importlib.reload(rerank_transformation_module) # Re-import after reload to get the fresh class - from litellm.llms.vertex_ai.rerank.transformation import VertexAIRerankConfig as FreshConfig + from litellm.llms.vertex_ai.rerank.transformation import ( + VertexAIRerankConfig as FreshConfig, + ) self.config = FreshConfig() self.model = "semantic-ranker-default@latest" - @patch('litellm.llms.vertex_ai.rerank.transformation.VertexAIRerankConfig._ensure_access_token') - def test_end_to_end_rerank_flow(self, mock_ensure_access_token): - """Test complete rerank flow from request to response.""" - # Mock authentication - mock_ensure_access_token.return_value = ("test-access-token", "test-project-123") - + def test_end_to_end_rerank_flow(self): + """ + Test complete rerank flow from request to response. + + Uses instance-level mocking to avoid class-reference issues caused by + importlib.reload(litellm) in conftest.py. + """ + # Mock authentication at instance level + mock_ensure_access_token = MagicMock( + return_value=("test-access-token", "test-project-123") + ) + self.config._ensure_access_token = mock_ensure_access_token + # Test documents documents = [ "Gemini is a cutting edge large language model created by Google.", @@ -38,43 +43,40 @@ class TestVertexAIRerankIntegration: "Google's Gemini AI model represents a significant advancement in artificial intelligence technology." ] query = "What is Google Gemini?" - + # Step 1: Test request transformation - with patch.object(self.config, 'get_vertex_ai_credentials', return_value=None), \ - patch.object(self.config, 'get_vertex_ai_project', return_value="test-project-123"): - - # Validate environment - headers = self.config.validate_environment( - headers={}, - model=self.model, - api_key=None - ) - - # Transform request - request_data = self.config.transform_rerank_request( - model=self.model, - optional_rerank_params={ - "query": query, - "documents": documents, - "top_n": 2, - "return_documents": True - }, - headers=headers - ) - - # Verify request structure - assert request_data["model"] == self.model - assert request_data["query"] == query - assert request_data["topN"] == 2 - assert request_data["ignoreRecordDetailsInResponse"] == False - assert len(request_data["records"]) == 4 - - # Verify record structure - for i, record in enumerate(request_data["records"]): - assert record["id"] == str(i) # 0-based indexing - assert "title" in record - assert "content" in record - assert record["content"] == documents[i] + # Validate environment + headers = self.config.validate_environment( + headers={}, + model=self.model, + api_key=None + ) + + # Transform request + request_data = self.config.transform_rerank_request( + model=self.model, + optional_rerank_params={ + "query": query, + "documents": documents, + "top_n": 2, + "return_documents": True + }, + headers=headers + ) + + # Verify request structure + assert request_data["model"] == self.model + assert request_data["query"] == query + assert request_data["topN"] == 2 + assert request_data["ignoreRecordDetailsInResponse"] == False + assert len(request_data["records"]) == 4 + + # Verify record structure + for i, record in enumerate(request_data["records"]): + assert record["id"] == str(i) # 0-based indexing + assert "title" in record + assert "content" in record + assert record["content"] == documents[i] # Step 2: Test response transformation # Mock Vertex AI Discovery Engine response diff --git a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py index 2e631054143..5bf2cb97fa9 100644 --- a/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/rerank/test_vertex_ai_rerank_transformation.py @@ -41,17 +41,16 @@ class TestVertexAIRerankTransform: for var, value in self._saved_env.items(): os.environ[var] = value - @patch('litellm.llms.vertex_ai.rerank.transformation.get_secret_str') - @patch('litellm.llms.vertex_ai.rerank.transformation.VertexAIRerankConfig._ensure_access_token') - def test_get_complete_url(self, mock_ensure_access_token, mock_get_secret_str): - """Test URL generation for Vertex AI Discovery Engine rerank API.""" - # Mock _ensure_access_token to return (token, project_id) - mock_ensure_access_token.return_value = ("mock-token", None) + def test_get_complete_url(self): + """ + Test URL generation for Vertex AI Discovery Engine rerank API. - # Mock get_secret_str to return the environment variable value - def mock_get_secret(key): - return os.environ.get(key) - mock_get_secret_str.side_effect = mock_get_secret + Uses instance-level mocking to avoid class-reference issues caused by + importlib.reload(litellm) in conftest.py. + """ + # Mock _ensure_access_token at instance level to return (token, project_id) + mock_ensure_access_token = MagicMock(return_value=("mock-token", None)) + self.config._ensure_access_token = mock_ensure_access_token # Test with project ID from environment with patch.dict(os.environ, {"VERTEXAI_PROJECT": "test-project-123"}): @@ -96,34 +95,31 @@ class TestVertexAIRerankTransform: finally: litellm.vertex_project = original_project - @patch('litellm.llms.vertex_ai.rerank.transformation.get_secret_str') - @patch('litellm.llms.vertex_ai.rerank.transformation.VertexAIRerankConfig._ensure_access_token') - def test_validate_environment(self, mock_ensure_access_token, mock_get_secret_str): - """Test environment validation and header setup.""" - # Mock the authentication - mock_ensure_access_token.return_value = ("test-access-token", "test-project-123") + def test_validate_environment(self): + """ + Test environment validation and header setup. - # Mock get_secret_str to return the environment variable value - def mock_get_secret(key): - return os.environ.get(key) - mock_get_secret_str.side_effect = mock_get_secret + Uses instance-level mocking to avoid class-reference issues caused by + importlib.reload(litellm) in conftest.py. + """ + # Mock the authentication at instance level + mock_ensure_access_token = MagicMock( + return_value=("test-access-token", "test-project-123") + ) + self.config._ensure_access_token = mock_ensure_access_token - # Mock the credential and project methods - with patch.object(self.config, 'get_vertex_ai_credentials', return_value=None), \ - patch.object(self.config, 'get_vertex_ai_project', return_value="test-project-123"): + headers = self.config.validate_environment( + headers={}, + model=self.model, + api_key=None + ) - headers = self.config.validate_environment( - headers={}, - model=self.model, - api_key=None - ) - - expected_headers = { - "Authorization": "Bearer test-access-token", - "Content-Type": "application/json", - "X-Goog-User-Project": "test-project-123" - } - assert headers == expected_headers + expected_headers = { + "Authorization": "Bearer test-access-token", + "Content-Type": "application/json", + "X-Goog-User-Project": "test-project-123" + } + assert headers == expected_headers def test_transform_rerank_request_basic(self): """Test basic request transformation for Vertex AI Discovery Engine format.""" @@ -457,33 +453,40 @@ class TestVertexAIRerankTransform: assert params["top_n"] == 2 assert params["return_documents"] == True - @patch('litellm.llms.vertex_ai.rerank.transformation.VertexAIRerankConfig._ensure_access_token') - def test_validate_environment_with_optional_params(self, mock_ensure_access_token): - """Test that validate_environment accepts and uses optional_params for credentials.""" - # Mock the authentication - mock_ensure_access_token.return_value = ("test-access-token", "test-project-123") - + def test_validate_environment_with_optional_params(self): + """ + Test that validate_environment accepts and uses optional_params for credentials. + + Uses instance-level mocking to avoid class-reference issues caused by + importlib.reload(litellm) in conftest.py. + """ + # Mock the authentication at instance level + mock_ensure_access_token = MagicMock( + return_value=("test-access-token", "test-project-123") + ) + self.config._ensure_access_token = mock_ensure_access_token + optional_params = { "vertex_credentials": "path/to/credentials.json", "vertex_project": "custom-project-id", "query": "test query", "documents": ["doc1"] } - + headers = self.config.validate_environment( headers={}, model=self.model, api_key=None, optional_params=optional_params ) - + # Verify that _ensure_access_token was called with the credentials from optional_params mock_ensure_access_token.assert_called_once() call_args = mock_ensure_access_token.call_args # The first call argument should be credentials (which will be the value from optional_params) # We can't check the exact value easily due to how get_vertex_ai_credentials pops values, # but we can verify the headers were set correctly - + expected_headers = { "Authorization": "Bearer test-access-token", "Content-Type": "application/json", diff --git a/tests/test_litellm/llms/watsonx/rerank/__init__.py b/tests/test_litellm/llms/watsonx/rerank/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/test_litellm/llms/watsonx/rerank/test_watsonx_rerank.py b/tests/test_litellm/llms/watsonx/rerank/test_watsonx_rerank.py new file mode 100644 index 00000000000..4a2edd9810b --- /dev/null +++ b/tests/test_litellm/llms/watsonx/rerank/test_watsonx_rerank.py @@ -0,0 +1,224 @@ +""" +Tests for IBM watsonx.ai rerank transformation functionality. +""" +import json +import re +import uuid +from unittest.mock import MagicMock + +import httpx +import pytest + +from litellm.llms.watsonx.common_utils import ( + WatsonXAIError, +) +from litellm.llms.watsonx.rerank.transformation import IBMWatsonXRerankConfig +from litellm.types.rerank import RerankResponse + + +class TestIBMWatsonXRerankTransform: + def setup_method(self): + self.config = IBMWatsonXRerankConfig() + self.model = "watsonx/cross-encoder/ms-marco-minilm-l-12-v2" + + def test_get_complete_url(self): + """Test URL generation for IBM watsonx.ai rerank API.""" + + api_base = "https://us-south.ml.cloud.ibm.com" + model = "watsonx/cross-encoder/ms-marco-minilm-l-12-v2" + url = self.config.get_complete_url(api_base, model) + assert url == "https://us-south.ml.cloud.ibm.com/ml/v1/text/rerank?version=2024-03-13" + + def test_map_cohere_rerank_params_basic(self): + """Test basic parameter mapping for IBM watsonx.ai rerank.""" + params = self.config.map_cohere_rerank_params( + non_default_params={ + "query": "hello", + "documents": ["hello", "world"], + "top_n": 2, + "return_documents": True, + "max_tokens_per_doc": 100, + }, + model="test", + drop_params=False, + query="hello", + documents=["hello", "world"], + ) + assert params["query"] == "hello" + assert params["inputs"] == [{"text": "hello"}, {"text": "world"}] + assert params["parameters"]["return_options"]["top_n"] == 2 + assert params["parameters"]["return_options"]["inputs"] is True + assert params["parameters"]["truncate_input_tokens"] == 100 + + def test_transform_rerank_request(self): + """Test request transformation for IBM watsonx.ai format.""" + optional_params = { + "query": "What is the capital of France?", + "documents": [ + "Paris is the capital of France.", + "France is a country in Europe.", + ], + "top_n": 2, + "return_documents": True, + "project_id": uuid.uuid4(), + } + + request_body = self.config.transform_rerank_request( + model="cross-encoder/ms-marco-minilm-l-12-v2", optional_rerank_params=optional_params, headers={} + ) + + assert request_body["model_id"] == "cross-encoder/ms-marco-minilm-l-12-v2" + assert request_body["project_id"] is not None + assert request_body["query"] == "What is the capital of France?" + assert request_body["documents"] == optional_params["documents"] + assert request_body["top_n"] == 2 + assert request_body["return_documents"] is True + + def test_transform_rerank_response_success(self): + """Test successful response transformation.""" + # Mock IBM watsonx.ai response format + response_data = { + "model_id": self.model, + "results": [ + { + "index": 0, + "score": 6.53515625, + "input": {"text": "Python is great for beginners due to simple syntax."}, + }, + {"index": 1, "score": -7.1875, "input": {"text": "JavaScript runs in browsers and is versatile."}}, + ], + "input_token_count": 62, + } + + # Create mock httpx response + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = response_data + mock_response.status_code = 200 + mock_response.headers = {} + + # Create mock logging object + mock_logging = MagicMock() + + model_response = RerankResponse() + + result = self.config.transform_rerank_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=mock_logging, + ) + + # Verify response structure + # IBM watsonx.ai doesn't return "id", so it uses "model" as the id + assert result.id == "watsonx/cross-encoder/ms-marco-minilm-l-12-v2" + assert len(result.results) == 2 + assert result.results[0]["index"] == 0 + assert result.results[0]["relevance_score"] == 6.53515625 + assert result.results[0]["document"]["text"] == "Python is great for beginners due to simple syntax." + assert result.results[1]["index"] == 1 + assert result.results[1]["relevance_score"] == -7.1875 + assert result.results[1]["document"]["text"] == "JavaScript runs in browsers and is versatile." + + # # Verify metadata + assert result.meta["tokens"]["input_tokens"] == 62 + + def test_transform_rerank_response_without_documents(self): + """Test response transformation when return_documents is False.""" + response_data = { + "model_id": self.model, + "results": [ + { + "index": 0, + "score": 6.53515625, + }, + { + "index": 1, + "score": -7.1875, + }, + ], + "input_token_count": 62, + } + + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = response_data + mock_response.status_code = 200 + mock_response.headers = {} + + mock_logging = MagicMock() + model_response = RerankResponse() + + result = self.config.transform_rerank_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=mock_logging, + ) + + # Verify response structure + # IBM watsonx.ai doesn't return "id", so it uses "model" as the id + assert result.id == "watsonx/cross-encoder/ms-marco-minilm-l-12-v2" + assert len(result.results) == 2 + + assert result.results[0]["index"] == 0 + assert result.results[0]["relevance_score"] == 6.53515625 + assert "document" not in result.results[0] + + assert result.results[1]["index"] == 1 + assert result.results[1]["relevance_score"] == -7.1875 + assert "document" not in result.results[1] + + def test_transform_rerank_response_missing_results(self): + """Test that missing results raises ValueError.""" + response_data = { + "model": self.model, + "usage": {"total_tokens": 10}, + } + + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.return_value = response_data + mock_response.status_code = 200 + mock_response.headers = {} + + mock_logging = MagicMock() + model_response = RerankResponse() + + expected_error_msg = re.escape("No results found") + + with pytest.raises(ValueError, match=expected_error_msg): + self.config.transform_rerank_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=mock_logging, + ) + + def test_transform_rerank_response_invalid_json(self): + """Test error handling for invalid JSON response.""" + mock_response = MagicMock(spec=httpx.Response) + mock_response.json.side_effect = json.JSONDecodeError("Invalid JSON", "doc", 0) + mock_response.text = "Invalid JSON response" + mock_response.status_code = 500 + mock_response.headers = {} + + mock_logging = MagicMock() + model_response = RerankResponse() + + expected_error_msg = re.escape("Failed to parse response") + + with pytest.raises(Exception, match=expected_error_msg): + self.config.transform_rerank_response( + model=self.model, + raw_response=mock_response, + model_response=model_response, + logging_obj=mock_logging, + ) + + def test_get_supported_cohere_rerank_params(self): + """Test getting supported parameters for IBM watsonx.ai rerank.""" + supported_params = self.config.get_supported_cohere_rerank_params(self.model) + assert "query" in supported_params + assert "documents" in supported_params + assert "top_n" in supported_params + assert "return_documents" in supported_params + assert "max_tokens_per_doc" in supported_params + assert len(supported_params) == 5 diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py index 0fb299a57cd..de2037793c5 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py @@ -230,7 +230,7 @@ class TestWrapSendWithDebugHeaders: ) message = {"type": "http.response.start", "status": 200, "headers": []} - asyncio.get_event_loop().run_until_complete(wrapped(message)) + asyncio.run(wrapped(message)) assert len(captured) == 1 headers = dict(captured[0]["headers"]) @@ -247,6 +247,6 @@ class TestWrapSendWithDebugHeaders: ) body_msg = {"type": "http.response.body", "body": b"hello"} - asyncio.get_event_loop().run_until_complete(wrapped(body_msg)) + asyncio.run(wrapped(body_msg)) assert captured[0] == body_msg diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index 1a50cacd308..464e5238325 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -39,7 +39,16 @@ def _reload_mcp_manager_module(): "litellm.proxy._experimental.mcp_server.mcp_server_manager" ] importlib.reload(utils_module) - return importlib.reload(manager_module) + reloaded = importlib.reload(manager_module) + # After reload, server.py still holds a stale reference to the old + # global_mcp_server_manager. Update it so tests that exercise server.py + # functions (e.g. _get_tools_from_mcp_servers) use the fresh instance. + server_module = sys.modules.get( + "litellm.proxy._experimental.mcp_server.server" + ) + if server_module is not None and hasattr(server_module, "global_mcp_server_manager"): + server_module.global_mcp_server_manager = reloaded.global_mcp_server_manager + return reloaded class TestMCPServerManager: diff --git a/tests/test_litellm/proxy/auth/test_model_checks.py b/tests/test_litellm/proxy/auth/test_model_checks.py index 789af480e72..193b014f03d 100644 --- a/tests/test_litellm/proxy/auth/test_model_checks.py +++ b/tests/test_litellm/proxy/auth/test_model_checks.py @@ -62,3 +62,27 @@ def test_get_complete_model_list_order(key_models, team_models, proxy_model_list infer_model_from_keys=False, llm_router=Router(model_list=model_list), ) == expected + + +def test_get_complete_model_list_byok_wildcard_expansion(): + """ + Test that wildcard models (e.g., openai/*) are expanded when the router has + no deployment for them - BYOK case where team has openai/* but proxy has + no openai config. + """ + from litellm.proxy.auth.model_checks import get_complete_model_list + from litellm import Router + + # Router with empty model_list - no openai/* deployment (BYOK scenario) + result = get_complete_model_list( + key_models=[], + team_models=["openai/*"], + proxy_model_list=[], + user_model=None, + infer_model_from_keys=False, + llm_router=Router(model_list=[]), + ) + # Should expand openai/* to actual OpenAI models + assert len(result) > 0 + assert all(m.startswith("openai/") for m in result) + assert "openai/*" not in result diff --git a/tests/test_litellm/proxy/auth/test_object_permission_loading.py b/tests/test_litellm/proxy/auth/test_object_permission_loading.py index 54e4c82471e..4db969c95e0 100644 --- a/tests/test_litellm/proxy/auth/test_object_permission_loading.py +++ b/tests/test_litellm/proxy/auth/test_object_permission_loading.py @@ -44,6 +44,11 @@ async def test_get_key_object_loads_object_permission(): vector_stores=["store1"], ) + # Mock proxy_logging_obj to handle async service hooks + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock() + mock_proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() + # Mock get_object_permission to return the permission with patch( "litellm.proxy.auth.auth_checks.get_object_permission", @@ -51,6 +56,9 @@ async def test_get_key_object_loads_object_permission(): ), patch( "litellm.proxy.auth.auth_checks._cache_key_object", AsyncMock() + ), patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + mock_proxy_logging_obj ): result = await get_key_object( hashed_token="test_token_hash", @@ -84,9 +92,17 @@ async def test_get_key_object_no_permission_id(): } mock_prisma_client.get_data = AsyncMock(return_value=mock_token_data) + # Mock proxy_logging_obj to handle async service hooks + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock() + mock_proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() + with patch( "litellm.proxy.auth.auth_checks._cache_key_object", AsyncMock() + ), patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + mock_proxy_logging_obj ): result = await get_key_object( hashed_token="test_token_hash", @@ -124,6 +140,11 @@ async def test_get_team_object_loads_object_permission(): vector_stores=["team_store1"], ) + # Mock proxy_logging_obj to handle async service hooks + mock_proxy_logging_obj = MagicMock() + mock_proxy_logging_obj.service_logging_obj.async_service_success_hook = AsyncMock() + mock_proxy_logging_obj.service_logging_obj.async_service_failure_hook = AsyncMock() + with patch( "litellm.proxy.auth.auth_checks._get_team_db_check", AsyncMock(return_value=mock_team) @@ -138,6 +159,9 @@ async def test_get_team_object_loads_object_permission(): return_value=True ), patch( "litellm.proxy.auth.auth_checks._update_last_db_access_time" + ), patch( + "litellm.proxy.proxy_server.proxy_logging_obj", + mock_proxy_logging_obj ): result = await get_team_object( team_id="test_team", diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 00e348b5b7c..68abf2a1c93 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -300,28 +300,36 @@ async def test_proxy_admin_expired_key_from_cache(): mock_get_key_object.return_value = expired_token # Set attributes on proxy_server module (these are imported inside _user_api_key_auth_builder) - import litellm.proxy.proxy_server - - setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma_client) - setattr(litellm.proxy.proxy_server, "user_api_key_cache", mock_cache) - setattr(litellm.proxy.proxy_server, "proxy_logging_obj", mock_proxy_logging_obj) - setattr(litellm.proxy.proxy_server, "master_key", "sk-master-key") - setattr(litellm.proxy.proxy_server, "general_settings", {}) - setattr(litellm.proxy.proxy_server, "llm_model_list", []) - setattr(litellm.proxy.proxy_server, "llm_router", None) - setattr(litellm.proxy.proxy_server, "open_telemetry_logger", None) - setattr(litellm.proxy.proxy_server, "model_max_budget_limiter", MagicMock()) - setattr(litellm.proxy.proxy_server, "user_custom_auth", None) - setattr(litellm.proxy.proxy_server, "jwt_handler", None) - setattr(litellm.proxy.proxy_server, "litellm_proxy_admin_name", "admin") - + import litellm.proxy.proxy_server as _proxy_server_mod + + _attrs_to_set = { + "prisma_client": mock_prisma_client, + "user_api_key_cache": mock_cache, + "proxy_logging_obj": mock_proxy_logging_obj, + "master_key": "sk-master-key", + "general_settings": {}, + "llm_model_list": [], + "llm_router": None, + "open_telemetry_logger": None, + "model_max_budget_limiter": MagicMock(), + "user_custom_auth": None, + "jwt_handler": None, + "litellm_proxy_admin_name": "admin", + } + _original_values = { + attr: getattr(_proxy_server_mod, attr, None) + for attr in _attrs_to_set + } try: - + for attr, val in _attrs_to_set.items(): + setattr(_proxy_server_mod, attr, val) + + # Create a mock request request = Request(scope={"type": "http"}) request._url = URL(url="/chat/completions") request_data = {} - + # Call the auth builder - should raise ProxyException for expired key # Note: api_key needs "Bearer " prefix for get_api_key() to process it correctly with pytest.raises(ProxyException) as exc_info: @@ -334,7 +342,7 @@ async def test_proxy_admin_expired_key_from_cache(): azure_apim_header=None, request_data=request_data, ) - + # Verify that ProxyException was raised with expired_key type assert hasattr(exc_info.value, "type"), "Exception should have 'type' attribute" assert exc_info.value.type == ProxyErrorTypes.expired_key, ( @@ -343,7 +351,7 @@ async def test_proxy_admin_expired_key_from_cache(): assert "Expired Key" in str(exc_info.value.message), ( f"Exception message should mention 'Expired Key', got: {exc_info.value.message}" ) - + # Verify that the param field does NOT leak the full API key (Issue #18731) # The param should be abbreviated like "sk-...XXXX" not the full plaintext key assert exc_info.value.param is not None, "Exception should have 'param' attribute" @@ -354,7 +362,7 @@ async def test_proxy_admin_expired_key_from_cache(): assert exc_info.value.param.startswith("sk-..."), ( f"Param should be abbreviated to 'sk-...XXXX' format. Got: {exc_info.value.param}" ) - + # Verify that cache deletion was called mock_delete_cache.assert_called_once() call_args = mock_delete_cache.call_args @@ -362,8 +370,9 @@ async def test_proxy_admin_expired_key_from_cache(): "Cache deletion should be called with the hashed key" ) finally: - # Clean up - restore original values if needed - pass + # Restore all module-level attributes so subsequent tests are not affected + for attr, val in _original_values.items(): + setattr(_proxy_server_mod, attr, val) diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py index 7d2b6e84de7..a3c1fd9ea05 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_generic_guardrail_api.py @@ -13,7 +13,7 @@ import pytest import litellm from litellm import ModelResponse -from litellm.exceptions import GuardrailRaisedException +from litellm.exceptions import GuardrailRaisedException, Timeout from litellm._version import version as litellm_version from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.generic_guardrail_api import ( @@ -704,6 +704,104 @@ class TestErrorHandling: assert "Generic Guardrail API failed" in str(exc_info.value) + @pytest.mark.asyncio + async def test_network_error_defaults_to_fail_closed_when_unreachable_fallback_not_set( + self, mock_request_data_input + ): + """Test default behavior is fail_closed when unreachable_fallback is omitted""" + guardrail = GenericGuardrailAPI( + api_base="https://api.test.guardrail.com", + headers={"Authorization": "Bearer test-key"}, + ) + + with patch.object( + guardrail.async_handler, + "post", + side_effect=httpx.RequestError("Connection failed", request=MagicMock()), + ): + with pytest.raises(Exception) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["test"]}, + request_data=mock_request_data_input, + input_type="request", + ) + + assert "Generic Guardrail API failed" in str(exc_info.value) + + @pytest.mark.asyncio + async def test_network_error_fail_open_allows_flow(self, mock_request_data_input): + """Test network error handling allows flow when unreachable_fallback=fail_open""" + guardrail = GenericGuardrailAPI( + api_base="https://api.test.guardrail.com", + headers={"Authorization": "Bearer test-key"}, + unreachable_fallback="fail_open", + ) + + with patch.object( + guardrail.async_handler, + "post", + side_effect=httpx.RequestError("Connection failed", request=MagicMock()), + ): + result = await guardrail.apply_guardrail( + inputs={"texts": ["test"]}, + request_data=mock_request_data_input, + input_type="request", + ) + + assert result.get("texts") == ["test"] + + @pytest.mark.asyncio + async def test_503_fail_open_allows_flow(self, mock_request_data_input): + """Test HTTP 503 allows flow when unreachable_fallback=fail_open""" + guardrail = GenericGuardrailAPI( + api_base="https://api.test.guardrail.com", + headers={"Authorization": "Bearer test-key"}, + unreachable_fallback="fail_open", + ) + + with patch.object( + guardrail.async_handler, + "post", + side_effect=httpx.HTTPStatusError( + "Service Unavailable", + request=MagicMock(), + response=MagicMock(status_code=503), + ), + ): + result = await guardrail.apply_guardrail( + inputs={"texts": ["test"]}, + request_data=mock_request_data_input, + input_type="request", + ) + + assert result.get("texts") == ["test"] + + @pytest.mark.asyncio + async def test_timeout_fail_open_allows_flow(self, mock_request_data_input): + """Test litellm.Timeout allows flow when unreachable_fallback=fail_open""" + guardrail = GenericGuardrailAPI( + api_base="https://api.test.guardrail.com", + headers={"Authorization": "Bearer test-key"}, + unreachable_fallback="fail_open", + ) + + with patch.object( + guardrail.async_handler, + "post", + side_effect=Timeout( + message="Connection timed out", + model="default-model-name", + llm_provider="litellm-httpx-handler", + ), + ): + result = await guardrail.apply_guardrail( + inputs={"texts": ["test"]}, + request_data=mock_request_data_input, + input_type="request", + ) + + assert result.get("texts") == ["test"] + class TestMultimodalSupport: """Test multimodal (image) message handling and serialization""" @@ -830,4 +928,4 @@ class TestMultimodalSupport: # Verify serialization succeeded call_args = mock_post.call_args json_payload = call_args.kwargs["json"] - assert isinstance(json_payload["structured_messages"], list) \ No newline at end of file + assert isinstance(json_payload["structured_messages"], list) diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 9a417f3566c..839885bc752 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -642,12 +642,9 @@ async def test_new_user_default_teams_flow(mocker): assert response.key == "sk-test-token-123" finally: - # Restore original default params - if original_default_params is not None: - litellm.default_internal_user_params = original_default_params - else: - if hasattr(litellm, "default_internal_user_params"): - delattr(litellm, "default_internal_user_params") + # Restore original default params (always assign, never delattr — the attribute + # is defined in litellm/__init__.py and delattr-ing it breaks parallel tests) + litellm.default_internal_user_params = original_default_params def test_update_internal_new_user_params_proxy_admin_role(): @@ -694,12 +691,7 @@ def test_update_internal_new_user_params_proxy_admin_role(): assert result["user_role"] == LitellmUserRoles.PROXY_ADMIN.value finally: - # Restore original default params - if original_default_params is not None: - litellm.default_internal_user_params = original_default_params - else: - if hasattr(litellm, "default_internal_user_params"): - delattr(litellm, "default_internal_user_params") + litellm.default_internal_user_params = original_default_params def test_update_internal_new_user_params_no_role_specified(): @@ -735,12 +727,7 @@ def test_update_internal_new_user_params_no_role_specified(): assert result["user_email"] == "user@example.com" finally: - # Restore original default params - if original_default_params is not None: - litellm.default_internal_user_params = original_default_params - else: - if hasattr(litellm, "default_internal_user_params"): - delattr(litellm, "default_internal_user_params") + litellm.default_internal_user_params = original_default_params def test_update_internal_new_user_params_internal_user_role(): @@ -780,12 +767,7 @@ def test_update_internal_new_user_params_internal_user_role(): assert result["user_role"] == LitellmUserRoles.INTERNAL_USER.value finally: - # Restore original default params - if original_default_params is not None: - litellm.default_internal_user_params = original_default_params - else: - if hasattr(litellm, "default_internal_user_params"): - delattr(litellm, "default_internal_user_params") + litellm.default_internal_user_params = original_default_params @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index b3f7b211951..efa7d27ec47 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -5609,6 +5609,122 @@ async def test_validate_key_list_check_key_hash_not_found(): @pytest.mark.asyncio +@patch( + "litellm.proxy.management_endpoints.key_management_endpoints.rotate_mcp_server_credentials_master_key" +) +async def test_rotate_master_key_model_data_valid_for_prisma( + mock_rotate_mcp, +): + """ + Test that _rotate_master_key produces valid data for Prisma create_many(). + + Regression test for: master key rotation fails with Prisma validation error + because created_at/updated_at are None (non-nullable DateTime) and + litellm_params/model_info are JSON strings (create_many expects dicts). + """ + from unittest.mock import AsyncMock, MagicMock + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.key_management_endpoints import ( + _rotate_master_key, + ) + + # Setup mock prisma client + mock_prisma_client = AsyncMock() + mock_prisma_client.db = MagicMock() + + # Mock model table — return one model + mock_model = MagicMock() + mock_model.model_id = "model-1" + mock_model.model_name = "test-model" + mock_model.litellm_params = '{"model": "openai/gpt-4", "api_key": "sk-encrypted-old"}' + mock_model.model_info = '{"id": "model-1"}' + mock_model.created_by = "admin" + mock_model.updated_by = "admin" + mock_prisma_client.db.litellm_proxymodeltable.find_many = AsyncMock( + return_value=[mock_model] + ) + + # Mock transaction context manager + mock_tx = AsyncMock() + mock_tx.litellm_proxymodeltable = MagicMock() + mock_tx.litellm_proxymodeltable.delete_many = AsyncMock() + mock_tx.litellm_proxymodeltable.create_many = AsyncMock() + mock_prisma_client.db.tx = MagicMock(return_value=AsyncMock( + __aenter__=AsyncMock(return_value=mock_tx), + __aexit__=AsyncMock(return_value=False), + )) + + # Mock config table — no env vars + mock_prisma_client.db.litellm_config.find_many = AsyncMock(return_value=[]) + + # Mock credentials table — no credentials + mock_prisma_client.db.litellm_credentialstable.find_many = AsyncMock( + return_value=[] + ) + + # Mock MCP rotation + mock_rotate_mcp.return_value = None + + # Mock proxy_config + mock_proxy_config = MagicMock() + mock_proxy_config.decrypt_model_list_from_db.return_value = [ + { + "model_name": "test-model", + "litellm_params": { + "model": "openai/gpt-4", + "api_key": "sk-decrypted-key", + }, + "model_info": {"id": "model-1"}, + } + ] + + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-1234", + user_id="test-user", + ) + + with patch( + "litellm.proxy.proxy_server.proxy_config", + mock_proxy_config, + ): + await _rotate_master_key( + prisma_client=mock_prisma_client, + user_api_key_dict=user_api_key_dict, + current_master_key="sk-old-master-key", + new_master_key="sk-new-master-key", + ) + + # Verify create_many was called + mock_tx.litellm_proxymodeltable.create_many.assert_called_once() + + # Get the data passed to create_many + call_args = mock_tx.litellm_proxymodeltable.create_many.call_args + created_models = call_args.kwargs.get("data") or call_args[1].get("data") + + assert len(created_models) == 1 + model_data = created_models[0] + + # Verify timestamps are NOT present (Prisma @default(now()) should apply) + assert "created_at" not in model_data, ( + "created_at should be excluded so Prisma @default(now()) applies" + ) + assert "updated_at" not in model_data, ( + "updated_at should be excluded so Prisma @default(now()) applies" + ) + + # Verify litellm_params and model_info are prisma.Json wrappers, NOT JSON strings + import prisma + + assert isinstance(model_data["litellm_params"], prisma.Json), ( + f"litellm_params should be prisma.Json for create_many(), got {type(model_data['litellm_params'])}" + ) + assert isinstance(model_data["model_info"], prisma.Json), ( + f"model_info should be prisma.Json for create_many(), got {type(model_data['model_info'])}" + ) + + # Verify delete_many was called inside the transaction (before create_many) + mock_tx.litellm_proxymodeltable.delete_many.assert_called_once() async def test_default_key_generate_params_duration(monkeypatch): """ Test that default_key_generate_params with 'duration' is applied diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 09b78335054..b022adae4e0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -3682,12 +3682,9 @@ async def test_role_mappings_override_default_internal_user_params(): # The models will be applied when new_user processes the request finally: - # Restore original default_internal_user_params - if original_default_params is not None: - litellm.default_internal_user_params = original_default_params - else: - if hasattr(litellm, "default_internal_user_params"): - delattr(litellm, "default_internal_user_params") + # Restore original default_internal_user_params (always assign, never delattr — + # the attribute is defined in litellm/__init__.py and delattr-ing it breaks parallel tests) + litellm.default_internal_user_params = original_default_params class TestSSOReadinessEndpoint: diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 452db3902c0..a8172e85910 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -1014,91 +1014,95 @@ async def test_add_litellm_metadata_from_request_headers(): # Set up test logger litellm._turn_on_debug() test_logger = TestCustomLogger() + original_callbacks = litellm.callbacks litellm.callbacks = [test_logger] - # Prepare test data (ensure no streaming, add mock_response and api_key to route to litellm.acompletion) - headers = {"x-litellm-spend-logs-metadata": '{"user_id": "12345", "project_id": "proj_abc", "request_type": "chat_completion", "timestamp": "2025-09-02T10:30:00Z"}'} - data = {"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], "stream": False, "mock_response": "Hi", "api_key": "fake-key"} - - # Create mock request with headers - mock_request = MagicMock(spec=Request) - mock_request.headers = headers - mock_request.url.path = "/chat/completions" - - # Create mock response - mock_fastapi_response = MagicMock(spec=Response) - - # Create mock user API key dict - mock_user_api_key_dict = UserAPIKeyAuth( - api_key="test-key", - user_id="test-user", - org_id="test-org" - ) - - # Create mock proxy logging object - mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) - - # Create async functions for the hooks - async def mock_during_call_hook(*args, **kwargs): - return None - - async def mock_pre_call_hook(*args, **kwargs): - return data - - async def mock_post_call_success_hook(*args, **kwargs): - # Return the response unchanged - return kwargs.get('response', args[2] if len(args) > 2 else None) - - mock_proxy_logging_obj.during_call_hook = mock_during_call_hook - mock_proxy_logging_obj.pre_call_hook = mock_pre_call_hook - mock_proxy_logging_obj.post_call_success_hook = mock_post_call_success_hook - - # Create mock proxy config - mock_proxy_config = MagicMock() - - # Create mock general settings - general_settings = {} - - # Create mock select_data_generator with correct signature - def mock_select_data_generator(response=None, user_api_key_dict=None, request_data=None): - async def mock_generator(): - yield "data: " + json.dumps({"choices": [{"delta": {"content": "Hello"}}]}) + "\n\n" - yield "data: [DONE]\n\n" - return mock_generator() - - # Create the processor - processor = ProxyBaseLLMRequestProcessing(data=data) - - # Call base_process_llm_request (it will use the mock_response="Hi" parameter) - result = await processor.base_process_llm_request( - request=mock_request, - fastapi_response=mock_fastapi_response, - user_api_key_dict=mock_user_api_key_dict, - route_type="acompletion", - proxy_logging_obj=mock_proxy_logging_obj, - general_settings=general_settings, - proxy_config=mock_proxy_config, - select_data_generator=mock_select_data_generator, - llm_router=None, - model="gpt-4", - is_streaming_request=False - ) - - # Sleep for 3 seconds to allow logging to complete - await asyncio.sleep(3) - - # Check if standard_logging_object was set - assert test_logger.standard_logging_object is not None, "standard_logging_object should be populated after LLM request" - - # Verify the logging object contains expected metadata - standard_logging_obj = test_logger.standard_logging_object + try: + # Prepare test data (ensure no streaming, add mock_response and api_key to route to litellm.acompletion) + headers = {"x-litellm-spend-logs-metadata": '{"user_id": "12345", "project_id": "proj_abc", "request_type": "chat_completion", "timestamp": "2025-09-02T10:30:00Z"}'} + data = {"model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], "stream": False, "mock_response": "Hi", "api_key": "fake-key"} - print(f"Standard logging object captured: {json.dumps(standard_logging_obj, indent=4, default=str)}") + # Create mock request with headers + mock_request = MagicMock(spec=Request) + mock_request.headers = headers + mock_request.url.path = "/chat/completions" + + # Create mock response + mock_fastapi_response = MagicMock(spec=Response) + + # Create mock user API key dict + mock_user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + user_id="test-user", + org_id="test-org" + ) + + # Create mock proxy logging object + mock_proxy_logging_obj = MagicMock(spec=ProxyLogging) + + # Create async functions for the hooks + async def mock_during_call_hook(*args, **kwargs): + return None + + async def mock_pre_call_hook(*args, **kwargs): + return data + + async def mock_post_call_success_hook(*args, **kwargs): + # Return the response unchanged + return kwargs.get('response', args[2] if len(args) > 2 else None) + + mock_proxy_logging_obj.during_call_hook = mock_during_call_hook + mock_proxy_logging_obj.pre_call_hook = mock_pre_call_hook + mock_proxy_logging_obj.post_call_success_hook = mock_post_call_success_hook + + # Create mock proxy config + mock_proxy_config = MagicMock() + + # Create mock general settings + general_settings = {} + + # Create mock select_data_generator with correct signature + def mock_select_data_generator(response=None, user_api_key_dict=None, request_data=None): + async def mock_generator(): + yield "data: " + json.dumps({"choices": [{"delta": {"content": "Hello"}}]}) + "\n\n" + yield "data: [DONE]\n\n" + return mock_generator() + + # Create the processor + processor = ProxyBaseLLMRequestProcessing(data=data) + + # Call base_process_llm_request (it will use the mock_response="Hi" parameter) + result = await processor.base_process_llm_request( + request=mock_request, + fastapi_response=mock_fastapi_response, + user_api_key_dict=mock_user_api_key_dict, + route_type="acompletion", + proxy_logging_obj=mock_proxy_logging_obj, + general_settings=general_settings, + proxy_config=mock_proxy_config, + select_data_generator=mock_select_data_generator, + llm_router=None, + model="gpt-4", + is_streaming_request=False + ) + + # Sleep for 3 seconds to allow logging to complete + await asyncio.sleep(3) + + # Check if standard_logging_object was set + assert test_logger.standard_logging_object is not None, "standard_logging_object should be populated after LLM request" + + # Verify the logging object contains expected metadata + standard_logging_obj = test_logger.standard_logging_object + + print(f"Standard logging object captured: {json.dumps(standard_logging_obj, indent=4, default=str)}") + + SPEND_LOGS_METADATA = standard_logging_obj["metadata"]["spend_logs_metadata"] + assert SPEND_LOGS_METADATA == dict(json.loads(headers["x-litellm-spend-logs-metadata"])), "spend_logs_metadata should be the same as the headers" + finally: + litellm.callbacks = original_callbacks - SPEND_LOGS_METADATA = standard_logging_obj["metadata"]["spend_logs_metadata"] - assert SPEND_LOGS_METADATA == dict(json.loads(headers["x-litellm-spend-logs-metadata"])), "spend_logs_metadata should be the same as the headers" - def test_get_internal_user_header_from_mapping_returns_expected_header(): mappings = [ diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index 2696867d017..f5d5a11ea11 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -25,6 +25,7 @@ sys.path.insert( import litellm from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.proxy_server import app, initialize +from litellm.utils import _invalidate_model_cost_lowercase_map example_embedding_result = { "object": "list", @@ -1743,30 +1744,39 @@ class TestPriceDataReloadAPI: def test_reload_model_cost_map_admin_access(self, client_with_auth): """Test that admin users can access the reload endpoint""" - with patch( - "litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map" - ) as mock_get_map: - mock_get_map.return_value = { - "gpt-3.5-turbo": {"input_cost_per_token": 0.001} - } - # Mock the database connection - with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: - mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) + # Save the original model_cost so the endpoint's direct assignment + # (litellm.model_cost = new_model_cost_map) does not contaminate + # subsequent tests running in the same worker process. + original_model_cost = litellm.model_cost.copy() + try: + with patch( + "litellm.litellm_core_utils.get_model_cost_map.get_model_cost_map" + ) as mock_get_map: + mock_get_map.return_value = { + "gpt-3.5-turbo": {"input_cost_per_token": 0.001} + } + # Mock the database connection + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.db.litellm_config.upsert = AsyncMock(return_value=None) - response = client_with_auth.post("/reload/model_cost_map") + response = client_with_auth.post("/reload/model_cost_map") - assert response.status_code == 200 - data = response.json() - assert data["status"] == "success" - assert "message" in data - assert "timestamp" in data - assert "models_count" in data - # The new implementation immediately reloads and returns the count - assert ( - "Price data reloaded successfully! 1 models updated." - in data["message"] - ) - assert data["models_count"] == 1 + assert response.status_code == 200 + data = response.json() + assert data["status"] == "success" + assert "message" in data + assert "timestamp" in data + assert "models_count" in data + # The new implementation immediately reloads and returns the count + assert ( + "Price data reloaded successfully! 1 models updated." + in data["message"] + ) + assert data["models_count"] == 1 + finally: + # Restore the full model cost map so subsequent tests are not affected + litellm.model_cost = original_model_cost + _invalidate_model_cost_lowercase_map() def test_reload_model_cost_map_non_admin_access(self, client_with_auth): """Test that non-admin users cannot access the reload endpoint""" diff --git a/tests/test_litellm/responses/test_metadata_codex_callback.py b/tests/test_litellm/responses/test_metadata_codex_callback.py new file mode 100644 index 00000000000..21c0c644521 --- /dev/null +++ b/tests/test_litellm/responses/test_metadata_codex_callback.py @@ -0,0 +1,176 @@ +""" +Test that metadata is passed to custom callbacks during chat completion calls to codex models. + +Fixes issue: Metadata is no longer passed to custom callback during chat completion +calls to codex models (#21204) + +Codex models (gpt-5.1-codex, gpt-5.2-codex) use mode=responses and route through +responses_api_bridge. The bridge converts metadata to litellm_metadata. This test +verifies metadata is preserved for custom callbacks via kwargs['litellm_params']['metadata']. +""" + +import asyncio +import os +import sys +from typing import Optional +from unittest.mock import AsyncMock, patch + +sys.path.insert(0, os.path.abspath("../../..")) + +import pytest + +import litellm +from litellm.integrations.custom_logger import CustomLogger + + +def _make_mock_http_response(response_dict: dict): + """Create a mock HTTP response that returns response_dict from .json().""" + + class MockResponse: + def __init__(self, json_data, status_code=200): + self._json_data = json_data + self.status_code = status_code + self.text = str(json_data) + self.headers = {} + + def json(self): + return self._json_data + + return MockResponse(response_dict, 200) + + +class MetadataCaptureCallback(CustomLogger): + """Custom callback that captures kwargs passed to async_log_success_event.""" + + def __init__(self): + self.captured_kwargs: Optional[dict] = None + + async def async_log_success_event( + self, kwargs, response_obj, start_time, end_time + ): + self.captured_kwargs = kwargs + + +@pytest.mark.asyncio +async def test_metadata_passed_to_custom_callback_codex_models(): + """ + Test that metadata passed to completion() is available in custom callback + when using codex models (responses API bridge path). + + Codex models have mode=responses and route through responses_api_bridge, + which passes litellm_metadata. The fix ensures this is preserved as + litellm_params.metadata for callback compatibility. + """ + from litellm.types.llms.openai import ResponsesAPIResponse + + mock_response = ResponsesAPIResponse.model_construct( + id="resp-test", + created_at=0, + output=[ + { + "type": "message", + "id": "msg-1", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "Hello!"}], + } + ], + object="response", + model="gpt-5.1-codex", + status="completed", + usage={ + "input_tokens": 5, + "output_tokens": 10, + "total_tokens": 15, + }, + ) + + test_metadata = {"foo": "bar", "trace_id": "test-123"} + callback = MetadataCaptureCallback() + original_callbacks = litellm.callbacks.copy() if litellm.callbacks else [] + litellm.callbacks = [callback] + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _make_mock_http_response( + mock_response.model_dump() + ) + # gpt-5.1-codex has mode=responses - routes through responses bridge + await litellm.acompletion( + model="gpt-5.1-codex", + messages=[{"role": "user", "content": "Hello"}], + metadata=test_metadata, + ) + + await asyncio.sleep(1) + + assert callback.captured_kwargs is not None, "Callback should have been invoked" + + litellm_params = callback.captured_kwargs.get("litellm_params", {}) + metadata = litellm_params.get("metadata") or {} + + assert "foo" in metadata, "metadata['foo'] should be accessible in callback" + assert metadata["foo"] == "bar" + assert metadata.get("trace_id") == "test-123" + + +@pytest.mark.asyncio +async def test_metadata_passed_via_litellm_metadata_responses_api(): + """ + Test that when calling responses() directly with litellm_metadata, + metadata is preserved for custom callbacks. + + Uses HTTP mock since mock_response returns early before update_environment_variables. + """ + from litellm.types.llms.openai import ResponsesAPIResponse + + mock_response = ResponsesAPIResponse.model_construct( + id="resp-test-2", + created_at=0, + output=[ + { + "type": "message", + "id": "msg-2", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "Hi there!"}], + } + ], + object="response", + model="gpt-4o", + status="completed", + usage={ + "input_tokens": 2, + "output_tokens": 3, + "total_tokens": 5, + }, + ) + + test_metadata = {"request_id": "req-456"} + callback = MetadataCaptureCallback() + litellm.callbacks = [callback] + + with patch( + "litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post", + new_callable=AsyncMock, + ) as mock_post: + mock_post.return_value = _make_mock_http_response( + mock_response.model_dump() + ) + await litellm.aresponses( + model="gpt-4o", + input="hi", + litellm_metadata=test_metadata, + ) + + await asyncio.sleep(1) + + assert callback.captured_kwargs is not None + + litellm_params = callback.captured_kwargs.get("litellm_params", {}) + metadata = litellm_params.get("metadata") or {} + + assert "request_id" in metadata + assert metadata["request_id"] == "req-456" 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 new file mode 100644 index 00000000000..e500ad3ca6e --- /dev/null +++ b/tests/test_litellm/router_utils/pre_call_checks/test_deployment_affinity_check.py @@ -0,0 +1,659 @@ +import asyncio +import os +import sys +from unittest.mock import AsyncMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../..")) + +import json + +import litellm +from litellm.caching.dual_cache import DualCache +from litellm.router_utils.pre_call_checks.deployment_affinity_check import ( + DeploymentAffinityCheck, +) + + +class MockResponse: + def __init__(self, json_data, status_code): + self._json_data = json_data + self.status_code = status_code + self.text = json.dumps(json_data) + self.headers = {} + + def json(self): + return self._json_data + + +@pytest.mark.asyncio +async def test_async_user_key_affinity_routes_to_same_deployment(): + """ + When deployment_affinity is enabled, subsequent requests from the same user key + should route to the same deployment (even if the routing strategy would pick another). + """ + mock_response_data = { + "id": "resp_mock-resp-123", + "object": "response", + "created_at": 1741476542, + "status": "completed", + "model": "azure/computer-use-preview", + "output": [ + { + "type": "message", + "id": "msg_123", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "Hello there!", "annotations": []}], + } + ], + "parallel_tool_calls": True, + "usage": { + "input_tokens": 5, + "output_tokens": 10, + "total_tokens": 15, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + "text": {"format": {"type": "text"}}, + "error": None, + "incomplete_details": None, + "instructions": None, + "metadata": {}, + "temperature": 1.0, + "tool_choice": "auto", + "tools": [], + "top_p": 1.0, + "max_output_tokens": None, + "previous_response_id": None, + "reasoning": {"effort": None, "summary": None}, + "truncation": "disabled", + "user": None, + } + + router = litellm.Router( + model_list=[ + { + "model_name": "azure-computer-use-preview", + "litellm_params": { + "model": "azure/computer-use-preview-1", + "api_key": "mock-api-key-1", + "api_version": "mock-api-version", + "api_base": "https://mock-endpoint-1.openai.azure.com", + }, + # Required for stable affinity scoping across multiple Azure deployments + "model_info": {"base_model": "computer-use-preview"}, + }, + { + "model_name": "azure-computer-use-preview", + "litellm_params": { + "model": "azure/computer-use-preview-2", + "api_key": "mock-api-key-2", + "api_version": "mock-api-version-2", + "api_base": "https://mock-endpoint-2.openai.azure.com", + }, + "model_info": {"base_model": "computer-use-preview"}, + }, + ], + optional_pre_call_checks=["deployment_affinity"], + ) + + model_group = "azure-computer-use-preview" + user_api_key_hash = "test-user-key-1" + + # Deterministic routing: first selection uses seq[0], second selection attempts seq[1] + # unless the list has been filtered to length=1 by deployment affinity. + 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) + + first_response = await router.aresponses( + model=model_group, + input="Hello, how are you?", + truncation="auto", + litellm_metadata={"user_api_key_hash": user_api_key_hash}, + ) + first_model_id = first_response._hidden_params["model_id"] + + # If affinity works, second request should be pinned to the same deployment + # even though deterministic_choice would pick the other deployment when len(seq)>1. + second_response = await router.aresponses( + model=model_group, + input="Follow-up question", + truncation="auto", + litellm_metadata={"user_api_key_hash": user_api_key_hash}, + ) + assert second_response._hidden_params["model_id"] == first_model_id + + +@pytest.mark.asyncio +async def test_async_user_key_affinity_routes_with_model_group_alias(): + """ + When Router model_group_alias is used, the requested model group (alias) can differ + from the internally-routed model group. Deployment affinity should still stick. + """ + mock_response_data = { + "id": "resp_mock-resp-alias", + "object": "response", + "created_at": 1741476542, + "status": "completed", + "model": "azure/computer-use-preview", + "output": [ + { + "type": "message", + "id": "msg_alias", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "Alias 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, + } + + canonical_model_group = "azure-computer-use-preview" + alias_model_group = "azure-computer-use-preview-alias" + user_api_key_hash = "test-user-key-alias" + + router = litellm.Router( + model_list=[ + { + "model_name": canonical_model_group, + "litellm_params": { + "model": "azure/computer-use-preview-1", + "api_key": "mock-api-key-1", + "api_version": "mock-api-version", + "api_base": "https://mock-endpoint-1.openai.azure.com", + }, + "model_info": {"base_model": "computer-use-preview"}, + }, + { + "model_name": canonical_model_group, + "litellm_params": { + "model": "azure/computer-use-preview-2", + "api_key": "mock-api-key-2", + "api_version": "mock-api-version-2", + "api_base": "https://mock-endpoint-2.openai.azure.com", + }, + "model_info": {"base_model": "computer-use-preview"}, + }, + ], + model_group_alias={alias_model_group: canonical_model_group}, + optional_pre_call_checks=["deployment_affinity"], + ) + + 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) + + first_response = await router.aresponses( + model=alias_model_group, + input="Hello", + truncation="auto", + litellm_metadata={"user_api_key_hash": user_api_key_hash}, + ) + first_model_id = first_response._hidden_params["model_id"] + + second_response = await router.aresponses( + model=alias_model_group, + input="Follow-up", + truncation="auto", + litellm_metadata={"user_api_key_hash": user_api_key_hash}, + ) + assert second_response._hidden_params["model_id"] == first_model_id + + +@pytest.mark.asyncio +async def test_async_previous_response_id_priority_over_user_key_affinity(): + """ + If both deployment_affinity and responses_api_deployment_check are enabled, + `previous_response_id` routing should take priority over user-key affinity. + """ + mock_response_data = { + "id": "resp_mock-resp-456", + "object": "response", + "created_at": 1741476542, + "status": "completed", + "model": "azure/computer-use-preview", + "output": [ + { + "type": "message", + "id": "msg_123", + "status": "completed", + "role": "assistant", + "content": [ + { + "type": "output_text", + "text": "I'm doing well, thank you for asking!", + "annotations": [], + } + ], + } + ], + "parallel_tool_calls": True, + "usage": { + "input_tokens": 10, + "output_tokens": 20, + "total_tokens": 30, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + "text": {"format": {"type": "text"}}, + "error": None, + "incomplete_details": None, + "instructions": None, + "metadata": {}, + "temperature": 1.0, + "tool_choice": "auto", + "tools": [], + "top_p": 1.0, + "max_output_tokens": None, + "previous_response_id": None, + "reasoning": {"effort": None, "summary": None}, + "truncation": "disabled", + "user": None, + } + + router = litellm.Router( + model_list=[ + { + "model_name": "azure-computer-use-preview", + "litellm_params": { + "model": "azure/computer-use-preview-1", + "api_key": "mock-api-key-1", + "api_version": "mock-api-version", + "api_base": "https://mock-endpoint-1.openai.azure.com", + }, + "model_info": {"base_model": "computer-use-preview"}, + }, + { + "model_name": "azure-computer-use-preview", + "litellm_params": { + "model": "azure/computer-use-preview-2", + "api_key": "mock-api-key-2", + "api_version": "mock-api-version-2", + "api_base": "https://mock-endpoint-2.openai.azure.com", + }, + "model_info": {"base_model": "computer-use-preview"}, + }, + ], + optional_pre_call_checks=[ + "deployment_affinity", + "responses_api_deployment_check", + ], + ) + + model_group = "azure-computer-use-preview" + user_api_key_hash = "test-user-key-1" + + 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=lambda seq: seq[0], + ): + mock_post.return_value = MockResponse(mock_response_data, 200) + + first_response = await router.aresponses( + model=model_group, + input="Hello, how are you?", + truncation="auto", + litellm_metadata={"user_api_key_hash": user_api_key_hash}, + ) + first_model_id = first_response._hidden_params["model_id"] + first_response_id = first_response.id + + all_model_ids = router.get_model_ids(model_name=model_group) + other_model_id = next(mid for mid in all_model_ids if mid != first_model_id) + + # Force user-key affinity to point to the OTHER deployment + affinity_cache_key = DeploymentAffinityCheck.get_affinity_cache_key( + model_group=model_group, + user_key=user_api_key_hash, + ) + await router.cache.async_set_cache(affinity_cache_key, {"model_id": other_model_id}, ttl=3600) + + # Even though user-key affinity points elsewhere, previous_response_id should pin + # to the deployment that created the original response. + follow_up = await router.aresponses( + model=model_group, + input="Follow-up question", + truncation="auto", + previous_response_id=first_response_id, + litellm_metadata={"user_api_key_hash": user_api_key_hash}, + ) + assert follow_up._hidden_params["model_id"] == first_model_id + + +@pytest.mark.asyncio +async def test_async_user_parameter_does_not_trigger_deployment_affinity(): + """ + The OpenAI `user` parameter identifies the *end-user* (not the API key), and should + not be used as an affinity key. + """ + mock_response_data = { + "id": "resp_mock-resp-sdk", + "object": "response", + "created_at": 1741476542, + "status": "completed", + "model": "azure/computer-use-preview", + "output": [ + { + "type": "message", + "id": "msg_sdk", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "SDK 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": "azure-sdk-test", + "litellm_params": { + "model": "azure/sdk-1", + "api_key": "mock", + "api_base": "https://mock1.openai.azure.com", + }, + "model_info": {"base_model": "sdk-test"}, + }, + { + "model_name": "azure-sdk-test", + "litellm_params": { + "model": "azure/sdk-2", + "api_key": "mock", + "api_base": "https://mock2.openai.azure.com", + }, + "model_info": {"base_model": "sdk-test"}, + }, + ], + optional_pre_call_checks=["deployment_affinity"], + ) + + model_group = "azure-sdk-test" + user_id = "sdk-user-123" + + 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) + + # First call with 'user' parameter (end-user id) + first_response = await router.aresponses( + model=model_group, + input="Hi", + user=user_id, + ) + first_model_id = first_response._hidden_params["model_id"] + + # Second call with same 'user' parameter should NOT be pinned by affinity + second_response = await router.aresponses( + model=model_group, + input="Follow-up", + user=user_id, + ) + assert second_response._hidden_params["model_id"] != first_model_id + + +@pytest.mark.asyncio +async def test_async_pre_call_hook_uses_model_map_key_scope(): + """ + Deployment affinity caching uses (user_api_key_hash, model_map_key) -> model_id. + """ + + cache = AsyncMock() + cache.async_set_cache = AsyncMock() + + callback = DeploymentAffinityCheck( + cache=cache, + ttl_seconds=123, + enable_user_key_affinity=True, + enable_responses_api_affinity=False, + ) + + kwargs = { + "model_info": {"id": "model-id-123"}, + "litellm_metadata": { + "user_api_key_hash": "user-key-abc", + "deployment_model_name": "claude-sonnet-4-5@20250929", + }, + } + + await callback.async_pre_call_deployment_hook(kwargs=kwargs, call_type=None) + + expected_cache_key = DeploymentAffinityCheck.get_affinity_cache_key( + model_group="claude-sonnet-4-5@20250929", + user_key="user-key-abc", + ) + cache.async_set_cache.assert_called_once_with( + expected_cache_key, + {"model_id": "model-id-123"}, + ttl=123, + ) + + +@pytest.mark.asyncio +async def test_async_filter_deployments_uses_stable_model_map_key_for_affinity_scope(): + """ + When a stable model-map key can be derived from the deployment set, affinity should + be scoped to that key (this helps stickiness across aliases). + + This is intentionally tested at the callback level (not via Router), to validate the + cache key selection logic deterministically. + """ + + user_key = "user-key-abc" + stable_model_map_key = "claude-sonnet-4-5@20250929" + + cache = AsyncMock() + cache.async_get_cache = AsyncMock() + + callback = DeploymentAffinityCheck( + cache=cache, + ttl_seconds=123, + enable_user_key_affinity=True, + enable_responses_api_affinity=False, + ) + + healthy_deployments = [ + { + "model_name": stable_model_map_key, + "litellm_params": {"model": f"vertex_ai/{stable_model_map_key}"}, + "model_info": {"id": "deployment-1"}, + }, + { + "model_name": stable_model_map_key, + "litellm_params": {"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"}, + "model_info": {"id": "deployment-2"}, + }, + ] + + expected_cache_key = DeploymentAffinityCheck.get_affinity_cache_key( + model_group=stable_model_map_key, + user_key=user_key, + ) + + async def get_cache_side_effect(*, key: str): + if key == expected_cache_key: + return {"model_id": "deployment-2"} + return None + + cache.async_get_cache.side_effect = get_cache_side_effect + + filtered = await callback.async_filter_deployments( + model="some-router-model-group", + healthy_deployments=healthy_deployments, + messages=None, + request_kwargs={"metadata": {"user_api_key_hash": user_key, "model_group": "alias-group"}}, + parent_otel_span=None, + ) + + assert len(filtered) == 1 + assert filtered[0]["model_info"]["id"] == "deployment-2" + + +@pytest.mark.asyncio +async def test_async_filter_deployments_falls_back_when_cached_deployment_is_unhealthy(): + """ + If affinity cache points to a deployment that's no longer healthy, callback should + return all healthy deployments so router can pick an available one. + """ + + user_key = "user-key-unhealthy" + stable_model_map_key = "claude-sonnet-4-5@20250929" + + cache = AsyncMock() + cache.async_get_cache = AsyncMock(return_value={"model_id": "stale-deployment"}) + + callback = DeploymentAffinityCheck( + cache=cache, + ttl_seconds=123, + enable_user_key_affinity=True, + enable_responses_api_affinity=False, + ) + + healthy_deployments = [ + { + "model_name": stable_model_map_key, + "litellm_params": {"model": f"vertex_ai/{stable_model_map_key}"}, + "model_info": {"id": "deployment-1"}, + }, + { + "model_name": stable_model_map_key, + "litellm_params": {"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"}, + "model_info": {"id": "deployment-2"}, + }, + ] + + filtered = await callback.async_filter_deployments( + model="some-router-model-group", + healthy_deployments=healthy_deployments, + messages=None, + request_kwargs={"metadata": {"user_api_key_hash": user_key}}, + parent_otel_span=None, + ) + + assert filtered == healthy_deployments + + +@pytest.mark.asyncio +async def test_async_user_key_affinity_ttl_expiry_allows_reroute(): + """ + After affinity TTL expires, cached pinning should no longer filter deployments. + """ + + callback = DeploymentAffinityCheck( + cache=DualCache(), + ttl_seconds=1, + enable_user_key_affinity=True, + enable_responses_api_affinity=False, + ) + + user_key = "ttl-user-key" + stable_model_map_key = "claude-sonnet-4-5@20250929" + healthy_deployments = [ + { + "model_name": stable_model_map_key, + "litellm_params": {"model": f"vertex_ai/{stable_model_map_key}"}, + "model_info": {"id": "deployment-1"}, + }, + { + "model_name": stable_model_map_key, + "litellm_params": {"model": f"bedrock/global.anthropic.{stable_model_map_key}-v1:0"}, + "model_info": {"id": "deployment-2"}, + }, + ] + + 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, + ) + + pinned = await callback.async_filter_deployments( + model="some-router-model-group", + healthy_deployments=healthy_deployments, + messages=None, + request_kwargs={"metadata": {"user_api_key_hash": user_key}}, + parent_otel_span=None, + ) + assert len(pinned) == 1 + assert pinned[0]["model_info"]["id"] == "deployment-1" + + await asyncio.sleep(1.2) + + after_ttl_expiry = await callback.async_filter_deployments( + model="some-router-model-group", + healthy_deployments=healthy_deployments, + messages=None, + request_kwargs={"metadata": {"user_api_key_hash": user_key}}, + parent_otel_span=None, + ) + assert after_ttl_expiry == healthy_deployments + + +def test_cache_key_does_not_double_hash_user_api_key_hash(): + """ + Proxy typically provides `metadata.user_api_key_hash` as a SHA-256 hex string. + The affinity cache key should not hash it again. + """ + + user_api_key_hash = "b95b015b66dd02a1c14e1e0a8729211f8ee53ec962658764f4cf58546c2c68e1" + key = DeploymentAffinityCheck.get_affinity_cache_key( + model_group="any-model-group", + user_key=user_api_key_hash, + ) + assert key.endswith(user_api_key_hash) diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index 74f5cf9bdd7..c2c20485b5e 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -1600,6 +1600,56 @@ def test_completion_cost_service_tier_priority(): ), "Costs from params and usage should be similar (both flex)" +def test_completion_cost_service_tier_for_bedrock(): + """Test that Bedrock cost calculation applies service_tier-specific pricing.""" + from litellm import completion_cost + + os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" + litellm.model_cost = litellm.get_model_cost_map(url="") + + model = "bedrock/us-east-1/test-bedrock-service-tier-cost-model" + litellm.register_model( + model_cost={ + model: { + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + "input_cost_per_token_priority": 0.01, + "output_cost_per_token_priority": 0.02, + "input_cost_per_token_flex": 0.0005, + "output_cost_per_token_flex": 0.001, + "litellm_provider": "bedrock", + "max_tokens": 8192, + } + } + ) + + usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150) + response = ModelResponse(usage=usage, model=model) + + default_cost = completion_cost( + completion_response=response, + model=model, + custom_llm_provider="bedrock", + ) + + priority_cost = completion_cost( + completion_response=response, + model=model, + custom_llm_provider="bedrock", + optional_params={"service_tier": "priority"}, + ) + + response_with_flex_tier = ModelResponse(usage=usage, model=model) + setattr(response_with_flex_tier, "service_tier", "flex") + flex_cost = completion_cost( + completion_response=response_with_flex_tier, + model=model, + custom_llm_provider="bedrock", + ) + + assert priority_cost > default_cost > flex_cost > 0 + + def test_gemini_cache_tokens_details_no_negative_values(): """ Test for Issue #18750: Negative text_tokens with Gemini caching diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 7374a605798..7b29a4d90aa 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -580,6 +580,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "annotation_cost_per_page": {"type": "number"}, "ocr_cost_per_page": {"type": "number"}, "code_interpreter_cost_per_session": {"type": "number"}, + "inference_geo": {"type": "string"}, "litellm_provider": {"type": "string"}, "max_audio_length_hours": {"type": "number"}, "max_audio_per_prompt": {"type": "number"}, diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx index d08d0f3af13..46dca6039a7 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.test.tsx @@ -72,6 +72,37 @@ describe("CreateKey", () => { expect(screen.getByRole("button", { name: /create new key/i })).toBeInTheDocument(); }); + it("should display 'AI APIs' label for the llm_api key type option", async () => { + renderWithProviders(); + + act(() => { + fireEvent.click(screen.getByRole("button", { name: /create new key/i })); + }); + + await waitFor(() => { + expect(screen.getByText("Key Type")).toBeInTheDocument(); + }); + + // Open the Key Type dropdown + const keyTypeSection = screen.getByText("Key Type").closest(".ant-form-item")!; + const selectElement = keyTypeSection.querySelector(".ant-select-selector")!; + act(() => { + fireEvent.mouseDown(selectElement); + }); + + await waitFor(() => { + // Verify "AI APIs" appears as an option + const options = document.querySelectorAll(".ant-select-item-option"); + const optionTexts = Array.from(options).map((el) => el.textContent); + const hasAIAPIs = optionTexts.some((text) => text?.includes("AI APIs")); + expect(hasAIAPIs).toBe(true); + + // Verify old "LLM API" label does NOT appear + const hasLLMAPI = optionTexts.some((text) => text?.includes("LLM API")); + expect(hasLLMAPI).toBe(false); + }); + }); + it("should include access_group_ids in keyCreateCall payload when access groups are selected", async () => { renderWithProviders(); diff --git a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx index 7937ebef666..9a99870c263 100644 --- a/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx +++ b/ui/litellm-dashboard/src/components/organisms/create_key_button.tsx @@ -728,15 +728,15 @@ const CreateKey: React.FC = ({ team, teams, data, addKey }) => { Default - Can call LLM API + Management routes + Can call AI APIs + Management routes - + - LLM API + AI APIs - Can call only LLM API routes (chat/completions, embeddings, etc.) + Can call only AI API routes (chat/completions, embeddings, etc.) diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx index bbe3c206cee..49a98b699e8 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.test.tsx @@ -494,7 +494,50 @@ describe("KeyEditView", () => { }); }); - it("should disable cancel button during submission", async () => { + it("should display 'AI APIs' label for the llm_api key type option", async () => { + const keyDataWithLlmApiRoutes = { + ...MOCK_KEY_DATA, + allowed_routes: ["llm_api_routes"], + }; + + renderWithProviders( + {}} + onSubmit={async () => {}} + accessToken={""} + userID={""} + userRole={""} + premiumUser={false} + />, + ); + + await waitFor(() => { + expect(screen.getByText("Key Type")).toBeInTheDocument(); + }); + + // The selected key type label should show "AI APIs" (not "LLM API") + const keyTypeSection = screen.getByText("Key Type").closest(".ant-form-item")!; + expect(keyTypeSection).toBeInTheDocument(); + + // Open the dropdown to see all options + const selectElement = keyTypeSection.querySelector(".ant-select-selector")!; + await userEvent.click(selectElement); + + await waitFor(() => { + // Verify "AI APIs" appears as an option label + const options = document.querySelectorAll(".ant-select-item-option"); + const optionTexts = Array.from(options).map((el) => el.textContent); + const hasAIAPIs = optionTexts.some((text) => text?.includes("AI APIs")); + expect(hasAIAPIs).toBe(true); + + // Verify old "LLM API" label does NOT appear + const hasLLMAPI = optionTexts.some((text) => text?.includes("LLM API")); + expect(hasLLMAPI).toBe(false); + }); + }); + + it("should display cancel button during submission", async () => { let resolveSubmit: (() => void) | undefined; const submitPromise = new Promise((resolve) => { resolveSubmit = resolve; diff --git a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx index 1ba3cc7f7b4..71e00e16542 100644 --- a/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx +++ b/ui/litellm-dashboard/src/components/templates/key_edit_view.tsx @@ -345,15 +345,15 @@ export function KeyEditView({ Default - Can call LLM API + Management routes + Can call AI APIs + Management routes - + - LLM API + AI APIs - Can call only LLM API routes (chat/completions, embeddings, etc.) + Can call only AI API routes (chat/completions, embeddings, etc.)