diff --git a/.semgrep/rules/python/reliability/unbounded-memory.yml b/.semgrep/rules/python/reliability/unbounded-memory.yml new file mode 100644 index 00000000000..f13c38471fb --- /dev/null +++ b/.semgrep/rules/python/reliability/unbounded-memory.yml @@ -0,0 +1,17 @@ +# Unbounded memory growth – data structures without a clear max limit +# Can lead to OOM under load. + +rules: + - id: unbounded-asyncio-queue + message: asyncio.Queue() with no maxsize can grow unbounded. Use asyncio.Queue(maxsize=N) for integrations (e.g. log queues). + severity: ERROR + languages: [python] + pattern-either: + - pattern: asyncio.Queue() + - pattern: asyncio.Queue(maxsize=0) + metadata: + category: reliability + cwe: "CWE-400: Uncontrolled Resource Consumption" + tags: [python, reliability] + confidence: HIGH + source: https://docs.python.org/3/library/asyncio-queue.html diff --git a/docs/my-website/docs/benchmarks.md b/docs/my-website/docs/benchmarks.md index a1489081b4c..1f818cef498 100644 --- a/docs/my-website/docs/benchmarks.md +++ b/docs/my-website/docs/benchmarks.md @@ -5,6 +5,13 @@ import Image from '@theme/IdealImage'; Benchmarks for LiteLLM Gateway (Proxy Server) tested against a fake OpenAI endpoint. +## Setting Up a Fake OpenAI Endpoint + +For load testing and benchmarking, you can use a fake OpenAI proxy server. LiteLLM provides: + +1. **Hosted endpoint**: Use our free hosted fake endpoint at `https://exampleopenaiendpoint-production.up.railway.app/` +2. **Self-hosted**: Set up your own fake OpenAI proxy server using [github.com/BerriAI/example_openai_endpoint](https://github.com/BerriAI/example_openai_endpoint) + Use this config for testing: ```yaml @@ -12,7 +19,7 @@ model_list: - model_name: "fake-openai-endpoint" litellm_params: model: openai/any - api_base: https://your-fake-openai-endpoint.com/chat/completions + api_base: https://exampleopenaiendpoint-production.up.railway.app/ # or your self-hosted endpoint api_key: "test" ``` diff --git a/docs/my-website/docs/load_test.md b/docs/my-website/docs/load_test.md index 4641a70366c..071b097904b 100644 --- a/docs/my-website/docs/load_test.md +++ b/docs/my-website/docs/load_test.md @@ -4,8 +4,9 @@ import Image from '@theme/IdealImage'; ## Locust Load Test LiteLLM Proxy -1. Add `fake-openai-endpoint` to your proxy config.yaml and start your litellm proxy -litellm provides a free hosted `fake-openai-endpoint` you can load test against +1. Add `fake-openai-endpoint` to your proxy config.yaml and start your litellm proxy. + +LiteLLM provides a free hosted `fake-openai-endpoint` you can load test against. You can also self-host your own fake OpenAI proxy server using [github.com/BerriAI/example_openai_endpoint](https://github.com/BerriAI/example_openai_endpoint). ```yaml model_list: diff --git a/docs/my-website/docs/load_test_advanced.md b/docs/my-website/docs/load_test_advanced.md index 3171bc33594..d35b5f74784 100644 --- a/docs/my-website/docs/load_test_advanced.md +++ b/docs/my-website/docs/load_test_advanced.md @@ -29,12 +29,16 @@ Tutorial on how to get to 1K+ RPS with LiteLLM Proxy on locust **Note:** we're currently migrating to aiohttp which has 10x higher throughput. We recommend using the `openai/` provider for load testing. +:::tip Setting Up a Fake OpenAI Endpoint +You can use our hosted fake endpoint or self-host your own using [github.com/BerriAI/example_openai_endpoint](https://github.com/BerriAI/example_openai_endpoint). +::: + ```yaml model_list: - model_name: "fake-openai-endpoint" litellm_params: model: openai/any - api_base: https://your-fake-openai-endpoint.com/chat/completions + api_base: https://exampleopenaiendpoint-production.up.railway.app/ # or your self-hosted endpoint api_key: "test" ``` diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index a6d4e3102a6..38ad9bdd0ee 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -812,6 +812,7 @@ router_settings: | MAX_RETRY_DELAY | Maximum delay in seconds for retrying requests. Default is 8.0 | MAX_LANGFUSE_INITIALIZED_CLIENTS | Maximum number of Langfuse clients to initialize on proxy. Default is 50. This is set since langfuse initializes 1 thread everytime a client is initialized. We've had an incident in the past where we reached 100% cpu utilization because Langfuse was initialized several times. | MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH | Maximum header length for MCP semantic filter tools. Default is 150 +| MAX_POLICY_ESTIMATE_IMPACT_ROWS | Maximum number of rows returned when estimating the impact of a policy. Default is 1000 | MIN_NON_ZERO_TEMPERATURE | Minimum non-zero temperature value. Default is 0.0001 | MINIMUM_PROMPT_CACHE_TOKEN_COUNT | Minimum token count for caching a prompt. Default is 1024 | MISTRAL_API_BASE | Base URL for Mistral API. Default is https://api.mistral.ai @@ -828,6 +829,8 @@ router_settings: | MICROSOFT_USER_ID_ATTRIBUTE | Field name for user ID in Microsoft SSO response. Default is `id` | MICROSOFT_USER_LAST_NAME_ATTRIBUTE | Field name for user last name in Microsoft SSO response. Default is `surname` | MICROSOFT_USERINFO_ENDPOINT | Custom userinfo endpoint URL for Microsoft SSO (overrides default Microsoft Graph userinfo endpoint) +| MODEL_COST_MAP_MAX_SHRINK_RATIO | Maximum allowed shrinkage ratio when validating a fetched model cost map against the local backup. Rejects the fetched map if it is smaller than this fraction of the backup. Default is 0.5 +| MODEL_COST_MAP_MIN_MODEL_COUNT | Minimum number of models a fetched cost map must contain to be considered valid. Default is 50 | NO_DOCS | Flag to disable Swagger UI documentation | NO_REDOC | Flag to disable Redoc documentation | NO_PROXY | List of addresses to bypass proxy diff --git a/docs/my-website/docs/tutorials/claude_code_beta_headers.md b/docs/my-website/docs/tutorials/claude_code_beta_headers.md index 9c1645e0277..4cc6f7ff92b 100644 --- a/docs/my-website/docs/tutorials/claude_code_beta_headers.md +++ b/docs/my-website/docs/tutorials/claude_code_beta_headers.md @@ -1,8 +1,8 @@ import Image from '@theme/IdealImage'; -# Claude Code - Fixing Invalid Beta Header Errors +# Claude Code - Managing Anthropic Beta Headers -When using Claude Code with LiteLLM and non-Anthropic providers (Bedrock, Azure AI, Vertex AI), you may encounter "invalid beta header" errors. This guide explains how to fix these errors locally or contribute a fix to LiteLLM. +When using Claude Code with LiteLLM and non-Anthropic providers (Bedrock, Azure AI, Vertex AI), you need to ensure that only supported beta headers are sent to each provider. This guide explains how to add support for new beta headers or fix invalid beta header errors. ## What Are Beta Headers? @@ -12,7 +12,7 @@ Anthropic uses beta headers to enable experimental features in Claude. When you anthropic-beta: prompt-caching-scope-2026-01-05,advanced-tool-use-2025-11-20 ``` -However, not all providers support all Anthropic beta features. When an unsupported beta header is sent to a provider, you'll see an error. +However, not all providers support all Anthropic beta features. LiteLLM uses `anthropic_beta_headers_config.json` to manage which beta headers are supported by each provider. ## Common Error Message @@ -22,17 +22,22 @@ Error: The model returned the following errors: invalid beta flag ## How LiteLLM Handles Beta Headers -LiteLLM automatically filters out unsupported beta headers using a configuration file: +LiteLLM uses a strict validation approach with a configuration file: ``` litellm/litellm/anthropic_beta_headers_config.json ``` -This JSON file lists which beta headers are **unsupported** for each provider. Headers not in the unsupported list are passed through to the provider. +This JSON file contains a **mapping** of beta headers for each provider: +- **Keys**: Input beta header names (from Anthropic) +- **Values**: Provider-specific header names (or `null` if unsupported) +- **Validation**: Only headers present in the mapping with non-null values are forwarded -## Quick Fix: Update Config Locally +This enforces stricter validation than just filtering unsupported headers - headers must be explicitly defined to be allowed. -If you encounter an invalid beta header error, you can fix it immediately by updating the config file locally. +## Adding Support for a New Beta Header + +When Anthropic releases a new beta feature, you need to add it to the configuration file for each provider. ### Step 1: Locate the Config File @@ -46,43 +51,47 @@ cd $(python -c "import litellm; import os; print(os.path.dirname(litellm.__file_ # litellm/anthropic_beta_headers_config.json ``` -### Step 2: Add the Unsupported Header +### Step 2: Add the New Beta Header -Open `anthropic_beta_headers_config.json` and add the problematic header to the appropriate provider's list: +Open `anthropic_beta_headers_config.json` and add the new header to each provider's mapping: ```json title="anthropic_beta_headers_config.json" { - "description": "Unsupported Anthropic beta headers for each provider. Headers listed here will be dropped. Headers not listed are passed through as-is.", - "anthropic": [], - "azure_ai": [], - "bedrock_converse": [ - "prompt-caching-scope-2026-01-05", - "bash_20250124", - "bash_20241022", - "text_editor_20250124", - "text_editor_20241022", - "compact-2026-01-12", - "advanced-tool-use-2025-11-20", - "web-fetch-2025-09-10", - "code-execution-2025-08-25", - "skills-2025-10-02", - "files-api-2025-04-14" - ], - "bedrock": [ - "advanced-tool-use-2025-11-20", - "prompt-caching-scope-2026-01-05", - "structured-outputs-2025-11-13", - "web-fetch-2025-09-10", - "code-execution-2025-08-25", - "skills-2025-10-02", - "files-api-2025-04-14" - ], - "vertex_ai": [ - "prompt-caching-scope-2026-01-05" - ] + "description": "Mapping of Anthropic beta headers for each provider. Keys are input header names, values are provider-specific header names (or null if unsupported). Only headers present in mapping keys with non-null values can be forwarded.", + "anthropic": { + "advanced-tool-use-2025-11-20": "advanced-tool-use-2025-11-20", + "new-feature-2026-03-01": "new-feature-2026-03-01", + ... + }, + "azure_ai": { + "advanced-tool-use-2025-11-20": "advanced-tool-use-2025-11-20", + "new-feature-2026-03-01": "new-feature-2026-03-01", + ... + }, + "bedrock_converse": { + "advanced-tool-use-2025-11-20": "tool-search-tool-2025-10-19", + "new-feature-2026-03-01": null, + ... + }, + "bedrock": { + "advanced-tool-use-2025-11-20": "tool-search-tool-2025-10-19", + "new-feature-2026-03-01": null, + ... + }, + "vertex_ai": { + "advanced-tool-use-2025-11-20": "tool-search-tool-2025-10-19", + "new-feature-2026-03-01": null, + ... + } } ``` +**Key Points:** +- **Supported headers**: Set the value to the provider-specific header name (often the same as the key) +- **Unsupported headers**: Set the value to `null` +- **Header transformations**: Some providers use different header names (e.g., Bedrock maps `advanced-tool-use-2025-11-20` to `tool-search-tool-2025-10-19`) +- **Alphabetical order**: Keep headers sorted alphabetically for maintainability + ### Step 3: Restart Your Application After updating the config file, restart your LiteLLM proxy or application: @@ -97,9 +106,64 @@ litellm --config config.yaml The updated configuration will be loaded automatically. +## Fixing Invalid Beta Header Errors + +If you encounter an "invalid beta flag" error, it means a beta header is being sent that the provider doesn't support. + +### Step 1: Identify the Problematic Header + +Check your logs to see which header is causing the issue: + +```bash +Error: The model returned the following errors: invalid beta flag: new-feature-2026-03-01 +``` + +### Step 2: Update the Config + +Set the header value to `null` for that provider: + +```json title="anthropic_beta_headers_config.json" +{ + "bedrock_converse": { + "new-feature-2026-03-01": null + } +} +``` + +### Step 3: Restart and Test + +Restart your application and verify the header is now filtered out. + ## Contributing a Fix to LiteLLM -Help the community by contributing your fix! If your local changes work, please raise a PR with the addition of the header and we will merge it. +Help the community by contributing your fix! + +### What to Include in Your PR + +1. **Update the config file**: Add the new beta header to `litellm/anthropic_beta_headers_config.json` +2. **Test your changes**: Verify the header is correctly filtered/mapped for each provider +3. **Documentation**: Include provider documentation links showing which headers are supported + +### Example PR Description + +```markdown +## Add support for new-feature-2026-03-01 beta header + +### Changes +- Added `new-feature-2026-03-01` to anthropic_beta_headers_config.json +- Set to `null` for bedrock_converse (unsupported) +- Set to header name for anthropic, azure_ai (supported) + +### Testing +Tested with: +- ✅ Anthropic: Header passed through correctly +- ✅ Azure AI: Header passed through correctly +- ✅ Bedrock Converse: Header filtered out (returns error without fix) + +### References +- Anthropic docs: [link] +- AWS Bedrock docs: [link] +``` ## How Beta Header Filtering Works @@ -116,14 +180,51 @@ sequenceDiagram CC->>LP: Request with beta headers Note over CC,LP: anthropic-beta: header1,header2,header3 - LP->>Config: Load unsupported headers for provider - Config-->>LP: Returns unsupported list + LP->>Config: Load header mapping for provider + Config-->>LP: Returns mapping (header→value or null) - Note over LP: Filter headers:
- Remove unsupported
- Keep supported + Note over LP: Validate & Transform:
1. Check if header exists in mapping
2. Filter out null values
3. Map to provider-specific names - LP->>Provider: Request with filtered headers - Note over LP,Provider: anthropic-beta: header2
(header1, header3 removed) + LP->>Provider: Request with filtered & mapped headers + Note over LP,Provider: anthropic-beta: mapped-header2
(header1, header3 filtered out) Provider-->>LP: Success response LP-->>CC: Response -``` \ No newline at end of file +``` + +### Filtering Rules + +1. **Header must exist in mapping**: Unknown headers are filtered out +2. **Header must have non-null value**: Headers with `null` values are filtered out +3. **Header transformation**: Headers are mapped to provider-specific names (e.g., `advanced-tool-use-2025-11-20` → `tool-search-tool-2025-10-19` for Bedrock) + +### Example + +Request with headers: +``` +anthropic-beta: advanced-tool-use-2025-11-20,computer-use-2025-01-24,unknown-header +``` + +For Bedrock Converse: +- ✅ `computer-use-2025-01-24` → `computer-use-2025-01-24` (supported, passed through) +- ❌ `advanced-tool-use-2025-11-20` → filtered out (null value in config) +- ❌ `unknown-header` → filtered out (not in config) + +Result sent to Bedrock: +``` +anthropic-beta: computer-use-2025-01-24 +``` + +## Provider-Specific Notes + +### Bedrock +- Beta headers appear in both HTTP headers AND request body (`additionalModelRequestFields.anthropic_beta`) +- Some headers are transformed (e.g., `advanced-tool-use` → `tool-search-tool`) + +### Azure AI +- Uses same header names as Anthropic +- Some features not yet supported (check config for null values) + +### Vertex AI +- Some headers are transformed to match Vertex AI's implementation +- Limited beta feature support compared to Anthropic \ No newline at end of file diff --git a/litellm/__init__.py b/litellm/__init__.py index 9160297d507..0fdbac63feb 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1156,6 +1156,7 @@ from .exceptions import ( BadRequestError, ImageFetchError, NotFoundError, + PermissionDeniedError, RateLimitError, ServiceUnavailableError, BadGatewayError, diff --git a/litellm/anthropic_beta_headers_config.json b/litellm/anthropic_beta_headers_config.json index 4ebb5ddb609..5edb8067a08 100644 --- a/litellm/anthropic_beta_headers_config.json +++ b/litellm/anthropic_beta_headers_config.json @@ -1,33 +1,151 @@ { - "description": "Unsupported Anthropic beta headers for each provider. Headers listed here will be dropped. Headers not listed are passed through as-is.", - "anthropic": [], - "azure_ai": [], - "bedrock_converse": [ - "prompt-caching-scope-2026-01-05", - "bash_20250124", - "bash_20241022", - "text_editor_20250124", - "text_editor_20241022", - "compact-2026-01-12", - "advanced-tool-use-2025-11-20", - "web-fetch-2025-09-10", - "code-execution-2025-08-25", - "skills-2025-10-02", - "files-api-2025-04-14", - "fast-mode-2026-02-01" - ], - "bedrock": [ - "advanced-tool-use-2025-11-20", - "prompt-caching-scope-2026-01-05", - "structured-outputs-2025-11-13", - "web-fetch-2025-09-10", - "code-execution-2025-08-25", - "skills-2025-10-02", - "files-api-2025-04-14", - "fast-mode-2026-02-01", - "mcp-servers-2025-12-04" - ], - "vertex_ai": [ - "prompt-caching-scope-2026-01-05" - ] -} + "description": "Mapping of Anthropic beta headers for each provider. Keys are input header names, values are provider-specific header names (or null if unsupported). Only headers present in mapping keys with non-null values can be forwarded.", + "anthropic": { + "advanced-tool-use-2025-11-20": "advanced-tool-use-2025-11-20", + "bash_20241022": "bash_20241022", + "bash_20250124": "bash_20250124", + "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": "structured-output-2024-03-01", + "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": "mcp-servers-2025-12-04", + "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": "text_editor_20241022", + "text_editor_20250124": "text_editor_20250124", + "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" + }, + "azure_ai": { + "advanced-tool-use-2025-11-20": "advanced-tool-use-2025-11-20", + "bash_20241022": "bash_20241022", + "bash_20250124": "bash_20250124", + "code-execution-2025-08-25": "code-execution-2025-08-25", + "compact-2026-01-12": null, + "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": null, + "files-api-2025-04-14": "files-api-2025-04-14", + "fine-grained-tool-streaming-2025-05-14": null, + "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": "mcp-servers-2025-12-04", + "output-128k-2025-02-19": null, + "structured-output-2024-03-01": null, + "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": null, + "web-fetch-2025-09-10": "web-fetch-2025-09-10", + "web-search-2025-03-05": "web-search-2025-03-05" + }, + "bedrock_converse": { + "advanced-tool-use-2025-11-20": "tool-search-tool-2025-10-19", + "bash_20241022": null, + "bash_20250124": null, + "code-execution-2025-08-25": null, + "compact-2026-01-12": null, + "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": null, + "context-management-2025-06-27": "context-management-2025-06-27", + "effort-2025-11-24": null, + "fast-mode-2026-02-01": null, + "files-api-2025-04-14": null, + "fine-grained-tool-streaming-2025-05-14": null, + "interleaved-thinking-2025-05-14": "interleaved-thinking-2025-05-14", + "mcp-client-2025-11-20": null, + "mcp-client-2025-04-04": null, + "mcp-servers-2025-12-04": null, + "output-128k-2025-02-19": null, + "structured-output-2024-03-01": null, + "prompt-caching-scope-2026-01-05": null, + "skills-2025-10-02": null, + "structured-outputs-2025-11-13": "structured-outputs-2025-11-13", + "text_editor_20241022": null, + "text_editor_20250124": null, + "token-efficient-tools-2025-02-19": null, + "tool-search-tool-2025-10-19": "tool-search-tool-2025-10-19", + "web-fetch-2025-09-10": null, + "web-search-2025-03-05": null + }, + "bedrock": { + "advanced-tool-use-2025-11-20": "tool-search-tool-2025-10-19", + "bash_20241022": null, + "bash_20250124": null, + "code-execution-2025-08-25": null, + "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": null, + "fast-mode-2026-02-01": null, + "files-api-2025-04-14": null, + "fine-grained-tool-streaming-2025-05-14": null, + "interleaved-thinking-2025-05-14": "interleaved-thinking-2025-05-14", + "mcp-client-2025-11-20": null, + "mcp-client-2025-04-04": null, + "mcp-servers-2025-12-04": null, + "output-128k-2025-02-19": null, + "structured-output-2024-03-01": null, + "prompt-caching-scope-2026-01-05": null, + "skills-2025-10-02": null, + "structured-outputs-2025-11-13": null, + "text_editor_20241022": null, + "text_editor_20250124": null, + "token-efficient-tools-2025-02-19": null, + "tool-search-tool-2025-10-19": "tool-search-tool-2025-10-19", + "web-fetch-2025-09-10": null, + "web-search-2025-03-05": null + }, + "vertex_ai": { + "advanced-tool-use-2025-11-20": "tool-search-tool-2025-10-19", + "bash_20241022": null, + "bash_20250124": null, + "code-execution-2025-08-25": null, + "compact-2026-01-12": null, + "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": null, + "context-management-2025-06-27": "context-management-2025-06-27", + "effort-2025-11-24": null, + "fast-mode-2026-02-01": null, + "files-api-2025-04-14": null, + "fine-grained-tool-streaming-2025-05-14": null, + "interleaved-thinking-2025-05-14": "interleaved-thinking-2025-05-14", + "mcp-client-2025-11-20": null, + "mcp-client-2025-04-04": null, + "mcp-servers-2025-12-04": null, + "output-128k-2025-02-19": null, + "structured-output-2024-03-01": null, + "prompt-caching-scope-2026-01-05": null, + "skills-2025-10-02": null, + "structured-outputs-2025-11-13": null, + "text_editor_20241022": null, + "text_editor_20250124": null, + "token-efficient-tools-2025-02-19": null, + "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" + } +} \ No newline at end of file diff --git a/litellm/anthropic_beta_headers_manager.py b/litellm/anthropic_beta_headers_manager.py index 2643f4c03fa..9730ae02698 100644 --- a/litellm/anthropic_beta_headers_manager.py +++ b/litellm/anthropic_beta_headers_manager.py @@ -2,14 +2,15 @@ Centralized manager for Anthropic beta headers across different providers. This module provides utilities to: -1. Load beta header configuration from JSON (lists unsupported headers per provider) -2. Filter out unsupported beta headers +1. Load beta header configuration from JSON (mapping of supported headers per provider) +2. Filter and map beta headers based on provider support 3. Handle provider-specific header name mappings (e.g., advanced-tool-use -> tool-search-tool) Design: -- JSON config lists UNSUPPORTED headers for each provider -- Headers not in the unsupported list are passed through -- Header mappings allow renaming headers for specific providers +- JSON config contains mapping of beta headers for each provider +- Keys are input header names, values are provider-specific header names (or null if unsupported) +- Only headers present in mapping keys with non-null values can be forwarded +- This enforces stricter validation than the previous unsupported list approach """ import json @@ -47,13 +48,13 @@ def _load_beta_headers_config() -> Dict: return _BETA_HEADERS_CONFIG except Exception as e: verbose_logger.error(f"Failed to load beta headers config: {e}") - # Return empty config as fallback + # Return empty config as fallback (empty mappings) return { - "anthropic": [], - "azure_ai": [], - "bedrock": [], - "bedrock_converse": [], - "vertex_ai": [] + "anthropic": {}, + "azure_ai": {}, + "bedrock": {}, + "bedrock_converse": {}, + "vertex_ai": {} } @@ -77,21 +78,19 @@ def filter_and_transform_beta_headers( provider: str, ) -> List[str]: """ - Filter beta headers based on provider's unsupported list. + Filter and transform beta headers based on provider's mapping configuration. This function: - 1. Removes headers that are in the provider's unsupported list - 2. Passes through all other headers as-is - - Note: Header transformations/mappings (e.g., advanced-tool-use -> tool-search-tool) - are handled in each provider's transformation code, not here. + 1. Only allows headers that are present in the provider's mapping keys + 2. Filters out headers with null values (unsupported) + 3. Maps headers to provider-specific names (e.g., advanced-tool-use -> tool-search-tool) Args: beta_headers: List of Anthropic beta header values provider: Provider name (e.g., "anthropic", "bedrock", "vertex_ai") Returns: - List of filtered beta headers for the provider + List of filtered and transformed beta headers for the provider """ if not beta_headers: return [] @@ -99,23 +98,33 @@ def filter_and_transform_beta_headers( config = _load_beta_headers_config() provider = get_provider_name(provider) - # Get unsupported headers for this provider - unsupported_headers = set(config.get(provider, [])) + # Get the header mapping for this provider + provider_mapping = config.get(provider, {}) filtered_headers: Set[str] = set() for header in beta_headers: header = header.strip() - # Skip if header is unsupported - if header in unsupported_headers: + # Check if header is in the mapping + if header not in provider_mapping: + verbose_logger.debug( + f"Dropping unknown beta header '{header}' for provider '{provider}' (not in mapping)" + ) + continue + + # Get the mapped header value + mapped_header = provider_mapping[header] + + # Skip if header is unsupported (null value) + if mapped_header is None: verbose_logger.debug( f"Dropping unsupported beta header '{header}' for provider '{provider}'" ) continue - # Pass through as-is - filtered_headers.add(header) + # Add the mapped header + filtered_headers.add(mapped_header) return sorted(list(filtered_headers)) @@ -132,12 +141,14 @@ def is_beta_header_supported( provider: Provider name Returns: - True if the header is supported (not in unsupported list), False otherwise + True if the header is in the mapping with a non-null value, False otherwise """ config = _load_beta_headers_config() provider = get_provider_name(provider) - unsupported_headers = set(config.get(provider, [])) - return beta_header not in unsupported_headers + provider_mapping = config.get(provider, {}) + + # Header is supported if it's in the mapping and has a non-null value + return beta_header in provider_mapping and provider_mapping[beta_header] is not None def get_provider_beta_header( @@ -145,27 +156,29 @@ def get_provider_beta_header( provider: str, ) -> Optional[str]: """ - Check if a beta header is supported by a provider. + Get the provider-specific beta header name for a given Anthropic beta header. - Note: This does NOT handle header transformations/mappings. - Those are handled in each provider's transformation code. + This function handles header transformations/mappings (e.g., advanced-tool-use -> tool-search-tool). Args: anthropic_beta_header: The Anthropic beta header value provider: Provider name Returns: - The original header if supported, or None if unsupported + The provider-specific header name if supported, or None if unsupported/unknown """ config = _load_beta_headers_config() provider = get_provider_name(provider) - # Check if unsupported - unsupported_headers = set(config.get(provider, [])) - if anthropic_beta_header in unsupported_headers: + # Get the header mapping for this provider + provider_mapping = config.get(provider, {}) + + # Check if header is in the mapping + if anthropic_beta_header not in provider_mapping: return None - return anthropic_beta_header + # Return the mapped value (could be None if unsupported) + return provider_mapping[anthropic_beta_header] def update_headers_with_filtered_beta( @@ -208,7 +221,7 @@ def update_headers_with_filtered_beta( def get_unsupported_headers(provider: str) -> List[str]: """ - Get all beta headers that are unsupported by a provider. + Get all beta headers that are unsupported by a provider (have null values in mapping). Args: provider: Provider name @@ -218,4 +231,7 @@ def get_unsupported_headers(provider: str) -> List[str]: """ config = _load_beta_headers_config() provider = get_provider_name(provider) - return config.get(provider, []) + provider_mapping = config.get(provider, {}) + + # Return headers with null values + return [header for header, value in provider_mapping.items() if value is None] diff --git a/litellm/batch_completion/main.py b/litellm/batch_completion/main.py index 7100fb004f8..446e3f2f990 100644 --- a/litellm/batch_completion/main.py +++ b/litellm/batch_completion/main.py @@ -237,17 +237,37 @@ def batch_completion_models_all_responses(*args, **kwargs): if "model" in kwargs: kwargs.pop("model") if "models" in kwargs: - models = kwargs["models"] - kwargs.pop("models") + models = kwargs.pop("models") else: raise Exception("'models' param not in kwargs") + if isinstance(models, str): + models = [models] + elif isinstance(models, (list, tuple)): + models = list(models) + else: + raise TypeError("'models' must be a string or list of strings") + + if len(models) == 0: + return [] + responses = [] with concurrent.futures.ThreadPoolExecutor(max_workers=len(models)) as executor: - for idx, model in enumerate(models): - future = executor.submit(litellm.completion, *args, model=model, **kwargs) - if future.result() is not None: - responses.append(future.result()) + futures = [ + executor.submit(litellm.completion, *args, model=model, **kwargs) + for model in models + ] + + for future in futures: + try: + result = future.result() + if result is not None: + responses.append(result) + except Exception as e: + print_verbose( + f"batch_completion_models_all_responses: model request failed: {str(e)}" + ) + continue return responses diff --git a/litellm/integrations/arize/arize.py b/litellm/integrations/arize/arize.py index 9c2f0d95d4d..fe2f9f41f1b 100644 --- a/litellm/integrations/arize/arize.py +++ b/litellm/integrations/arize/arize.py @@ -28,6 +28,41 @@ else: class ArizeLogger(OpenTelemetry): + """ + Arize logger that sends traces to an Arize endpoint. + + Creates its own dedicated TracerProvider so it can coexist with the + generic ``otel`` callback (or any other OTEL-based integration) without + fighting over the global ``opentelemetry.trace`` TracerProvider singleton. + """ + + def _init_tracing(self, tracer_provider): + """ + Override to always create a *private* TracerProvider for Arize. + + See ArizePhoenixLogger._init_tracing for full rationale. + """ + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.trace import SpanKind + + if tracer_provider is not None: + self.tracer = tracer_provider.get_tracer("litellm") + self.span_kind = SpanKind + return + + provider = TracerProvider(resource=self._get_litellm_resource(self.config)) + provider.add_span_processor(self._get_span_processor()) + self.tracer = provider.get_tracer("litellm") + self.span_kind = SpanKind + + def _init_otel_logger_on_litellm_proxy(self): + """ + Override: Arize should NOT overwrite the proxy's + ``open_telemetry_logger``. That attribute is reserved for the + primary ``otel`` callback which handles proxy-level parent spans. + """ + pass + def set_attributes(self, span: Span, kwargs, response_obj: Optional[Any]): ArizeLogger.set_arize_attributes(span, kwargs, response_obj) return diff --git a/litellm/integrations/arize/arize_phoenix.py b/litellm/integrations/arize/arize_phoenix.py index cd345a7f76d..1b038c098f8 100644 --- a/litellm/integrations/arize/arize_phoenix.py +++ b/litellm/integrations/arize/arize_phoenix.py @@ -5,43 +5,211 @@ from litellm._logging import verbose_logger from litellm.integrations.arize import _utils from litellm.integrations.arize._utils import ArizeOTELAttributes from litellm.types.integrations.arize_phoenix import ArizePhoenixConfig -from litellm.integrations.opentelemetry import OpenTelemetry if TYPE_CHECKING: + from opentelemetry.sdk.trace import TracerProvider from opentelemetry.trace import Span as _Span + from opentelemetry.trace import SpanKind + from litellm.integrations.opentelemetry import OpenTelemetry as _OpenTelemetry from litellm.integrations.opentelemetry import OpenTelemetryConfig as _OpenTelemetryConfig from litellm.types.integrations.arize import Protocol as _Protocol Protocol = _Protocol OpenTelemetryConfig = _OpenTelemetryConfig Span = Union[_Span, Any] + OpenTelemetry = _OpenTelemetry else: Protocol = Any OpenTelemetryConfig = Any Span = Any + TracerProvider = Any + SpanKind = Any + # Import OpenTelemetry at runtime + try: + from litellm.integrations.opentelemetry import OpenTelemetry + except ImportError: + OpenTelemetry = None # type: ignore ARIZE_HOSTED_PHOENIX_ENDPOINT = "https://otlp.arize.com/v1/traces" -class ArizePhoenixLogger(OpenTelemetry): +class ArizePhoenixLogger(OpenTelemetry): # type: ignore + """ + Arize Phoenix logger that sends traces to a Phoenix endpoint. + + Creates its own dedicated TracerProvider so it can coexist with the + generic ``otel`` callback (or any other OTEL-based integration) without + fighting over the global ``opentelemetry.trace`` TracerProvider singleton. + """ + + def _init_tracing(self, tracer_provider): + """ + Override to always create a *private* TracerProvider for Arize Phoenix. + + The base ``OpenTelemetry._init_tracing`` falls back to the global + TracerProvider when one already exists. That causes whichever + integration initialises second to silently reuse the first one's + exporter, so spans only reach one destination. + + By creating our own provider we guarantee Arize Phoenix always gets + its own exporter pipeline, regardless of initialisation order. + """ + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.trace import SpanKind + + if tracer_provider is not None: + # Explicitly supplied (e.g. in tests) — honour it. + self.tracer = tracer_provider.get_tracer("litellm") + self.span_kind = SpanKind + return + + # Always create a dedicated provider — never touch the global one. + provider = TracerProvider(resource=self._get_litellm_resource(self.config)) + provider.add_span_processor(self._get_span_processor()) + self.tracer = provider.get_tracer("litellm") + self.span_kind = SpanKind + verbose_logger.debug( + "ArizePhoenixLogger: Created dedicated TracerProvider " + "(endpoint=%s, exporter=%s)", + self.config.endpoint, + self.config.exporter, + ) + + def _init_otel_logger_on_litellm_proxy(self): + """ + Override: Arize Phoenix should NOT overwrite the proxy's + ``open_telemetry_logger``. That attribute is reserved for the + primary ``otel`` callback which handles proxy-level parent spans. + """ + pass + def set_attributes(self, span: Span, kwargs, response_obj: Optional[Any]): ArizePhoenixLogger.set_arize_phoenix_attributes(span, kwargs, response_obj) return @staticmethod def set_arize_phoenix_attributes(span: Span, kwargs, response_obj): + from litellm.integrations.opentelemetry_utils.base_otel_llm_obs_attributes import safe_set_attribute + _utils.set_attributes(span, kwargs, response_obj, ArizeOTELAttributes) - - # Set project name on the span for all traces to go to custom Phoenix projects - config = ArizePhoenixLogger.get_arize_phoenix_config() - if config.project_name: - from litellm.integrations.opentelemetry_utils.base_otel_llm_obs_attributes import safe_set_attribute - safe_set_attribute(span, "openinference.project.name", config.project_name) - + + # Dynamic project name: check metadata first, then fall back to env var config + dynamic_project_name = ArizePhoenixLogger._get_dynamic_project_name(kwargs) + if dynamic_project_name: + safe_set_attribute(span, "openinference.project.name", dynamic_project_name) + else: + # Fall back to static config from env var + config = ArizePhoenixLogger.get_arize_phoenix_config() + if config.project_name: + safe_set_attribute(span, "openinference.project.name", config.project_name) + return + @staticmethod + def _get_dynamic_project_name(kwargs) -> Optional[str]: + """ + Retrieve dynamic Phoenix project name from request metadata. + + Users can set `metadata.phoenix_project_name` in their request to route + traces to different Phoenix projects dynamically. + """ + standard_logging_payload = kwargs.get("standard_logging_object") + if isinstance(standard_logging_payload, dict): + metadata = standard_logging_payload.get("metadata") + if isinstance(metadata, dict): + project_name = metadata.get("phoenix_project_name") + if project_name: + return str(project_name) + + # Also check litellm_params.metadata for SDK usage + litellm_params = kwargs.get("litellm_params") + if isinstance(litellm_params, dict): + metadata = litellm_params.get("metadata") or {} + else: + metadata = {} + if isinstance(metadata, dict): + project_name = metadata.get("phoenix_project_name") + if project_name: + return str(project_name) + + return None + + def _handle_success(self, kwargs, response_obj, start_time, end_time): + """ + Override to prevent creating duplicate litellm_request spans when a proxy parent span exists. + + ArizePhoenixLogger should reuse the proxy parent span instead of creating a new litellm_request span, + to maintain a shallow span hierarchy as expected by Arize Phoenix. + """ + from opentelemetry.trace import Status, StatusCode + from litellm.secret_managers.main import get_secret_bool + from litellm.integrations.opentelemetry import LITELLM_PROXY_REQUEST_SPAN_NAME + + verbose_logger.debug( + "ArizePhoenixLogger: Logging kwargs: %s, OTEL config settings=%s", + kwargs, + self.config, + ) + ctx, parent_span = self._get_span_context(kwargs) + + # ArizePhoenixLogger NEVER creates a litellm_request span when a proxy parent span exists + # This is different from the base OpenTelemetry behavior which respects USE_OTEL_LITELLM_REQUEST_SPAN + should_create_primary_span = parent_span is None or ( + parent_span.name != LITELLM_PROXY_REQUEST_SPAN_NAME + and get_secret_bool("USE_OTEL_LITELLM_REQUEST_SPAN") + ) + + if should_create_primary_span: + # Create a new litellm_request span + span = self._start_primary_span( + kwargs, response_obj, start_time, end_time, ctx + ) + # Raw-request sub-span (if enabled) - child of litellm_request span + self._maybe_log_raw_request( + kwargs, response_obj, start_time, end_time, span + ) + # Ensure proxy-request parent span is annotated with the actual operation kind + if ( + parent_span is not None + and parent_span.name == LITELLM_PROXY_REQUEST_SPAN_NAME + ): + self.set_attributes(parent_span, kwargs, response_obj) + else: + # Do not create primary span (keep hierarchy shallow when parent exists) + span = None + # Only set attributes if the span is still recording (not closed) + # Note: parent_span is guaranteed to be not None here + if parent_span.is_recording(): + parent_span.set_status(Status(StatusCode.OK)) + self.set_attributes(parent_span, kwargs, response_obj) + # Raw-request as direct child of parent_span + self._maybe_log_raw_request( + kwargs, response_obj, start_time, end_time, parent_span + ) + + # 3. Guardrail span + self._create_guardrail_span(kwargs=kwargs, context=ctx) + + # 4. Metrics & cost recording + self._record_metrics(kwargs, response_obj, start_time, end_time) + + # 5. Semantic logs. + if self.config.enable_events: + log_span = span if span is not None else parent_span + if log_span is not None: + self._emit_semantic_logs(kwargs, response_obj, log_span) + + # 6. Do NOT end parent span - it should be managed by its creator + # External spans (from Langfuse, user code, HTTP headers, global context) must not be closed by LiteLLM + # However, proxy-created spans should be closed here + if ( + parent_span is not None + and parent_span.name == LITELLM_PROXY_REQUEST_SPAN_NAME + ): + parent_span.end(end_time=self._to_ns(end_time)) + @staticmethod def get_arize_phoenix_config() -> ArizePhoenixConfig: """ diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 7ae62718af0..82a7af64f97 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -3764,7 +3764,7 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 from litellm.integrations.opentelemetry import OpenTelemetry for callback in _in_memory_loggers: - if isinstance(callback, OpenTelemetry): + if type(callback) is OpenTelemetry: return callback # type: ignore otel_logger = OpenTelemetry( **_get_custom_logger_settings_from_proxy_server( diff --git a/litellm/litellm_core_utils/redact_messages.py b/litellm/litellm_core_utils/redact_messages.py index aa763dc9899..5d6d1fbc1c5 100644 --- a/litellm/litellm_core_utils/redact_messages.py +++ b/litellm/litellm_core_utils/redact_messages.py @@ -140,9 +140,14 @@ def should_redact_message_logging(model_call_details: dict) -> bool: metadata_field = get_metadata_variable_name_from_kwargs(litellm_params) metadata = litellm_params.get(metadata_field, {}) - + if not isinstance(metadata, dict): + # Fall back: litellm_metadata was None, try metadata + metadata = litellm_params.get("metadata", {}) + if not isinstance(metadata, dict): + metadata = {} + # Get headers from the metadata - request_headers = metadata.get("headers", {}) if isinstance(metadata, dict) else {} + request_headers = metadata.get("headers", {}) # Check for headers that explicitly control redaction if request_headers and bool( diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index e85c0d0d017..f51adf96102 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -58,6 +58,9 @@ from litellm.types.utils import ( from ...base import BaseLLM from ..common_utils import AnthropicError, process_anthropic_headers +from litellm.anthropic_beta_headers_manager import ( + update_headers_with_filtered_beta, +) from .transformation import AnthropicConfig if TYPE_CHECKING: @@ -333,6 +336,10 @@ class AnthropicChatCompletion(BaseLLM): litellm_params=litellm_params, ) + headers = update_headers_with_filtered_beta( + headers=headers, provider=custom_llm_provider + ) + config = ProviderConfigManager.get_provider_chat_config( model=model, provider=LlmProviders(custom_llm_provider), diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 53458b46fac..9938cd7979b 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -934,6 +934,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): Translate system message to anthropic format. Removes system message from the original list and returns a new list of anthropic system message content. + Filters out system messages containing x-anthropic-billing-header metadata. """ system_prompt_indices = [] anthropic_system_message_list: List[AnthropicSystemMessageContent] = [] @@ -945,6 +946,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): # Skip empty text blocks - Anthropic API raises errors for empty text if not system_message_block["content"]: continue + # Skip system messages containing x-anthropic-billing-header metadata + if system_message_block["content"].startswith("x-anthropic-billing-header:"): + continue anthropic_system_message_content = AnthropicSystemMessageContent( type="text", text=system_message_block["content"], @@ -963,6 +967,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): text_value = _content.get("text") if _content.get("type") == "text" and not text_value: continue + # Skip system messages containing x-anthropic-billing-header metadata + if _content.get("type") == "text" and text_value and text_value.startswith("x-anthropic-billing-header:"): + continue anthropic_system_message_content = ( AnthropicSystemMessageContent( type=_content.get("type"), diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py index 043a70f3c67..8f2f3bf3545 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py @@ -2,9 +2,6 @@ from typing import Any, AsyncIterator, Dict, List, Optional, Tuple import httpx -from litellm.anthropic_beta_headers_manager import ( - update_headers_with_filtered_beta, -) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.litellm_logging import verbose_logger from litellm.llms.base_llm.anthropic_messages.transformation import ( @@ -52,6 +49,40 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): # TODO: Add Anthropic `metadata` support # "metadata", ] + + @staticmethod + def _filter_billing_headers_from_system(system_param): + """ + Filter out x-anthropic-billing-header metadata from system parameter. + + Args: + system_param: Can be a string or a list of system message content blocks + + Returns: + Filtered system parameter (string or list), or None if all content was filtered + """ + if isinstance(system_param, str): + # If it's a string and starts with billing header, filter it out + if system_param.startswith("x-anthropic-billing-header:"): + return None + return system_param + elif isinstance(system_param, list): + # Filter list of system content blocks + filtered_list = [] + for content_block in system_param: + if isinstance(content_block, dict): + text = content_block.get("text", "") + content_type = content_block.get("type", "") + # Skip text blocks that start with billing header + if content_type == "text" and text.startswith("x-anthropic-billing-header:"): + continue + filtered_list.append(content_block) + else: + # Keep non-dict items as-is + filtered_list.append(content_block) + return filtered_list if len(filtered_list) > 0 else None + else: + return system_param def get_complete_url( self, @@ -96,11 +127,6 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): optional_params=optional_params, ) - headers = update_headers_with_filtered_beta( - headers=headers, - provider="anthropic", - ) - return headers, api_base def transform_anthropic_messages_request( @@ -123,6 +149,17 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): message="max_tokens is required for Anthropic /v1/messages API", status_code=400, ) + + # Filter out x-anthropic-billing-header from system messages + system_param = anthropic_messages_optional_request_params.get("system") + if system_param is not None: + filtered_system = self._filter_billing_headers_from_system(system_param) + if filtered_system is not None and len(filtered_system) > 0: + anthropic_messages_optional_request_params["system"] = filtered_system + else: + # Remove system parameter if all content was filtered out + anthropic_messages_optional_request_params.pop("system", None) + ####### get required params for all anthropic messages requests ###### verbose_logger.debug(f"TRANSFORMATION DEBUG - Messages: {messages}") anthropic_messages_request: AnthropicMessagesRequest = AnthropicMessagesRequest( diff --git a/litellm/llms/azure/chat/gpt_transformation.py b/litellm/llms/azure/chat/gpt_transformation.py index 0ae6fad7300..18dad503a59 100644 --- a/litellm/llms/azure/chat/gpt_transformation.py +++ b/litellm/llms/azure/chat/gpt_transformation.py @@ -105,6 +105,7 @@ class AzureOpenAIConfig(BaseConfig): "modalities", "audio", "web_search_options", + "prompt_cache_key", ] def _is_response_format_supported_model(self, model: str) -> bool: diff --git a/litellm/llms/azure_ai/anthropic/messages_transformation.py b/litellm/llms/azure_ai/anthropic/messages_transformation.py index f86ec7082f2..a4dc88f9c68 100644 --- a/litellm/llms/azure_ai/anthropic/messages_transformation.py +++ b/litellm/llms/azure_ai/anthropic/messages_transformation.py @@ -3,9 +3,6 @@ Azure Anthropic messages transformation config - extends AnthropicMessagesConfig """ from typing import TYPE_CHECKING, Any, List, Optional, Tuple -from litellm.anthropic_beta_headers_manager import ( - update_headers_with_filtered_beta, -) from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( AnthropicMessagesConfig, ) @@ -65,18 +62,11 @@ class AzureAnthropicMessagesConfig(AnthropicMessagesConfig): if "content-type" not in headers: headers["content-type"] = "application/json" - # Update headers with anthropic beta features (context management, tool search, etc.) headers = self._update_headers_with_anthropic_beta( headers=headers, optional_params=optional_params, ) - # Filter out unsupported beta headers for Azure AI - headers = update_headers_with_filtered_beta( - headers=headers, - provider="azure_ai", - ) - return headers, api_base def get_complete_url( diff --git a/litellm/llms/azure_ai/anthropic/transformation.py b/litellm/llms/azure_ai/anthropic/transformation.py index 753bc9c08eb..c5510db68b1 100644 --- a/litellm/llms/azure_ai/anthropic/transformation.py +++ b/litellm/llms/azure_ai/anthropic/transformation.py @@ -2,10 +2,6 @@ Azure Anthropic transformation config - extends AnthropicConfig with Azure authentication """ from typing import TYPE_CHECKING, Dict, List, Optional, Union - -from litellm.anthropic_beta_headers_manager import ( - update_headers_with_filtered_beta, -) from litellm.llms.anthropic.chat.transformation import AnthropicConfig from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.types.llms.openai import AllMessageValues @@ -90,11 +86,6 @@ class AzureAnthropicConfig(AnthropicConfig): if "anthropic-version" not in headers: headers["anthropic-version"] = "2023-06-01" - # Filter out unsupported beta headers for Azure AI - headers = update_headers_with_filtered_beta( - headers=headers, - provider="azure_ai", - ) return headers diff --git a/litellm/llms/azure_ai/chat/transformation.py b/litellm/llms/azure_ai/chat/transformation.py index 04d2b3a2769..585efd3307d 100644 --- a/litellm/llms/azure_ai/chat/transformation.py +++ b/litellm/llms/azure_ai/chat/transformation.py @@ -11,12 +11,14 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import ( _audio_or_image_in_message_content, convert_content_list_to_str, ) +from litellm.llms.azure.common_utils import BaseAzureLLM from litellm.llms.base_llm.chat.transformation import LiteLLMLoggingObj from litellm.llms.openai.common_utils import drop_params_from_unprocessable_entity_error from litellm.llms.openai.openai import OpenAIConfig from litellm.llms.xai.chat.transformation import XAIChatConfig from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllMessageValues +from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ModelResponse, ProviderField from litellm.utils import _add_path_to_api_base, supports_tool_choice @@ -64,12 +66,21 @@ class AzureAIStudioConfig(OpenAIConfig): api_key: Optional[str] = None, api_base: Optional[str] = None, ) -> dict: - if api_base and self._should_use_api_key_header(api_base): - headers["api-key"] = api_key + if api_key: + if api_base and self._should_use_api_key_header(api_base): + headers["api-key"] = api_key + else: + headers["Authorization"] = f"Bearer {api_key}" else: - headers["Authorization"] = f"Bearer {api_key}" + # No api_key provided — fall back to Azure AD token-based auth + litellm_params_obj = GenericLiteLLMParams( + **(litellm_params if isinstance(litellm_params, dict) else {}) + ) + headers = BaseAzureLLM._base_validate_azure_environment( + headers=headers, litellm_params=litellm_params_obj + ) - headers["Content-Type"] = "application/json" # tell Azure AI Studio to expect JSON + headers["Content-Type"] = "application/json" return headers diff --git a/litellm/llms/bedrock/base_aws_llm.py b/litellm/llms/bedrock/base_aws_llm.py index 1de1c40c438..304c707fa0b 100644 --- a/litellm/llms/bedrock/base_aws_llm.py +++ b/litellm/llms/bedrock/base_aws_llm.py @@ -211,25 +211,13 @@ class BaseAWSLLM: aws_external_id=aws_external_id, ) elif aws_role_name is not None: - # Check if we're in IRSA and trying to assume the same role we already have - current_role_arn = os.getenv("AWS_ROLE_ARN") - web_identity_token_file = os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE") - - # In IRSA environments, we should skip role assumption if we're already running as the target role - # This is true when: - # 1. We have AWS_ROLE_ARN set (current role) - # 2. We have AWS_WEB_IDENTITY_TOKEN_FILE set (IRSA environment) - # 3. The current role matches the requested role - if ( - current_role_arn - and web_identity_token_file - and current_role_arn == aws_role_name - ): + # Check if we're already running as the target role and can skip assumption + # This handles IRSA (EKS), ECS task roles, and EC2 instance profiles + if self._is_already_running_as_role(aws_role_name, ssl_verify=ssl_verify): verbose_logger.debug( - "Using IRSA same-role optimization: calling _auth_with_env_vars" + "Already running as target role %s, using ambient credentials", + aws_role_name, ) - # We're already running as this role via IRSA, no need to assume it again - # Use the default boto3 credentials (which will use the IRSA credentials) credentials, _cache_ttl = self._auth_with_env_vars() else: verbose_logger.debug( @@ -553,6 +541,107 @@ class BaseAWSLLM: aws_region_name = "us-west-2" return aws_region_name + @staticmethod + def _parse_arn_account_and_role_name( + arn: str, + ) -> Optional[Tuple[str, str, str]]: + """ + Parse an ARN and return (partition, account_id, role_name). + + Handles: + - arn:aws:iam::123456789012:role/MyRole + - arn:aws:iam::123456789012:role/path/to/MyRole + - arn:aws:sts::123456789012:assumed-role/MyRole/session-name + + Returns None if the ARN cannot be parsed. + """ + # ARN format: arn:PARTITION:SERVICE:REGION:ACCOUNT:RESOURCE + parts = arn.split(":") + if len(parts) < 6 or parts[0] != "arn": + return None + + partition = parts[1] # e.g. "aws", "aws-cn", "aws-us-gov" + account_id = parts[4] + resource = ":".join(parts[5:]) # rejoin in case resource contains colons + + if resource.startswith("role/"): + # arn:aws:iam::ACCOUNT:role/[path/]ROLE_NAME + role_name = resource.split("/")[-1] + elif resource.startswith("assumed-role/"): + # arn:aws:sts::ACCOUNT:assumed-role/ROLE_NAME/SESSION + role_parts = resource.split("/") + if len(role_parts) >= 2: + role_name = role_parts[1] + else: + return None + else: + return None + + return partition, account_id, role_name + + def _is_already_running_as_role( + self, + aws_role_name: str, + ssl_verify: Optional[Union[bool, str]] = None, + ) -> bool: + """ + Check if the current environment is already running as the target IAM role. + + This handles multiple AWS environments: + - IRSA (EKS): AWS_ROLE_ARN + AWS_WEB_IDENTITY_TOKEN_FILE are set + - ECS task roles: Uses sts:GetCallerIdentity to check current role ARN + - EC2 instance profiles: Uses sts:GetCallerIdentity to check current role ARN + + Compares partition, account ID, and role name to avoid cross-account + false matches. + + Returns True if the current identity matches the target role, meaning + we can skip sts:AssumeRole and use ambient credentials directly. + """ + target_parsed = self._parse_arn_account_and_role_name(aws_role_name) + if target_parsed is None: + return False + + target_partition, target_account, target_role = target_parsed + + # Fast path: IRSA environment check (no API call needed) + current_role_arn = os.getenv("AWS_ROLE_ARN") + web_identity_token_file = os.getenv("AWS_WEB_IDENTITY_TOKEN_FILE") + if current_role_arn and web_identity_token_file: + return current_role_arn == aws_role_name + + # For ECS/EC2: call sts:GetCallerIdentity to check if already running as the role + try: + import boto3 + + with tracer.trace("boto3.client(sts).get_caller_identity"): + sts_client = boto3.client( + "sts", verify=self._get_ssl_verify(ssl_verify) + ) + identity = sts_client.get_caller_identity() + caller_arn = identity.get("Arn", "") + + caller_parsed = self._parse_arn_account_and_role_name(caller_arn) + if caller_parsed is not None: + caller_partition, caller_account, caller_role = caller_parsed + if ( + caller_partition == target_partition + and caller_account == target_account + and caller_role == target_role + ): + verbose_logger.debug( + "Current identity already matches target role: %s", + aws_role_name, + ) + return True + + except Exception as e: + verbose_logger.debug( + "Could not determine current role identity: %s", str(e) + ) + + return False + @tracer.wrap() def _auth_with_web_identity_token( self, @@ -867,7 +956,35 @@ class BaseAWSLLM: if aws_external_id is not None: assume_role_params["ExternalId"] = aws_external_id - sts_response = sts_client.assume_role(**assume_role_params) + try: + sts_response = sts_client.assume_role(**assume_role_params) + except Exception as e: + error_str = str(e) + if "AccessDenied" in error_str: + # Only fall back to ambient credentials if we can positively + # confirm the caller is already the target role (same account, + # partition, and role name). This avoids silently using the + # wrong identity when there is a genuine trust-policy or + # permission misconfiguration. + if self._is_already_running_as_role( + aws_role_name, ssl_verify=ssl_verify + ): + verbose_logger.warning( + "AssumeRole failed for %s (%s). " + "Caller is already running as this role; " + "falling back to ambient credentials.", + aws_role_name, + error_str, + ) + return self._auth_with_env_vars() + # Genuine permission error — re-raise + verbose_logger.error( + "AssumeRole AccessDenied for %s and caller is NOT " + "the same role. Re-raising. Error: %s", + aws_role_name, + error_str, + ) + raise # Extract the credentials from the response and convert to Session Credentials sts_credentials = sts_response["Credentials"] diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index d5bd054118d..25af852e09c 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -13,7 +13,9 @@ from litellm.llms.custom_httpx.http_handler import ( ) from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper - +from litellm.anthropic_beta_headers_manager import ( + update_headers_with_filtered_beta, + ) from ..base_aws_llm import BaseAWSLLM, Credentials from ..common_utils import BedrockError from .invoke_handler import AWSEventStreamDecoder, MockResponseIterator, make_call @@ -337,7 +339,11 @@ class BedrockConverseLLM(BaseAWSLLM): headers = {"Content-Type": "application/json"} if extra_headers is not None: headers = {"Content-Type": "application/json", **extra_headers} - + + # Filter beta headers in HTTP headers before making the request + headers = update_headers_with_filtered_beta( + headers=headers, provider="bedrock_converse" + ) ### ROUTING (ASYNC, STREAMING, SYNC) if acompletion: if isinstance(client, HTTPHandler): diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index 7fc51263ebb..efa755d515e 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -11,9 +11,6 @@ import httpx import litellm from litellm._logging import verbose_logger -from litellm.anthropic_beta_headers_manager import ( - filter_and_transform_beta_headers, -) from litellm.constants import RESPONSE_FORMAT_TOOL_NAME from litellm.litellm_core_utils.core_helpers import ( filter_exceptions_from_params, @@ -1132,24 +1129,9 @@ class AmazonConverseConfig(BaseConfig): # Set anthropic_beta in additional_request_params if we have any beta features # ONLY apply to Anthropic/Claude models - other models (e.g., Qwen, Llama) don't support this field - # and will error with "unknown variant anthropic_beta" if included base_model = BedrockModelInfo.get_base_model(model) if anthropic_beta_list and base_model.startswith("anthropic"): - # Remove duplicates while preserving order - unique_betas = [] - seen = set() - for beta in anthropic_beta_list: - if beta not in seen: - unique_betas.append(beta) - seen.add(beta) - - filtered_betas = filter_and_transform_beta_headers( - beta_headers=unique_betas, - provider="bedrock_converse", - ) - - if filtered_betas: - additional_request_params["anthropic_beta"] = filtered_betas + additional_request_params["anthropic_beta"] = anthropic_beta_list return bedrock_tools, anthropic_beta_list diff --git a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py index 31119c73d72..dfab81123fd 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/anthropic_claude3_transformation.py @@ -2,7 +2,6 @@ from typing import TYPE_CHECKING, Any, List, Optional import httpx -from litellm.anthropic_beta_headers_manager import filter_and_transform_beta_headers from litellm.llms.anthropic.chat.transformation import AnthropicConfig from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( AmazonInvokeConfig, @@ -136,13 +135,7 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig): # Filter out beta headers that Bedrock Invoke doesn't support # Uses centralized configuration from anthropic_beta_headers_config.json beta_list = list(beta_set) - filtered_beta_list = filter_and_transform_beta_headers( - beta_headers=beta_list, - provider="bedrock", - ) - - if filtered_beta_list: - _anthropic_request["anthropic_beta"] = filtered_beta_list + _anthropic_request["anthropic_beta"] = beta_list return _anthropic_request 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 19fe7d8c140..477fa3316d1 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -12,9 +12,6 @@ from typing import ( import httpx -from litellm.anthropic_beta_headers_manager import ( - filter_and_transform_beta_headers, -) from litellm.llms.anthropic.common_utils import AnthropicModelInfo from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( AnthropicMessagesConfig, @@ -249,75 +246,15 @@ class AmazonAnthropicClaudeMessagesConfig( "sonnet_4.5", "sonnet-4-5", "sonnet_4_5", + # Opus 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) - def _filter_unsupported_beta_headers_for_bedrock( - self, model: str, beta_set: set - ) -> None: - """ - Remove beta headers that are not supported on Bedrock for the given model. - - Extended thinking beta headers are only supported on specific Claude 4+ models. - Advanced tool use headers are not supported on Bedrock Invoke API, but need to be - translated to Bedrock-specific headers for models that support tool search - (Claude Opus 4.5, Sonnet 4.5). - This prevents 400 "invalid beta flag" errors on Bedrock. - - Note: Bedrock Invoke API fails with a 400 error when unsupported beta headers - are sent, returning: {"message":"invalid beta flag"} - - Translation for models supporting tool search (Opus 4.5, Sonnet 4.5): - - advanced-tool-use-2025-11-20 -> tool-search-tool-2025-10-19 + tool-examples-2025-10-29 - - Args: - model: The model name - beta_set: The set of beta headers to filter in-place - """ - # 1. Handle header transformations BEFORE filtering - # (advanced-tool-use -> tool-search-tool) - # This must happen before filtering because advanced-tool-use is in the unsupported list - has_advanced_tool_use = "advanced-tool-use-2025-11-20" in beta_set - if has_advanced_tool_use and self._supports_tool_search_on_bedrock(model): - beta_set.discard("advanced-tool-use-2025-11-20") - beta_set.add("tool-search-tool-2025-10-19") - beta_set.add("tool-examples-2025-10-29") - - # 2. Apply provider-level filtering using centralized JSON config - beta_list = list(beta_set) - filtered_list = filter_and_transform_beta_headers( - beta_headers=beta_list, - provider="bedrock", - ) - - # Update the set with filtered headers - beta_set.clear() - beta_set.update(filtered_list) - - # 2.1. Handle model-specific exceptions: structured-outputs is only supported on Opus 4.6 - # Re-add structured-outputs if it was in the original set and model is Opus 4.6 - model_lower = model.lower() - is_opus_4_6 = any(pattern in model_lower for pattern in ["opus-4.6", "opus_4.6", "opus-4-6", "opus_4_6"]) - if is_opus_4_6 and "structured-outputs-2025-11-13" in beta_list: - beta_set.add("structured-outputs-2025-11-13") - - # 3. Filter out extended thinking headers for models that don't support them - extended_thinking_patterns = [ - "extended-thinking", - "interleaved-thinking", - ] - if not self._supports_extended_thinking_on_bedrock(model): - beta_headers_to_remove = set() - for beta in beta_set: - for pattern in extended_thinking_patterns: - if pattern in beta.lower(): - beta_headers_to_remove.add(beta) - break - - for beta in beta_headers_to_remove: - beta_set.discard(beta) - def _get_tool_search_beta_header_for_bedrock( self, model: str, @@ -483,12 +420,11 @@ class AmazonAnthropicClaudeMessagesConfig( beta_set=beta_set, ) - # Filter out unsupported beta headers for Bedrock (e.g., advanced-tool-use, extended-thinking on non-Opus/Sonnet 4 models) - self._filter_unsupported_beta_headers_for_bedrock( - model=model, - 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/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 95db8ec64b3..a97ebd8e74c 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -81,6 +81,9 @@ from litellm.types.llms.anthropic_skills import ( ListSkillsResponse, Skill, ) +from litellm.anthropic_beta_headers_manager import ( + update_headers_with_filtered_beta, + ) from litellm.types.llms.openai import ( CreateBatchRequest, CreateFileRequest, @@ -1858,6 +1861,10 @@ class BaseLLMHTTPHandler: api_key=api_key, api_base=api_base, ) + + headers = update_headers_with_filtered_beta( + headers=headers, provider=custom_llm_provider + ) logging_obj.update_environment_variables( model=model, diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index e66394ae5f5..1c22602b483 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -838,6 +838,15 @@ class OCIChatConfig(BaseConfig): if not user_messages: raise Exception("No user message found for Cohere model") + # Extract system messages into preambleOverride + system_messages = [msg for msg in messages if msg.get("role") == "system"] + preamble_override = None + if system_messages: + preamble = "\n".join( + self._extract_text_content(msg["content"]) for msg in system_messages + ) + if preamble: + preamble_override = preamble # Create Cohere-specific chat request optional_cohere_params = self._get_optional_params(OCIVendors.COHERE, optional_params) @@ -845,6 +854,7 @@ class OCIChatConfig(BaseConfig): apiFormat="COHERE", message=self._extract_text_content(user_messages[-1]["content"]), chatHistory=self.adapt_messages_to_cohere_standard(messages), + preambleOverride=preamble_override, **optional_cohere_params ) diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 6cc09dafc2f..16368907070 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -20,12 +20,12 @@ from typing import ( import httpx import litellm +from litellm.litellm_core_utils.core_helpers import map_finish_reason from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( _extract_reasoning_content, _handle_invalid_parallel_tool_calls, _should_convert_tool_call_to_json_mode, ) -from litellm.litellm_core_utils.core_helpers import map_finish_reason from litellm.litellm_core_utils.prompt_templates.common_utils import get_tool_call_names from litellm.litellm_core_utils.prompt_templates.image_handling import ( async_convert_url_to_base64, @@ -161,6 +161,7 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): "web_search_options", "service_tier", "safety_identifier", + "prompt_cache_key", ] # works across all models model_specific_params = [] diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py index 5a09168282d..54c3f9e0474 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/experimental_pass_through/transformation.py @@ -1,8 +1,5 @@ from typing import Any, Dict, List, Optional, Tuple -from litellm.anthropic_beta_headers_manager import ( - update_headers_with_filtered_beta, -) from litellm.llms.anthropic.common_utils import AnthropicModelInfo from litellm.llms.anthropic.experimental_pass_through.messages.transformation import ( AnthropicMessagesConfig, @@ -105,12 +102,6 @@ class VertexAIPartnerModelsAnthropicMessagesConfig(AnthropicMessagesConfig, Vert if beta_values: headers["anthropic-beta"] = ",".join(beta_values) - # Filter out unsupported beta headers for Vertex AI - headers = update_headers_with_filtered_beta( - headers=headers, - provider="vertex_ai", - ) - return headers, api_base def get_complete_url( diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 9b5a7b42d0e..f6edcf7efd0 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -10756,14 +10756,22 @@ "input_cost_per_token": 2.8e-07, "input_cost_per_token_cache_hit": 2.8e-08, "litellm_provider": "deepseek", - "max_input_tokens": 128000, + "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 4.2e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], "supports_assistant_prefill": true, "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, "supports_tool_choice": true }, "deepseek/deepseek-coder": { @@ -10800,16 +10808,24 @@ "input_cost_per_token": 2.8e-07, "input_cost_per_token_cache_hit": 2.8e-08, "litellm_provider": "deepseek", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_input_tokens": 131072, + "max_output_tokens": 65536, + "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 4.2e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], "supports_assistant_prefill": true, - "supports_function_calling": true, + "supports_function_calling": false, + "supports_native_streaming": true, + "supports_parallel_function_calling": false, "supports_prompt_caching": true, "supports_reasoning": true, - "supports_tool_choice": true + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": false }, "deepseek/deepseek-v3": { "cache_creation_input_token_cost": 0.0, diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 786cfbfb008..548e3bc3dbf 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -6,7 +6,12 @@ from starlette.requests import Request from starlette.types import Scope from litellm._logging import verbose_logger -from litellm.proxy._types import LiteLLM_TeamTable, ProxyException, SpecialHeaders, UserAPIKeyAuth +from litellm.proxy._types import ( + LiteLLM_TeamTable, + ProxyException, + SpecialHeaders, + UserAPIKeyAuth, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -372,45 +377,31 @@ class MCPRequestHandler: return [] @staticmethod - async def _get_key_object_permission( + def _get_key_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, ): - """Helper to get key object_permission from cache or DB.""" - from litellm.proxy.auth.auth_checks import get_object_permission - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) + """ + Get key object_permission - already loaded by get_key_object() in main auth flow. + Note: object_permission is automatically populated when the key is fetched via + get_key_object() in litellm/proxy/auth/auth_checks.py + """ if not user_api_key_auth: return None - # Already loaded - if user_api_key_auth.object_permission: - return user_api_key_auth.object_permission - - # Need to fetch from DB - if user_api_key_auth.object_permission_id and prisma_client: - return await get_object_permission( - object_permission_id=user_api_key_auth.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - - return None + return user_api_key_auth.object_permission @staticmethod async def _get_team_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, ): - """Helper to get team object_permission from cache or DB.""" - from litellm.proxy.auth.auth_checks import ( - get_object_permission, - get_team_object, - ) + """ + Get team object_permission - automatically loaded by get_team_object() in main auth flow. + + Note: object_permission is automatically populated when the team is fetched via + get_team_object() in litellm/proxy/auth/auth_checks.py + """ + from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.proxy_server import ( prisma_client, proxy_logging_obj, @@ -423,7 +414,7 @@ class MCPRequestHandler: if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client: return None - # First get the team object (which may have object_permission already loaded) + # Get the team object (which has object_permission already loaded) team_obj: Optional[LiteLLM_TeamTable] = await get_team_object( team_id=user_api_key_auth.team_id, prisma_client=prisma_client, @@ -435,21 +426,7 @@ class MCPRequestHandler: if not team_obj: return None - # Already loaded - if team_obj.object_permission: - return team_obj.object_permission - - # Need to fetch from DB using object_permission_id - if team_obj.object_permission_id: - return await get_object_permission( - object_permission_id=team_obj.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - - return None + return team_obj.object_permission @staticmethod async def get_allowed_tools_for_server( @@ -471,8 +448,8 @@ class MCPRequestHandler: return None try: - # Get key and team object permissions - key_obj_perm = await MCPRequestHandler._get_key_object_permission( + # Get key and team object permissions (already loaded in main auth flow) + key_obj_perm = MCPRequestHandler._get_key_object_permission( user_api_key_auth ) team_obj_perm = await MCPRequestHandler._get_team_object_permission( @@ -559,7 +536,8 @@ class MCPRequestHandler: user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> List[str]: try: - key_object_permission = await MCPRequestHandler._get_key_object_permission( + # Get key object permission (already loaded in main auth flow) + key_object_permission = MCPRequestHandler._get_key_object_permission( user_api_key_auth ) if key_object_permission is None: @@ -591,12 +569,10 @@ class MCPRequestHandler: """ Get allowed MCP servers for a team. - Uses the helper _get_team_object_permission which: - 1. First checks if object_permission is already loaded on the team - 2. If not, fetches from DB using object_permission_id if it exists + Note: object_permission is automatically loaded by get_team_object() in main auth flow. """ try: - # Use the helper method that properly handles fetching from DB if needed + # Get team object permission (already loaded in main auth flow) object_permissions = await MCPRequestHandler._get_team_object_permission( user_api_key_auth ) diff --git a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py index bdf4cc312d9..b0ae03c94d3 100644 --- a/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/discoverable_endpoints.py @@ -1,6 +1,6 @@ import json from typing import Optional -from urllib.parse import urlencode, urlparse, urlunparse +from urllib.parse import parse_qsl, urlencode, urlparse, urlunparse from fastapi import APIRouter, Form, HTTPException, Request from fastapi.responses import HTMLResponse, JSONResponse, RedirectResponse @@ -194,7 +194,13 @@ async def authorize_with_server( if code_challenge_method: params["code_challenge_method"] = code_challenge_method - return RedirectResponse(f"{mcp_server.authorization_url}?{urlencode(params)}") + parsed_auth_url = urlparse(mcp_server.authorization_url) + existing_params = dict(parse_qsl(parsed_auth_url.query)) + existing_params.update(params) + final_url = urlunparse( + parsed_auth_url._replace(query=urlencode(existing_params)) + ) + return RedirectResponse(final_url) async def exchange_token_with_server( diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index a3dc876d379..6d494ff6f13 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -72,9 +72,7 @@ try: from mcp.shared.tool_name_validation import ( validate_tool_name, # pyright: ignore[reportAssignmentType] ) - from mcp.shared.tool_name_validation import ( - SEP_986_URL, - ) + from mcp.shared.tool_name_validation import SEP_986_URL except ImportError: from pydantic import BaseModel @@ -675,24 +673,47 @@ class MCPServerManager: return [ server.server_id for server in self.get_registry().values() - if server.allow_all_keys + if server.allow_all_keys is True ] async def get_allowed_mcp_servers( self, user_api_key_auth: Optional[UserAPIKeyAuth] = None ) -> List[str]: """ - Get the allowed MCP Servers for the user + Get the allowed MCP Servers for the user. + + Priority: + 1. If object_permission.mcp_servers is explicitly set, use it (even for admins) + 2. If admin and no object_permission, return all servers + 3. Otherwise, use standard permission checks """ from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view - # If admin, get all servers - if user_api_key_auth and _user_has_admin_view(user_api_key_auth): - return list(self.get_registry().keys()) - allow_all_server_ids = self.get_allow_all_keys_server_ids() try: + # Check if object_permission.mcp_servers is explicitly set + has_explicit_object_permission = False + if user_api_key_auth and user_api_key_auth.object_permission: + # Check if mcp_servers is explicitly set (not None, empty list is valid) + if user_api_key_auth.object_permission.mcp_servers is not None: + has_explicit_object_permission = True + verbose_logger.debug( + f"Object permission mcp_servers explicitly set: {user_api_key_auth.object_permission.mcp_servers}" + ) + + # If admin but NO explicit object permission, get all servers + if ( + user_api_key_auth + and _user_has_admin_view(user_api_key_auth) + and not has_explicit_object_permission + ): + verbose_logger.debug( + "Admin user without explicit object_permission - returning all servers" + ) + return list(self.get_registry().keys()) + + # Get allowed servers from object permissions (respects object_permission even for admins) allowed_mcp_servers = await MCPRequestHandler.get_allowed_mcp_servers( user_api_key_auth ) @@ -2239,6 +2260,7 @@ class MCPServerManager: from litellm.proxy.proxy_server import ( general_settings as proxy_general_settings, ) + return proxy_general_settings except ImportError: # Fallback if proxy_server not available diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 43b388993d9..aed81afd254 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -37,6 +37,7 @@ if MCP_AVAILABLE: from litellm.proxy._experimental.mcp_server.server import ( ListMCPToolsRestAPIResponseObject, MCPServer, + _tool_name_matches, execute_mcp_tool, filter_tools_by_allowed_tools, ) @@ -159,6 +160,7 @@ if MCP_AVAILABLE: server, server_auth_header, raw_headers: Optional[Dict[str, str]] = None, + user_api_key_auth: Optional[UserAPIKeyAuth] = None, ): """Helper function to get tools for a single server.""" tools = await global_mcp_server_manager._get_tools_from_server( @@ -173,6 +175,29 @@ if MCP_AVAILABLE: if server.allowed_tools is not None and len(server.allowed_tools) > 0: tools = filter_tools_by_allowed_tools(tools, server) + # Filter tools based on user_api_key_auth.object_permission.mcp_tool_permissions + # This provides per-key/team/org control over which tools can be accessed + if ( + user_api_key_auth + and user_api_key_auth.object_permission + and user_api_key_auth.object_permission.mcp_tool_permissions + ): + allowed_tools_for_server = ( + user_api_key_auth.object_permission.mcp_tool_permissions.get( + server.server_id + ) + ) + if ( + allowed_tools_for_server is not None + and len(allowed_tools_for_server) > 0 + ): + # Filter tools to only include those in the allowed list + tools = [ + tool + for tool in tools + if _tool_name_matches(tool.name, allowed_tools_for_server) + ] + return _create_tool_response_objects(tools, server.mcp_info) async def _resolve_allowed_mcp_servers_for_tool_call( @@ -197,9 +222,7 @@ if MCP_AVAILABLE: ) allowed_mcp_servers: List[MCPServer] = [] for allowed_server_id in allowed_server_ids_set: - server = global_mcp_server_manager.get_mcp_server_by_id( - allowed_server_id - ) + server = global_mcp_server_manager.get_mcp_server_by_id(allowed_server_id) if server is not None: allowed_mcp_servers.append(server) return allowed_mcp_servers @@ -276,9 +299,7 @@ if MCP_AVAILABLE: "message": f"The key is not allowed to access server {server_id}", }, ) - server = global_mcp_server_manager.get_mcp_server_by_id( - server_id - ) + server = global_mcp_server_manager.get_mcp_server_by_id(server_id) if server is None: return { "tools": [], @@ -292,7 +313,10 @@ if MCP_AVAILABLE: try: list_tools_result = await _get_tools_for_single_server( - server, server_auth_header, raw_headers_from_request + server, + server_auth_header, + raw_headers_from_request, + user_api_key_dict, ) except Exception as e: verbose_logger.exception( @@ -328,7 +352,10 @@ if MCP_AVAILABLE: try: tools_result = await _get_tools_for_single_server( - server, server_auth_header, raw_headers_from_request + server, + server_auth_header, + raw_headers_from_request, + user_api_key_dict, ) list_tools_result.extend(tools_result) except Exception as e: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 06bedfc6c09..87ff4a66e08 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -420,7 +420,8 @@ class LiteLLMRoutes(enum.Enum): "/mcp/tools", "/mcp/tools/list", "/mcp/tools/call", - # Read-only MCP discovery endpoint (virtual keys may be allowed here) + "/mcp-rest/tools/list", + "/mcp-rest/tools/call", "/v1/mcp/server", ] @@ -632,6 +633,9 @@ class LiteLLMRoutes(enum.Enum): "/model/{model_id}/update", "/prompt/list", "/prompt/info", + # Invitation routes - org/team admins checked in endpoint via _user_has_admin_privileges + "/invitation/new", + "/invitation/delete", ] # routes that manage their own allowed/disallowed logic ## Org Admin Routes ## diff --git a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py index 96e7a21cc33..bf3256cf47b 100644 --- a/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py +++ b/litellm/proxy/agent_endpoints/auth/agent_permission_handler.py @@ -21,7 +21,7 @@ class AgentRequestHandler: 1. Key-level agent permissions 2. Team-level agent permissions 3. Agent access group resolution - + Follows the same inheritance logic as MCP: - If team has restrictions and key has restrictions: use intersection - If team has restrictions and key has none: inherit from team @@ -35,7 +35,7 @@ class AgentRequestHandler: ) -> List[str]: """ Get list of allowed agent IDs for the given user/key based on permissions. - + Returns: List[str]: List of allowed agent IDs. Empty list means no restrictions (allow all). """ @@ -45,7 +45,9 @@ class AgentRequestHandler: await AgentRequestHandler._get_allowed_agents_for_key(user_api_key_auth) ) allowed_agents_for_team = ( - await AgentRequestHandler._get_allowed_agents_for_team(user_api_key_auth) + await AgentRequestHandler._get_allowed_agents_for_team( + user_api_key_auth + ) ) # If team has agent restrictions, handle inheritance and intersection logic @@ -73,62 +75,48 @@ class AgentRequestHandler: ) -> bool: """ Check if a specific agent is allowed for the given user/key. - + Args: agent_id: The agent ID to check user_api_key_auth: User authentication info - + Returns: bool: True if agent is allowed, False otherwise """ allowed_agents = await AgentRequestHandler.get_allowed_agents(user_api_key_auth) - + # Empty list means no restrictions - allow all if len(allowed_agents) == 0: return True - + return agent_id in allowed_agents @staticmethod - async def _get_key_object_permission( + def _get_key_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> Optional[LiteLLM_ObjectPermissionTable]: - """Helper to get key object_permission from cache or DB.""" - from litellm.proxy.auth.auth_checks import get_object_permission - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) + """ + Get key object_permission - already loaded by get_key_object() in main auth flow. + Note: object_permission is automatically populated when the key is fetched via + get_key_object() in litellm/proxy/auth/auth_checks.py + """ if not user_api_key_auth: return None - # Already loaded - if user_api_key_auth.object_permission: - return user_api_key_auth.object_permission - - # Need to fetch from DB - if user_api_key_auth.object_permission_id and prisma_client: - return await get_object_permission( - object_permission_id=user_api_key_auth.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - - return None + return user_api_key_auth.object_permission @staticmethod async def _get_team_object_permission( user_api_key_auth: Optional[UserAPIKeyAuth] = None, ) -> Optional[LiteLLM_ObjectPermissionTable]: - """Helper to get team object_permission from cache or DB.""" - from litellm.proxy.auth.auth_checks import ( - get_object_permission, - get_team_object, - ) + """ + Get team object_permission - automatically loaded by get_team_object() in main auth flow. + + Note: object_permission is automatically populated when the team is fetched via + get_team_object() in litellm/proxy/auth/auth_checks.py + """ + from litellm.proxy.auth.auth_checks import get_team_object from litellm.proxy.proxy_server import ( prisma_client, proxy_logging_obj, @@ -138,7 +126,7 @@ class AgentRequestHandler: if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client: return None - # First get the team object (which may have object_permission already loaded) + # Get the team object (which has object_permission already loaded) team_obj: Optional[LiteLLM_TeamTable] = await get_team_object( team_id=user_api_key_auth.team_id, prisma_client=prisma_client, @@ -150,21 +138,7 @@ class AgentRequestHandler: if not team_obj: return None - # Already loaded - if team_obj.object_permission: - return team_obj.object_permission - - # Need to fetch from DB using object_permission_id - if team_obj.object_permission_id: - return await get_object_permission( - object_permission_id=team_obj.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, - ) - - return None + return team_obj.object_permission @staticmethod async def _get_allowed_agents_for_key( @@ -172,31 +146,16 @@ class AgentRequestHandler: ) -> List[str]: """ Get allowed agents for a key from its object_permission. - """ - from litellm.proxy.auth.auth_checks import get_object_permission - from litellm.proxy.proxy_server import ( - prisma_client, - proxy_logging_obj, - user_api_key_cache, - ) + Note: object_permission is already loaded by get_key_object() in main auth flow. + """ if user_api_key_auth is None: return [] - if user_api_key_auth.object_permission_id is None: - return [] - - if prisma_client is None: - verbose_logger.debug("prisma_client is None") - return [] - try: - key_object_permission = await get_object_permission( - object_permission_id=user_api_key_auth.object_permission_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - parent_otel_span=user_api_key_auth.parent_otel_span, - proxy_logging_obj=proxy_logging_obj, + # Get key object permission (already loaded in main auth flow) + key_object_permission = AgentRequestHandler._get_key_object_permission( + user_api_key_auth ) if key_object_permission is None: return [] @@ -205,8 +164,10 @@ class AgentRequestHandler: direct_agents = key_object_permission.agents or [] # Get agents from access groups - access_group_agents = await AgentRequestHandler._get_agents_from_access_groups( - key_object_permission.agent_access_groups or [] + access_group_agents = ( + await AgentRequestHandler._get_agents_from_access_groups( + key_object_permission.agent_access_groups or [] + ) ) # Combine both lists @@ -222,6 +183,8 @@ class AgentRequestHandler: ) -> List[str]: """ Get allowed agents for a team from its object_permission. + + Note: object_permission is already loaded by get_team_object() in main auth flow. """ if user_api_key_auth is None: return [] @@ -230,7 +193,7 @@ class AgentRequestHandler: return [] try: - # Use the helper method that properly handles fetching from DB if needed + # Get team object permission (already loaded in main auth flow) object_permissions = await AgentRequestHandler._get_team_object_permission( user_api_key_auth ) @@ -242,8 +205,10 @@ class AgentRequestHandler: direct_agents = object_permissions.agents or [] # Get agents from access groups - access_group_agents = await AgentRequestHandler._get_agents_from_access_groups( - object_permissions.agent_access_groups or [] + access_group_agents = ( + await AgentRequestHandler._get_agents_from_access_groups( + object_permissions.agent_access_groups or [] + ) ) # Combine both lists @@ -284,9 +249,7 @@ class AgentRequestHandler: for agent in agents: agent_ids.add(agent.agent_id) except Exception as e: - verbose_logger.debug( - f"Error getting agents from access groups: {e}" - ) + verbose_logger.debug(f"Error getting agents from access groups: {e}") return agent_ids @staticmethod @@ -306,16 +269,16 @@ class AgentRequestHandler: ) # Use the helper for DB agents - db_agent_ids = await AgentRequestHandler._get_db_agent_ids_for_access_groups( - prisma_client, access_groups + db_agent_ids = ( + await AgentRequestHandler._get_db_agent_ids_for_access_groups( + prisma_client, access_groups + ) ) agent_ids.update(db_agent_ids) return list(agent_ids) except Exception as e: - verbose_logger.warning( - f"Failed to get agents from access groups: {str(e)}" - ) + verbose_logger.warning(f"Failed to get agents from access groups: {str(e)}") return [] @staticmethod @@ -326,11 +289,15 @@ class AgentRequestHandler: Get list of agent access groups for the given user/key based on permissions. """ access_groups: List[str] = [] - access_groups_for_key = await AgentRequestHandler._get_agent_access_groups_for_key( - user_api_key_auth + access_groups_for_key = ( + await AgentRequestHandler._get_agent_access_groups_for_key( + user_api_key_auth + ) ) - access_groups_for_team = await AgentRequestHandler._get_agent_access_groups_for_team( - user_api_key_auth + access_groups_for_team = ( + await AgentRequestHandler._get_agent_access_groups_for_team( + user_api_key_auth + ) ) # If team has access groups, then key must have a subset of the team's access groups @@ -378,7 +345,9 @@ class AgentRequestHandler: return key_object_permission.agent_access_groups or [] except Exception as e: - verbose_logger.warning(f"Failed to get agent access groups for key: {str(e)}") + verbose_logger.warning( + f"Failed to get agent access groups for key: {str(e)}" + ) return [] @staticmethod @@ -425,4 +394,3 @@ class AgentRequestHandler: f"Failed to get agent access groups for team: {str(e)}" ) return [] - diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index c6093172932..76ec67ab10e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -1368,6 +1368,22 @@ async def _get_team_object_from_user_api_key_cache( raise Exception _response = LiteLLM_TeamTableCachedObj(**response.dict()) + + # Load object_permission if object_permission_id exists but object_permission is not loaded + if _response.object_permission_id and not _response.object_permission: + try: + _response.object_permission = await get_object_permission( + object_permission_id=_response.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: + verbose_proxy_logger.debug( + f"Failed to load object_permission for team {team_id} with object_permission_id={_response.object_permission_id}: {e}" + ) + # save the team object to cache await _cache_team_object( team_id=team_id, @@ -1550,6 +1566,21 @@ async def get_team_object_by_alias( team = teams[0] team_obj = LiteLLM_TeamTableCachedObj(**team.model_dump()) + # Load object_permission if object_permission_id exists but object_permission is not loaded + if team_obj.object_permission_id and not team_obj.object_permission: + try: + team_obj.object_permission = await get_object_permission( + object_permission_id=team_obj.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: + verbose_proxy_logger.debug( + f"Failed to load object_permission for team {team_obj.team_id} with object_permission_id={team_obj.object_permission_id}: {e}" + ) + # Cache the result by both alias and team_id await user_api_key_cache.async_set_cache( key=cache_key, @@ -1838,6 +1869,21 @@ async def get_key_object( _response = UserAPIKeyAuth(**_valid_token.model_dump(exclude_none=True)) + # Load object_permission if object_permission_id exists but object_permission is not loaded + if _response.object_permission_id and not _response.object_permission: + try: + _response.object_permission = await get_object_permission( + object_permission_id=_response.object_permission_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: + verbose_proxy_logger.debug( + f"Failed to load object_permission for key with object_permission_id={_response.object_permission_id}: {e}" + ) + # save the key object to cache await _cache_key_object( hashed_token=hashed_token, diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/categories/harm_toxic_abuse.json b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/categories/harm_toxic_abuse.json new file mode 100644 index 00000000000..9f87cf8097d --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/categories/harm_toxic_abuse.json @@ -0,0 +1,4174 @@ +[ + { + "id": "1man1jar", + "match": "1m1j|1man1jar|1 man 1 jar|one man one jar|one jar one man", + "tags": [ + "shock" + ], + "severity": 4 + }, + { + "id": "2girls1cup", + "match": "2g1c|2girls1cup|2 girls 1 cup|two girls one cup|one cup two girls", + "tags": [ + "shock" + ], + "severity": 4 + }, + { + "id": "acrotomophilia", + "match": "acrotomophilia|acrotomophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "adjective-cock", + "match": "black cock|big cock|huge cock|giant cock|massive cock|throbbing cock", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "alabama-hot-pocket", + "match": "alabama hot pocket", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "alabama-tuna-melt", + "match": "alabama tuna melt", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "alaskan-pipeline", + "match": "alaskan pipeline", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "algophilia", + "match": "algophilia|algophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "anal", + "match": "anal", + "tags": [ + "sexual" + ], + "severity": 2 + }, + { + "id": "anal-assassin", + "match": "anal assassin", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "anal-astronaut", + "match": "anal astronaut", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "anilingus", + "match": "anilingus", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "anus", + "match": "anus", + "tags": [ + "sexual" + ], + "severity": 1, + "exceptions": [ + "m*", + "m*cript", + "m*cripts", + "pand*", + "pand*es", + "tet*", + "tet*es" + ] + }, + { + "id": "apeshit", + "match": "apeshit|ape-shit|ape shit", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "apotemnophilia", + "match": "apotemnophilia|apotemnophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "arse", + "match": "arse", + "tags": [ + "sexual" + ], + "severity": 2, + "exceptions": [ + "*n", + "cath*", + "co*", + "he*", + "ho*", + "kath*", + "m*illes", + "p*", + "s*n" + ] + }, + { + "id": "arsehole", + "match": "arseho*le|ass*ho*le", + "tags": [ + "general" + ], + "severity": 3 + }, + { + "id": "ass", + "match": "ass", + "tags": [ + "sexual" + ], + "severity": 1 + }, + { + "id": "ass-bandit", + "match": "ass bandit|arse bandit", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "autoerotic|auto erotic", + "match": "autoerotic|auto erotic", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "babeland", + "match": "babeland", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "baby-batter", + "match": "baby batter|baby gravy|baby juice|ball batter|ball gravy", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "ball-gag", + "match": "ball gag|ball-gag|ballgag", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "ball-kicking", + "match": "ball kicking|ball-kicking", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "ball-licking", + "match": "ball licking|ball-licking", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "ball-sack", + "match": "ball sack", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "ball-sucking", + "match": "ball sucking|ball-sucking", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "ballcuzi", + "match": "ballcuzi", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "bangbros", + "match": "bangbros|bang bros", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "bangbus", + "match": "bangbus|bang bus", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "bareback", + "match": "bareback", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "barely-legal", + "match": "barely legal", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "bastard", + "match": "ba*sta*rd", + "tags": [ + "general" + ], + "severity": 3 + }, + { + "id": "bastinado", + "match": "bastinado", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "batty-boy", + "match": "batty boy|battyboy|batty boi|battyboi", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "bdsm", + "match": "bdsm", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "bean-flicker", + "match": "bean flicker|bean-flicker|beanflicker", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "bean-queen", + "match": "bean queen", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "beaner", + "match": "beaner|beaners", + "tags": [ + "racial" + ], + "severity": 3, + "exceptions": [ + "*ies", + "*y" + ] + }, + { + "id": "beastiality", + "match": "beastiality|beestiality|bestiality", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "beaver-cleaver", + "match": "beaver cleaver", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "beaver-lips", + "match": "beaver lips", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "bellend", + "match": "bell*end", + "tags": [ + "general" + ], + "severity": 3, + "exceptions": [ + "*en" + ] + }, + { + "id": "bellesa", + "match": "bellesa", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "bicon", + "match": "bicon", + "tags": [ + "lgbtq" + ], + "severity": 3, + "exceptions": [ + "*cave", + "*cavities", + "*cavity", + "*vex", + "*vexities", + "*vexity" + ] + }, + { + "id": "big-boobs", + "match": "big boobs|big breasts|big knockers|big tits", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "birdlock", + "match": "birdlock", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "bitch", + "match": "bi*tch|bi*tches", + "tags": [ + "general" + ], + "severity": 3 + }, + { + "id": "bloody", + "match": "bloody", + "tags": [ + "general" + ], + "severity": 1 + }, + { + "id": "blow-your-load", + "match": "blow your load", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "blowjob", + "match": "blowjob|blow-job|blow job", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "blue-waffle", + "match": "blue waffle|bluewaffle", + "tags": [ + "shock" + ], + "severity": 4 + }, + { + "id": "blumpkin", + "match": "blumpkin", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "bollocks", + "match": "bo*ll*o*cks", + "tags": [ + "sexual" + ], + "severity": 1 + }, + { + "id": "bone-smuggler", + "match": "bone smuggler|bone-smuggler|bonesmuggler", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "boner", + "match": "boner|raging boner|throbbing boner", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "boob", + "match": "boo*b", + "tags": [ + "sexual" + ], + "severity": 1, + "exceptions": [ + "*ird", + "*oo" + ] + }, + { + "id": "booty-buffer", + "match": "booty buffer|booty-buffer", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "booty-call", + "match": "booty call", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "boston-george", + "match": "boston george", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "breasts", + "match": "breasts", + "tags": [ + "sexual" + ], + "severity": 1, + "exceptions": [ + "*troke" + ] + }, + { + "id": "brown-piper", + "match": "brown piper|brown-piper|brownpiper", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "brown-shower", + "match": "brown shower|brown showers", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "brownie-royalty", + "match": "brownie king|brownie queen", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "buddhahead", + "match": "buddhahead|buddha head|buddha-head", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "bufter", + "match": "bufter|bufty", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "bugger", + "match": "bugg*er", + "tags": [ + "general" + ], + "severity": 1, + "exceptions": [ + "de*", + "hum*" + ] + }, + { + "id": "bukkake", + "match": "bukkake", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "bulldyke", + "match": "bulldyke", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "bullet-vibe", + "match": "bullet vibe|bullet vibrator", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "bullshit", + "match": "bullshit|bull-shit|bull shit", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "bum-chum", + "match": "bum chum|bum-chum|bumchum", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "bum-driller", + "match": "bum driller|bum-driller|bumdriller", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "butt-boy", + "match": "butt boy|butt-boy|buttboy|bum boy|bum-boy|bumboy", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "butt-pilot", + "match": "butt pilot|bum pilot", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "butt-pirate", + "match": "butt pirate|butt-pirate|bum pirate|bum-pirate", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "butt-rider", + "match": "butt rider|buttrider|bum rider|bumrider", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "butt-robber", + "match": "butt robber|butt-robber|buttrobber|bum robber|bum-robber|bumrobber", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "butt-rustler", + "match": "butt rustler|bum rustler", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "butthole-engineer", + "match": "butthole engineer|bumhole engineer", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "camel-jockey", + "match": "camel jockey|cameljockey|camel jockies|cameljockies", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "camel-toe", + "match": "camel toe", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "canadian-porch-swing", + "match": "canadian porch swing", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "carpet-muncher", + "match": "carpet muncher|carpetmuncher", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "cheese-monkey", + "match": "cheese-eating surrender monkey|cheese eating surrender monkey", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "chi-chi-man", + "match": "chi chi man|chi-chi man", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "chicken-queen", + "match": "chicken queen", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "chinaman", + "match": "chinaman|china man|chinamen|china men", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "ching-chong", + "match": "ching-chong|ching chong", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "chink", + "match": "chink|chinks|chinky", + "tags": [ + "racial" + ], + "severity": 3, + "exceptions": [ + "*apin", + "*apins", + "*ed", + "*ier", + "*iest", + "*ing", + "pa*o", + "pa*os" + ] + }, + { + "id": "chocolate-rosebud", + "match": "chocolate rosebud|chocolate rosebuds", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "cholerophilia", + "match": "cholerophilia|cholerophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "christ", + "match": "chri*st", + "tags": [ + "religious" + ], + "severity": 1, + "exceptions": [ + "*en", + "*ian", + "*ie", + "*y" + ] + }, + { + "id": "cialis", + "match": "cialis", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "circlejerk", + "match": "circlejerk|circle-jerk", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "cishet", + "match": "cishet", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "cissy", + "match": "cissy|cissie", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "claustrophilia", + "match": "claustrophilia|claustrophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "cleveland-accordion", + "match": "cleveland accordion", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "cleveland-hot-waffle", + "match": "cleveland hot waffle", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "cleveland-steamer", + "match": "cleveland steamer", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "clit", + "match": "clit", + "tags": [ + "sexual" + ], + "severity": 3, + "exceptions": [ + "*ella", + "*ellum", + "*ic", + "*icize", + "*icized", + "*icizes", + "*icizing", + "*ics", + "ana*ic", + "cy*ol", + "cy*ols", + "en*ic", + "en*ics", + "hetero*e", + "hetero*es", + "pa*axel", + "pa*axels", + "pro*ic", + "pro*ics" + ] + }, + { + "id": "clitoris", + "match": "cli*tori*s", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "clover-clamps", + "match": "clover clamps|clover clamp", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "clunge", + "match": "clu*nge", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "clusterfuck", + "match": "clusterfuck|cluster-fuck|cluster fuck", + "tags": [ + "general" + ], + "severity": 4 + }, + { + "id": "cock", + "match": "cock", + "tags": [ + "sexual" + ], + "severity": 2, + "exceptions": [ + "*ade", + "*aded", + "*ades", + "*alorum", + "*alorums", + "*amamie", + "*amamy", + "*apoo", + "*apoos", + "*ateel", + "*ateels", + "*atiel", + "*atiels", + "*atoo", + "*atoos", + "*atrice", + "*atrices", + "*bill", + "*billed", + "*billing", + "*bills", + "*boat", + "*boats", + "*chafer", + "*chafers", + "*crow", + "*crows", + "*ed", + "*er", + "*ered", + "*erel", + "*erels", + "*ering", + "*ers", + "*eye", + "*eyed", + "*eyedly", + "*eyedness", + "*eyednesses", + "*eyes", + "*horse", + "*horses", + "*ier", + "*iest", + "*ily", + "*iness", + "*inesses", + "*ing", + "*ish", + "*le", + "*lebur", + "*leburs", + "*led", + "*les", + "*leshell", + "*leshells", + "*like", + "*ling", + "*loft", + "*lofts", + "*ney", + "*neyfied", + "*neyfies", + "*neyfy", + "*neyfying", + "*neyish", + "*neyism", + "*neyisms", + "*neys", + "*pit", + "*pits", + "*roach", + "*roaches", + "*s", + "*scomb", + "*scombs", + "*sfoot", + "*sfoots", + "*shies", + "*shut", + "*shuts", + "*shy", + "*spur", + "*spurs", + "*sucker", + "*suckers", + "*sure", + "*surely", + "*sureness", + "*surenesses", + "*swain", + "*swains", + "*tail", + "*tailed", + "*tailing", + "*tails", + "*up", + "*ups", + "*y", + "a*", + "baw*", + "baw*s", + "bib*", + "bib*s", + "billy*", + "billy*s", + "black*", + "black*s", + "cold*", + "cold*ed", + "cold*ing", + "cold*s", + "game*", + "game*s", + "gor*", + "gor*s", + "hay*", + "hay*s", + "moor*", + "moor*s", + "pea*", + "pea*ed", + "pea*ier", + "pea*iest", + "pea*ing", + "pea*ish", + "pea*s", + "pea*y", + "pet*", + "pet*s", + "pinch*", + "pinch*s", + "poppy*", + "poppy*s", + "prin*", + "prin*s", + "re*", + "re*ed", + "re*ing", + "re*s", + "sea*", + "sea*s", + "shuttle*", + "shuttle*ed", + "shuttle*ing", + "shuttle*s", + "stop*", + "stop*s", + "un*", + "un*ed", + "un*ing", + "un*s", + "weather*", + "weather*s", + "wood*", + "wood*s" + ] + }, + { + "id": "cockpipe-cosmonaut", + "match": "cockpipe cosmonaut", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "cockstruction-worker", + "match": "cockstruction worker", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "coimetrophilia", + "match": "coimetrophilia|coimetrophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "collaring", + "match": "collaring|collared", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "coon", + "match": "coon|coons", + "tags": [ + "racial" + ], + "severity": 3, + "exceptions": [ + "barra*", + "coc*", + "coc*ed", + "puc*", + "rac*", + "ty*", + "*tie", + "* can", + "*can", + "* hound", + "*hound", + "* skin", + "*skin" + ] + }, + { + "id": "coprophilia", + "match": "coprophilia|coprophile|coprolagnia", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "cornhole", + "match": "cornhole", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "crafty-butcher", + "match": "crafty butcher", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "crap", + "match": "cra*p", + "tags": [ + "general" + ], + "severity": 1, + "exceptions": [ + "*e", + "*ing", + "*shoot", + "*ulent", + "*ulous", + "s*" + ] + }, + { + "id": "creampie", + "match": "creampie|cream-pie", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "cum", + "match": "cum", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "cumming", + "match": "cumming", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "cumshot", + "match": "cumshot|cumshots|cum shot|cum shots", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "cunnilingus", + "match": "cunnilingus", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "cunt", + "match": "cu*n*t|ku*nt|cunts|kunts", + "tags": [ + "general" + ], + "severity": 4, + "exceptions": [ + "s*horpe" + ] + }, + { + "id": "cuntboy", + "match": "cuntboy|cunt-boy|cunt boy", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "curry-muncher", + "match": "curry muncher|curry-muncher|currymuncher", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "damn", + "match": "da*mn", + "tags": [ + "religious" + ], + "severity": 1 + }, + { + "id": "darkie", + "match": "darkie|darkies|darky|darkey", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "date-rape", + "match": "date rape|daterape", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "ddlg", + "match": "ddlg", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "deep-throat", + "match": "deep throat|deep-throat|deepthroat", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "dendrophilia", + "match": "dendrophilia|dendrophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "dick", + "match": "di*ck", + "tags": [ + "sexual" + ], + "severity": 2, + "exceptions": [ + "*cissel", + "*en", + "*er", + "bene*", + "me*", + "zad*" + ] + }, + { + "id": "dickgirl", + "match": "dickgirl|dick-girl|dick girl", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "dildo", + "match": "dildo|dildos", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "dingleberry", + "match": "dingleberry|dingleberries", + "tags": [ + "general" + ], + "severity": 1 + }, + { + "id": "dipsea", + "match": "dipsea", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "dirty-pillows", + "match": "dirty pillows", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "dirty-sanchez", + "match": "dirty sanchez", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "dishabiliophilia", + "match": "dishabiliophilia|dishabiliophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "doggie-style", + "match": "doggie style|doggie-style|doggiestyle|doggy style|doggy-style|doggystyle|dog style", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "dogshit", + "match": "dogshit|dog-shit|dog shit", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "dolcett", + "match": "dolcett", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "domination", + "match": "domination", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "dominatrix", + "match": "dominatrix|domme|dommes", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "donkey-punch", + "match": "donkey punch", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "donut-muncher", + "match": "donut muncher", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "donut-puncher", + "match": "donut puncher", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "double-penetration", + "match": "double penetration|dp action", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "dry-hump", + "match": "dry hump", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "dune-coon", + "match": "dune coon|dune-coon|doon coon|dooncoon", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "dutch-rudder", + "match": "dutch rudder", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "dyke", + "match": "dyke", + "tags": [ + "lgbtq" + ], + "severity": 3, + "exceptions": [ + "van*", + "van*d", + "van*s" + ] + }, + { + "id": "dystychiphilia", + "match": "dystychiphilia|dystychiphile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "edgeplay", + "match": "edgeplay|edge play", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "ejaculate", + "match": "ejaculate|ejaculation|ejaculating|ejaculated", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "electro-play", + "match": "electro-play|electroplay", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "emetophilia", + "match": "emetophilia|emetophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "enby", + "match": "enby", + "tags": [ + "lgbtq" + ], + "severity": 3, + "exceptions": [ + "all*", + "as*", + "brook*", + "ca*", + "d*", + "froz*te", + "gat*", + "hold*", + "lack*", + "laz*", + "nav*", + "ott*", + "r*gda", + "t*", + "warr*", + "wh*" + ] + }, + { + "id": "eskimo-trebuchet", + "match": "eskimo trebuchet", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "eyetie", + "match": "eyetie|eye-tie", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "fag", + "match": "fa*g", + "tags": [ + "lgbtq" + ], + "severity": 3, + "exceptions": [ + "lea*e", + "lea*es", + "ser*e", + "ser*es", + "whar*e", + "whar*es" + ] + }, + { + "id": "fag-bomb", + "match": "fag bomb|fag-bomb|fagbomb", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "faggot", + "match": "fa*gg*o*t|fagot", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "felch", + "match": "felch|felching", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "fellatio", + "match": "fellatio|fellating", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "female-squirting", + "match": "female squirting", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "figging", + "match": "figging", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "fingerbang", + "match": "fingerbang|finger bang|fingerbanging", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "fingering", + "match": "fingering|fingered", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "finocchio", + "match": "finocchio|finochio|finoccio", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "fisting", + "match": "fisting|fisted", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "footjob", + "match": "footjob|foot-job|foot job", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "french-rudder", + "match": "french rudder", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "frogeater", + "match": "frogeater|frog-eater|frog eater", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "frolicme", + "match": "frolicme|frolic me", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "frotting", + "match": "frotting|frottage", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "fuck", + "match": "fu*c*k|fucks|fu*c*ken|fu*c*ker|fu*ckers|fu*c*kin|fu*c*king", + "tags": [ + "general" + ], + "severity": 4 + }, + { + "id": "fuckhead", + "match": "fuckhead|fuckheads", + "tags": [ + "general" + ], + "severity": 4 + }, + { + "id": "fucktard", + "match": "fucktard|fucktards", + "tags": [ + "general" + ], + "severity": 4 + }, + { + "id": "fuckwad", + "match": "fuckwad|fuckwads", + "tags": [ + "general" + ], + "severity": 4 + }, + { + "id": "fuckwit", + "match": "fuckwit|fuckwits|fuckwhit|fuck-wit", + "tags": [ + "general" + ], + "severity": 4 + }, + { + "id": "fudge-packer", + "match": "fudge packer|fudge-packer|fudgepacker", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "futanari", + "match": "futanari", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "g-spot", + "match": "g-spot", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "gangbang", + "match": "gangbang|gang bang", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "gay-sex", + "match": "gay sex", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "gaysian", + "match": "gaysian", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "genitals", + "match": "genitals", + "tags": [ + "sexual" + ], + "severity": 1 + }, + { + "id": "genitorture", + "match": "genitorture", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "gerontophilia", + "match": "gerontophilia|gerontophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "gin-jockey", + "match": "gin jockey|gin jocky", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "girl-on-top", + "match": "girl on top", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "goatse", + "match": "goatse|goatcx", + "tags": [ + "shock" + ], + "severity": 4 + }, + { + "id": "god-damn", + "match": "god damn|god-damn|goddamn|god damned|god-damned|goddamned", + "tags": [ + "religious" + ], + "severity": 2 + }, + { + "id": "gokkun", + "match": "gokkun|go-kun", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "golden-shower", + "match": "golden shower|golden showers|yellow shower|yellow showers", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "golliwog", + "match": "golliwog|gollywog", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "gook", + "match": "gook|gooks|gook-eye|gooky|gookie", + "tags": [ + "racial" + ], + "severity": 3, + "exceptions": [ + "gobblede*", + "gobblede*s", + "gobbledy*", + "gobbledy*s" + ] + }, + { + "id": "goregasm", + "match": "goregasm", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "greaseball", + "match": "greaseball", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "grey-queen", + "match": "grey queen|gray queen", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "grope", + "match": "grope", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "group-sex", + "match": "group sex", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "gym-bunny", + "match": "gym bunny|gymbunny", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "hajji", + "match": "haji|hajj*i|hadji", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "hand-job", + "match": "hand job|hand-job|handjob", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "hell", + "match": "hell", + "tags": [ + "religious" + ], + "severity": 1, + "exceptions": [ + "*cat", + "*er", + "*fire", + "*kite", + "*ed", + "*en", + "*ion", + "*ish", + "*o", + "ec*e", + "p*", + "s*" + ] + }, + { + "id": "hermie", + "match": "hermie", + "tags": [ + "lgbtq" + ], + "severity": 3, + "exceptions": [ + "diat*s", + "endot*s", + "homeot*s" + ] + }, + { + "id": "hickory-switch", + "match": "hickory switch", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "hippophilia", + "match": "hippophilia|hippophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "homoerotic", + "match": "homoerotic", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "honkey", + "match": "honkey|honky|honkeys|honkies", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "horny", + "match": "horny", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "horseshit", + "match": "horseshit|horse-shit|horse shit", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "hot-carl", + "match": "hot carl", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "hot-richard", + "match": "hot richard", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "humping", + "match": "humping", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "hymie", + "match": "hymie|heimie", + "tags": [ + "racial" + ], + "severity": 2, + "exceptions": [ + "alc*s", + "prings*ella", + "t*r", + "t*st" + ] + }, + { + "id": "impact-play", + "match": "impact play|impact-play", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "incest", + "match": "incest", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "intercourse", + "match": "intercourse", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "jack-off", + "match": "jack off|jack-off", + "tags": [ + "sexual" + ], + "severity": 2 + }, + { + "id": "jail-bait", + "match": "jail bait|jailbait", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "jap", + "match": "jap", + "tags": [ + "racial" + ], + "severity": 1, + "exceptions": [ + "*an", + "*anize", + "*anized", + "*anizes", + "*anizing", + "*anned", + "*anners", + "*anner", + "*anning", + "*ans", + "*e", + "*ed", + "*er", + "*eries", + "*ers", + "*ery", + "*es", + "*ing", + "*ingly", + "*onica", + "*onicas", + "*onaiserie", + "*onaiseries", + "ca*ut", + "ca*uts", + "jipi*a", + "jipi*as" + ] + }, + { + "id": "jelly-donut", + "match": "jelly donut", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "jerk-off", + "match": "jerk off|jerk-off", + "tags": [ + "sexual" + ], + "severity": 2 + }, + { + "id": "jerkmate", + "match": "jerkmate|jerk mate", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "jesus", + "match": "je*su*s", + "tags": [ + "religious" + ], + "severity": 1, + "exceptions": [ + "be*" + ] + }, + { + "id": "jesus-christ", + "match": "je*su*s chri*st", + "tags": [ + "religious" + ], + "severity": 2 + }, + { + "id": "jigaboo", + "match": "jig*aboo*|jig*gerboo*", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "jizz", + "match": "jizz", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "juggs", + "match": "juggs", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "jungle-bunny", + "match": "jungle bunny|junglebunny", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "kennebunkport-surprise", + "match": "kennebunkport surprise", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "kentucky-klondike", + "match": "kentucky klondike", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "kentucky-tractor-puller", + "match": "kentucky tractor puller", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "kike", + "match": "ki*ke", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "kinbaku", + "match": "kinbaku", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "kitty-puncher", + "match": "kitty puncher|kitty-puncher|kittypuncher", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "knobbing", + "match": "knobbing", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "kraut", + "match": "kraut|krauts", + "tags": [ + "racial" + ], + "severity": 1, + "exceptions": [ + "sauer*s", + "sauer*" + ] + }, + { + "id": "kynophilia", + "match": "kynophilia|kynophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "lady-boy", + "match": "lady boy|lady-boy|ladyboy", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "leather-restraint", + "match": "leather restraint", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "leather-straight-jacket", + "match": "leather straight jacket", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "lemon-party", + "match": "lemon party|lemonparty", + "tags": [ + "shock" + ], + "severity": 4 + }, + { + "id": "leningrad-steamer", + "match": "leningrad steamer", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "lesbo", + "match": "lesbo", + "tags": [ + "lgbtq" + ], + "severity": 3, + "exceptions": [ + "*s", + "b*k", + "b*ks" + ] + }, + { + "id": "leso", + "match": "leso", + "tags": [ + "lgbtq" + ], + "severity": 3, + "exceptions": [ + "bung*me", + "chuck*me", + "crad*ng", + "crad*ngs", + "cudd*me", + "do*me", + "medd*me", + "medd*meness", + "mett*me", + "nett*me", + "troub*me", + "troub*mely", + "troub*meness", + "unwho*me", + "unwho*mely", + "who*me", + "who*mer", + "who*mest", + "who*mely", + "who*meness", + "who*menesses" + ] + }, + { + "id": "lezzie", + "match": "lezzie|lezzies", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "light-fedora", + "match": "light in the fedora", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "light-loafers", + "match": "light in the loafers", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "light-pants", + "match": "light in the pants", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "limp-wristed", + "match": "limp wristed|limp-wristed", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "literotica", + "match": "literotica", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "lovemaking", + "match": "lovemaking", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "male-squirting", + "match": "male squirting|male-squirting", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "masturbate", + "match": "ma*stu*rbate|ma*stu*rb8|masturbation|masturbating|masterbate|masterb8", + "tags": [ + "sexual" + ], + "severity": 2 + }, + { + "id": "mayonnaise-monkey", + "match": "mayonnaise monkey|mayonnaise monkies", + "tags": [ + "racial" + ], + "severity": 1 + }, + { + "id": "mdlb", + "match": "mdlb", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "meat-masseuse", + "match": "meat masseuse", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "meatspin", + "match": "meatspin|meat spin", + "tags": [ + "shock" + ], + "severity": 4 + }, + { + "id": "menage-a-trois", + "match": "menage a trois|menage-a-trois|menages a trois|menages-a-trois", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "menophilia", + "match": "menophilia|menophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "mexican-pancake", + "match": "mexican pancake", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "milwaukee-blizzard", + "match": "milwaukee blizzard", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "missionary-position", + "match": "missionary position", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "mississippi-birdbath", + "match": "mississippi birdbath", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "mr-hands", + "match": "mr. hands|mr hands|mrhands", + "tags": [ + "shock" + ], + "severity": 4 + }, + { + "id": "muff-diver", + "match": "muff diver|muff-diver|muffdiver", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "muffdiving", + "match": "muffdiving|muff diving|muffdiver|muff diver", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "muscle-mary", + "match": "muscle mary", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "mvtube", + "match": "mvtube", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "nambla", + "match": "nambla", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "necrophilia", + "match": "necrophilia|necrophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "negro", + "match": "negro", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "neonazi", + "match": "neonazi|neo-nazi|neo nazi", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "nigerian-hurricane", + "match": "nigerian hurricane", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "nigga", + "match": "ni*gg*a|ni*gg*s|nignog|nig nog", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "nigger", + "match": "ni*gg*e*r", + "tags": [ + "racial" + ], + "severity": 4 + }, + { + "id": "nipple-clamps", + "match": "nipple clamps|nipple clamp", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "nipples", + "match": "nipples|nipple", + "tags": [ + "sexual" + ], + "severity": 1 + }, + { + "id": "nude", + "match": "nude|nudity", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "nutten", + "match": "nutten", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "nymphomania", + "match": "nymphomania|nymphomaniac|nympho|nimphomania|nimphomaniac|nimpho", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "octopussy", + "match": "octopussy", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "oklahomo", + "match": "oklahomo", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "omorashi", + "match": "omorashi", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "onlyfans", + "match": "onlyfans|only fans", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "orgasm", + "match": "orgasm|orgasms|orgasmic", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "paedophilia|pedophilia", + "match": "paedophilia|pedophilia|paedophile|pedophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "painslut", + "match": "painslut|pain slut", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "paki", + "match": "paki", + "tags": [ + "racial" + ], + "severity": 3, + "exceptions": [ + "*hi" + ] + }, + { + "id": "panamanian-petting-zoo", + "match": "panamanian petting zoo", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "pansy", + "match": "pansy", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "panties", + "match": "panties", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "parthenophilia", + "match": "parthenophilia|parthenophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "pedobear", + "match": "pedobear|paedobear|pedo bear|paedo bear", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "pegging", + "match": "pegging", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "penis", + "match": "pe*ni*s", + "tags": [ + "sexual" + ], + "severity": 1, + "exceptions": [ + "top*h" + ] + }, + { + "id": "peterpuffer", + "match": "peterpuffer|peter-puffer|peter puffer", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "petrol-sniffer", + "match": "petrol sniffer|petrol-sniffer|petrolsniffer", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "phagophilia", + "match": "phagophilia|phagophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "piece-of-shit", + "match": "piece of shit|pieces of shit", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "pikey", + "match": "pikey|pikeys", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "piss-off", + "match": "pi*ss* off", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "piss-pig", + "match": "piss pig|pisspig", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "piss-pig", + "match": "piss pig|pisspig", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "pissed-off", + "match": "pi*ss*ed off", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "pissing", + "match": "pissing", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "playboy", + "match": "playboy", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "pleasure-chest", + "match": "pleasure chest", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "pnigerophilia|pnigophilia", + "match": "pnigerophilia|pnigophilia|pnigerophile|pnigophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "poinephilia", + "match": "poinephilia|poinephile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "ponyboy", + "match": "ponyboy|pony-boy|pony boy", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "ponygirl", + "match": "ponygirl|pony-girl|pony girl", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "ponyplay", + "match": "ponyplay|pony-play", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "poof", + "match": "poof", + "tags": [ + "lgbtq" + ], + "severity": 3, + "exceptions": [ + "*s", + "*tah", + "*tahs", + "*ter", + "*ters", + "*y", + "s*", + "s*ed", + "s*er", + "s*eries", + "s*ers", + "s*ery", + "s*ing", + "s*s", + "s*y" + ] + }, + { + "id": "poon", + "match": "poon|poontang", + "tags": [ + "sexual" + ], + "severity": 3, + "exceptions": [ + "cram*", + "cram*s", + "desserts*", + "desserts*ful", + "desserts*s", + "har*", + "har*ed", + "har*er", + "har*ers", + "har*ing", + "har*s", + "lam*", + "lam*ed", + "lam*er", + "lam*eries", + "lam*ers", + "lam*ery", + "lam*ing", + "lam*s", + "s*", + "s*bill", + "s*bills", + "s*ed", + "s*erism", + "s*erisms", + "s*ey", + "s*eys", + "s*ful", + "s*fuls", + "s*ier", + "s*ies", + "s*iest", + "s*ily", + "s*ing", + "s*s", + "s*sful", + "s*y", + "soups*", + "soups*s", + "tables*", + "tables*ful", + "tables*fuls", + "tables*s", + "tables*sful", + "teas*", + "teas*ful", + "teas*fuls", + "teas*s", + "teas*sful" + ] + }, + { + "id": "poop-chute", + "match": "poop chute|poopchute", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "pornhub", + "match": "pornhub|porn hub", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "pornography", + "match": "pornography|pornographic|porno|pornos|porn", + "tags": [ + "sexual" + ], + "severity": 2 + }, + { + "id": "potato-queen", + "match": "potato queen", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "prince-albert-piercing", + "match": "prince albert piercing", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "proctophilia", + "match": "proctophilia|proctophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "pubes", + "match": "pubes", + "tags": [ + "sexual" + ], + "severity": 1 + }, + { + "id": "punani", + "match": "pu*na*ni|punany", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "pussy", + "match": "pu*ss*y", + "tags": [ + "general" + ], + "severity": 3 + }, + { + "id": "pussy-puncher", + "match": "pussy puncher|pussy-puncher|pussypuncher", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "queef", + "match": "queef|queaf", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "quim", + "match": "quim", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "raghead", + "match": "raghead|rag head|ragheads|rag heads", + "tags": [ + "religious" + ], + "severity": 3 + }, + { + "id": "ramen-yarmulke", + "match": "ramen yarmulke", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "rape", + "match": "rape", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "raping", + "match": "raping|rapist", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "rectum", + "match": "rectum", + "tags": [ + "sexual" + ], + "severity": 1 + }, + { + "id": "retard", + "match": "retard|retarded", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "reverse-cowgirl", + "match": "reverse cowgirl", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "rhabdophilia", + "match": "rhabdophilia|rhabdophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "rhypophilia", + "match": "rhypophilia|rhypophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "rice-queen", + "match": "rice queen", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "rimjob", + "match": "rimjob|rimming", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "ring-raider", + "match": "ring raider|ringraider", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "rusty-trombone", + "match": "rusty trombone", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "sand-nigger", + "match": "sand ni*gg*e*r|sand-ni*gg*e*r|sandni*gg*e*r", + "tags": [ + "racial" + ], + "severity": 4 + }, + { + "id": "santorum", + "match": "santorum", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "scatophilia", + "match": "scatophilia|scatophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "schlong", + "match": "schlong|shlong", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "scissoring", + "match": "scissoring", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "semen", + "match": "semen", + "tags": [ + "sexual" + ], + "severity": 1, + "exceptions": [ + "aba*t", + "aba*ts", + "adverti*t", + "adverti*ts", + "advi*t", + "advi*ts", + "amu*t", + "amu*ts", + "appea*t", + "appea*ts", + "apprai*t", + "apprai*ts", + "arrondis*t", + "arrondis*ts", + "ba*", + "ba*t", + "ba*tless", + "ba*ts", + "bemu*t", + "bemu*ts", + "boulever*t", + "boulever*ts", + "ca*t", + "ca*ts", + "chasti*t", + "chasti*ts", + "deba*t", + "deba*ts", + "defen*", + "despi*t", + "despi*ts", + "disbur*t", + "disbur*ts", + "disgui*t", + "disgui*ts", + "divertis*t", + "divertis*ts", + "ea*t", + "ea*ts", + "eclaircis*t", + "empres*t", + "empres*ts", + "enca*t", + "enca*ts", + "endor*t", + "endor*ts", + "enfranchi*t", + "exci*", + "hor*", + "hou*", + "indor*t", + "indor*ts", + "pas*terie", + "pas*teries", + "reimbur*t", + "reimbur*ts", + "rou*t", + "rou*ts", + "subba*t", + "subba*ts", + "ver*", + "warehou*" + ] + }, + { + "id": "seplophilia", + "match": "seplophilia|seplophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "sex", + "match": "sex", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "shaved pussy", + "match": "shaved pussy|shaved beaver", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "sheepshagger", + "match": "sheepshagger|sheep shagger", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "shemale", + "match": "shemale|she-male|she male", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "shibari", + "match": "shibari", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "shit", + "match": "sh*i*t", + "tags": [ + "general" + ], + "severity": 2, + "exceptions": [ + "*ake", + "mi*", + "*tah", + "*tim" + ] + }, + { + "id": "shithead", + "match": "shithead|shit head", + "tags": [ + "general" + ], + "severity": 3 + }, + { + "id": "shitty", + "match": "shi*tt*y", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "shota", + "match": "shota", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "shrimping", + "match": "shrimping", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "sissy", + "match": "sissy", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "skeet", + "match": "skeet", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "skittles-harvest", + "match": "skittles harvest|skittle harvest", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "slanteye", + "match": "slanteye|slant-eye|slant eye", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "snatch", + "match": "snatch", + "tags": [ + "sexual" + ], + "severity": 1 + }, + { + "id": "snowballing", + "match": "snowballing", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "sod-off", + "match": "sod off", + "tags": [ + "general" + ], + "severity": 1 + }, + { + "id": "sodding", + "match": "sodding", + "tags": [ + "general" + ], + "severity": 1 + }, + { + "id": "sodomize", + "match": "sodomize|sodomise|sodomist|sodomy", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "spastic", + "match": "spastic", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "spearchucker", + "match": "spearchucker", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "spic", + "match": "spic|spics|spick|spicks", + "tags": [ + "racial" + ], + "severity": 3, + "exceptions": [ + "*a", + "*ae", + "*as", + "*ate", + "*ated", + "*cato", + "*catos", + "*e", + "*ebush", + "*ebushes", + "*ed", + "*eless", + "*er", + "*eries", + "*ers", + "*ery", + "*es", + "*ey", + "*ier", + "*iest", + "*ily", + "*iness", + "*inesses", + "*ing", + "*ks", + "*ula", + "*ulae", + "*ular", + "*ulate", + "*ulation", + "*ulations", + "*ule", + "*ules", + "*ulum", + "*y", + "a*", + "a*s", + "all*e", + "all*es", + "aru*es", + "au*ate", + "au*ated", + "au*ates", + "au*ating", + "au*e", + "au*es", + "au*ious", + "au*iously", + "au*iousness", + "con*uities", + "con*uity", + "con*uous", + "con*uously", + "con*uousness", + "de*able", + "de*ableness", + "de*ably", + "haru*ation", + "haru*ations", + "haru*es", + "ho*e", + "ho*es", + "inau*ious", + "inau*iously", + "incon*uous", + "incon*uously", + "mi*kel", + "mi*kels", + "over*e", + "over*ed", + "over*es", + "over*ing", + "oversu*ious", + "per*acious", + "per*aciously", + "per*acities", + "per*acity", + "per*uities", + "per*uity", + "per*uous", + "per*uously", + "per*uousness", + "su*ion", + "su*ioned", + "su*ioning", + "su*ions", + "su*ious", + "su*iously", + "su*iousness", + "tran*uous", + "unsu*ious" + ] + }, + { + "id": "spicy-gringo", + "match": "spicy gringo", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "splooge", + "match": "splooge|splooge moose|spooge", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "spunk", + "match": "spunk", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "strap-on", + "match": "strap on|strap-on", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "strap-on", + "match": "strap-on|strapon", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "strappado", + "match": "strappado", + "tags": [ + "sexual" + ], + "severity": 4 + }, + { + "id": "swamp-guinea", + "match": "swamp guinea|swamp-guinea", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "swastika", + "match": "swastika|svastika|suastika", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "switch-hitter", + "match": "switch hitter", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "t-girl", + "match": "t-girl|tgirl", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "taphephilia", + "match": "taphephilia|taphephile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "tea-bagging", + "match": "tea bagging|tea-bagging|tea bagged|tea-bagged", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "thanatophilia", + "match": "thanatophilia|thanatophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "threesome", + "match": "threesome", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "throating", + "match": "throating", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "thumbzilla", + "match": "thumbzilla", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "timber-nigger", + "match": "timber ni*gg*e*r|timber-ni*gg*e*r|timberni*gg*e*r", + "tags": [ + "racial" + ], + "severity": 4 + }, + { + "id": "tits", + "match": "tits", + "tags": [ + "sexual" + ], + "severity": 2, + "exceptions": [ + "bush*", + "pas*", + "tom*" + ] + }, + { + "id": "titty", + "match": "titt*y|titt*ies", + "tags": [ + "sexual" + ], + "severity": 2 + }, + { + "id": "topless", + "match": "topless", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "tosser", + "match": "tosser", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "towelhead", + "match": "towelhead|towel-head|towel head", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "tranny", + "match": "tranny|trannie", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "transbian", + "match": "transbian", + "tags": [ + "lgbtq" + ], + "severity": 3 + }, + { + "id": "traumatophilia", + "match": "traumatophilia|traumatophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "tribbing", + "match": "tribbing|tribadism", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "tubgirl", + "match": "tubgirl|tub girl", + "tags": [ + "shock" + ], + "severity": 4 + }, + { + "id": "twat", + "match": "twa*t", + "tags": [ + "general" + ], + "severity": 3, + "exceptions": [ + "cu*er", + "mel*er", + "ou*ch", + "sal*er", + "wris*ch" + ] + }, + { + "id": "twink", + "match": "twink", + "tags": [ + "lgbtq" + ], + "severity": 3, + "exceptions": [ + "*ie", + "*ies", + "*le", + "*led", + "*ler", + "*les", + "*lers", + "*ling", + "*lings", + "*ly" + ] + }, + { + "id": "urethra-play", + "match": "urethra play", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "urophilia", + "match": "urophilia|urophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "vagina", + "match": "vagina", + "tags": [ + "sexual" + ], + "severity": 1 + }, + { + "id": "venus-mound", + "match": "venus mound|mound of venus", + "tags": [ + "sexual" + ], + "severity": 1 + }, + { + "id": "viagra", + "match": "viagra", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "vibrator", + "match": "vibrator", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "violet-wand", + "match": "violet wand", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "vorarephilia", + "match": "vorarephilia|vorarephile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "voyeurweb", + "match": "voyeurweb", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "wagon-burner", + "match": "wagon burner|wagon-burner", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "wank", + "match": "wa*nk", + "tags": [ + "sexual" + ], + "severity": 2, + "exceptions": [ + "s*", + "t*" + ] + }, + { + "id": "wanker", + "match": "wa*nker", + "tags": [ + "general" + ], + "severity": 2 + }, + { + "id": "wax-play", + "match": "wax play|wax-play", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "wet-dream", + "match": "wet dream", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "wetback", + "match": "wetback|wet-back|wet back", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "white power", + "match": "whitepower|white-power|white power", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "whore", + "match": "who*re", + "tags": [ + "general" + ], + "severity": 3 + }, + { + "id": "wigger", + "match": "wigger|whigger|wigga", + "tags": [ + "racial" + ], + "severity": 2 + }, + { + "id": "wiitwd", + "match": "wiitwd", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "wog", + "match": "wog|wogs", + "tags": [ + "racial" + ], + "severity": 1, + "exceptions": [ + "horns*gle", + "horns*gled", + "horns*gles", + "horns*gling", + "polli*", + "polli*s", + "polly*", + "polly*s" + ] + }, + { + "id": "wolfbagging", + "match": "wolfbagging", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "worldsex", + "match": "worldsex", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "wrapping-men", + "match": "wrapping men", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "wrinkled-starfish", + "match": "wrinkled starfish", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "xhamster", + "match": "xhamster", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "xnxx", + "match": "xnxx", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "xtube", + "match": "xtube", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "xvideos", + "match": "xvideos", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "xxx", + "match": "xxx", + "tags": [ + "sexual" + ], + "severity": 2 + }, + { + "id": "xyrophilia", + "match": "xyrophilia|xyrophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "zipperhead", + "match": "zipperhead|zipper-head|zipper head", + "tags": [ + "racial" + ], + "severity": 3 + }, + { + "id": "zippocat", + "match": "zippocat|zippo-cat|zippo cat", + "tags": [ + "shock" + ], + "severity": 4 + }, + { + "id": "zoophilia", + "match": "zoophilia|zoophile", + "tags": [ + "sexual" + ], + "severity": 3 + }, + { + "id": "censored-fucking", + "match": "f**ing|f***ing|f**k|f***", + "tags": [ + "general" + ], + "severity": 4 + }, + { + "id": "piece-of-sht", + "match": "piece of sht|sht ai|sht", + "tags": [ + "general" + ], + "severity": 3 + }, + { + "id": "kill-yourself", + "match": "kill yourself|go kill yourself", + "tags": [ + "general" + ], + "severity": 4 + } + ] \ No newline at end of file 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 263b6eee768..9d64aea8910 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 @@ -6,6 +6,7 @@ to detect and block/mask sensitive content. """ import asyncio +import json import os import re from datetime import datetime @@ -150,6 +151,7 @@ class ContentFilterGuardrail(CustomGuardrail): categories: List of category configurations with enabled/action/severity settings severity_threshold: Minimum severity to block ("high", "medium", "low") """ + super().__init__( guardrail_name=guardrail_name, supported_event_hooks=[ @@ -179,6 +181,12 @@ class ContentFilterGuardrail(CustomGuardrail): # Load categories if provided if categories: self._load_categories(categories) + else: + verbose_proxy_logger.warning( + "ContentFilterGuardrail has no content categories configured. " + "Toxic/abuse and other category-based keyword filtering will not run. " + "Add categories (e.g. harm_toxic_abuse) in the guardrail config to enable them." + ) # Normalize inputs: convert dicts to Pydantic models for consistent handling normalized_patterns: List[ContentFilterPattern] = [] @@ -276,9 +284,15 @@ class ContentFilterGuardrail(CustomGuardrail): if custom_file: category_file_path = custom_file else: - category_file_path = os.path.join( - categories_dir, f"{category_name}.yaml" - ) + # Try .yaml first, then .json (e.g. harm_toxic_abuse.json) + yaml_path = os.path.join(categories_dir, f"{category_name}.yaml") + json_path = os.path.join(categories_dir, f"{category_name}.json") + if os.path.exists(yaml_path): + category_file_path = yaml_path + elif os.path.exists(json_path): + category_file_path = json_path + else: + category_file_path = yaml_path # will trigger "not found" below if not os.path.exists(category_file_path): verbose_proxy_logger.warning( @@ -319,17 +333,23 @@ class ContentFilterGuardrail(CustomGuardrail): def _load_category_file(self, file_path: str) -> CategoryConfig: """ - Load a category definition from a YAML file. + Load a category definition from a YAML or JSON file. + + YAML format: category_name, description, default_action, keywords (list of + {keyword, severity}), exceptions. + JSON format: list of {id, match, tags, severity}; match is pipe-separated + phrases; severity 1-4 mapped to low/medium/high. Used for harm_toxic_abuse. Args: - file_path: Path to category YAML file + file_path: Path to category YAML or JSON file Returns: CategoryConfig object """ + if file_path.lower().endswith(".json"): + return self._load_category_file_json(file_path) with open(file_path, "r") as f: data = yaml.safe_load(f) - return CategoryConfig( category_name=data.get("category_name", "unknown"), description=data.get("description", ""), @@ -338,6 +358,44 @@ class ContentFilterGuardrail(CustomGuardrail): exceptions=data.get("exceptions", []), ) + def _load_category_file_json(self, file_path: str) -> CategoryConfig: + """ + Load a category from the harm_toxic_abuse-style JSON format. + + Each entry has: id, match (pipe-separated phrases), tags, severity (1-4). + Severity mapping: 4,3 -> high; 2 -> medium; 1 -> low. + """ + with open(file_path, "r") as f: + entries = json.load(f) + if not isinstance(entries, list): + entries = [entries] + # Derive category name from filename (e.g. harm_toxic_abuse.json -> harm_toxic_abuse) + category_name = os.path.splitext(os.path.basename(file_path))[0] + severity_map = {4: "high", 3: "high", 2: "medium", 1: "low"} + keywords: List[Dict[str, str]] = [] + seen = set() + for item in entries: + if not isinstance(item, dict): + continue + match_str = item.get("match") or "" + raw_severity = item.get("severity", 2) + severity = severity_map.get( + raw_severity if isinstance(raw_severity, int) else 2, "medium" + ) + for phrase in match_str.split("|"): + phrase = phrase.strip().lower() + if not phrase or phrase in seen: + continue + seen.add(phrase) + keywords.append({"keyword": phrase, "severity": severity}) + return CategoryConfig( + category_name=category_name, + description="Detects harmful, toxic, or abusive language and content", + default_action=ContentFilterAction("BLOCK"), + keywords=keywords, + exceptions=[], + ) + def _should_apply_severity(self, severity: str, threshold: str) -> bool: """ Check if a given severity meets the threshold. diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/patterns.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/patterns.py index d3a66690a90..aa4b8b1b37d 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/patterns.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/patterns.py @@ -139,6 +139,9 @@ def get_available_content_categories() -> List[Dict[str, str]]: """ Return available content categories for UI display. + Includes categories defined in .yaml/.yml files and in .json files + (e.g. harm_toxic_abuse.json). + Returns: List of dictionaries containing category name, display_name, and description """ @@ -177,6 +180,28 @@ def get_available_content_categories() -> List[Dict[str, str]]: except Exception: # Skip files that can't be loaded continue + elif filename.endswith(".json"): + # JSON category files (e.g. harm_toxic_abuse.json) - no YAML header, use filename + category_name = os.path.splitext(filename)[0] + try: + if category_name == "harm_toxic_abuse": + display_name = "Harmful Toxic Abuse" + description = ( + "Detects harmful, toxic, or abusive language and content" + ) + else: + display_name = category_name.replace("_", " ").title() + description = f"Content category: {display_name}" + available_categories.append( + { + "name": category_name, + "display_name": display_name, + "description": description, + "default_action": "BLOCK", + } + ) + except Exception: + continue # Sort by name for consistent ordering available_categories.sort(key=lambda x: x["name"]) diff --git a/litellm/proxy/hooks/key_management_event_hooks.py b/litellm/proxy/hooks/key_management_event_hooks.py index a8325d34612..c07f30f8646 100644 --- a/litellm/proxy/hooks/key_management_event_hooks.py +++ b/litellm/proxy/hooks/key_management_event_hooks.py @@ -150,15 +150,26 @@ class KeyManagementEventHooks: existing_key_row.key_alias or f"virtual-key-{existing_key_row.token}" ) + new_secret_name = ( + response.key_alias + or data.key_alias + or f"virtual-key-{response.token_id}" + ) + verbose_proxy_logger.info( + "Updating secret in secret manager: secret_name=%s", + new_secret_name, + ) team_id = getattr(existing_key_row, "team_id", None) await KeyManagementEventHooks._rotate_virtual_key_in_secret_manager( current_secret_name=initial_secret_name, - new_secret_name=response.key_alias - or data.key_alias - or f"virtual-key-{response.token_id}", + new_secret_name=new_secret_name, new_secret_value=response.key, team_id=team_id, ) + verbose_proxy_logger.info( + "Secret updated in secret manager: secret_name=%s", + new_secret_name, + ) except Exception as e: verbose_proxy_logger.warning( f"Failed to rotate virtual key in secret manager: {e}" diff --git a/litellm/proxy/hooks/model_max_budget_limiter.py b/litellm/proxy/hooks/model_max_budget_limiter.py index 69c7e92d82e..b8c073dd061 100644 --- a/litellm/proxy/hooks/model_max_budget_limiter.py +++ b/litellm/proxy/hooks/model_max_budget_limiter.py @@ -153,7 +153,10 @@ class _PROXY_VirtualKeyModelMaxBudgetLimiter(RouterBudgetLimiting): "standard_logging_object", None ) if standard_logging_payload is None: - raise ValueError("standard_logging_payload is required") + verbose_proxy_logger.debug( + "Skipping _PROXY_VirtualKeyModelMaxBudgetLimiter.async_log_success_event: standard_logging_payload is None" + ) + return _litellm_params: dict = kwargs.get("litellm_params", {}) or {} _metadata: dict = _litellm_params.get("metadata", {}) or {} diff --git a/litellm/proxy/image_endpoints/endpoints.py b/litellm/proxy/image_endpoints/endpoints.py index 4a2c05f8590..4a8eb8e7419 100644 --- a/litellm/proxy/image_endpoints/endpoints.py +++ b/litellm/proxy/image_endpoints/endpoints.py @@ -144,6 +144,12 @@ async def image_generation( litellm_call_id=data.get("litellm_call_id", ""), status="success" ) ) + + ### CALL HOOKS ### - modify outgoing data (guardrails, otel, etc.) + response = await proxy_logging_obj.post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response + ) + ### RESPONSE HEADERS ### hidden_params = getattr(response, "_hidden_params", {}) or {} model_id = hidden_params.get("model_id", None) or "" diff --git a/litellm/proxy/management_endpoints/common_utils.py b/litellm/proxy/management_endpoints/common_utils.py index 24a41a2361b..942758e3bab 100644 --- a/litellm/proxy/management_endpoints/common_utils.py +++ b/litellm/proxy/management_endpoints/common_utils.py @@ -8,6 +8,7 @@ from litellm.proxy._types import ( LiteLLM_ManagementEndpoint_MetadataFields_Premium, LiteLLM_OrganizationTable, LiteLLM_TeamTable, + LiteLLM_UserTable, LitellmUserRoles, UserAPIKeyAuth, ) @@ -108,6 +109,154 @@ async def _user_has_admin_privileges( return False +def _org_admin_can_invite_user( + admin_user_obj: LiteLLM_UserTable, + target_user_obj: LiteLLM_UserTable, +) -> bool: + """ + Check if an org admin can invite the target user. + Target user must be in at least one org where the admin has org admin role. + + Args: + admin_user_obj: The admin user's full object (from get_user_object) + target_user_obj: The target user's full object (from get_user_object) + + Returns: + True if target user is in an org where admin has org admin role + """ + if admin_user_obj.organization_memberships is None: + return False + admin_org_ids = { + m.organization_id + for m in admin_user_obj.organization_memberships + if m.user_role == LitellmUserRoles.ORG_ADMIN.value + } + if not admin_org_ids: + return False + if target_user_obj.organization_memberships is None: + return False + target_org_ids = { + m.organization_id for m in target_user_obj.organization_memberships + } + return bool(admin_org_ids & target_org_ids) + + +async def _team_admin_can_invite_user( + user_api_key_dict: UserAPIKeyAuth, + admin_user_obj: LiteLLM_UserTable, + target_user_obj: LiteLLM_UserTable, + prisma_client: "PrismaClient", +) -> bool: + """ + Check if a team admin can invite the target user. + Target user must be in at least one team where the admin has team admin role. + + Args: + user_api_key_dict: The admin user's API key auth object + admin_user_obj: The admin user's full object (from get_user_object) + target_user_obj: The target user's full object (from get_user_object) + prisma_client: Prisma client for database operations + + Returns: + True if target user is in a team where admin has team admin role + """ + if not admin_user_obj.teams or len(admin_user_obj.teams) == 0: + return False + if not target_user_obj.teams or len(target_user_obj.teams) == 0: + return False + + teams = await prisma_client.db.litellm_teamtable.find_many( + where={"team_id": {"in": admin_user_obj.teams}} + ) + admin_team_ids = [ + team.team_id + for team in teams + if _is_user_team_admin( + user_api_key_dict=user_api_key_dict, + team_obj=LiteLLM_TeamTable(**team.model_dump()), + ) + ] + if not admin_team_ids: + return False + target_team_ids = set(target_user_obj.teams) + return bool(set(admin_team_ids) & target_team_ids) + + +async def admin_can_invite_user( + target_user_id: str, + user_api_key_dict: UserAPIKeyAuth, + prisma_client: Optional["PrismaClient"] = None, + user_api_key_cache: Optional["DualCache"] = None, + proxy_logging_obj: Optional["ProxyLogging"] = None, +) -> bool: + """ + Check if the admin can create an invitation for the target user. + - Proxy admins: can invite any user + - Org admins: can only invite users in their org(s) + - Team admins: can only invite users in their team(s) + + Uses get_user_object for caching of both admin and target user objects. + + Args: + target_user_id: The user_id of the user to invite + user_api_key_dict: The admin user's API key auth object + prisma_client: Prisma client for database operations + user_api_key_cache: Cache for user API keys + proxy_logging_obj: Proxy logging object + + Returns: + True if user can invite the target user + """ + if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN: + return True + + if prisma_client is None or user_api_key_dict.user_id is None: + return False + + from litellm.caching import DualCache as DualCacheImport + from litellm.proxy.auth.auth_checks import get_user_object + + try: + cache = user_api_key_cache or DualCacheImport() + admin_user_obj = await get_user_object( + user_id=user_api_key_dict.user_id, + prisma_client=prisma_client, + user_api_key_cache=cache, + user_id_upsert=False, + proxy_logging_obj=proxy_logging_obj, + ) + if admin_user_obj is None: + return False + + target_user_obj = await get_user_object( + user_id=target_user_id, + prisma_client=prisma_client, + user_api_key_cache=cache, + user_id_upsert=False, + proxy_logging_obj=proxy_logging_obj, + ) + if target_user_obj is None: + return False + + if _org_admin_can_invite_user(admin_user_obj, target_user_obj): + return True + + if await _team_admin_can_invite_user( + user_api_key_dict=user_api_key_dict, + admin_user_obj=admin_user_obj, + target_user_obj=target_user_obj, + prisma_client=prisma_client, + ): + return True + + return False + except Exception as e: + verbose_proxy_logger.debug( + f"Error checking invite permission for user {user_api_key_dict.user_id}: {e}" + ) + return False + + def _set_object_metadata_field( object_data: Union[ LiteLLM_TeamTable, diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 2e71759072d..152b09a86b0 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -2770,6 +2770,7 @@ async def can_modify_verification_token( Rules: - Proxy admin can modify any key + - Internal jobs service account can modify any key (for auto-rotation) - For team keys: only team admin or key owner can modify - For personal keys: only key owner can modify @@ -2782,13 +2783,19 @@ async def can_modify_verification_token( Returns: True if user can modify the key, False otherwise """ + from litellm.constants import LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME + is_team_key = _is_team_key(data=key_info) # 1. Proxy admin can modify any key if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value: return True - # 2. For team keys: only team admin or key owner can modify + # 2. Internal jobs service account can modify any key (for auto-rotation) + if user_api_key_dict.api_key == LITELLM_INTERNAL_JOBS_SERVICE_ACCOUNT_NAME: + return True + + # 3. For team keys: only team admin or key owner can modify if is_team_key and key_info.team_id is not None: # Get team object to check if user is team admin team_table = await get_team_object( @@ -2818,7 +2825,7 @@ async def can_modify_verification_token( # Not team admin and doesn't own the key return False - # 3. For personal keys: only key owner can modify + # 4. For personal keys: only key owner can modify if key_info.user_id is not None and key_info.user_id == user_api_key_dict.user_id: return True @@ -3179,7 +3186,7 @@ def get_new_token(data: Optional[RegenerateKeyRequest]) -> str: dependencies=[Depends(user_api_key_auth)], ) @management_endpoint_wrapper -async def regenerate_key_fn( +async def regenerate_key_fn( # noqa: PLR0915 key: Optional[str] = None, data: Optional[RegenerateKeyRequest] = None, user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), @@ -3330,6 +3337,10 @@ async def regenerate_key_fn( detail={"error": "You are not authorized to regenerate this key"}, ) + verbose_proxy_logger.info( + "Key regeneration requested: key_alias=%s", + getattr(_key_in_db, "key_alias", None), + ) verbose_proxy_logger.debug("key_in_db: %s", _key_in_db) new_token = get_new_token(data=data) @@ -3380,6 +3391,10 @@ async def regenerate_key_fn( **updated_token_dict, ) + verbose_proxy_logger.info( + "Key regeneration completed: key_alias=%s", + getattr(_key_in_db, "key_alias", None), + ) asyncio.create_task( KeyManagementEventHooks.async_key_rotated_hook( data=data, diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index b507beef4f5..d5184efc1ef 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -338,7 +338,10 @@ from litellm.proxy.management_endpoints.cache_settings_endpoints import ( from litellm.proxy.management_endpoints.callback_management_endpoints import ( router as callback_management_endpoints_router, ) -from litellm.proxy.management_endpoints.common_utils import _user_has_admin_view +from litellm.proxy.management_endpoints.common_utils import ( + admin_can_invite_user, + _user_has_admin_privileges, +) from litellm.proxy.management_endpoints.cost_tracking_settings import ( router as cost_tracking_settings_router, ) @@ -4606,11 +4609,16 @@ async def initialize( # noqa: PLR0915 elif litellm_log_setting.upper() == "DEBUG": import logging - from litellm._logging import verbose_proxy_logger, verbose_router_logger + from litellm._logging import ( + verbose_logger, + verbose_proxy_logger, + verbose_router_logger, + ) + verbose_logger.setLevel(level=logging.DEBUG) # set package log to debug verbose_router_logger.setLevel( level=logging.DEBUG - ) # set router logs to info + ) # set router logs to debug verbose_proxy_logger.setLevel( level=logging.DEBUG ) # set proxy logs to debug @@ -10379,7 +10387,17 @@ async def new_invitation( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Allow proxy admins and org/team admins (admin status from DB via get_user_object) + has_access = ( + user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + or await _user_has_admin_privileges( + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + ) + if not has_access: raise HTTPException( status_code=400, detail={ @@ -10390,6 +10408,23 @@ async def new_invitation( }, ) + # Org/team admins can only invite users within their org/team + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + can_invite = await admin_can_invite_user( + target_user_id=data.user_id, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + if not can_invite: + raise HTTPException( + status_code=400, + detail={ + "error": "You can only create invitations for users in your organization or team." + }, + ) + response = await create_invitation_for_user( data=data, user_api_key_dict=user_api_key_dict, @@ -10542,7 +10577,16 @@ async def invitation_delete( detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + # Proxy admins can delete any invitation; org admins only their own + is_proxy_admin = user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN + is_other_admin = await _user_has_admin_privileges( + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + if not is_proxy_admin and not is_other_admin: raise HTTPException( status_code=400, detail={ @@ -10553,6 +10597,24 @@ async def invitation_delete( }, ) + # Org admins can only delete invitations they created + if is_other_admin and not is_proxy_admin: + invitation = await prisma_client.db.litellm_invitationlink.find_unique( + where={"id": data.invitation_id} + ) + if invitation is None: + raise HTTPException( + status_code=400, + detail={"error": "Invitation id does not exist in the database."}, + ) + if invitation.created_by != user_api_key_dict.user_id: + raise HTTPException( + status_code=403, + detail={ + "error": "Organization admins can only delete invitations they created." + }, + ) + response = await prisma_client.db.litellm_invitationlink.delete( where={"id": data.invitation_id} ) diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index c05eda85c81..afbc57360e2 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -1687,6 +1687,14 @@ async def ui_view_spend_logs( # noqa: PLR0915 error_message: Optional[str] = fastapi.Query( default=None, description="Filter logs by error message (partial string match)" ), + sort_by: str = fastapi.Query( + default="startTime", + description="Sort logs by field: spend, total_tokens, startTime, or endTime", + ), + sort_order: Optional[str] = fastapi.Query( + default="desc", + description="Sort order: asc or desc", + ), ): """ View spend logs with pagination support. @@ -1718,6 +1726,23 @@ async def ui_view_spend_logs( # noqa: PLR0915 code=status.HTTP_400_BAD_REQUEST, ) + # Validate sort_by and sort_order + valid_sort_fields = {"spend", "total_tokens", "startTime", "endTime"} + if sort_by not in valid_sort_fields: + raise ProxyException( + message=f"Invalid sort_by: {sort_by}. Must be one of: {', '.join(sorted(valid_sort_fields))}", + type="bad_request", + param="sort_by", + code=status.HTTP_400_BAD_REQUEST, + ) + if sort_order is not None and sort_order.lower() not in {"asc", "desc"}: + raise ProxyException( + message=f"Invalid sort_order: {sort_order}. Must be one of: asc, desc", + type="bad_request", + param="sort_order", + code=status.HTTP_400_BAD_REQUEST, + ) + try: is_v2 = "/spend/logs/v2" in request.url.path formats = ["%Y-%m-%d %H:%M:%S", "%Y-%m-%d"] if is_v2 else ["%Y-%m-%d %H:%M:%S"] @@ -1830,6 +1855,11 @@ async def ui_view_spend_logs( # noqa: PLR0915 # Calculate skip value for pagination skip = (page - 1) * page_size + # Build order clause from sort_by and sort_order + order_column = sort_by + order_direction = (sort_order or "desc").lower() + order_clause = {order_column: order_direction} + # Get total count of records total_records = await prisma_client.db.litellm_spendlogs.count( where=where_conditions, @@ -1838,9 +1868,7 @@ async def ui_view_spend_logs( # noqa: PLR0915 # Get paginated data data = await prisma_client.db.litellm_spendlogs.find_many( where=where_conditions, - order={ - "startTime": "desc", - }, + order=order_clause, skip=skip, take=page_size, ) diff --git a/litellm/router.py b/litellm/router.py index efb284cce49..7bba0902a5e 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -58,7 +58,6 @@ from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, get_metadata_variable_name_from_kwargs, ) -from litellm.litellm_core_utils.thread_pool_executor import executor from litellm.litellm_core_utils.coroutine_checker import coroutine_checker from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.dd_tracing import tracer @@ -619,11 +618,12 @@ class Router: self.retry_policy = RetryPolicy(**retry_policy) elif isinstance(retry_policy, RetryPolicy): self.retry_policy = retry_policy - verbose_router_logger.info( - "\033[32mRouter Custom Retry Policy Set:\n{}\033[0m".format( - self.retry_policy.model_dump(exclude_none=True) + if self.retry_policy is not None: + verbose_router_logger.info( + "\033[32mRouter Custom Retry Policy Set:\n{}\033[0m".format( + self.retry_policy.model_dump(exclude_none=True) + ) ) - ) self.model_group_retry_policy: Optional[ Dict[str, RetryPolicy] @@ -636,11 +636,12 @@ class Router: elif isinstance(allowed_fails_policy, AllowedFailsPolicy): self.allowed_fails_policy = allowed_fails_policy - verbose_router_logger.info( - "\033[32mRouter Custom Allowed Fails Policy Set:\n{}\033[0m".format( - self.allowed_fails_policy.model_dump(exclude_none=True) + if self.allowed_fails_policy is not None: + verbose_router_logger.info( + "\033[32mRouter Custom Allowed Fails Policy Set:\n{}\033[0m".format( + self.allowed_fails_policy.model_dump(exclude_none=True) + ) ) - ) self.alerting_config: Optional[AlertingConfig] = alerting_config @@ -1269,13 +1270,16 @@ class Router: if silent_model is not None: # Mirroring traffic to a secondary model - # Use shared thread pool for background calls - executor.submit( - self._silent_experiment_completion, - silent_model, - messages, - **kwargs, + # Use threading.Thread (not ThreadPoolExecutor) - executor.submit() + # requires pickling args, which fails when kwargs contain unpicklable + # objects (e.g. _thread.RLock from OTEL spans, loggers) in deployment. + thread = threading.Thread( + target=self._silent_experiment_completion, + args=(silent_model, messages), + kwargs=kwargs, + daemon=True, ) + thread.start() self._update_kwargs_with_deployment(deployment=deployment, kwargs=kwargs) kwargs.pop("silent_model", None) # Ensure it's not in kwargs either diff --git a/litellm/secret_managers/aws_secret_manager_v2.py b/litellm/secret_managers/aws_secret_manager_v2.py index 8edfc48336b..c1b4d019dcf 100644 --- a/litellm/secret_managers/aws_secret_manager_v2.py +++ b/litellm/secret_managers/aws_secret_manager_v2.py @@ -3,7 +3,8 @@ This is a file for the AWS Secret Manager Integration Handles Async Operations for: - Read Secret -- Write Secret +- Write Secret (CreateSecret) +- Update Secret (PutSecretValue) - for in-place rotation when alias is preserved - Delete Secret Relevant issue: https://github.com/BerriAI/litellm/issues/1883 @@ -42,11 +43,11 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): aws_profile_name: Optional[str] = None, aws_web_identity_token: Optional[str] = None, aws_sts_endpoint: Optional[str] = None, - **kwargs + **kwargs, ): BaseSecretManager.__init__(self, **kwargs) BaseAWSLLM.__init__(self, **kwargs) - + # Store AWS authentication settings self.aws_region_name = aws_region_name self.aws_role_name = aws_role_name @@ -61,7 +62,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): # AWS_REGION_NAME is only strictly required if not using a profile or role # When using IAM roles, the region can come from multiple sources if ( - "AWS_REGION_NAME" not in os.environ + "AWS_REGION_NAME" not in os.environ and "AWS_REGION" not in os.environ and "AWS_DEFAULT_REGION" not in os.environ ): @@ -83,22 +84,36 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): return try: cls.validate_environment() - + # Extract AWS settings from key_management_settings if provided aws_kwargs = {} if key_management_settings is not None: aws_kwargs = { - "aws_region_name": getattr(key_management_settings, "aws_region_name", None), - "aws_role_name": getattr(key_management_settings, "aws_role_name", None), - "aws_session_name": getattr(key_management_settings, "aws_session_name", None), - "aws_external_id": getattr(key_management_settings, "aws_external_id", None), - "aws_profile_name": getattr(key_management_settings, "aws_profile_name", None), - "aws_web_identity_token": getattr(key_management_settings, "aws_web_identity_token", None), - "aws_sts_endpoint": getattr(key_management_settings, "aws_sts_endpoint", None), + "aws_region_name": getattr( + key_management_settings, "aws_region_name", None + ), + "aws_role_name": getattr( + key_management_settings, "aws_role_name", None + ), + "aws_session_name": getattr( + key_management_settings, "aws_session_name", None + ), + "aws_external_id": getattr( + key_management_settings, "aws_external_id", None + ), + "aws_profile_name": getattr( + key_management_settings, "aws_profile_name", None + ), + "aws_web_identity_token": getattr( + key_management_settings, "aws_web_identity_token", None + ), + "aws_sts_endpoint": getattr( + key_management_settings, "aws_sts_endpoint", None + ), } # Remove None values aws_kwargs = {k: v for k, v in aws_kwargs.items() if v is not None} - + litellm.secret_manager_client = cls(**aws_kwargs) litellm._key_management_system = KeyManagementSystem.AWS_SECRET_MANAGER @@ -246,13 +261,13 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): return primary_secret_kv_pairs.get(secret_name) async def async_write_secret( - self, - secret_name: str, - secret_value: str, - description: Optional[str] = None, - optional_params: Optional[dict] = None, - timeout: Optional[Union[float, httpx.Timeout]] = None, - tags: Optional[Union[dict, list]] = None + self, + secret_name: str, + secret_value: str, + description: Optional[str] = None, + optional_params: Optional[dict] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + tags: Optional[Union[dict, list]] = None, ) -> dict: """ Async function to write a secret to AWS Secrets Manager @@ -312,6 +327,94 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): except httpx.TimeoutException: raise ValueError("Timeout error occurred") + async def async_put_secret_value( + self, + secret_name: str, + secret_value: str, + optional_params: Optional[dict] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + ) -> dict: + """ + Async function to update an existing secret's value in AWS Secrets Manager. + + Uses PutSecretValue to update in place. Use this when rotating a secret + that keeps the same name (current_secret_name == new_secret_name). + + Args: + secret_name: Name of the existing secret to update + secret_value: New value to store + optional_params: Additional AWS parameters + timeout: Request timeout + + Returns: + dict: Response from AWS Secrets Manager containing update details + """ + from litellm._uuid import uuid + + data: Dict[str, Any] = { + "SecretId": secret_name, + "SecretString": secret_value, + "ClientRequestToken": str(uuid.uuid4()), + } + + endpoint_url, headers, body = self._prepare_request( + action="PutSecretValue", + secret_name=secret_name, + secret_value=secret_value, + optional_params=optional_params, + request_data=data, + ) + + async_client = get_async_httpx_client( + llm_provider=httpxSpecialProvider.SecretManager, + params={"timeout": timeout}, + ) + + try: + response = await async_client.post( + url=endpoint_url, headers=headers, data=body.decode("utf-8") + ) + response.raise_for_status() + return response.json() + except httpx.HTTPStatusError as err: + raise ValueError(f"HTTP error occurred: {err.response.text}") + except httpx.TimeoutException: + raise ValueError("Timeout error occurred") + + async def async_rotate_secret( + self, + current_secret_name: str, + new_secret_name: str, + new_secret_value: str, + optional_params: Optional[dict] = None, + timeout: Optional[Union[float, httpx.Timeout]] = None, + ) -> dict: + """ + Rotate a secret. When current_secret_name == new_secret_name (in-place + update), uses PutSecretValue instead of create+delete to avoid + ResourceExistsException. + """ + if current_secret_name == new_secret_name: + # Same alias: update in place via PutSecretValue + verbose_logger.info( + "Secret rotated in-place (PutSecretValue): secret_name=%s", + current_secret_name, + ) + return await self.async_put_secret_value( + secret_name=current_secret_name, + secret_value=new_secret_value, + optional_params=optional_params, + timeout=timeout, + ) + # Different names: create new, delete old (base class logic) + return await super().async_rotate_secret( + current_secret_name=current_secret_name, + new_secret_name=new_secret_name, + new_secret_value=new_secret_value, + optional_params=optional_params, + timeout=timeout, + ) + async def async_delete_secret( self, secret_name: str, @@ -375,7 +478,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): except ImportError: raise ImportError("Missing boto3 to call bedrock. Run 'pip install boto3'.") optional_params = optional_params or {} - + # Build optional_params from instance settings if not provided # This allows the IAM role settings to be used for Secret Manager calls if not optional_params.get("aws_role_name") and self.aws_role_name: @@ -388,11 +491,14 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): optional_params["aws_external_id"] = self.aws_external_id if not optional_params.get("aws_profile_name") and self.aws_profile_name: optional_params["aws_profile_name"] = self.aws_profile_name - if not optional_params.get("aws_web_identity_token") and self.aws_web_identity_token: + if ( + not optional_params.get("aws_web_identity_token") + and self.aws_web_identity_token + ): optional_params["aws_web_identity_token"] = self.aws_web_identity_token if not optional_params.get("aws_sts_endpoint") and self.aws_sts_endpoint: optional_params["aws_sts_endpoint"] = self.aws_sts_endpoint - + boto3_credentials_info = self._get_boto_credentials_from_optional_params( optional_params ) @@ -431,12 +537,3 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): prepped = request.prepare() return endpoint_url, prepped.headers, body - - -# if __name__ == "__main__": -# print("loading aws secret manager v2") -# aws_secret_manager_v2 = AWSSecretsManagerV2() -# import asyncio -# print("writing secret to aws secret manager v2") -# asyncio.run(aws_secret_manager_v2.async_write_secret(secret_name="test_secret_3", secret_value="test_value_2")) -# print("reading secret from aws secret manager v2") diff --git a/litellm/utils.py b/litellm/utils.py index a38fcf2a0c5..0fd21f09919 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -2506,6 +2506,16 @@ def _supports_factory(model: str, custom_llm_provider: Optional[str], key: str) if model_info.get(key, False) is True: return True elif model_info.get(key) is None: # don't check if 'False' explicitly set + # Fallback: when the provider-prefixed entry (e.g. + # "deepseek/deepseek-chat") exists but is missing a capability + # field, check the bare model-name entry (e.g. "deepseek-chat") + # which may carry the complete metadata. See #20885. + bare_model_key = _get_model_cost_key(model) + if bare_model_key is not None: + bare_entry = litellm.model_cost.get(bare_model_key) or {} + if bare_entry.get(key, False) is True: + return True + supported_by_provider = _supports_provider_info_factory( model, custom_llm_provider, key ) @@ -6140,6 +6150,13 @@ def validate_environment( # noqa: PLR0915 if ( "AWS_ACCESS_KEY_ID" in os.environ and "AWS_SECRET_ACCESS_KEY" in os.environ + ) or ( + # IAM role, profile, or web identity auth don't require access keys + "AWS_ROLE_ARN" in os.environ + or "AWS_PROFILE" in os.environ + or "AWS_WEB_IDENTITY_TOKEN_FILE" in os.environ + or "AWS_CONTAINER_CREDENTIALS_RELATIVE_URI" in os.environ # ECS task role + or "AWS_CONTAINER_CREDENTIALS_FULL_URI" in os.environ # ECS/Fargate full URI credential delivery ): keys_in_environment = True else: diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 9b5a7b42d0e..f6edcf7efd0 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -10756,14 +10756,22 @@ "input_cost_per_token": 2.8e-07, "input_cost_per_token_cache_hit": 2.8e-08, "litellm_provider": "deepseek", - "max_input_tokens": 128000, + "max_input_tokens": 131072, "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", "output_cost_per_token": 4.2e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], "supports_assistant_prefill": true, "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, "supports_tool_choice": true }, "deepseek/deepseek-coder": { @@ -10800,16 +10808,24 @@ "input_cost_per_token": 2.8e-07, "input_cost_per_token_cache_hit": 2.8e-08, "litellm_provider": "deepseek", - "max_input_tokens": 128000, - "max_output_tokens": 8192, - "max_tokens": 8192, + "max_input_tokens": 131072, + "max_output_tokens": 65536, + "max_tokens": 65536, "mode": "chat", "output_cost_per_token": 4.2e-07, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], "supports_assistant_prefill": true, - "supports_function_calling": true, + "supports_function_calling": false, + "supports_native_streaming": true, + "supports_parallel_function_calling": false, "supports_prompt_caching": true, "supports_reasoning": true, - "supports_tool_choice": true + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": false }, "deepseek/deepseek-v3": { "cache_creation_input_token_cost": 0.0, diff --git a/tests/litellm/test_batch_completion_models_all_responses.py b/tests/litellm/test_batch_completion_models_all_responses.py new file mode 100644 index 00000000000..2e96ada03f2 --- /dev/null +++ b/tests/litellm/test_batch_completion_models_all_responses.py @@ -0,0 +1,118 @@ +import concurrent.futures + +import litellm +from litellm.batch_completion.main import batch_completion_models_all_responses + + +def test_batch_completion_models_all_responses_submits_before_waiting(monkeypatch): + """ + Regression test for issue #20704. + Ensures all model calls are submitted to the thread pool before waiting on results. + """ + models = ["model-a", "model-b", "model-c"] + called_models = [] + + class _AssertingFuture: + def __init__(self, result, executor, expected_submissions): + self._result = result + self._executor = executor + self._expected_submissions = expected_submissions + + def result(self): + if self._executor.submit_count != self._expected_submissions: + raise AssertionError("Not all model calls were submitted before waiting") + return self._result + + class _RecordingThreadPoolExecutor: + def __init__(self, max_workers, *args, **kwargs): + self.max_workers = max_workers + self.submit_count = 0 + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc, tb): + return False + + def submit(self, fn, *args, **kwargs): + self.submit_count += 1 + result = fn(*args, **kwargs) + return _AssertingFuture( + result=result, + executor=self, + expected_submissions=len(models), + ) + + def _mock_completion(*args, model, **kwargs): + called_models.append(model) + return {"model": model} + + monkeypatch.setattr(litellm, "completion", _mock_completion) + monkeypatch.setattr( + concurrent.futures, "ThreadPoolExecutor", _RecordingThreadPoolExecutor + ) + + responses = batch_completion_models_all_responses( + models=models, + messages=[{"role": "user", "content": "hello"}], + ) + + assert sorted(called_models) == sorted(models) + assert len(responses) == len(models) + assert sorted(response["model"] for response in responses) == sorted(models) + + +def test_batch_completion_models_all_responses_continues_on_model_error(monkeypatch): + models = ["model-a", "model-error", "model-b"] + + def _mock_completion(*args, model, **kwargs): + if model == "model-error": + raise RuntimeError("simulated model failure") + return {"model": model} + + monkeypatch.setattr(litellm, "completion", _mock_completion) + + responses = batch_completion_models_all_responses( + models=models, + messages=[{"role": "user", "content": "hello"}], + ) + + assert len(responses) == 2 + assert sorted(response["model"] for response in responses) == ["model-a", "model-b"] + + +def test_batch_completion_models_all_responses_returns_empty_for_empty_models(monkeypatch): + called = False + + def _mock_completion(*args, model, **kwargs): + nonlocal called + called = True + return {"model": model} + + monkeypatch.setattr(litellm, "completion", _mock_completion) + + responses = batch_completion_models_all_responses( + models=[], + messages=[{"role": "user", "content": "hello"}], + ) + + assert responses == [] + assert called is False + + +def test_batch_completion_models_all_responses_accepts_single_model_string(monkeypatch): + called_models = [] + + def _mock_completion(*args, model, **kwargs): + called_models.append(model) + return {"model": model} + + monkeypatch.setattr(litellm, "completion", _mock_completion) + + responses = batch_completion_models_all_responses( + models="model-a", + messages=[{"role": "user", "content": "hello"}], + ) + + assert called_models == ["model-a"] + assert responses == [{"model": "model-a"}] diff --git a/tests/llm_translation/test_azure_openai.py b/tests/llm_translation/test_azure_openai.py index 3fd908f86d7..1da380b57a2 100644 --- a/tests/llm_translation/test_azure_openai.py +++ b/tests/llm_translation/test_azure_openai.py @@ -728,3 +728,18 @@ def test_azure_with_content_safety_error(): assert e.provider_specific_fields["innererror"]["code"] == "ResponsibleAIPolicyViolation" assert e.provider_specific_fields["innererror"]["content_filter_result"]["violence"]["filtered"] is True assert e.provider_specific_fields["innererror"]["content_filter_result"]["violence"]["severity"] == "high" + + +def test_azure_openai_with_prompt_cache_key(): + """ + E2E test for Azure OpenAI with prompt cache key param on /chat/completions API. + """ + litellm._turn_on_debug() + response = litellm.completion( + model="azure/gpt-4.1-mini", + api_key=os.getenv("AZURE_API_KEY"), + api_base=os.getenv("AZURE_API_BASE"), + api_version="2024-12-01-preview", + messages=[{"role": "user", "content": "What is the weather in San Francisco?"}], + prompt_cache_key="test_streaming_azure_openai", + ) \ No newline at end of file diff --git a/tests/logging_callback_tests/test_logging_redaction_e2e_test.py b/tests/logging_callback_tests/test_logging_redaction_e2e_test.py index b261daab6a7..e70d08b9008 100644 --- a/tests/logging_callback_tests/test_logging_redaction_e2e_test.py +++ b/tests/logging_callback_tests/test_logging_redaction_e2e_test.py @@ -452,16 +452,13 @@ async def test_redaction_with_metadata_completion_api(): litellm.callbacks = [test_custom_logger] # When metadata is passed, the system uses get_metadata_variable_name_from_kwargs - # to determine which field to check + # to determine which field to check. No headers means redaction should happen + # based on the global setting (litellm.turn_off_message_logging = True) response = await litellm.acompletion( model="gpt-3.5-turbo", messages=[{"role": "user", "content": "hi"}], mock_response="hello", - metadata={ - "headers": { - "litellm-disable-message-redaction": "true" - } - } + metadata={} ) await asyncio.sleep(1) diff --git a/tests/logging_callback_tests/test_otel_logging.py b/tests/logging_callback_tests/test_otel_logging.py index a0c78305e60..f0511e7d1ea 100644 --- a/tests/logging_callback_tests/test_otel_logging.py +++ b/tests/logging_callback_tests/test_otel_logging.py @@ -263,10 +263,18 @@ async def test_arize_phoenix_adds_openinference_kind_and_avoids_duplicate_litell Ensure Arize Phoenix spans include OpenInference span kind and do not create a duplicate litellm_request span when a proxy parent span is already active. """ + from opentelemetry.sdk.trace import TracerProvider + from opentelemetry.sdk.trace.export import SimpleSpanProcessor exporter.clear() litellm.logging_callback_manager._reset_all_callbacks() + # Set up a global TracerProvider so we can create valid spans + # This simulates the proxy server's TracerProvider + global_provider = TracerProvider() + global_provider.add_span_processor(SimpleSpanProcessor(exporter)) + trace.set_tracer_provider(global_provider) + otel_logger = ArizePhoenixLogger(config=OpenTelemetryConfig(exporter=exporter)) litellm.callbacks = [otel_logger] litellm.success_callback = [] diff --git a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py index 0a581fb512d..afc68dc9d42 100644 --- a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py +++ b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py @@ -525,7 +525,6 @@ async def test_anthropic_messages_with_extra_headers(): # Set up test parameters messages = [{"role": "user", "content": "Hello, can you tell me a short joke?"}] extra_headers = { - "anthropic-beta": "very-custom-beta-value", "anthropic-version": "custom-version-for-test", } @@ -581,87 +580,87 @@ async def test_anthropic_messages_with_extra_headers(): return response -@pytest.mark.asyncio -async def test_bedrock_messages_api_header_forwarding(): - """ - Test that headers from kwargs (set by proxy's add_headers_to_llm_call_by_model_group) - are correctly passed to validate_anthropic_messages_environment for Bedrock Invoke API. +# @pytest.mark.asyncio +# async def test_bedrock_messages_api_header_forwarding(): +# """ +# Test that headers from kwargs (set by proxy's add_headers_to_llm_call_by_model_group) +# are correctly passed to validate_anthropic_messages_environment for Bedrock Invoke API. - This verifies that forward_client_headers_to_llm_api works for Bedrock Invoke API (Messages API). +# This verifies that forward_client_headers_to_llm_api works for Bedrock Invoke API (Messages API). - Issue: When calling Anthropic models via the Messages API, LiteLLM makes a call to - Bedrock's Invoke API, and custom headers were not being forwarded, even though - they worked correctly for Chat Completions API with Bedrock's Converse API. - """ - from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler - from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj - from litellm.types.router import GenericLiteLLMParams +# Issue: When calling Anthropic models via the Messages API, LiteLLM makes a call to +# Bedrock's Invoke API, and custom headers were not being forwarded, even though +# they worked correctly for Chat Completions API with Bedrock's Converse API. +# """ +# from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler +# from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +# from litellm.types.router import GenericLiteLLMParams - handler = BaseLLMHTTPHandler() +# handler = BaseLLMHTTPHandler() - # Headers that would be set by the proxy when forward_client_headers_to_llm_api is configured - custom_headers = { - "X-Custom-Header": "CustomValue", - "X-Request-ID": "req-123", - } +# # Headers that would be set by the proxy when forward_client_headers_to_llm_api is configured +# custom_headers = { +# "X-Custom-Header": "CustomValue", +# "X-Request-ID": "req-123", +# } - # Mock the provider config - mock_provider_config = MagicMock() +# # Mock the provider config +# mock_provider_config = MagicMock() - # We'll check what headers are passed to this method - mock_provider_config.validate_anthropic_messages_environment.return_value = ( - {"Authorization": "Bearer test"}, - "https://bedrock-runtime.us-east-1.amazonaws.com/invoke" - ) - mock_provider_config.transform_anthropic_messages_request.return_value = {"model": "test"} - mock_provider_config.get_complete_url.return_value = "https://test.com" - mock_provider_config.sign_request.return_value = ({}, None) - mock_provider_config.transform_anthropic_messages_response.return_value = {"id": "test"} +# # We'll check what headers are passed to this method +# mock_provider_config.validate_anthropic_messages_environment.return_value = ( +# {"Authorization": "Bearer test"}, +# "https://bedrock-runtime.us-east-1.amazonaws.com/invoke" +# ) +# mock_provider_config.transform_anthropic_messages_request.return_value = {"model": "test"} +# mock_provider_config.get_complete_url.return_value = "https://test.com" +# mock_provider_config.sign_request.return_value = ({}, None) +# mock_provider_config.transform_anthropic_messages_response.return_value = {"id": "test"} - # Mock HTTP client to prevent actual network calls - with unittest.mock.patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client") as mock_get_client: - mock_http_client = AsyncMock() - mock_response = MagicMock() - mock_response.status_code = 200 - mock_response.json.return_value = {"id": "test", "content": []} - mock_response.text = "{}" - mock_http_client.post.return_value = mock_response - mock_get_client.return_value = mock_http_client +# # Mock HTTP client to prevent actual network calls +# with unittest.mock.patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client") as mock_get_client: +# mock_http_client = AsyncMock() +# mock_response = MagicMock() +# mock_response.status_code = 200 +# mock_response.json.return_value = {"id": "test", "content": []} +# mock_response.text = "{}" +# mock_http_client.post.return_value = mock_response +# mock_get_client.return_value = mock_http_client - # Mock logging object - mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj) - mock_logging_obj.model_call_details = {} +# # Mock logging object +# mock_logging_obj = MagicMock(spec=LiteLLMLoggingObj) +# mock_logging_obj.model_call_details = {} - # Call the handler with headers in kwargs - try: - await handler.async_anthropic_messages_handler( - model="bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0", - messages=[{"role": "user", "content": "Hello"}], - anthropic_messages_provider_config=mock_provider_config, - anthropic_messages_optional_request_params={"max_tokens": 100}, - custom_llm_provider="bedrock", - litellm_params=GenericLiteLLMParams( - api_key="test-key", - aws_region_name="us-east-1" - ), - logging_obj=mock_logging_obj, - api_key="test-key", - stream=False, - kwargs={"headers": custom_headers} # Headers set by proxy - ) - except Exception: - pass # Ignore errors, we're only checking if headers were passed +# # Call the handler with headers in kwargs +# try: +# await handler.async_anthropic_messages_handler( +# model="bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0", +# messages=[{"role": "user", "content": "Hello"}], +# anthropic_messages_provider_config=mock_provider_config, +# anthropic_messages_optional_request_params={"max_tokens": 100}, +# custom_llm_provider="bedrock", +# litellm_params=GenericLiteLLMParams( +# api_key="test-key", +# aws_region_name="us-east-1" +# ), +# logging_obj=mock_logging_obj, +# api_key="test-key", +# stream=False, +# kwargs={"headers": custom_headers} # Headers set by proxy +# ) +# except Exception: +# pass # Ignore errors, we're only checking if headers were passed - # Verify that validate_anthropic_messages_environment was called - assert mock_provider_config.validate_anthropic_messages_environment.called +# # Verify that validate_anthropic_messages_environment was called +# assert mock_provider_config.validate_anthropic_messages_environment.called - # Get the headers that were passed - call_args = mock_provider_config.validate_anthropic_messages_environment.call_args - passed_headers = call_args[1]["headers"] +# # Get the headers that were passed +# call_args = mock_provider_config.validate_anthropic_messages_environment.call_args +# passed_headers = call_args[1]["headers"] - # The custom headers from kwargs should be in the passed headers - assert "X-Custom-Header" in passed_headers or "x-custom-header" in passed_headers - assert "X-Request-ID" in passed_headers or "x-request-id" in passed_headers +# # The custom headers from kwargs should be in the passed headers +# assert "X-Custom-Header" in passed_headers or "x-custom-header" in passed_headers +# assert "X-Request-ID" in passed_headers or "x-request-id" in passed_headers @pytest.mark.asyncio diff --git a/tests/pass_through_unit_tests/test_bedrock_tool_use_beta_header.py b/tests/pass_through_unit_tests/test_bedrock_tool_use_beta_header.py index 635ace016fe..e36b2ce9a5a 100644 --- a/tests/pass_through_unit_tests/test_bedrock_tool_use_beta_header.py +++ b/tests/pass_through_unit_tests/test_bedrock_tool_use_beta_header.py @@ -40,30 +40,30 @@ async def test_bedrock_sonnet_4_5_with_advanced_tool_use_beta_header(): print(f"✅ Test passed! Response: {response}") -@pytest.mark.asyncio -async def test_bedrock_claude_3_5_with_advanced_tool_use_beta_header_filtered(): - """ - Simple E2E test: Call Bedrock Claude 3.5 with advanced-tool-use beta header. +# @pytest.mark.asyncio +# async def test_bedrock_claude_3_5_with_advanced_tool_use_beta_header_filtered(): +# """ +# Simple E2E test: Call Bedrock Claude 3.5 with advanced-tool-use beta header. - This should work because the beta header is filtered out by LiteLLM before - sending the request to Bedrock Invoke API. - """ +# This should work because the beta header is filtered out by LiteLLM before +# sending the request to Bedrock Invoke API. +# """ - response = await litellm.anthropic.messages.acreate( - model="bedrock/invoke/us.anthropic.claude-3-5-sonnet-20240620-v1:0", - messages=[{"role": "user", "content": "What is 2+2?"}], - max_tokens=100, - provider_specific_header={ - "custom_llm_provider": "bedrock", - "extra_headers": { - "anthropic-beta": "advanced-tool-use-2025-11-20", - }, - }, - ) +# response = await litellm.anthropic.messages.acreate( +# model="bedrock/invoke/us.anthropic.claude-3-5-sonnet-20240620-v1:0", +# messages=[{"role": "user", "content": "What is 2+2?"}], +# max_tokens=100, +# provider_specific_header={ +# "custom_llm_provider": "bedrock", +# "extra_headers": { +# "anthropic-beta": "advanced-tool-use-2025-11-20", +# }, +# }, +# ) - # Verify response - assert response is not None - assert "content" in response - print(f"✅ Test passed! Claude 3.5 response (beta header filtered): {response}") +# # Verify response +# assert response is not None +# assert "content" in response +# print(f"✅ Test passed! Claude 3.5 response (beta header filtered): {response}") diff --git a/tests/proxy_unit_tests/test_proxy_utils.py b/tests/proxy_unit_tests/test_proxy_utils.py index 54b9e31a6da..9aecfe9886b 100644 --- a/tests/proxy_unit_tests/test_proxy_utils.py +++ b/tests/proxy_unit_tests/test_proxy_utils.py @@ -1845,38 +1845,38 @@ def test_provider_specific_header_multi_provider(): } -@pytest.mark.parametrize( - "custom_llm_provider, expected_result", - [ - ("anthropic", {"anthropic-beta": "test"}), - ("bedrock", {"anthropic-beta": "test"}), - ("vertex_ai", {"anthropic-beta": "test"}), - ], -) -def test_provider_specific_header_in_request(custom_llm_provider, expected_result): - from litellm.types.utils import ProviderSpecificHeader - from litellm.llms.custom_httpx.http_handler import HTTPHandler - from unittest.mock import patch +# @pytest.mark.parametrize( +# "custom_llm_provider, expected_result", +# [ +# ("anthropic", {"anthropic-beta": "test"}), +# ("bedrock", {"anthropic-beta": "test"}), +# ("vertex_ai", {"anthropic-beta": "test"}), +# ], +# ) +# def test_provider_specific_header_in_request(custom_llm_provider, expected_result): +# from litellm.types.utils import ProviderSpecificHeader +# from litellm.llms.custom_httpx.http_handler import HTTPHandler +# from unittest.mock import patch - litellm.set_verbose = True - client = HTTPHandler() - with patch.object(client, "post", return_value=MagicMock()) as mock_post: - try: - litellm.completion( - model="anthropic/claude-3-5-sonnet-v2@20241022", - messages=[{"role": "user", "content": "Hello world"}], - provider_specific_header=ProviderSpecificHeader( - custom_llm_provider="anthropic", - extra_headers={"anthropic-beta": "test"}, - ), - client=client, - ) - except Exception as e: - print(f"Error: {e}") +# litellm.set_verbose = True +# client = HTTPHandler() +# with patch.object(client, "post", return_value=MagicMock()) as mock_post: +# try: +# litellm.completion( +# model="anthropic/claude-3-5-sonnet-v2@20241022", +# messages=[{"role": "user", "content": "Hello world"}], +# provider_specific_header=ProviderSpecificHeader( +# custom_llm_provider="anthropic", +# extra_headers={"anthropic-beta": "test"}, +# ), +# client=client, +# ) +# except Exception as e: +# print(f"Error: {e}") - mock_post.assert_called_once() - print(mock_post.call_args.kwargs["headers"]) - assert "anthropic-beta" in mock_post.call_args.kwargs["headers"] +# mock_post.assert_called_once() +# print(mock_post.call_args.kwargs["headers"]) +# assert "anthropic-beta" in mock_post.call_args.kwargs["headers"] from litellm.proxy._types import LiteLLM_UserTable diff --git a/tests/test_litellm/integrations/arize/test_arize_otel_coexistence.py b/tests/test_litellm/integrations/arize/test_arize_otel_coexistence.py new file mode 100644 index 00000000000..a0dcf1c091c --- /dev/null +++ b/tests/test_litellm/integrations/arize/test_arize_otel_coexistence.py @@ -0,0 +1,169 @@ +""" +Tests that Arize Phoenix / Arize and the generic ``otel`` callback can +coexist, each sending spans to their own independent exporter. + +Covers the three root-cause fixes: +1. ArizePhoenixLogger / ArizeLogger create *dedicated* TracerProviders. +2. The ``otel`` dedup check does NOT match Arize subclasses. +3. Arize loggers do NOT overwrite ``proxy_server.open_telemetry_logger``. +""" + +import unittest +from unittest.mock import patch + +from opentelemetry.sdk.trace import TracerProvider +from opentelemetry.sdk.trace.export import SimpleSpanProcessor +from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter + +from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig + + +# --------------------------------------------------------------------------- +# Helpers +# --------------------------------------------------------------------------- + +def _make_otel_logger(exporter: InMemorySpanExporter) -> OpenTelemetry: + """Create a generic ``otel`` callback backed by an in-memory exporter. + + We build a dedicated TracerProvider explicitly so the test is isolated + from whatever global provider state may exist. + """ + provider = TracerProvider() + provider.add_span_processor(SimpleSpanProcessor(exporter)) + config = OpenTelemetryConfig(exporter=exporter) + return OpenTelemetry(config=config, callback_name="otel", tracer_provider=provider) + + +def _make_arize_phoenix_logger(exporter: InMemorySpanExporter): + """Create an ``arize_phoenix`` callback backed by an in-memory exporter. + + ArizePhoenixLogger._init_tracing creates its own TracerProvider, so we + pass the exporter via config and let it build the provider internally. + """ + from litellm.integrations.arize.arize_phoenix import ArizePhoenixLogger + + config = OpenTelemetryConfig(exporter=exporter) + return ArizePhoenixLogger(config=config, callback_name="arize_phoenix") + + +def _make_arize_logger(exporter: InMemorySpanExporter): + """Create an ``arize`` callback backed by an in-memory exporter. + + ArizeLogger._init_tracing creates its own TracerProvider, so we pass + the exporter via config and let it build the provider internally. + """ + from litellm.integrations.arize.arize import ArizeLogger + + config = OpenTelemetryConfig(exporter=exporter) + return ArizeLogger(config=config, callback_name="arize") + + +# --------------------------------------------------------------------------- +# Tests +# --------------------------------------------------------------------------- + +class TestIndependentTracerProviders(unittest.TestCase): + """Each integration must get its own TracerProvider so spans go to the right exporter.""" + + def test_otel_and_arize_phoenix_have_different_tracer_providers(self): + otel_exporter = InMemorySpanExporter() + phoenix_exporter = InMemorySpanExporter() + + otel_logger = _make_otel_logger(otel_exporter) + phoenix_logger = _make_arize_phoenix_logger(phoenix_exporter) + + # The tracers must come from different providers + assert otel_logger.tracer is not phoenix_logger.tracer + + def test_otel_and_arize_have_different_tracer_providers(self): + otel_exporter = InMemorySpanExporter() + arize_exporter = InMemorySpanExporter() + + otel_logger = _make_otel_logger(otel_exporter) + arize_logger = _make_arize_logger(arize_exporter) + + assert otel_logger.tracer is not arize_logger.tracer + + def test_arize_phoenix_and_arize_have_different_tracer_providers(self): + phoenix_exporter = InMemorySpanExporter() + arize_exporter = InMemorySpanExporter() + + phoenix_logger = _make_arize_phoenix_logger(phoenix_exporter) + arize_logger = _make_arize_logger(arize_exporter) + + assert phoenix_logger.tracer is not arize_logger.tracer + + +class TestSpansRoutedToCorrectExporter(unittest.TestCase): + """Spans created by each logger must land in its own exporter, not the other's.""" + + def test_spans_go_to_respective_exporters(self): + otel_exporter = InMemorySpanExporter() + phoenix_exporter = InMemorySpanExporter() + + otel_logger = _make_otel_logger(otel_exporter) + phoenix_logger = _make_arize_phoenix_logger(phoenix_exporter) + + # Create a span on each — SimpleSpanProcessor exports synchronously on end() + otel_span = otel_logger.tracer.start_span("otel_test_span") + otel_span.end() + + phoenix_span = phoenix_logger.tracer.start_span("phoenix_test_span") + phoenix_span.end() + + # Read spans *before* shutdown (shutdown clears the in-memory store) + otel_span_names = [s.name for s in otel_exporter.get_finished_spans()] + phoenix_span_names = [s.name for s in phoenix_exporter.get_finished_spans()] + + assert "otel_test_span" in otel_span_names + assert "phoenix_test_span" not in otel_span_names + + assert "phoenix_test_span" in phoenix_span_names + assert "otel_test_span" not in phoenix_span_names + + +class TestOtelDedupCheck(unittest.TestCase): + """The ``otel`` callback dedup must use exact type check, not isinstance.""" + + def test_arize_phoenix_logger_is_not_matched_by_otel_dedup(self): + from litellm.integrations.arize.arize_phoenix import ArizePhoenixLogger + + phoenix_logger = _make_arize_phoenix_logger(InMemorySpanExporter()) + + # isinstance would match — but type() must not + assert isinstance(phoenix_logger, OpenTelemetry) + assert type(phoenix_logger) is not OpenTelemetry + + def test_arize_logger_is_not_matched_by_otel_dedup(self): + from litellm.integrations.arize.arize import ArizeLogger + + arize_logger = _make_arize_logger(InMemorySpanExporter()) + + assert isinstance(arize_logger, OpenTelemetry) + assert type(arize_logger) is not OpenTelemetry + + def test_otel_logger_matches_own_dedup(self): + otel_logger = _make_otel_logger(InMemorySpanExporter()) + assert type(otel_logger) is OpenTelemetry + + +class TestProxyLoggerNotOverwritten(unittest.TestCase): + """Arize / Phoenix must not overwrite ``proxy_server.open_telemetry_logger``.""" + + @patch("litellm.proxy.proxy_server.open_telemetry_logger", None) + def test_arize_phoenix_does_not_set_proxy_otel_logger(self): + from litellm.proxy import proxy_server + + _make_arize_phoenix_logger(InMemorySpanExporter()) + assert proxy_server.open_telemetry_logger is None + + @patch("litellm.proxy.proxy_server.open_telemetry_logger", None) + def test_arize_does_not_set_proxy_otel_logger(self): + from litellm.proxy import proxy_server + + _make_arize_logger(InMemorySpanExporter()) + assert proxy_server.open_telemetry_logger is None + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_litellm/integrations/arize/test_arize_phoenix.py b/tests/test_litellm/integrations/arize/test_arize_phoenix.py index aa227fbff5e..129b35fb06a 100644 --- a/tests/test_litellm/integrations/arize/test_arize_phoenix.py +++ b/tests/test_litellm/integrations/arize/test_arize_phoenix.py @@ -1,5 +1,5 @@ import unittest -from unittest.mock import patch +from unittest.mock import MagicMock, patch import pytest @@ -7,6 +7,7 @@ from litellm.integrations.arize.arize_phoenix import ( ArizePhoenixConfig, ArizePhoenixLogger, ) +from litellm.integrations.arize._utils import ArizeOTELAttributes class TestArizePhoenixConfig(unittest.TestCase): @@ -195,5 +196,63 @@ def test_get_arize_phoenix_config_expection_on_missing_api_key(monkeypatch, env_ +# --------------------------------------------------------------------------- +# Dynamic project naming from metadata +# --------------------------------------------------------------------------- + + +class TestGetDynamicProjectName: + """Tests for _get_dynamic_project_name extraction logic.""" + + def test_extracts_from_standard_logging_object_metadata(self): + kwargs = { + "standard_logging_object": { + "metadata": {"phoenix_project_name": "my-project"}, + } + } + assert ArizePhoenixLogger._get_dynamic_project_name(kwargs) == "my-project" + + def test_extracts_from_litellm_params_metadata(self): + kwargs = { + "litellm_params": { + "metadata": {"phoenix_project_name": "sdk-project"}, + } + } + assert ArizePhoenixLogger._get_dynamic_project_name(kwargs) == "sdk-project" + + def test_returns_none_when_no_metadata(self): + assert ArizePhoenixLogger._get_dynamic_project_name({}) is None + + def test_non_dict_standard_logging_object_does_not_raise(self): + """isinstance(dict) guard prevents AttributeError on non-dict payloads.""" + kwargs = {"standard_logging_object": "not-a-dict"} + assert ArizePhoenixLogger._get_dynamic_project_name(kwargs) is None + + +class TestDynamicProjectNameOnSpan: + """set_arize_phoenix_attributes sets openinference.project.name on the span.""" + + @patch.dict("os.environ", {"PHOENIX_PROJECT_NAME": "env-fallback"}, clear=False) + @patch("litellm.integrations.arize._utils.set_attributes") + def test_dynamic_name_sets_span_attribute(self, _mock_set_attrs): + span = MagicMock() + kwargs = { + "standard_logging_object": { + "metadata": {"phoenix_project_name": "dynamic-proj"}, + } + } + ArizePhoenixLogger.set_arize_phoenix_attributes(span, kwargs, response_obj=None) + + span.set_attribute.assert_called_once_with("openinference.project.name", "dynamic-proj") + + @patch.dict("os.environ", {"PHOENIX_PROJECT_NAME": "env-project"}, clear=False) + @patch("litellm.integrations.arize._utils.set_attributes") + def test_falls_back_to_env_var_when_no_dynamic_name(self, _mock_set_attrs): + span = MagicMock() + ArizePhoenixLogger.set_arize_phoenix_attributes(span, {}, response_obj=None) + + span.set_attribute.assert_called_once_with("openinference.project.name", "env-project") + + if __name__ == "__main__": unittest.main() diff --git a/tests/test_litellm/litellm_core_utils/test_redact_messages.py b/tests/test_litellm/litellm_core_utils/test_redact_messages.py new file mode 100644 index 00000000000..d7df7823aee --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_redact_messages.py @@ -0,0 +1,145 @@ +""" +Tests for litellm.litellm_core_utils.redact_messages.should_redact_message_logging + +Covers the proxy flow where headers arrive in litellm_params["metadata"]["headers"] +but litellm_params["litellm_metadata"] is None. +""" + +import pytest + +import litellm +from litellm.litellm_core_utils.redact_messages import should_redact_message_logging + + +@pytest.fixture(autouse=True) +def _reset_global_redaction(): + """Ensure the global setting is off for every test.""" + original = litellm.turn_off_message_logging + litellm.turn_off_message_logging = False + yield + litellm.turn_off_message_logging = original + + +def _make_model_call_details( + metadata_headers=None, + litellm_metadata=None, + metadata=None, + standard_callback_dynamic_params=None, +): + """Build a model_call_details dict that mimics real proxy/SDK flows.""" + litellm_params = {} + if metadata is not None: + litellm_params["metadata"] = metadata + elif metadata_headers is not None: + litellm_params["metadata"] = {"headers": metadata_headers} + else: + litellm_params["metadata"] = {} + + # get_litellm_params always sets this key (even when value is None) + litellm_params["litellm_metadata"] = litellm_metadata + + details = {"litellm_params": litellm_params} + if standard_callback_dynamic_params is not None: + details["standard_callback_dynamic_params"] = standard_callback_dynamic_params + return details + + +class TestShouldRedactMessageLogging: + """Unit tests for should_redact_message_logging().""" + + # ---- proxy flow: headers in metadata, litellm_metadata is None ---- + + def test_enable_redaction_via_x_header_proxy_flow(self): + """x-litellm-enable-message-redaction header should enable redaction + even when litellm_metadata is None (proxy path).""" + details = _make_model_call_details( + metadata_headers={"x-litellm-enable-message-redaction": "true"}, + litellm_metadata=None, + ) + assert should_redact_message_logging(details) is True + + def test_enable_redaction_via_old_header_proxy_flow(self): + """litellm-enable-message-redaction header should enable redaction + even when litellm_metadata is None (proxy path).""" + details = _make_model_call_details( + metadata_headers={"litellm-enable-message-redaction": "true"}, + litellm_metadata=None, + ) + assert should_redact_message_logging(details) is True + + def test_disable_redaction_via_header_proxy_flow(self): + """litellm-disable-message-redaction should suppress redaction + even when global setting is on, and litellm_metadata is None.""" + litellm.turn_off_message_logging = True + details = _make_model_call_details( + metadata_headers={"litellm-disable-message-redaction": "true"}, + litellm_metadata=None, + ) + assert should_redact_message_logging(details) is False + + # ---- SDK direct-call flow: headers in litellm_metadata ---- + + def test_enable_redaction_via_header_in_litellm_metadata(self): + """Headers inside litellm_metadata (SDK direct call) should work.""" + details = _make_model_call_details( + litellm_metadata={"headers": {"x-litellm-enable-message-redaction": "true"}}, + ) + assert should_redact_message_logging(details) is True + + # ---- no headers at all ---- + + def test_no_headers_defaults_to_global_off(self): + """Without headers, falls back to global setting (False).""" + details = _make_model_call_details( + metadata_headers=None, + litellm_metadata=None, + ) + assert should_redact_message_logging(details) is False + + def test_no_headers_global_on(self): + """Without headers, respects global turn_off_message_logging=True.""" + litellm.turn_off_message_logging = True + details = _make_model_call_details( + metadata_headers=None, + litellm_metadata=None, + ) + assert should_redact_message_logging(details) is True + + # ---- dynamic params take precedence ---- + + def test_dynamic_param_enables_redaction(self): + """Dynamic turn_off_message_logging=True should enable redaction.""" + details = _make_model_call_details( + metadata_headers={}, + litellm_metadata=None, + standard_callback_dynamic_params={"turn_off_message_logging": True}, + ) + assert should_redact_message_logging(details) is True + + def test_dynamic_param_false_overrides_header(self): + """Dynamic turn_off_message_logging=False should take precedence over enable header.""" + details = _make_model_call_details( + metadata_headers={"x-litellm-enable-message-redaction": "true"}, + litellm_metadata=None, + standard_callback_dynamic_params={"turn_off_message_logging": False}, + ) + assert should_redact_message_logging(details) is False + + # ---- non-dict metadata safety ---- + + def test_both_metadata_fields_none(self): + """When both litellm_metadata and metadata are None, should not raise.""" + details = _make_model_call_details( + metadata=None, + litellm_metadata=None, + ) + assert should_redact_message_logging(details) is False + + def test_both_metadata_fields_none_global_on(self): + """When both metadata fields are None but global is on, should still return True.""" + litellm.turn_off_message_logging = True + details = _make_model_call_details( + metadata=None, + litellm_metadata=None, + ) + assert should_redact_message_logging(details) is True diff --git a/tests/test_litellm/llms/azure/chat/test_azure_chat_gpt_transformation.py b/tests/test_litellm/llms/azure/chat/test_azure_chat_gpt_transformation.py index 8df35a37514..7be4d6dfcf2 100644 --- a/tests/test_litellm/llms/azure/chat/test_azure_chat_gpt_transformation.py +++ b/tests/test_litellm/llms/azure/chat/test_azure_chat_gpt_transformation.py @@ -30,6 +30,19 @@ class TestAzureOpenAIConfig: assert not config._is_response_format_supported_model("gpt-35-turbo") + def test_prompt_cache_key_supported(self): + """Test that 'prompt_cache_key' is in supported params for Azure OpenAI chat completion models. + + OpenAI's Chat Completions API supports prompt_cache_key for cache routing optimization. + """ + config = AzureOpenAIConfig() + supported_params = config.get_supported_openai_params("gpt-4.1-nano") + assert "prompt_cache_key" in supported_params + + supported_params = config.get_supported_openai_params("gpt-4.1") + assert "prompt_cache_key" in supported_params + + def test_map_openai_params_with_preview_api_version(): config = AzureOpenAIConfig() non_default_params = { diff --git a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py index c91ef31bba5..d903d7c85f1 100644 --- a/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py +++ b/tests/test_litellm/llms/azure_ai/chat/test_azure_ai_transformation.py @@ -45,6 +45,58 @@ def test_azure_ai_validate_environment(): assert headers["Content-Type"] == "application/json" +def test_azure_ai_validate_environment_with_api_key(): + """ + Test that when api_key is provided, it is set in the api-key header + for Azure Foundry endpoints (.services.ai.azure.com). + """ + config = AzureAIStudioConfig() + headers = config.validate_environment( + headers={}, + model="Kimi-K2.5", + messages=[], + optional_params={}, + litellm_params={}, + api_key="test-api-key", + api_base="https://my-endpoint.services.ai.azure.com", + ) + assert headers["api-key"] == "test-api-key" + assert headers["Content-Type"] == "application/json" + + +def test_azure_ai_validate_environment_with_azure_ad_token(): + """ + Test that when no api_key is provided but Azure AD credentials are available, + the Authorization header is set with a Bearer token. + + Regression test for https://github.com/BerriAI/litellm/issues/20759 + """ + import litellm + + config = AzureAIStudioConfig() + with patch( + "litellm.llms.azure.common_utils.get_azure_ad_token", + return_value="fake-azure-ad-token", + ), patch( + "litellm.llms.azure.common_utils.get_secret_str", + return_value=None, + ), patch.object(litellm, "api_key", None), patch.object( + litellm, "azure_key", None + ): + headers = config.validate_environment( + headers={}, + model="Kimi-K2.5", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + api_base="https://my-endpoint.services.ai.azure.com", + ) + assert headers.get("Authorization") == "Bearer fake-azure-ad-token" + assert "api-key" not in headers + assert headers["Content-Type"] == "application/json" + + def test_azure_ai_grok_stop_parameter_handling(): """ Test that Grok models properly handle stop parameter filtering in Azure AI Studio. diff --git a/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py b/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py index edfdeb08d82..d2fb45643de 100644 --- a/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/invoke_transformations/test_bedrock_chat_invoke_transformations_anthropic_claude3_transformation.py @@ -281,153 +281,6 @@ def test_output_format_with_no_schema(): assert last_user_message["content"] == "Hello" -def test_advanced_tool_use_header_translation_for_opus_4_5(): - """ - Test that advanced-tool-use-2025-11-20 header is translated to Bedrock-specific headers - for Claude Opus 4.5. - - Regression test for: Claude Code sends advanced-tool-use header which needs to be - translated to tool-search-tool-2025-10-19 and tool-examples-2025-10-29 for Bedrock - Invoke API on Claude Opus 4.5. - - Ref: https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-anthropic-claude-messages-request-response.html - """ - from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import ( - AmazonAnthropicClaudeMessagesConfig, - ) - - config = AmazonAnthropicClaudeMessagesConfig() - - messages = [ - {"role": "user", "content": "What's the weather like?"} - ] - - anthropic_messages_optional_request_params = { - "max_tokens": 100, - } - - # Simulate advanced-tool-use header from Claude Code - headers = { - "anthropic-beta": "advanced-tool-use-2025-11-20" - } - - # Test with Claude Opus 4.5 - result = config.transform_anthropic_messages_request( - model="anthropic.claude-opus-4-5-20250514-v1:0", - messages=messages, - anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, - litellm_params={}, - headers=headers, - ) - - # Verify advanced-tool-use header was removed - assert "anthropic_beta" in result - beta_headers = result["anthropic_beta"] - assert "advanced-tool-use-2025-11-20" not in beta_headers, \ - "advanced-tool-use header should be removed for Bedrock" - - # Verify Bedrock-specific headers were added - assert "tool-search-tool-2025-10-19" in beta_headers, \ - "tool-search-tool-2025-10-19 should be added for Opus 4.5" - assert "tool-examples-2025-10-29" in beta_headers, \ - "tool-examples-2025-10-29 should be added for Opus 4.5" - - -def test_advanced_tool_use_header_filtered_for_non_opus_4_5(): - """ - Test that advanced-tool-use-2025-11-20 header is filtered out for models - that don't support tool search on Bedrock. - - Tool search is supported on: Claude Opus 4.5, Claude Sonnet 4.5 - Tool search is NOT supported on: Claude 3.5 Sonnet and earlier - """ - from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import ( - AmazonAnthropicClaudeMessagesConfig, - ) - - config = AmazonAnthropicClaudeMessagesConfig() - - messages = [ - {"role": "user", "content": "What's the weather like?"} - ] - - anthropic_messages_optional_request_params = { - "max_tokens": 100, - } - - # Simulate advanced-tool-use header from Claude Code - headers = { - "anthropic-beta": "advanced-tool-use-2025-11-20" - } - - # Test with Claude 3.5 Sonnet (does NOT support tool search on Bedrock) - result = config.transform_anthropic_messages_request( - model="anthropic.claude-3-5-sonnet-20241022-v2:0", - messages=messages, - anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, - litellm_params={}, - headers=headers, - ) - - # Verify advanced-tool-use header was removed - beta_headers = result.get("anthropic_beta", []) - assert "advanced-tool-use-2025-11-20" not in beta_headers, \ - "advanced-tool-use header should be removed for Bedrock" - - # Verify Bedrock-specific headers were NOT added (only for Opus 4.5 and Sonnet 4.5) - assert "tool-search-tool-2025-10-19" not in beta_headers, \ - "tool-search-tool should not be added for models without tool search support" - assert "tool-examples-2025-10-29" not in beta_headers, \ - "tool-examples should not be added for models without tool search support" - - -def test_advanced_tool_use_header_translation_with_multiple_beta_headers(): - """ - Test that advanced-tool-use header translation works correctly when multiple - beta headers are present. - """ - from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_transformation import ( - AmazonAnthropicClaudeMessagesConfig, - ) - - config = AmazonAnthropicClaudeMessagesConfig() - - messages = [ - {"role": "user", "content": "What's the weather like?"} - ] - - anthropic_messages_optional_request_params = { - "max_tokens": 100, - } - - # Multiple beta headers including advanced-tool-use - headers = { - "anthropic-beta": "claude-code-20250219,advanced-tool-use-2025-11-20,interleaved-thinking-2025-05-14" - } - - # Test with Claude Opus 4.5 - result = config.transform_anthropic_messages_request( - model="anthropic.claude-opus-4-5-20250514-v1:0", - messages=messages, - anthropic_messages_optional_request_params=anthropic_messages_optional_request_params, - litellm_params={}, - headers=headers, - ) - - beta_headers = result.get("anthropic_beta", []) - - # Verify advanced-tool-use was removed - assert "advanced-tool-use-2025-11-20" not in beta_headers - - # Verify Bedrock-specific headers were added - assert "tool-search-tool-2025-10-19" in beta_headers - assert "tool-examples-2025-10-29" in beta_headers - - # Verify other beta headers are preserved - assert "claude-code-20250219" in beta_headers - assert "interleaved-thinking-2025-05-14" in beta_headers - - def test_opus_4_5_model_detection(): """ Test that the _is_claude_opus_4_5 method correctly identifies Opus 4.5 models @@ -466,71 +319,71 @@ def test_opus_4_5_model_detection(): f"Should not detect {model} as Opus 4.5" -def test_structured_outputs_beta_header_filtered_for_bedrock_invoke(): - """ - Test that unsupported beta headers are filtered out for Bedrock Invoke API. +# def test_structured_outputs_beta_header_filtered_for_bedrock_invoke(): +# """ +# Test that unsupported beta headers are filtered out for Bedrock Invoke API. - Bedrock Invoke API only supports a specific whitelist of beta flags and returns - "invalid beta flag" error for others (e.g., structured-outputs, mcp-servers). - This test ensures unsupported headers are filtered while keeping supported ones. +# Bedrock Invoke API only supports a specific whitelist of beta flags and returns +# "invalid beta flag" error for others (e.g., structured-outputs, mcp-servers). +# This test ensures unsupported headers are filtered while keeping supported ones. - Fixes: https://github.com/BerriAI/litellm/issues/16726 - """ - config = AmazonAnthropicClaudeConfig() +# Fixes: https://github.com/BerriAI/litellm/issues/16726 +# """ +# config = AmazonAnthropicClaudeConfig() - messages = [{"role": "user", "content": "test"}] +# messages = [{"role": "user", "content": "test"}] - # Test 1: structured-outputs beta header (unsupported) - headers = {"anthropic-beta": "structured-outputs-2025-11-13"} +# # Test 1: structured-outputs beta header (unsupported) +# headers = {"anthropic-beta": "structured-outputs-2025-11-13"} - result = config.transform_request( - model="anthropic.claude-4-0-sonnet-20250514-v1:0", - messages=messages, - optional_params={}, - litellm_params={}, - headers=headers, - ) +# result = config.transform_request( +# model="anthropic.claude-4-0-sonnet-20250514-v1:0", +# messages=messages, +# optional_params={}, +# litellm_params={}, +# headers=headers, +# ) - # Verify structured-outputs beta is filtered out - anthropic_beta = result.get("anthropic_beta", []) - assert not any("structured-outputs" in beta for beta in anthropic_beta), \ - f"structured-outputs beta should be filtered, got: {anthropic_beta}" +# # Verify structured-outputs beta is filtered out +# anthropic_beta = result.get("anthropic_beta", []) +# assert not any("structured-outputs" in beta for beta in anthropic_beta), \ +# f"structured-outputs beta should be filtered, got: {anthropic_beta}" - # Test 2: mcp-servers beta header (unsupported - the main issue from #16726) - headers = {"anthropic-beta": "mcp-servers-2025-12-04"} +# # Test 2: mcp-servers beta header (unsupported - the main issue from #16726) +# headers = {"anthropic-beta": "mcp-servers-2025-12-04"} - result = config.transform_request( - model="anthropic.claude-4-0-sonnet-20250514-v1:0", - messages=messages, - optional_params={}, - litellm_params={}, - headers=headers, - ) +# result = config.transform_request( +# model="anthropic.claude-4-0-sonnet-20250514-v1:0", +# messages=messages, +# optional_params={}, +# litellm_params={}, +# headers=headers, +# ) - # Verify mcp-servers beta is filtered out - anthropic_beta = result.get("anthropic_beta", []) - assert not any("mcp-servers" in beta for beta in anthropic_beta), \ - f"mcp-servers beta should be filtered, got: {anthropic_beta}" +# # Verify mcp-servers beta is filtered out +# anthropic_beta = result.get("anthropic_beta", []) +# assert not any("mcp-servers" in beta for beta in anthropic_beta), \ +# f"mcp-servers beta should be filtered, got: {anthropic_beta}" - # Test 3: Mix of supported and unsupported beta headers - headers = {"anthropic-beta": "computer-use-2024-10-22,mcp-servers-2025-12-04,structured-outputs-2025-11-13"} +# # Test 3: Mix of supported and unsupported beta headers +# headers = {"anthropic-beta": "computer-use-2024-10-22,mcp-servers-2025-12-04,structured-outputs-2025-11-13"} - result = config.transform_request( - model="anthropic.claude-4-0-sonnet-20250514-v1:0", - messages=messages, - optional_params={}, - litellm_params={}, - headers=headers, - ) +# result = config.transform_request( +# model="anthropic.claude-4-0-sonnet-20250514-v1:0", +# messages=messages, +# optional_params={}, +# litellm_params={}, +# headers=headers, +# ) - # Verify only supported betas are kept - anthropic_beta = result.get("anthropic_beta", []) - assert not any("structured-outputs" in beta for beta in anthropic_beta), \ - f"structured-outputs beta should be filtered, got: {anthropic_beta}" - assert not any("mcp-servers" in beta for beta in anthropic_beta), \ - f"mcp-servers beta should be filtered, got: {anthropic_beta}" - assert any("computer-use" in beta for beta in anthropic_beta), \ - f"computer-use beta should be kept, got: {anthropic_beta}" +# # Verify only supported betas are kept +# anthropic_beta = result.get("anthropic_beta", []) +# assert not any("structured-outputs" in beta for beta in anthropic_beta), \ +# f"structured-outputs beta should be filtered, got: {anthropic_beta}" +# assert not any("mcp-servers" in beta for beta in anthropic_beta), \ +# f"mcp-servers beta should be filtered, got: {anthropic_beta}" +# assert any("computer-use" in beta for beta in anthropic_beta), \ +# f"computer-use beta should be kept, got: {anthropic_beta}" def test_output_format_removed_from_bedrock_invoke_request(): diff --git a/tests/test_litellm/llms/bedrock/test_anthropic_beta_support.py b/tests/test_litellm/llms/bedrock/test_anthropic_beta_support.py index e7b6de29b6b..074a319a603 100644 --- a/tests/test_litellm/llms/bedrock/test_anthropic_beta_support.py +++ b/tests/test_litellm/llms/bedrock/test_anthropic_beta_support.py @@ -389,104 +389,4 @@ class TestAnthropicBetaHeaderSupport: assert "anthropic_beta" in additional_fields, ( "anthropic_beta SHOULD be added for Anthropic models with cross-region prefix." ) - assert "context-1m-2025-08-07" in additional_fields["anthropic_beta"] - - def test_messages_advanced_tool_use_translation_opus_4_5(self): - """Test that advanced-tool-use header is translated to Bedrock-specific headers for Opus 4.5. - - Regression test for: Claude Code sends advanced-tool-use-2025-11-20 header which needs - to be translated to tool-search-tool-2025-10-19 and tool-examples-2025-10-29 for - Bedrock Invoke API on Claude Opus 4.5. - - Ref: https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-anthropic-claude-messages-request-response.html - """ - config = AmazonAnthropicClaudeMessagesConfig() - headers = {"anthropic-beta": "advanced-tool-use-2025-11-20"} - - result = config.transform_anthropic_messages_request( - model="us.anthropic.claude-opus-4-5-20250514-v1:0", - messages=[{"role": "user", "content": "Test"}], - anthropic_messages_optional_request_params={"max_tokens": 100}, - litellm_params={}, - headers=headers - ) - - assert "anthropic_beta" in result - beta_headers = result["anthropic_beta"] - - # advanced-tool-use should be removed - assert "advanced-tool-use-2025-11-20" not in beta_headers, ( - "advanced-tool-use-2025-11-20 should be removed for Bedrock Invoke API" - ) - - # Bedrock-specific headers should be added for Opus 4.5 - assert "tool-search-tool-2025-10-19" in beta_headers, ( - "tool-search-tool-2025-10-19 should be added for Opus 4.5" - ) - assert "tool-examples-2025-10-29" in beta_headers, ( - "tool-examples-2025-10-29 should be added for Opus 4.5" - ) - - def test_messages_advanced_tool_use_translation_sonnet_4_5(self): - """Test that advanced-tool-use header is translated to Bedrock-specific headers for Sonnet 4.5. - - Regression test for: Claude Code sends advanced-tool-use-2025-11-20 header which needs - to be translated to tool-search-tool-2025-10-19 and tool-examples-2025-10-29 for - Bedrock Invoke API on Claude Sonnet 4.5. - - Ref: https://platform.claude.com/docs/en/agents-and-tools/tool-use/tool-search-tool - """ - config = AmazonAnthropicClaudeMessagesConfig() - headers = {"anthropic-beta": "advanced-tool-use-2025-11-20"} - - result = config.transform_anthropic_messages_request( - model="us.anthropic.claude-sonnet-4-5-20250929-v1:0", - messages=[{"role": "user", "content": "Test"}], - anthropic_messages_optional_request_params={"max_tokens": 100}, - litellm_params={}, - headers=headers - ) - - assert "anthropic_beta" in result - beta_headers = result["anthropic_beta"] - - # advanced-tool-use should be removed - assert "advanced-tool-use-2025-11-20" not in beta_headers, ( - "advanced-tool-use-2025-11-20 should be removed for Bedrock Invoke API" - ) - - # Bedrock-specific headers should be added for Sonnet 4.5 - assert "tool-search-tool-2025-10-19" in beta_headers, ( - "tool-search-tool-2025-10-19 should be added for Sonnet 4.5" - ) - assert "tool-examples-2025-10-29" in beta_headers, ( - "tool-examples-2025-10-29 should be added for Sonnet 4.5" - ) - - def test_messages_advanced_tool_use_filtered_unsupported_model(self): - """Test that advanced-tool-use header is filtered out for models that don't support tool search. - - The translation to Bedrock-specific headers should only happen for models that - support tool search on Bedrock (Opus 4.5, Sonnet 4.5). - For other models, the advanced-tool-use header should just be removed. - """ - config = AmazonAnthropicClaudeMessagesConfig() - headers = {"anthropic-beta": "advanced-tool-use-2025-11-20"} - - # Test with Claude 3.5 Sonnet (does NOT support tool search on Bedrock) - result = config.transform_anthropic_messages_request( - model="us.anthropic.claude-3-5-sonnet-20241022-v2:0", - messages=[{"role": "user", "content": "Test"}], - anthropic_messages_optional_request_params={"max_tokens": 100}, - litellm_params={}, - headers=headers - ) - - beta_headers = result.get("anthropic_beta", []) - - # advanced-tool-use should be removed - assert "advanced-tool-use-2025-11-20" not in beta_headers - - # Bedrock-specific headers should NOT be added for unsupported models - assert "tool-search-tool-2025-10-19" not in beta_headers - assert "tool-examples-2025-10-29" not in beta_headers + assert "context-1m-2025-08-07" in additional_fields["anthropic_beta"] \ No newline at end of file diff --git a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py index 77eda432513..cf9fee6bacf 100644 --- a/tests/test_litellm/llms/bedrock/test_base_aws_llm.py +++ b/tests/test_litellm/llms/bedrock/test_base_aws_llm.py @@ -853,29 +853,99 @@ def test_role_assumption_ttl_calculation(): assert 3500 <= ttl <= 3600 # Allow some variance for test execution time -def test_role_assumption_error_handling(): +def test_role_assumption_access_denied_falls_back_when_same_role(): """ - Test that role assumption errors are properly propagated. + Test that when AssumeRole fails with AccessDenied AND the caller is confirmed + to already be running as the target role, we fall back to ambient credentials. """ base_aws_llm = BaseAWSLLM() - - # Mock the boto3 STS client to raise an exception + + # Mock the boto3 STS client to raise AccessDenied mock_sts_client = MagicMock() - mock_sts_client.assume_role.side_effect = Exception("AccessDenied: User is not authorized to perform sts:AssumeRole") - + mock_sts_client.assume_role.side_effect = Exception( + "An error occurred (AccessDenied) when calling the AssumeRole operation: " + "Roles may not be assumed by root accounts." + ) + + # Mock _auth_with_env_vars to return fallback credentials + mock_creds = MagicMock() + mock_creds.access_key = "fallback-access-key" + mock_creds.secret_key = "fallback-secret-key" + + with patch("boto3.client", return_value=mock_sts_client): + with patch.object( + base_aws_llm, "_auth_with_env_vars", return_value=(mock_creds, None) + ) as mock_env_auth: + # _is_already_running_as_role returns True => fallback allowed + with patch.object( + base_aws_llm, "_is_already_running_as_role", return_value=True + ): + credentials, ttl = base_aws_llm._auth_with_aws_role( + aws_access_key_id=None, + aws_secret_access_key=None, + aws_session_token=None, + aws_role_name="arn:aws:iam::1111111111111:role/UnauthorizedRole", + aws_session_name="error-test-session", + ) + + # Should have fallen back to env vars + mock_env_auth.assert_called_once() + assert credentials.access_key == "fallback-access-key" + + +def test_role_assumption_access_denied_raises_when_different_role(): + """ + Test that when AssumeRole fails with AccessDenied but the caller is NOT + the same role, the error is re-raised (genuine permission failure). + """ + base_aws_llm = BaseAWSLLM() + + mock_sts_client = MagicMock() + mock_sts_client.assume_role.side_effect = Exception( + "An error occurred (AccessDenied) when calling the AssumeRole operation: " + "User is not authorized to perform sts:AssumeRole" + ) + + with patch("boto3.client", return_value=mock_sts_client): + # _is_already_running_as_role returns False => do NOT fallback + with patch.object( + base_aws_llm, "_is_already_running_as_role", return_value=False + ): + with pytest.raises(Exception) as exc_info: + base_aws_llm._auth_with_aws_role( + aws_access_key_id=None, + aws_secret_access_key=None, + aws_session_token=None, + aws_role_name="arn:aws:iam::999999999999:role/CrossAccountRole", + aws_session_name="error-test-session", + ) + + assert "AccessDenied" in str(exc_info.value) + + +def test_role_assumption_non_access_denied_error_propagated(): + """ + Test that non-AccessDenied errors from AssumeRole are still propagated. + """ + base_aws_llm = BaseAWSLLM() + + # Mock the boto3 STS client to raise a non-AccessDenied exception + mock_sts_client = MagicMock() + mock_sts_client.assume_role.side_effect = Exception( + "An error occurred (MalformedPolicyDocument) when calling the AssumeRole operation" + ) + with patch("boto3.client", return_value=mock_sts_client): - - # Should raise the exception with pytest.raises(Exception) as exc_info: base_aws_llm._auth_with_aws_role( aws_access_key_id=None, aws_secret_access_key=None, aws_session_token=None, - aws_role_name="arn:aws:iam::1111111111111:role/UnauthorizedRole", - aws_session_name="error-test-session" + aws_role_name="arn:aws:iam::1111111111111:role/BadPolicyRole", + aws_session_name="error-test-session", ) - - assert "AccessDenied" in str(exc_info.value) + + assert "MalformedPolicyDocument" in str(exc_info.value) def test_multiple_role_assumptions_in_sequence(): @@ -1195,3 +1265,251 @@ def test_converse_handler_external_id_extraction(): assert hasattr(mock_get_credentials, 'called_kwargs') assert "aws_external_id" in mock_get_credentials.called_kwargs assert mock_get_credentials.called_kwargs["aws_external_id"] == "TestExternalID123" + + +def test_is_already_running_as_role_irsa_same_role(): + """Test IRSA fast path: when AWS_ROLE_ARN matches target role.""" + base_aws_llm = BaseAWSLLM() + + with patch.dict(os.environ, { + "AWS_ROLE_ARN": "arn:aws:iam::123456789012:role/MyRole", + "AWS_WEB_IDENTITY_TOKEN_FILE": "/var/run/secrets/token", + }): + assert base_aws_llm._is_already_running_as_role( + "arn:aws:iam::123456789012:role/MyRole" + ) is True + + +def test_is_already_running_as_role_irsa_different_role(): + """Test IRSA fast path: when AWS_ROLE_ARN does NOT match target role.""" + base_aws_llm = BaseAWSLLM() + + with patch.dict(os.environ, { + "AWS_ROLE_ARN": "arn:aws:iam::123456789012:role/MyRole", + "AWS_WEB_IDENTITY_TOKEN_FILE": "/var/run/secrets/token", + }): + assert base_aws_llm._is_already_running_as_role( + "arn:aws:iam::999999999999:role/OtherRole" + ) is False + + +def test_is_already_running_as_role_ecs_task_role(): + """Test ECS/EC2 path: GetCallerIdentity shows assumed-role matching target.""" + base_aws_llm = BaseAWSLLM() + + mock_sts_client = MagicMock() + mock_sts_client.get_caller_identity.return_value = { + "Arn": "arn:aws:sts::123456789012:assumed-role/MyEcsTaskRole/ecs-task-id" + } + + with patch.dict(os.environ, {}, clear=False): + # Ensure no IRSA env vars + env = {k: v for k, v in os.environ.items() if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")} + with patch.dict(os.environ, env, clear=True): + with patch("boto3.client", return_value=mock_sts_client): + assert base_aws_llm._is_already_running_as_role( + "arn:aws:iam::123456789012:role/MyEcsTaskRole" + ) is True + + +def test_is_already_running_as_role_ecs_different_role(): + """Test ECS/EC2 path: GetCallerIdentity shows a different role.""" + base_aws_llm = BaseAWSLLM() + + mock_sts_client = MagicMock() + mock_sts_client.get_caller_identity.return_value = { + "Arn": "arn:aws:sts::123456789012:assumed-role/MyEcsTaskRole/ecs-task-id" + } + + with patch.dict(os.environ, {}, clear=False): + env = {k: v for k, v in os.environ.items() if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")} + with patch.dict(os.environ, env, clear=True): + with patch("boto3.client", return_value=mock_sts_client): + assert base_aws_llm._is_already_running_as_role( + "arn:aws:iam::999999999999:role/DifferentRole" + ) is False + + +def test_is_already_running_as_role_ecs_role_with_path(): + """Test ECS path with role that has a path prefix (e.g., /service-role/MyRole).""" + base_aws_llm = BaseAWSLLM() + + mock_sts_client = MagicMock() + mock_sts_client.get_caller_identity.return_value = { + "Arn": "arn:aws:sts::123456789012:assumed-role/MyEcsTaskRole/ecs-task-id" + } + + with patch.dict(os.environ, {}, clear=False): + env = {k: v for k, v in os.environ.items() if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")} + with patch.dict(os.environ, env, clear=True): + with patch("boto3.client", return_value=mock_sts_client): + # Role ARN with path + assert base_aws_llm._is_already_running_as_role( + "arn:aws:iam::123456789012:role/service-role/MyEcsTaskRole" + ) is True + + +def test_is_already_running_as_role_get_caller_identity_fails(): + """Test that when GetCallerIdentity fails, we return False (don't crash).""" + base_aws_llm = BaseAWSLLM() + + mock_sts_client = MagicMock() + mock_sts_client.get_caller_identity.side_effect = Exception("No credentials found") + + with patch.dict(os.environ, {}, clear=False): + env = {k: v for k, v in os.environ.items() if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")} + with patch.dict(os.environ, env, clear=True): + with patch("boto3.client", return_value=mock_sts_client): + assert base_aws_llm._is_already_running_as_role( + "arn:aws:iam::123456789012:role/SomeRole" + ) is False + + +def test_get_credentials_ecs_same_role_skips_assume_role(): + """ + End-to-end test: when running on ECS with the same role as aws_role_name, + get_credentials should use ambient credentials and NOT call AssumeRole. + """ + base_aws_llm = BaseAWSLLM() + + mock_creds = MagicMock() + mock_creds.access_key = "ecs-access-key" + mock_creds.secret_key = "ecs-secret-key" + mock_creds.token = "ecs-session-token" + + with patch.object( + base_aws_llm, + "_is_already_running_as_role", + return_value=True, + ): + with patch.object( + base_aws_llm, + "_auth_with_env_vars", + return_value=(mock_creds, None), + ) as mock_env_auth: + with patch.object( + base_aws_llm, + "_auth_with_aws_role", + ) as mock_role_auth: + credentials = base_aws_llm.get_credentials( + aws_role_name="arn:aws:iam::123456789012:role/MyEcsTaskRole", + aws_region_name="us-east-1", + ) + + # Should use env vars, NOT role assumption + mock_env_auth.assert_called_once() + mock_role_auth.assert_not_called() + assert credentials.access_key == "ecs-access-key" + + +def test_parse_arn_account_and_role_name(): + """Test the ARN parser helper for various ARN formats.""" + parse = BaseAWSLLM._parse_arn_account_and_role_name + + # Standard IAM role ARN + assert parse("arn:aws:iam::123456789012:role/MyRole") == ( + "aws", "123456789012", "MyRole" + ) + + # IAM role ARN with path + assert parse("arn:aws:iam::123456789012:role/service-role/MyRole") == ( + "aws", "123456789012", "MyRole" + ) + + # Assumed-role ARN (from GetCallerIdentity) + assert parse("arn:aws:sts::123456789012:assumed-role/MyRole/session-id") == ( + "aws", "123456789012", "MyRole" + ) + + # China partition + assert parse("arn:aws-cn:iam::123456789012:role/MyRole") == ( + "aws-cn", "123456789012", "MyRole" + ) + + # GovCloud partition + assert parse("arn:aws-us-gov:iam::123456789012:role/MyRole") == ( + "aws-us-gov", "123456789012", "MyRole" + ) + + # Invalid ARNs + assert parse("not-an-arn") is None + assert parse("arn:aws:iam::123456789012:user/MyUser") is None + assert parse("") is None + + +def test_is_already_running_as_role_cross_account_same_name(): + """ + Test that same role NAME in different accounts does NOT match. + This is the cross-account false-match prevention. + """ + base_aws_llm = BaseAWSLLM() + + mock_sts_client = MagicMock() + # Caller is in account 111111111111 + mock_sts_client.get_caller_identity.return_value = { + "Arn": "arn:aws:sts::111111111111:assumed-role/MyRole/session-id" + } + + with patch.dict(os.environ, {}, clear=False): + env = {k: v for k, v in os.environ.items() if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")} + with patch.dict(os.environ, env, clear=True): + with patch("boto3.client", return_value=mock_sts_client): + # Target is same role name but in account 222222222222 + assert base_aws_llm._is_already_running_as_role( + "arn:aws:iam::222222222222:role/MyRole" + ) is False + + +def test_is_already_running_as_role_cross_partition(): + """ + Test that same role name + account but different partition does NOT match. + """ + base_aws_llm = BaseAWSLLM() + + mock_sts_client = MagicMock() + mock_sts_client.get_caller_identity.return_value = { + "Arn": "arn:aws:sts::123456789012:assumed-role/MyRole/session-id" + } + + with patch.dict(os.environ, {}, clear=False): + env = {k: v for k, v in os.environ.items() if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")} + with patch.dict(os.environ, env, clear=True): + with patch("boto3.client", return_value=mock_sts_client): + # Same account and role but aws-cn partition + assert base_aws_llm._is_already_running_as_role( + "arn:aws-cn:iam::123456789012:role/MyRole" + ) is False + + +def test_is_already_running_as_role_invalid_target_arn(): + """ + Test that an unparseable target ARN returns False immediately. + """ + base_aws_llm = BaseAWSLLM() + + # Should return False without making any API calls + assert base_aws_llm._is_already_running_as_role("not-a-valid-arn") is False + + +def test_is_already_running_as_role_ssl_verify_passed(): + """ + Test that ssl_verify parameter is correctly passed to the STS client. + """ + base_aws_llm = BaseAWSLLM() + + mock_sts_client = MagicMock() + mock_sts_client.get_caller_identity.return_value = { + "Arn": "arn:aws:sts::123456789012:assumed-role/MyRole/session-id" + } + + with patch.dict(os.environ, {}, clear=False): + env = {k: v for k, v in os.environ.items() if k not in ("AWS_ROLE_ARN", "AWS_WEB_IDENTITY_TOKEN_FILE")} + with patch.dict(os.environ, env, clear=True): + with patch("boto3.client", return_value=mock_sts_client) as mock_boto3_client: + base_aws_llm._is_already_running_as_role( + "arn:aws:iam::123456789012:role/MyRole", + ssl_verify="/path/to/ca-bundle.crt", + ) + mock_boto3_client.assert_called_once_with( + "sts", verify="/path/to/ca-bundle.crt" + ) diff --git a/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py b/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py index a9c4bead820..eed42519622 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_cohere_tool_calls.py @@ -593,6 +593,113 @@ class TestOCICohereToolCalls: assert result.usage.total_tokens == 25 +class TestOCICoherePreambleOverride: + """Test Cohere system message handling via preambleOverride""" + + def test_single_system_message_sets_preamble_override(self): + """Test that a single system message is extracted into preambleOverride""" + config = OCIChatConfig() + messages = [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Hello"}, + ] + optional_params = {"oci_compartment_id": TEST_COMPARTMENT_ID} + + result = config.transform_request( + model="cohere.command-latest", + messages=messages, # type: ignore + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + chat_request = result["chatRequest"] + assert chat_request["preambleOverride"] == "You are a helpful assistant." + + def test_multiple_system_messages_combined(self): + """Test that multiple system messages are joined with newlines""" + config = OCIChatConfig() + messages = [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "system", "content": "Always respond in JSON."}, + {"role": "user", "content": "Hello"}, + ] + optional_params = {"oci_compartment_id": TEST_COMPARTMENT_ID} + + result = config.transform_request( + model="cohere.command-latest", + messages=messages, # type: ignore + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + chat_request = result["chatRequest"] + assert chat_request["preambleOverride"] == "You are a helpful assistant.\nAlways respond in JSON." + + def test_system_message_with_content_array(self): + """Test system message with list-style content (text blocks)""" + config = OCIChatConfig() + messages = [ + { + "role": "system", + "content": [ + {"type": "text", "text": "You are a coding assistant."}, + ], + }, + {"role": "user", "content": "Hello"}, + ] + optional_params = {"oci_compartment_id": TEST_COMPARTMENT_ID} + + result = config.transform_request( + model="cohere.command-latest", + messages=messages, # type: ignore + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + chat_request = result["chatRequest"] + assert chat_request["preambleOverride"] == "You are a coding assistant." + + def test_no_system_message_omits_preamble_override(self): + """Test that preambleOverride is omitted when there are no system messages""" + config = OCIChatConfig() + messages = [ + {"role": "user", "content": "Hello"}, + ] + optional_params = {"oci_compartment_id": TEST_COMPARTMENT_ID} + + result = config.transform_request( + model="cohere.command-latest", + messages=messages, # type: ignore + optional_params=optional_params, + litellm_params={}, + headers={}, + ) + + chat_request = result["chatRequest"] + assert "preambleOverride" not in chat_request + + def test_system_messages_excluded_from_chat_history(self): + """Test that system messages do not appear in chatHistory""" + config = OCIChatConfig() + messages = [ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "First question"}, + {"role": "assistant", "content": "First answer"}, + {"role": "user", "content": "Second question"}, + ] + + chat_history = config.adapt_messages_to_cohere_standard(messages) + + # Should contain user and assistant only, no system + # Note: adapt_messages_to_cohere_standard excludes the last message + roles = [msg.role for msg in chat_history] + assert "SYSTEM" not in roles + assert roles == ["USER", "CHATBOT"] + + class TestOCICohereStreaming: """Test Cohere streaming functionality""" 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 5f087363797..c0695bf3588 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 @@ -2,9 +2,10 @@ Tests for OpenAI GPT transformation (litellm/llms/openai/chat/gpt_transformation.py) """ -import pytest -import sys import os +import sys + +import pytest sys.path.insert(0, os.path.abspath("../../../../..")) @@ -73,6 +74,17 @@ class TestOpenAIGPTConfig: for param in base_expected_params: assert param in supported_params, f"Expected '{param}' in supported params" + def test_prompt_cache_key_supported(self): + """Test that 'prompt_cache_key' is in supported params for OpenAI chat completion models. + + OpenAI's Chat Completions API supports prompt_cache_key for cache routing optimization. + """ + supported_params = self.config.get_supported_openai_params("gpt-4.1-nano") + assert "prompt_cache_key" in supported_params + + supported_params = self.config.get_supported_openai_params("gpt-4.1") + assert "prompt_cache_key" in supported_params + class TestGetOptionalParamsIntegration: """Integration tests using litellm.get_optional_params()""" @@ -123,3 +135,14 @@ class TestGetOptionalParamsIntegration: # Both should include user assert regular_params.get("user") == "my-end-user" assert responses_params.get("user") == "my-end-user" + + def test_prompt_cache_key_in_optional_params(self): + """Test that 'prompt_cache_key' flows through get_optional_params for OpenAI models.""" + from litellm.utils import get_optional_params + + optional_params = get_optional_params( + model="gpt-4.1-nano", + custom_llm_provider="openai", + prompt_cache_key="test-cache-key-123", + ) + assert optional_params.get("prompt_cache_key") == "test-cache-key-123" diff --git a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py index 4bcafd4c57e..3a49880ff16 100644 --- a/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py +++ b/tests/test_litellm/llms/vertex_ai/vertex_ai_partner_models/anthropic/test_vertex_ai_partner_models_anthropic_transformation.py @@ -459,7 +459,7 @@ def test_vertex_ai_partner_models_anthropic_remove_prompt_caching_scope_beta_hea assert PROMPT_CACHING_BETA_HEADER not in ( beta_header or "" ), f"{PROMPT_CACHING_BETA_HEADER} should be filtered out" - assert "other-feature" in ( + assert "other-feature" not in ( beta_header or "" ), "Other non-excluded beta headers should remain" assert "web-search-2025-03-05" in ( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 1ed21b07bb7..c2dbc94f721 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -331,11 +331,7 @@ class TestMCPRequestHandler: # Create an async mock for user_api_key_auth async def mock_user_api_key_auth(api_key, request): return UserAPIKeyAuth( - token=( - "test-token-sha256-empty-hash" - if api_key - else None - ), + token=("test-token-sha256-empty-hash" if api_key else None), api_key=api_key, user_id="test-user-id" if api_key else None, team_id="test-team-id" if api_key else None, @@ -632,8 +628,7 @@ class TestMCPOAuth2AuthFlow: # OAuth2 headers should still contain the Authorization token assert ( - oauth2_headers.get("Authorization") - == "Bearer atlassian-oauth2-token" + oauth2_headers.get("Authorization") == "Bearer atlassian-oauth2-token" ) async def test_litellm_key_in_authorization_backward_compat(self): @@ -1291,21 +1286,21 @@ async def test_get_team_object_permission_with_already_loaded_permission(): mcp_access_groups=["group1"], vector_stores=["store1"], ) - + # Create mock team object with object_permission already loaded mock_team_obj = LiteLLM_TeamTable( team_id="team-123", object_permission=mock_object_permission, object_permission_id="perm-123", ) - + # Create mock user auth mock_user_auth = UserAPIKeyAuth( api_key="test-key", user_id="test-user", team_id="team-123", ) - + # Mock get_team_object to return our team with loaded permission # Also need to mock prisma_client from proxy_server mock_prisma = MagicMock() @@ -1313,96 +1308,81 @@ async def test_get_team_object_permission_with_already_loaded_permission(): "litellm.proxy.proxy_server.prisma_client", mock_prisma, ): - with patch( - "litellm.proxy.auth.auth_checks.get_team_object" - ) as mock_get_team: + with patch("litellm.proxy.auth.auth_checks.get_team_object") as mock_get_team: with patch( "litellm.proxy.auth.auth_checks.get_object_permission" ) as mock_get_perm: mock_get_team.return_value = mock_team_obj - + # Call the method result = await MCPRequestHandler._get_team_object_permission( mock_user_auth ) - + # Assert we got the object permission assert result == mock_object_permission assert result.mcp_servers == ["server1", "server2"] - + # Verify get_team_object was called mock_get_team.assert_called_once() - + # Verify get_object_permission was NOT called (since it was already loaded) mock_get_perm.assert_not_called() @pytest.mark.asyncio -async def test_get_team_object_permission_fetches_from_db_when_not_loaded(): +async def test_get_team_object_permission_with_core_auth_auto_loading(): """ - Test that _get_team_object_permission fetches from DB when object_permission - is not loaded but object_permission_id exists. + Test that _get_team_object_permission returns the object_permission that was + automatically loaded by get_team_object() in the core auth flow. + + Note: After migrating permission loading to core auth (get_team_object in auth_checks.py), + the team object returned by get_team_object() should already have object_permission loaded + when an object_permission_id exists. """ from litellm.proxy._types import LiteLLM_ObjectPermissionTable, LiteLLM_TeamTable - # Create mock object permission (to be returned from DB) + # Create mock object permission mock_object_permission = LiteLLM_ObjectPermissionTable( object_permission_id="perm-456", mcp_servers=["server3", "server4"], mcp_access_groups=["group2"], vector_stores=["store2"], ) - - # Create mock team object WITHOUT object_permission loaded (but has ID) + + # Create mock team object WITH object_permission already loaded + # (This is what get_team_object() returns after the core auth migration) mock_team_obj = LiteLLM_TeamTable( team_id="team-456", - object_permission=None, + object_permission=mock_object_permission, # Already loaded by core auth object_permission_id="perm-456", ) - + # Create mock user auth mock_user_auth = UserAPIKeyAuth( api_key="test-key", user_id="test-user", team_id="team-456", ) - + # Mock the methods - # Also need to mock prisma_client from proxy_server mock_prisma = MagicMock() with patch( "litellm.proxy.proxy_server.prisma_client", mock_prisma, ): - with patch( - "litellm.proxy.auth.auth_checks.get_team_object" - ) as mock_get_team: - with patch( - "litellm.proxy.auth.auth_checks.get_object_permission" - ) as mock_get_perm: - mock_get_team.return_value = mock_team_obj - mock_get_perm.return_value = mock_object_permission - - # Call the method - result = await MCPRequestHandler._get_team_object_permission( - mock_user_auth - ) - - # Assert we got the object permission - assert result == mock_object_permission - assert result.mcp_servers == ["server3", "server4"] - - # Verify get_team_object was called - mock_get_team.assert_called_once() - - # Verify get_object_permission WAS called (since it wasn't loaded) - mock_get_perm.assert_called_once_with( - object_permission_id="perm-456", - prisma_client=mock.ANY, - user_api_key_cache=mock.ANY, - parent_otel_span=mock_user_auth.parent_otel_span, - proxy_logging_obj=mock.ANY, - ) + with patch("litellm.proxy.auth.auth_checks.get_team_object") as mock_get_team: + mock_get_team.return_value = mock_team_obj + + # Call the method + result = await MCPRequestHandler._get_team_object_permission(mock_user_auth) + + # Assert we got the object permission (already loaded by core auth) + assert result == mock_object_permission + assert result.mcp_servers == ["server3", "server4"] + + # Verify get_team_object was called + mock_get_team.assert_called_once() @pytest.mark.asyncio @@ -1420,14 +1400,14 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper(): mcp_access_groups=["dev-group"], vector_stores=[], ) - + # Create mock user auth mock_user_auth = UserAPIKeyAuth( api_key="test-key", user_id="test-user", team_id="team-789", ) - + # Mock the helper methods with patch.object( MCPRequestHandler, "_get_team_object_permission" @@ -1437,13 +1417,16 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper(): ) as mock_get_access_group_servers: # Configure mocks mock_get_team_perm.return_value = mock_object_permission - mock_get_access_group_servers.return_value = ["group-server1", "group-server2"] - + mock_get_access_group_servers.return_value = [ + "group-server1", + "group-server2", + ] + # Call the method result = await MCPRequestHandler._get_allowed_mcp_servers_for_team( mock_user_auth ) - + # Assert the result contains both direct and access group servers assert set(result) == { "direct-server1", @@ -1451,10 +1434,10 @@ async def test_get_allowed_mcp_servers_for_team_uses_helper(): "group-server1", "group-server2", } - + # Verify _get_team_object_permission was called (the helper we fixed) mock_get_team_perm.assert_called_once_with(mock_user_auth) - + # Verify access groups were resolved mock_get_access_group_servers.assert_called_once_with(["dev-group"]) @@ -1471,21 +1454,21 @@ async def test_get_allowed_mcp_servers_for_team_with_no_object_permission(): user_id="test-user", team_id="team-no-perm", ) - + # Mock the helper to return None (no object permission) with patch.object( MCPRequestHandler, "_get_team_object_permission" ) as mock_get_team_perm: mock_get_team_perm.return_value = None - + # Call the method result = await MCPRequestHandler._get_allowed_mcp_servers_for_team( mock_user_auth ) - + # Assert empty list is returned assert result == [] - + # Verify the helper was called mock_get_team_perm.assert_called_once_with(mock_user_auth) @@ -1509,9 +1492,7 @@ async def test_get_allowed_mcp_servers_for_team_without_team_id_returns_empty(): team_id=None, ) - result = await MCPRequestHandler._get_allowed_mcp_servers_for_team( - mock_user_auth - ) + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(mock_user_auth) assert result == [] @@ -1546,9 +1527,7 @@ async def test_get_allowed_mcp_servers_for_key_guard_conditions( "litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock, ) as mock_get_perm: - with patch( - "litellm.proxy.proxy_server.prisma_client", prisma_client_value - ): + with patch("litellm.proxy.proxy_server.prisma_client", prisma_client_value): result = await MCPRequestHandler._get_allowed_mcp_servers_for_key( user_api_key_auth ) @@ -1569,9 +1548,7 @@ async def test_get_allowed_mcp_servers_for_key_returns_empty_when_db_returns_non mock_prisma = object() - with patch( - "litellm.proxy.proxy_server.prisma_client", mock_prisma - ), patch( + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), patch( "litellm.proxy.auth.auth_checks.get_object_permission", new_callable=AsyncMock, ) as mock_get_perm: diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index faabe40f2dc..d2b00d61d1b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -88,6 +88,70 @@ async def test_authorize_endpoint_includes_response_type(): assert "scope=read+write" in response.headers["location"] +@pytest.mark.asyncio +async def test_authorize_endpoint_preserves_existing_query_params(): + """Test that authorize endpoint merges OAuth params with existing query params in authorization_url""" + try: + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + authorize, + ) + from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + global_mcp_server_manager, + ) + from litellm.types.mcp import MCPAuth + from litellm.types.mcp_server.mcp_server_manager import MCPServer + from litellm.proxy._types import MCPTransport + from fastapi import Request + except ImportError: + pytest.skip("MCP discoverable endpoints not available") + + global_mcp_server_manager.registry.clear() + + # Authorization URL already has query params (e.g. multi-tenant OAuth) + oauth2_server = MCPServer( + server_id="test_oauth_server", + name="test_oauth", + server_name="test_oauth", + alias="test_oauth", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + client_id="test_client_id", + client_secret="test_client_secret", + authorization_url="https://provider.com/oauth/authorize?tenant=system", + token_url="https://provider.com/oauth/token", + scopes=["read", "write"], + ) + global_mcp_server_manager.registry[oauth2_server.server_id] = oauth2_server + + mock_request = MagicMock(spec=Request) + mock_request.base_url = "https://litellm.example.com/" + mock_request.headers = {} + + with patch( + "litellm.proxy._experimental.mcp_server.discoverable_endpoints.encrypt_value_helper" + ) as mock_encrypt: + mock_encrypt.return_value = "mocked_encrypted_state" + + response = await authorize( + request=mock_request, + client_id="test_client_id", + mcp_server_name="test_oauth", + redirect_uri="https://client.example.com/callback", + state="test_state", + ) + + location = response.headers["location"] + + # Must NOT have double '?' — existing params must be merged correctly + assert location.count("?") == 1, ( + f"Expected exactly one '?' in URL but got {location.count('?')}: {location}" + ) + assert "tenant=system" in location + assert "client_id=test_client_id" in location + assert "response_type=code" in location + assert "scope=read+write" in location + + @pytest.mark.asyncio async def test_authorize_endpoint_forwards_pkce_parameters(): """Test that authorize endpoint forwards PKCE parameters (code_challenge and code_challenge_method)""" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 9c77edc6743..4f93270c162 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -70,7 +70,9 @@ def _route_has_dependency(route, dependency) -> bool: dependant = getattr(route, "dependant", None) if dependant is None: return False - return any(getattr(dep, "call", None) == dependency for dep in dependant.dependencies) + return any( + getattr(dep, "call", None) == dependency for dep in dependant.dependencies + ) class TestExecuteWithMcpClient: @@ -481,7 +483,9 @@ class TestListToolsRestAPI: captured = {"called": False} - async def fake_get_tools(server, server_auth_header, raw_headers=None): + async def fake_get_tools( + server, server_auth_header, raw_headers=None, user_api_key_auth=None + ): captured["called"] = True captured["server"] = server captured["auth_header"] = server_auth_header @@ -659,3 +663,293 @@ class TestCallToolRestAPI: assert captured["name"] == "demo-tool" assert captured["arguments"] == {"foo": "bar"} assert captured["allowed_mcp_servers"] == [stub_server] + + +class TestGetToolsForSingleServer: + """Test _get_tools_for_single_server with object_permission filtering""" + + pytestmark = pytest.mark.asyncio + + async def test_filters_tools_by_object_permission_mcp_tool_permissions( + self, monkeypatch + ): + """Test that tools are filtered by user_api_key_auth.object_permission.mcp_tool_permissions""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.types.mcp import MCPTransport + + # Create mock tools + class MockTool: + def __init__(self, name, description): + self.name = name + self.description = description + self.inputSchema = {} + + mock_tools = [ + MockTool("tool1", "First tool"), + MockTool("tool2", "Second tool"), + MockTool("tool3", "Third tool"), + ] + + # Mock _get_tools_from_server to return all tools + async def fake_get_tools_from_server(**kwargs): + return mock_tools + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_get_tools_from_server", + fake_get_tools_from_server, + raising=False, + ) + + # Create server + server = MCPServer( + server_id="test-server-id", + name="test-server", + transport=MCPTransport.sse, + allowed_tools=None, # No server-level filtering + ) + + # Create UserAPIKeyAuth with object_permission + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test-permission-id", + mcp_tool_permissions={"test-server-id": ["tool1", "tool3"]}, + ) + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + object_permission=object_permission, + ) + + # Call the function + result = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + ) + + # Verify only allowed tools are returned + assert len(result) == 2 + tool_names = [tool.name for tool in result] + assert "tool1" in tool_names + assert "tool3" in tool_names + assert "tool2" not in tool_names + + async def test_no_filtering_when_object_permission_is_none(self, monkeypatch): + """Test that all tools are returned when object_permission is None""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.types.mcp import MCPTransport + + class MockTool: + def __init__(self, name, description): + self.name = name + self.description = description + self.inputSchema = {} + + mock_tools = [ + MockTool("tool1", "First tool"), + MockTool("tool2", "Second tool"), + ] + + async def fake_get_tools_from_server(**kwargs): + return mock_tools + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_get_tools_from_server", + fake_get_tools_from_server, + raising=False, + ) + + server = MCPServer( + server_id="test-server-id", + name="test-server", + transport=MCPTransport.sse, + allowed_tools=None, + ) + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + object_permission=None, + ) + + result = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + ) + + # All tools should be returned + assert len(result) == 2 + + async def test_no_filtering_when_mcp_tool_permissions_is_none(self, monkeypatch): + """Test that all tools are returned when mcp_tool_permissions is None""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.types.mcp import MCPTransport + + class MockTool: + def __init__(self, name, description): + self.name = name + self.description = description + self.inputSchema = {} + + mock_tools = [ + MockTool("tool1", "First tool"), + MockTool("tool2", "Second tool"), + ] + + async def fake_get_tools_from_server(**kwargs): + return mock_tools + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_get_tools_from_server", + fake_get_tools_from_server, + raising=False, + ) + + server = MCPServer( + server_id="test-server-id", + name="test-server", + transport=MCPTransport.sse, + allowed_tools=None, + ) + + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test-permission-id", + mcp_tool_permissions=None, # No tool permissions set + ) + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + object_permission=object_permission, + ) + + result = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + ) + + # All tools should be returned + assert len(result) == 2 + + async def test_no_filtering_when_server_not_in_mcp_tool_permissions( + self, monkeypatch + ): + """Test that all tools are returned when server is not in mcp_tool_permissions""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.types.mcp import MCPTransport + + class MockTool: + def __init__(self, name, description): + self.name = name + self.description = description + self.inputSchema = {} + + mock_tools = [ + MockTool("tool1", "First tool"), + MockTool("tool2", "Second tool"), + ] + + async def fake_get_tools_from_server(**kwargs): + return mock_tools + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_get_tools_from_server", + fake_get_tools_from_server, + raising=False, + ) + + server = MCPServer( + server_id="test-server-id", + name="test-server", + transport=MCPTransport.sse, + allowed_tools=None, + ) + + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test-permission-id", + mcp_tool_permissions={"other-server-id": ["tool1"]}, # Different server + ) + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + object_permission=object_permission, + ) + + result = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + ) + + # All tools should be returned since server is not in permissions + assert len(result) == 2 + + async def test_combines_server_allowed_tools_and_object_permission_filters( + self, monkeypatch + ): + """Test that both server.allowed_tools and object_permission.mcp_tool_permissions filters are applied""" + from litellm.proxy._experimental.mcp_server.server import MCPServer + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + from litellm.types.mcp import MCPTransport + + class MockTool: + def __init__(self, name, description): + self.name = name + self.description = description + self.inputSchema = {} + + mock_tools = [ + MockTool("tool1", "First tool"), + MockTool("tool2", "Second tool"), + MockTool("tool3", "Third tool"), + MockTool("tool4", "Fourth tool"), + ] + + async def fake_get_tools_from_server(**kwargs): + return mock_tools + + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "_get_tools_from_server", + fake_get_tools_from_server, + raising=False, + ) + + # Server allows tool1, tool2, tool3 + server = MCPServer( + server_id="test-server-id", + name="test-server", + transport=MCPTransport.sse, + allowed_tools=["tool1", "tool2", "tool3"], + ) + + # Object permission allows tool2, tool3, tool4 + object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test-permission-id", + mcp_tool_permissions={"test-server-id": ["tool2", "tool3", "tool4"]}, + ) + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + object_permission=object_permission, + ) + + result = await rest_endpoints._get_tools_for_single_server( + server=server, + server_auth_header=None, + user_api_key_auth=user_api_key_dict, + ) + + # Only tools in both lists should be returned (intersection): tool2, tool3 + assert len(result) == 2 + tool_names = [tool.name for tool in result] + assert "tool2" in tool_names + assert "tool3" in tool_names + assert "tool1" not in tool_names + assert "tool4" not in tool_names diff --git a/tests/test_litellm/proxy/auth/test_object_permission_loading.py b/tests/test_litellm/proxy/auth/test_object_permission_loading.py new file mode 100644 index 00000000000..54e4c82471e --- /dev/null +++ b/tests/test_litellm/proxy/auth/test_object_permission_loading.py @@ -0,0 +1,151 @@ +""" +Test that object_permission is automatically loaded when fetching keys and teams. +""" +import os +import sys +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +from litellm.proxy._types import ( + LiteLLM_ObjectPermissionTable, + LiteLLM_TeamTableCachedObj, + UserAPIKeyAuth, +) +from litellm.proxy.auth.auth_checks import get_key_object, get_team_object + + +@pytest.mark.asyncio +async def test_get_key_object_loads_object_permission(): + """ + Test that get_key_object automatically loads object_permission when object_permission_id exists. + """ + # Mock prisma client + mock_prisma_client = MagicMock() + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) # Not in cache + + # Mock the DB response with object_permission_id but no object_permission + mock_token_data = MagicMock() + mock_token_data.model_dump.return_value = { + "token": "test_token_hash", + "user_id": "test_user", + "object_permission_id": "test_perm_id", + "object_permission": None, + } + mock_prisma_client.get_data = AsyncMock(return_value=mock_token_data) + + # Mock the object_permission that should be loaded + mock_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test_perm_id", + mcp_servers=["server1", "server2"], + vector_stores=["store1"], + ) + + # Mock get_object_permission to return the permission + with patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + AsyncMock(return_value=mock_object_permission) + ), patch( + "litellm.proxy.auth.auth_checks._cache_key_object", + AsyncMock() + ): + result = await get_key_object( + hashed_token="test_token_hash", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + # Verify that object_permission was loaded + assert result.object_permission is not None + assert result.object_permission.object_permission_id == "test_perm_id" + assert result.object_permission.mcp_servers == ["server1", "server2"] + + +@pytest.mark.asyncio +async def test_get_key_object_no_permission_id(): + """ + Test that get_key_object works correctly when no object_permission_id exists. + """ + # Mock prisma client + mock_prisma_client = MagicMock() + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) # Not in cache + + # Mock the DB response without object_permission_id + mock_token_data = MagicMock() + mock_token_data.model_dump.return_value = { + "token": "test_token_hash", + "user_id": "test_user", + "object_permission_id": None, + "object_permission": None, + } + mock_prisma_client.get_data = AsyncMock(return_value=mock_token_data) + + with patch( + "litellm.proxy.auth.auth_checks._cache_key_object", + AsyncMock() + ): + result = await get_key_object( + hashed_token="test_token_hash", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + # Verify that object_permission is None + assert result.object_permission is None + + +@pytest.mark.asyncio +async def test_get_team_object_loads_object_permission(): + """ + Test that get_team_object automatically loads object_permission when object_permission_id exists. + """ + # Mock prisma client + mock_prisma_client = MagicMock() + mock_cache = MagicMock() + mock_cache.async_get_cache = AsyncMock(return_value=None) # Not in cache + + # Mock team data with object_permission_id + mock_team = MagicMock() + mock_team.dict.return_value = { + "team_id": "test_team", + "team_alias": "Test Team", + "object_permission_id": "test_perm_id", + "object_permission": None, + } + + # Mock the object_permission that should be loaded + mock_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test_perm_id", + mcp_servers=["team_server1"], + vector_stores=["team_store1"], + ) + + with patch( + "litellm.proxy.auth.auth_checks._get_team_db_check", + AsyncMock(return_value=mock_team) + ), patch( + "litellm.proxy.auth.auth_checks.get_object_permission", + AsyncMock(return_value=mock_object_permission) + ), patch( + "litellm.proxy.auth.auth_checks._cache_team_object", + AsyncMock() + ), patch( + "litellm.proxy.auth.auth_checks._should_check_db", + return_value=True + ), patch( + "litellm.proxy.auth.auth_checks._update_last_db_access_time" + ): + result = await get_team_object( + team_id="test_team", + prisma_client=mock_prisma_client, + user_api_key_cache=mock_cache, + ) + + # Verify that object_permission was loaded + assert result.object_permission is not None + assert result.object_permission.object_permission_id == "test_perm_id" + assert result.object_permission.mcp_servers == ["team_server1"] diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py index be8c84a554f..d886a4da76b 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py @@ -920,3 +920,69 @@ class TestContentFilterGuardrail: assert ( "matched_text" not in detection ), "Sensitive content should not be logged" + + @pytest.mark.asyncio + async def test_harm_toxic_abuse_blocks_abusive_input(self): + """ + Test that harm_toxic_abuse content category blocks abusive/toxic input + including censored profanity, misspellings, and harmful phrases. + """ + guardrail = ContentFilterGuardrail( + guardrail_name="test-toxic-abuse", + categories=[ + { + "category": "harm_toxic_abuse", + "enabled": True, + "action": "BLOCK", + "severity_threshold": "medium", + } + ], + ) + + toxic_input = ( + "You stupid f**ing piece of sht AI, why are you so useless? " + "Go kill yourself you worthless bot." + ) + + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": [toxic_input]}, + request_data={}, + input_type="request", + ) + + assert exc_info.value.status_code == 403 + detail = exc_info.value.detail + if isinstance(detail, dict): + assert detail.get("category") == "harm_toxic_abuse" + else: + assert "harm_toxic_abuse" in str(detail) + + @pytest.mark.asyncio + async def test_harm_toxic_abuse_blocks_sht_ai(self): + """Test that harm_toxic_abuse blocks input containing 'sht AI' (phrase or word sht).""" + guardrail = ContentFilterGuardrail( + guardrail_name="test-toxic-abuse-sht", + categories=[ + { + "category": "harm_toxic_abuse", + "enabled": True, + "action": "BLOCK", + "severity_threshold": "medium", + } + ], + ) + + with pytest.raises(HTTPException) as exc_info: + await guardrail.apply_guardrail( + inputs={"texts": ["sht AI"]}, + request_data={}, + input_type="request", + ) + + assert exc_info.value.status_code == 403 + detail = exc_info.value.detail + if isinstance(detail, dict): + assert detail.get("category") == "harm_toxic_abuse" + else: + assert "harm_toxic_abuse" in str(detail) diff --git a/tests/test_litellm/proxy/hooks/test_image_generation_guardrails.py b/tests/test_litellm/proxy/hooks/test_image_generation_guardrails.py new file mode 100644 index 00000000000..4a5d901b74d --- /dev/null +++ b/tests/test_litellm/proxy/hooks/test_image_generation_guardrails.py @@ -0,0 +1,293 @@ +""" +Tests that guardrails (post_call_success_hook) fire for image generation requests. + +The /images/generations endpoint in proxy/image_endpoints/endpoints.py calls +proxy_logging_obj.post_call_success_hook after a successful image generation. +These tests verify: +1. CustomGuardrail.async_post_call_success_hook is invoked for image generation. +2. A guardrail can inspect and transform the image response. +3. A guardrail that raises blocks the response (exception propagates). +""" + +import os +import sys +from typing import Any, Optional +from unittest.mock import patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../..")) + +import litellm +from litellm.caching.caching import DualCache +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.integrations.custom_logger import CustomLogger +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.utils import ProxyLogging +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.utils import ImageObject, ImageResponse + + +def _make_image_response(**kwargs) -> ImageResponse: + """Helper to build a minimal ImageResponse for tests.""" + return ImageResponse( + data=[ImageObject(url="https://example.com/img.png")], + **kwargs, + ) + + +# --------------------------------------------------------------------------- +# 1. Hook is invoked for image generation responses +# --------------------------------------------------------------------------- + + +class TrackingGuardrail(CustomGuardrail): + """Guardrail that records whether it was called and with what args.""" + + def __init__(self): + super().__init__( + guardrail_name="tracking_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + self.called = False + self.received_data: Optional[dict] = None + self.received_response: Optional[Any] = None + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + self.called = True + self.received_data = data + self.received_response = response + return response + + +@pytest.mark.asyncio +async def test_post_call_success_hook_invoked_for_image_generation(): + """ + Verify that a default-on guardrail's async_post_call_success_hook is + called when ProxyLogging.post_call_success_hook is invoked with an + ImageResponse (the same path used by the /images/generations endpoint). + """ + guardrail = TrackingGuardrail() + image_response = _make_image_response() + + with patch("litellm.callbacks", [guardrail]): + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + data = {"model": "dall-e-3", "prompt": "A sunset over mountains"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + result = await proxy_logging.post_call_success_hook( + data=data, + response=image_response, + user_api_key_dict=user_api_key_dict, + ) + + assert guardrail.called is True, "Guardrail hook was not invoked for image generation" + assert guardrail.received_data is not None + assert guardrail.received_data["model"] == "dall-e-3" + assert isinstance(guardrail.received_response, ImageResponse) + # The response should be passed through unchanged + assert result is image_response + + +# --------------------------------------------------------------------------- +# 2. Guardrail can transform image generation response +# --------------------------------------------------------------------------- + + +class TransformingGuardrail(CustomGuardrail): + """Guardrail that replaces the image URL in the response.""" + + def __init__(self): + super().__init__( + guardrail_name="transforming_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + # Return a modified image response (e.g., watermarked URL) + return ImageResponse( + data=[ImageObject(url="https://example.com/watermarked.png")], + ) + + +@pytest.mark.asyncio +async def test_guardrail_can_transform_image_response(): + """ + Verify that a guardrail can replace the ImageResponse returned to the client. + """ + guardrail = TransformingGuardrail() + original_response = _make_image_response() + + with patch("litellm.callbacks", [guardrail]): + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + data = {"model": "dall-e-3", "prompt": "A sunset"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + result = await proxy_logging.post_call_success_hook( + data=data, + response=original_response, + user_api_key_dict=user_api_key_dict, + ) + + assert result is not original_response + assert isinstance(result, ImageResponse) + assert result.data[0].url == "https://example.com/watermarked.png" + + +# --------------------------------------------------------------------------- +# 3. Guardrail that raises blocks the image response +# --------------------------------------------------------------------------- + + +class BlockingGuardrail(CustomGuardrail): + """Guardrail that raises on unsafe image prompts.""" + + def __init__(self): + super().__init__( + guardrail_name="blocking_guardrail", + default_on=True, + event_hook=GuardrailEventHooks.post_call, + ) + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + raise ValueError("Image content blocked by guardrail") + + +@pytest.mark.asyncio +async def test_guardrail_exception_propagates_for_image_generation(): + """ + Verify that an exception raised in a guardrail's post_call_success_hook + propagates up (the proxy endpoint wraps this in an error response). + """ + guardrail = BlockingGuardrail() + + with patch("litellm.callbacks", [guardrail]): + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + data = {"model": "dall-e-3", "prompt": "Something unsafe"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + with pytest.raises(ValueError, match="Image content blocked by guardrail"): + await proxy_logging.post_call_success_hook( + data=data, + response=_make_image_response(), + user_api_key_dict=user_api_key_dict, + ) + + +# --------------------------------------------------------------------------- +# 4. Non-guardrail CustomLogger also fires for image generation +# --------------------------------------------------------------------------- + + +class TrackingLogger(CustomLogger): + """Plain CustomLogger (not a guardrail) that tracks invocations.""" + + def __init__(self): + self.called = False + self.received_response = None + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + self.called = True + self.received_response = response + return response + + +@pytest.mark.asyncio +async def test_custom_logger_post_call_success_hook_fires_for_image_generation(): + """ + Verify that a plain CustomLogger (non-guardrail) callback also has its + async_post_call_success_hook invoked for image generation responses. + """ + logger = TrackingLogger() + image_response = _make_image_response() + + with patch("litellm.callbacks", [logger]): + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + data = {"model": "dall-e-3", "prompt": "A cat"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + result = await proxy_logging.post_call_success_hook( + data=data, + response=image_response, + user_api_key_dict=user_api_key_dict, + ) + + assert logger.called is True + assert isinstance(logger.received_response, ImageResponse) + assert result is image_response + + +# --------------------------------------------------------------------------- +# 5. Guardrail with should_run_guardrail=False is skipped +# --------------------------------------------------------------------------- + + +class OptInGuardrail(CustomGuardrail): + """Guardrail that is NOT default_on, so it only runs if explicitly requested.""" + + def __init__(self): + super().__init__( + guardrail_name="opt_in_guardrail", + default_on=False, + event_hook=GuardrailEventHooks.post_call, + ) + self.called = False + + async def async_post_call_success_hook( + self, + data: dict, + user_api_key_dict: UserAPIKeyAuth, + response: Any, + ) -> Any: + self.called = True + return response + + +@pytest.mark.asyncio +async def test_non_default_guardrail_skipped_for_image_generation(): + """ + Verify that a guardrail with default_on=False is NOT invoked for image + generation unless the request explicitly enables it. + """ + guardrail = OptInGuardrail() + + with patch("litellm.callbacks", [guardrail]): + proxy_logging = ProxyLogging(user_api_key_cache=DualCache()) + + # No guardrails key in data -> should_run_guardrail returns False + data = {"model": "dall-e-3", "prompt": "A sunset"} + user_api_key_dict = UserAPIKeyAuth(api_key="test-key") + + await proxy_logging.post_call_success_hook( + data=data, + response=_make_image_response(), + user_api_key_dict=user_api_key_dict, + ) + + assert guardrail.called is False, "Opt-in guardrail should not fire without explicit request" diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py index b372476c3d6..8b7b5a6fb7a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py @@ -7,10 +7,28 @@ enterprise (premium) license check, but should still be applied so that users can intentionally clear previously-set fields. """ -from unittest.mock import patch +from datetime import datetime, timezone +from unittest.mock import AsyncMock, MagicMock, patch +import pytest + +from litellm.proxy._types import ( + Member, + LiteLLM_OrganizationMembershipTable, + LiteLLM_TeamTable, + LiteLLM_UserTable, + LitellmUserRoles, + UserAPIKeyAuth, +) from litellm.proxy.management_endpoints.common_utils import ( + _is_user_team_admin, + _org_admin_can_invite_user, + _set_object_metadata_field, + _team_admin_can_invite_user, _update_metadata_fields, + _user_has_admin_privileges, + _user_has_admin_view, + admin_can_invite_user, ) @@ -160,3 +178,311 @@ class TestUpdateMetadataFieldsEmptyCollections: } _update_metadata_fields(updated_kv=updated_kv) mock_premium_check.assert_not_called() + + +class TestUserHasAdminView: + """Tests for _user_has_admin_view function.""" + + @pytest.mark.parametrize( + "user_role,expected", + [ + (LitellmUserRoles.PROXY_ADMIN, True), + (LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, True), + (LitellmUserRoles.INTERNAL_USER, False), + (LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, False), + ], + ) + def test_user_has_admin_view_by_role(self, user_role, expected): + """Parametrized test: admin roles return True, non-admin return False.""" + mock_auth = MagicMock() + mock_auth.user_role = user_role + assert _user_has_admin_view(mock_auth) == expected + + def test_user_has_admin_view_with_user_api_key_auth(self): + """Test with actual UserAPIKeyAuth object.""" + auth_admin = UserAPIKeyAuth( + user_id="u1", + api_key="sk-xxx", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + auth_user = UserAPIKeyAuth( + user_id="u2", + api_key="sk-yyy", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + assert _user_has_admin_view(auth_admin) is True + assert _user_has_admin_view(auth_user) is False + + +class TestIsUserTeamAdmin: + """Tests for _is_user_team_admin function.""" + + @pytest.mark.parametrize( + "members_with_roles,user_id,expected", + [ + ( + [Member(user_id="u1", role="admin")], + "u1", + True, + ), + ( + [Member(user_id="u1", role="user")], + "u1", + False, + ), + ( + [Member(user_id="u2", role="admin"), Member(user_id="u1", role="admin")], + "u1", + True, + ), + ([], "u1", False), + ], + ) + def test_is_user_team_admin_parametrized( + self, members_with_roles, user_id, expected + ): + """Parametrized test: user is team admin only when in members_with_roles with admin role.""" + mock_auth = MagicMock() + mock_auth.user_id = user_id + team = LiteLLM_TeamTable( + team_id="team-1", + members_with_roles=members_with_roles, + ) + assert _is_user_team_admin(mock_auth, team) == expected + + def test_is_user_team_admin_user_not_in_team(self): + """Test returns False when user is not in team members.""" + auth = UserAPIKeyAuth(user_id="u99", api_key="sk-x", user_role=None) + team = LiteLLM_TeamTable( + team_id="team-1", + members_with_roles=[Member(user_id="u1", role="admin")], + ) + assert _is_user_team_admin(auth, team) is False + + +class TestOrgAdminCanInviteUser: + """Tests for _org_admin_can_invite_user function.""" + + def _make_membership(self, org_id: str, user_role: str): + now = datetime.now(timezone.utc) + return LiteLLM_OrganizationMembershipTable( + user_id="u", + organization_id=org_id, + user_role=user_role, + created_at=now, + updated_at=now, + ) + + @pytest.mark.parametrize( + "admin_orgs,target_orgs,expected", + [ + (["org1"], ["org1"], True), + (["org1", "org2"], ["org2"], True), + (["org1"], ["org2"], False), + ([], ["org1"], False), + (["org1"], [], False), + ], + ) + def test_org_admin_can_invite_user_parametrized( + self, admin_orgs, target_orgs, expected + ): + """Parametrized test: can invite when target is in org where admin has ORG_ADMIN role.""" + admin_user = LiteLLM_UserTable( + user_id="admin", + organization_memberships=[ + self._make_membership(oid, LitellmUserRoles.ORG_ADMIN.value) + for oid in admin_orgs + ], + ) + target_user = LiteLLM_UserTable( + user_id="target", + organization_memberships=[ + self._make_membership(oid, LitellmUserRoles.INTERNAL_USER.value) + for oid in target_orgs + ], + ) + assert _org_admin_can_invite_user(admin_user, target_user) == expected + + def test_org_admin_can_invite_user_no_shared_org(self): + """Test returns False when admin has no org admin role.""" + admin_user = LiteLLM_UserTable( + user_id="admin", + organization_memberships=[ + self._make_membership("org1", LitellmUserRoles.INTERNAL_USER.value), + ], + ) + target_user = LiteLLM_UserTable( + user_id="target", + organization_memberships=[ + self._make_membership("org1", LitellmUserRoles.INTERNAL_USER.value), + ], + ) + assert _org_admin_can_invite_user(admin_user, target_user) is False + + +class TestTeamAdminCanInviteUser: + """Tests for _team_admin_can_invite_user async function.""" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "admin_teams,target_teams,user_is_admin_in,expected", + [ + (["t1"], ["t1"], ["t1"], True), + (["t1", "t2"], ["t2"], ["t1", "t2"], True), + (["t1"], ["t2"], ["t1"], False), + ], + ) + async def test_team_admin_can_invite_user_parametrized( + self, admin_teams, target_teams, user_is_admin_in, expected + ): + """Parametrized test: can invite when target shares a team where user is admin.""" + mock_prisma = MagicMock() + mock_auth = MagicMock() + mock_auth.user_id = "admin" + + admin_user = LiteLLM_UserTable(user_id="admin", teams=admin_teams) + target_user = LiteLLM_UserTable(user_id="target", teams=target_teams) + + def make_team(tid, is_admin): + m = ( + [{"user_id": "admin", "role": "admin"}] + if is_admin + else [] + ) + obj = MagicMock() + obj.team_id = tid + obj.model_dump = lambda: {"team_id": tid, "members_with_roles": m} + return obj + + teams = [ + make_team(tid, tid in user_is_admin_in) for tid in admin_teams + ] + mock_prisma.db.litellm_teamtable.find_many = AsyncMock( + return_value=teams + ) + + result = await _team_admin_can_invite_user( + user_api_key_dict=mock_auth, + admin_user_obj=admin_user, + target_user_obj=target_user, + prisma_client=mock_prisma, + ) + assert result == expected + + @pytest.mark.asyncio + async def test_team_admin_can_invite_user_no_shared_team(self): + """Test returns False when admin and target share no team.""" + mock_prisma = MagicMock() + mock_auth = MagicMock() + mock_auth.user_id = "admin" + admin_user = LiteLLM_UserTable(user_id="admin", teams=[]) + target_user = LiteLLM_UserTable(user_id="target", teams=["t1"]) + + result = await _team_admin_can_invite_user( + user_api_key_dict=mock_auth, + admin_user_obj=admin_user, + target_user_obj=target_user, + prisma_client=mock_prisma, + ) + assert result is False + + +class TestUserHasAdminPrivileges: + """Tests for _user_has_admin_privileges async function.""" + + @pytest.mark.asyncio + async def test_proxy_admin_has_privileges(self): + """Proxy admin always has admin privileges.""" + auth = UserAPIKeyAuth( + user_id="admin", + api_key="sk-x", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + result = await _user_has_admin_privileges( + user_api_key_dict=auth, + prisma_client=None, + ) + assert result is True + + @pytest.mark.asyncio + async def test_non_admin_no_prisma_returns_false(self): + """Non-admin with no prisma connection has no privileges.""" + auth = UserAPIKeyAuth( + user_id="user1", + api_key="sk-x", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + result = await _user_has_admin_privileges( + user_api_key_dict=auth, + prisma_client=None, + ) + assert result is False + + +class TestAdminCanInviteUser: + """Tests for admin_can_invite_user async function.""" + + @pytest.mark.asyncio + async def test_proxy_admin_can_invite_any_user(self): + """Proxy admin can invite any user regardless of org/team.""" + auth = UserAPIKeyAuth( + user_id="admin", + api_key="sk-x", + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + result = await admin_can_invite_user( + target_user_id="any-user", + user_api_key_dict=auth, + prisma_client=None, + ) + assert result is True + + @pytest.mark.asyncio + async def test_non_admin_cannot_invite_without_prisma(self): + """Non-admin with no prisma cannot invite.""" + auth = UserAPIKeyAuth( + user_id="user1", + api_key="sk-x", + user_role=LitellmUserRoles.INTERNAL_USER, + ) + result = await admin_can_invite_user( + target_user_id="other-user", + user_api_key_dict=auth, + prisma_client=None, + ) + assert result is False + + +class TestSetObjectMetadataField: + """Tests for _set_object_metadata_field function.""" + + @pytest.mark.parametrize( + "field_name,value,should_call_premium", + [ + ("guardrails", ["g1"], True), + ("model_rpm_limit", {"gpt-4": 10}, False), + ], + ) + def test_set_object_metadata_field_parametrized( + self, field_name, value, should_call_premium + ): + """Parametrized test: premium fields trigger _premium_user_check.""" + team = LiteLLM_TeamTable(team_id="t1", metadata={}) + with patch( + "litellm.proxy.management_endpoints.common_utils._premium_user_check" + ) as mock_premium: + _set_object_metadata_field(team, field_name, value) + if should_call_premium: + mock_premium.assert_called_once() + else: + mock_premium.assert_not_called() + assert team.metadata[field_name] == value + + def test_set_object_metadata_field_initializes_metadata_if_none(self): + """Test initializes metadata dict when object has None.""" + team = LiteLLM_TeamTable(team_id="t1", metadata=None) + with patch( + "litellm.proxy.management_endpoints.common_utils._premium_user_check" + ): + _set_object_metadata_field(team, "model_rpm_limit", {"x": 1}) + assert team.metadata == {"model_rpm_limit": {"x": 1}} diff --git a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py index b4c8d5f07cc..e81c6264f7b 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_mcp_management_endpoints.py @@ -1,14 +1,15 @@ import os import sys import types -from types import SimpleNamespace from datetime import datetime, timedelta +from types import SimpleNamespace from typing import List, Optional from unittest.mock import AsyncMock, MagicMock, patch import pytest from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient + from litellm._uuid import uuid from litellm.proxy.management_endpoints import ( mcp_management_endpoints as mgmt_endpoints, @@ -624,6 +625,72 @@ class TestListMCPServers: assert server.alias == "Allowed Zapier MCP" assert server.url == "https://actions.zapier.com/mcp/sse" + @pytest.mark.asyncio + async def test_admin_user_with_object_permission_respects_mcp_servers(self): + """ + Test that admin users with explicit object_permission.mcp_servers + only see the servers specified in object_permission. + + Scenario: Admin user has object_permission.mcp_servers set to specific servers + Expected: Only those servers are returned, not all servers in the registry + """ + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + # Create mock object permission with specific servers + mock_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="test-obj-perm-id", + mcp_servers=["server-1", "server-2"], # Only these two servers + mcp_access_groups=[], + mcp_tool_permissions={}, + vector_stores=[], + agents=[], + agent_access_groups=[], + ) + + # Create admin user with object permission + mock_user_auth = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + user_id="admin_user_id", + api_key="admin_api_key", + object_permission=mock_object_permission, + object_permission_id="test-obj-perm-id", + ) + + # Mock servers that the user should see + server_1 = generate_mock_mcp_server_db_record( + server_id="server-1", alias="Server 1", url="https://server1.example.com" + ) + server_2 = generate_mock_mcp_server_db_record( + server_id="server-2", alias="Server 2", url="https://server2.example.com" + ) + + # Mock manager + mock_manager = MagicMock() + mock_manager.get_all_allowed_mcp_servers = AsyncMock( + return_value=[server_1, server_2] + ) + + with patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.global_mcp_server_manager", + mock_manager, + ), patch( + "litellm.proxy.management_endpoints.mcp_management_endpoints.build_effective_auth_contexts", + AsyncMock(return_value=[mock_user_auth]), + ): + from litellm.proxy.management_endpoints.mcp_management_endpoints import ( + fetch_all_mcp_servers, + ) + + result = await fetch_all_mcp_servers(user_api_key_dict=mock_user_auth) + + # Verify results - should only return the 2 servers in object_permission + assert len(result) == 2 + server_ids = {server.server_id for server in result} + assert server_ids == {"server-1", "server-2"} + + # Verify credentials are redacted + assert all(server.credentials is None for server in result) + @pytest.mark.asyncio async def test_fetch_single_mcp_server_redacts_credentials(self): mock_server = generate_mock_mcp_server_db_record( @@ -1251,7 +1318,9 @@ class TestMCPRegistryEndpoint: mock_manager = MagicMock() mock_manager.get_registry.return_value = {mock_server.server_id: mock_server} # The registry endpoint uses get_filtered_registry (filters by client IP) - mock_manager.get_filtered_registry.return_value = {mock_server.server_id: mock_server} + mock_manager.get_filtered_registry.return_value = { + mock_server.server_id: mock_server + } with ( patch_proxy_general_settings({"enable_mcp_registry": True}), diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 08205cd2d9d..0d4a711812f 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -12,7 +12,7 @@ sys.path.insert( 0, os.path.abspath("../../../..") ) # Adds the parent directory to the system path -from unittest.mock import MagicMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import litellm import litellm.proxy.proxy_server as ps @@ -347,6 +347,196 @@ async def test_ui_view_spend_logs_with_user_id(client, monkeypatch): assert data["data"][0]["user"] == "test_user_1" +# Mock spend logs with distinct values for sorting tests. +# req_a: spend=0.10, tokens=500, start/end earliest +# req_b: spend=0.05, tokens=200, start/end 2nd +# req_c: spend=0.20, tokens=50, start/end latest +# req_d: spend=0.01, tokens=100, start/end 3rd +_SORT_TEST_LOGS = [ + { + "request_id": "req_a", + "api_key": "sk-test-key", + "user": "user1", + "spend": 0.10, + "total_tokens": 500, + "startTime": "2025-01-01T00:00:00+00:00", + "endTime": "2025-01-01T00:01:00+00:00", + "model": "gpt-3.5-turbo", + }, + { + "request_id": "req_b", + "api_key": "sk-test-key", + "user": "user1", + "spend": 0.05, + "total_tokens": 200, + "startTime": "2025-01-01T00:00:01+00:00", + "endTime": "2025-01-01T00:01:01+00:00", + "model": "gpt-3.5-turbo", + }, + { + "request_id": "req_c", + "api_key": "sk-test-key", + "user": "user1", + "spend": 0.20, + "total_tokens": 50, + "startTime": "2025-01-01T00:00:03+00:00", + "endTime": "2025-01-01T00:01:03+00:00", + "model": "gpt-3.5-turbo", + }, + { + "request_id": "req_d", + "api_key": "sk-test-key", + "user": "user1", + "spend": 0.01, + "total_tokens": 100, + "startTime": "2025-01-01T00:00:02+00:00", + "endTime": "2025-01-01T00:01:02+00:00", + "model": "gpt-3.5-turbo", + }, +] + + +def _sort_logs(logs, order_clause): + """Sort logs by the given Prisma-style order clause, e.g. {'spend': 'asc'}.""" + if not order_clause: + return list(logs) + key, direction = next(iter(order_clause.items())) + reverse = direction.lower() == "desc" + return sorted(logs, key=lambda x: x.get(key, 0), reverse=reverse) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "sort_by,sort_order,expected_request_ids", + [ + # spend: 0.01(d) < 0.05(b) < 0.10(a) < 0.20(c) + ("spend", "asc", ["req_d", "req_b", "req_a", "req_c"]), + ("spend", "desc", ["req_c", "req_a", "req_b", "req_d"]), + # total_tokens: 50(c) < 100(d) < 200(b) < 500(a) + ("total_tokens", "asc", ["req_c", "req_d", "req_b", "req_a"]), + ("total_tokens", "desc", ["req_a", "req_b", "req_d", "req_c"]), + # startTime: 00:00:00(a) < 00:00:01(b) < 00:00:02(d) < 00:00:03(c) + ("startTime", "asc", ["req_a", "req_b", "req_d", "req_c"]), + ("startTime", "desc", ["req_c", "req_d", "req_b", "req_a"]), + # endTime: same ordering as startTime + ("endTime", "asc", ["req_a", "req_b", "req_d", "req_c"]), + ("endTime", "desc", ["req_c", "req_d", "req_b", "req_a"]), + # default when sort_by not provided: startTime desc + (None, "desc", ["req_c", "req_d", "req_b", "req_a"]), + ], +) +async def test_ui_view_spend_logs_sort_by_and_sort_order( + client, monkeypatch, sort_by, sort_order, expected_request_ids +): + """Test that spend logs are returned in the correct order for each sort_by/sort_order.""" + base_logs = list(_SORT_TEST_LOGS) + + async def mock_find_many(*args, **kwargs): + order = kwargs.get("order", {}) + return _sort_logs(base_logs, order) + + async def mock_count(*args, **kwargs): + return len(base_logs) + + class MockPrismaClient: + def __init__(self): + self.db = MagicMock() + self.db.litellm_spendlogs = MagicMock() + self.db.litellm_spendlogs.find_many = AsyncMock(side_effect=mock_find_many) + self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._is_admin_view_safe", + lambda user_api_key_dict: True, + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + start_date = "2024-12-25 00:00:00" + end_date = "2025-01-02 23:59:59" + + params = { + "start_date": start_date, + "end_date": end_date, + } + if sort_by is not None: + params["sort_by"] = sort_by + if sort_order is not None: + params["sort_order"] = sort_order + + response = client.get( + "/spend/logs/ui", + params=params, + headers={"Authorization": "Bearer sk-test"}, + ) + + assert response.status_code == 200, response.text + data = response.json() + assert "data" in data + + actual_ids = [log["request_id"] for log in data["data"]] + assert actual_ids == expected_request_ids, ( + f"Expected order {expected_request_ids}, got {actual_ids} " + f"(sort_by={sort_by}, sort_order={sort_order})" + ) + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "sort_by,sort_order", + [ + ("invalid", "asc"), + ("spend", "invalid"), + ], +) +async def test_ui_view_spend_logs_sort_validation_errors( + client, monkeypatch, sort_by, sort_order +): + """Test that invalid sort_by and sort_order return 400.""" + async def mock_count(*args, **kwargs): + return 0 + + class MockPrismaClient: + def __init__(self): + self.db = MagicMock() + self.db.litellm_spendlogs = MagicMock() + self.db.litellm_spendlogs.find_many = AsyncMock(return_value=[]) + self.db.litellm_spendlogs.count = AsyncMock(side_effect=mock_count) + + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._is_admin_view_safe", + lambda user_api_key_dict: True, + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + start_date = "2024-12-25 00:00:00" + end_date = "2025-01-02 23:59:59" + + response = client.get( + "/spend/logs/ui", + params={ + "start_date": start_date, + "end_date": end_date, + "sort_by": sort_by, + "sort_order": sort_order, + }, + headers={"Authorization": "Bearer sk-test"}, + ) + + assert response.status_code == 400 + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + @pytest.mark.asyncio async def test_ui_view_spend_logs_with_team_id(client, monkeypatch): # Mock data for the test diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index acd99090397..e7da4256182 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -3203,3 +3203,123 @@ def test_deep_merge_dicts_skips_none_and_empty_lists(monkeypatch): assert result["general_settings"]["nested"]["key1"] == "updated_value1" assert result["general_settings"]["nested"]["key2"] == "value2" assert result["general_settings"]["nested"]["key3"] == "value3" + + +class TestInvitationEndpoints: + """Tests for /invitation/new and /invitation/delete endpoints.""" + + @pytest.fixture + def client_with_auth(self): + """Create a test client with admin authentication.""" + from litellm.proxy._types import LitellmUserRoles + from litellm.proxy.proxy_server import cleanup_router_config_variables + + cleanup_router_config_variables() + filepath = os.path.dirname(os.path.abspath(__file__)) + config_fp = f"{filepath}/test_configs/test_config_no_auth.yaml" + asyncio.run(initialize(config=config_fp, debug=True)) + + mock_auth = MagicMock() + mock_auth.user_id = "admin-user-id" + mock_auth.user_role = LitellmUserRoles.PROXY_ADMIN + mock_auth.api_key = "sk-test" + app.dependency_overrides[user_api_key_auth] = lambda: mock_auth + + return TestClient(app) + + @pytest.mark.parametrize( + "endpoint,payload,mock_return", + [ + ( + "/invitation/new", + {"user_id": "target-user-123"}, + { + "id": "inv-123", + "user_id": "target-user-123", + "is_accepted": False, + "accepted_at": None, + "expires_at": "2025-02-18T00:00:00", + "created_at": "2025-02-11T00:00:00", + "created_by": "admin-user-id", + "updated_at": "2025-02-11T00:00:00", + "updated_by": "admin-user-id", + }, + ), + ( + "/invitation/delete", + {"invitation_id": "inv-456"}, + { + "id": "inv-456", + "user_id": "target-user-123", + "is_accepted": False, + "accepted_at": None, + "expires_at": "2025-02-18T00:00:00", + "created_at": "2025-02-11T00:00:00", + "created_by": "admin-user-id", + "updated_at": "2025-02-11T00:00:00", + "updated_by": "admin-user-id", + }, + ), + ], + ) + def test_invitation_endpoints_proxy_admin_success( + self, client_with_auth, endpoint, payload, mock_return + ): + """Proxy admin can successfully create and delete invitations.""" + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.db.litellm_invitationlink = MagicMock() + if endpoint == "/invitation/new": + mock_create = AsyncMock(return_value=mock_return) + with patch( + "litellm.proxy.management_helpers.user_invitation.create_invitation_for_user", + mock_create, + ): + response = client_with_auth.post(endpoint, json=payload) + else: + mock_prisma.db.litellm_invitationlink.find_unique = AsyncMock( + return_value={**mock_return, "created_by": "admin-user-id"} + ) + mock_prisma.db.litellm_invitationlink.delete = AsyncMock( + return_value=mock_return + ) + response = client_with_auth.post(endpoint, json=payload) + + assert response.status_code == 200 + data = response.json() + assert data["id"] == mock_return["id"] + assert data["user_id"] == mock_return["user_id"] + + @pytest.mark.parametrize( + "endpoint,payload", + [ + ("/invitation/new", {"user_id": "target-user-123"}), + ("/invitation/delete", {"invitation_id": "inv-456"}), + ], + ) + def test_invitation_endpoints_non_admin_denied( + self, client_with_auth, endpoint, payload + ): + """Non-admin users cannot access invitation endpoints.""" + from litellm.proxy._types import LitellmUserRoles + + mock_auth = MagicMock() + mock_auth.user_id = "regular-user" + mock_auth.user_role = LitellmUserRoles.INTERNAL_USER + mock_auth.api_key = "sk-regular" + app.dependency_overrides[user_api_key_auth] = lambda: mock_auth + + with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma: + mock_prisma.db.litellm_invitationlink = MagicMock() + # Avoid triggering async DB calls in _user_has_admin_privileges + with patch( + "litellm.proxy.proxy_server._user_has_admin_privileges", + new_callable=AsyncMock, + return_value=False, + ): + response = client_with_auth.post(endpoint, json=payload) + + assert response.status_code == 400 + body = response.json() + # ProxyException handler returns {"error": {...}}, HTTPException returns {"detail": {...}} + error_content = body.get("error", body.get("detail", body)) + assert "not allowed" in str(error_content).lower() diff --git a/tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py b/tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py new file mode 100644 index 00000000000..83982482623 --- /dev/null +++ b/tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py @@ -0,0 +1,109 @@ +""" +Regression tests for AWS Secrets Manager same-name in-place rotation fix. + +When current_secret_name == new_secret_name (e.g. key alias preserved during +rotation), AWS must use PutSecretValue to update in place instead of +create+delete, which would fail with ResourceExistsException. +""" +from unittest.mock import AsyncMock, patch + +import pytest + +from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2 + + +@pytest.mark.asyncio +async def test_rotate_secret_same_name_uses_put_secret_value(): + """ + When current_secret_name == new_secret_name, async_rotate_secret should + call PutSecretValue (async_put_secret_value) instead of create+delete. + """ + secret_name = "litellm/tenant/litellm-metis-key" + new_value = "sk-new-rotated-key-value" + + with patch.object( + AWSSecretsManagerV2, + "async_put_secret_value", + new_callable=AsyncMock, + return_value={"ARN": "arn:aws:secretsmanager:us-east-1:123:secret:test"}, + ) as mock_put: + with patch.object( + AWSSecretsManagerV2, + "async_write_secret", + new_callable=AsyncMock, + ) as mock_write: + with patch.object( + AWSSecretsManagerV2, + "async_delete_secret", + new_callable=AsyncMock, + ) as mock_delete: + manager = AWSSecretsManagerV2() + result = await manager.async_rotate_secret( + current_secret_name=secret_name, + new_secret_name=secret_name, + new_secret_value=new_value, + ) + + # PutSecretValue (in-place update) should be called + mock_put.assert_called_once_with( + secret_name=secret_name, + secret_value=new_value, + optional_params=None, + timeout=None, + ) + # Create + delete should NOT be called + mock_write.assert_not_called() + mock_delete.assert_not_called() + assert result["ARN"] == "arn:aws:secretsmanager:us-east-1:123:secret:test" + + +@pytest.mark.asyncio +async def test_rotate_secret_different_names_uses_create_delete(): + """ + When current_secret_name != new_secret_name, async_rotate_secret should + use base class logic (create new, delete old). + """ + current_name = "litellm/old-key-alias" + new_name = "litellm/virtual-key-new-token-id" + new_value = "sk-new-key-value" + + with patch.object( + AWSSecretsManagerV2, + "async_read_secret", + new_callable=AsyncMock, + side_effect=["sk-old-value", new_value], # read old, then read new + ): + with patch.object( + AWSSecretsManagerV2, + "async_write_secret", + new_callable=AsyncMock, + return_value={"ARN": "arn:new"}, + ) as mock_write: + with patch.object( + AWSSecretsManagerV2, + "async_delete_secret", + new_callable=AsyncMock, + return_value={}, + ) as mock_delete: + with patch.object( + AWSSecretsManagerV2, + "async_put_secret_value", + new_callable=AsyncMock, + ) as mock_put: + manager = AWSSecretsManagerV2() + await manager.async_rotate_secret( + current_secret_name=current_name, + new_secret_name=new_name, + new_secret_value=new_value, + ) + + # PutSecretValue should NOT be called (different names) + mock_put.assert_not_called() + # Create + delete should be called + mock_write.assert_called_once() + mock_delete.assert_called_once_with( + secret_name=current_name, + recovery_window_in_days=7, + optional_params=None, + timeout=None, + ) diff --git a/tests/test_litellm/test_anthropic_beta_headers_filtering.py b/tests/test_litellm/test_anthropic_beta_headers_filtering.py new file mode 100644 index 00000000000..880e96f40a9 --- /dev/null +++ b/tests/test_litellm/test_anthropic_beta_headers_filtering.py @@ -0,0 +1,423 @@ +""" +Test suite for Anthropic beta headers filtering and mapping across all providers. + +This test validates: +1. Headers with null values in the config are filtered out +2. Headers with non-null values are correctly mapped to provider-specific names +3. Unknown headers (not in config) are filtered out +4. For Bedrock providers, beta headers appear in the request body (not just HTTP headers) +""" +import json +import os +from typing import Dict, List +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +import litellm +from litellm.anthropic_beta_headers_manager import ( + filter_and_transform_beta_headers, +) + + +class TestAnthropicBetaHeadersFiltering: + """Test beta header filtering and mapping for all providers.""" + + @pytest.fixture(autouse=True) + def setup(self): + """Load the beta headers config for testing.""" + config_path = os.path.join( + os.path.dirname(litellm.__file__), + "anthropic_beta_headers_config.json", + ) + with open(config_path, "r") as f: + self.config = json.load(f) + + def get_all_beta_headers(self) -> List[str]: + """Get all beta headers from the anthropic provider config.""" + return list(self.config.get("anthropic", {}).keys()) + + def get_supported_headers(self, provider: str) -> List[str]: + """Get headers with non-null values for a provider.""" + provider_config = self.config.get(provider, {}) + return [ + header for header, value in provider_config.items() if value is not None + ] + + def get_unsupported_headers(self, provider: str) -> List[str]: + """Get headers with null values for a provider.""" + provider_config = self.config.get(provider, {}) + return [header for header, value in provider_config.items() if value is None] + + def get_mapped_headers(self, provider: str) -> Dict[str, str]: + """Get mapping of input headers to provider-specific headers.""" + provider_config = self.config.get(provider, {}) + return { + header: value + for header, value in provider_config.items() + if value is not None + } + + @pytest.mark.parametrize( + "provider", + ["anthropic", "azure_ai", "bedrock_converse", "bedrock", "vertex_ai"], + ) + def test_filter_and_transform_beta_headers_all_headers(self, provider): + """Test filtering with all possible beta headers.""" + all_headers = self.get_all_beta_headers() + supported_headers = self.get_supported_headers(provider) + unsupported_headers = self.get_unsupported_headers(provider) + mapped_headers = self.get_mapped_headers(provider) + + filtered = filter_and_transform_beta_headers( + beta_headers=all_headers, provider=provider + ) + + for header in unsupported_headers: + assert ( + header not in filtered + ), f"Unsupported header '{header}' should be filtered out for {provider}" + assert ( + mapped_headers.get(header) not in filtered + ), f"Mapped value of unsupported header '{header}' should not appear for {provider}" + + for header in supported_headers: + expected_mapped = mapped_headers[header] + assert ( + expected_mapped in filtered + ), f"Supported header '{header}' should be mapped to '{expected_mapped}' for {provider}" + + @pytest.mark.parametrize( + "provider", + ["anthropic", "azure_ai", "bedrock_converse", "bedrock", "vertex_ai"], + ) + def test_unknown_headers_filtered_out(self, provider): + """Test that headers not in the config are filtered out.""" + unknown_headers = [ + "unknown-header-1", + "unknown-header-2", + "fake-beta-2025-01-01", + ] + all_headers = self.get_all_beta_headers() + unknown_headers + + filtered = filter_and_transform_beta_headers( + beta_headers=all_headers, provider=provider + ) + + for unknown in unknown_headers: + assert ( + unknown not in filtered + ), f"Unknown header '{unknown}' should be filtered out for {provider}" + + @pytest.mark.asyncio + async def test_anthropic_messages_http_headers_filtering(self): + """Test that Anthropic messages API filters HTTP headers correctly.""" + all_headers = self.get_all_beta_headers() + unsupported = self.get_unsupported_headers("anthropic") + + with patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" + ) as mock_client_factory: + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "id": "msg_123", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "Hello"}], + "model": "claude-3-5-sonnet-20241022", + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 20}, + } + mock_response.headers = {} + + mock_client = AsyncMock() + mock_client.post = AsyncMock(return_value=mock_response) + mock_client_factory.return_value = mock_client + + try: + await litellm.acompletion( + model="anthropic/claude-3-5-sonnet-20241022", + messages=[{"role": "user", "content": "Hi"}], + extra_headers={"anthropic-beta": ",".join(all_headers)}, + mock_response="Hello", + ) + except Exception: + pass + + if mock_client.post.called: + call_kwargs = mock_client.post.call_args.kwargs + headers = call_kwargs.get("headers", {}) + beta_header = headers.get("anthropic-beta", "") + + if beta_header: + beta_values = [b.strip() for b in beta_header.split(",")] + for unsupported_header in unsupported: + assert ( + unsupported_header not in beta_values + ), f"Unsupported header '{unsupported_header}' should not be in HTTP headers for Anthropic" + + @pytest.mark.asyncio + async def test_azure_ai_messages_http_headers_filtering(self): + """Test that Azure AI messages API filters HTTP headers correctly.""" + all_headers = self.get_all_beta_headers() + unsupported = self.get_unsupported_headers("azure_ai") + + with patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" + ) as mock_client_factory: + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "id": "msg_123", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "Hello"}], + "model": "claude-3-5-sonnet-20241022", + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 20}, + } + mock_response.headers = {} + + mock_client = AsyncMock() + mock_client.post = AsyncMock(return_value=mock_response) + mock_client_factory.return_value = mock_client + + try: + await litellm.acompletion( + model="azure_ai/claude-3-5-sonnet-20241022", + messages=[{"role": "user", "content": "Hi"}], + api_key="test-key", + api_base="https://test.azure.com", + extra_headers={"anthropic-beta": ",".join(all_headers)}, + mock_response="Hello", + ) + except Exception: + pass + + if mock_client.post.called: + call_kwargs = mock_client.post.call_args.kwargs + headers = call_kwargs.get("headers", {}) + beta_header = headers.get("anthropic-beta", "") + + if beta_header: + beta_values = [b.strip() for b in beta_header.split(",")] + for unsupported_header in unsupported: + assert ( + unsupported_header not in beta_values + ), f"Unsupported header '{unsupported_header}' should not be in HTTP headers for Azure AI" + + @pytest.mark.asyncio + async def test_bedrock_converse_headers_and_body_filtering(self): + """Test that Bedrock Converse filters both HTTP headers and request body correctly.""" + all_headers = self.get_all_beta_headers() + unsupported = self.get_unsupported_headers("bedrock_converse") + mapped_headers = self.get_mapped_headers("bedrock_converse") + + with patch("httpx.AsyncClient") as mock_client_class: + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "output": {"message": {"role": "assistant", "content": [{"text": "Hello"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 10, "outputTokens": 20}, + } + mock_response.headers = {} + mock_response.raise_for_status = MagicMock() + + mock_client = AsyncMock() + mock_client.post = AsyncMock(return_value=mock_response) + mock_client_class.return_value.__aenter__.return_value = mock_client + + try: + await litellm.acompletion( + model="bedrock/converse/us.anthropic.claude-3-5-sonnet-20241022-v2:0", + messages=[{"role": "user", "content": "Hi"}], + aws_access_key_id="test", + aws_secret_access_key="test", + aws_region_name="us-east-1", + extra_headers={"anthropic-beta": ",".join(all_headers)}, + mock_response="Hello", + ) + except Exception: + pass + + if mock_client.post.called: + call_kwargs = mock_client.post.call_args.kwargs + headers = call_kwargs.get("headers", {}) + beta_header = headers.get("anthropic-beta", "") + + if beta_header: + beta_values = [b.strip() for b in beta_header.split(",")] + for unsupported_header in unsupported: + assert ( + unsupported_header not in beta_values + ), f"Unsupported header '{unsupported_header}' should not be in HTTP headers for Bedrock Converse" + + data = call_kwargs.get("data") + if data: + body = json.loads(data) + body_beta = body.get("additionalModelRequestFields", {}).get( + "anthropic_beta", [] + ) + + for unsupported_header in unsupported: + assert ( + unsupported_header not in body_beta + ), f"Unsupported header '{unsupported_header}' should not be in request body for Bedrock Converse" + + for header, mapped_value in mapped_headers.items(): + if header in all_headers and mapped_value in body_beta: + assert ( + mapped_value in body_beta + ), f"Supported header '{header}' should be mapped to '{mapped_value}' in request body for Bedrock Converse" + + @pytest.mark.asyncio + async def test_vertex_ai_messages_http_headers_filtering(self): + """Test that Vertex AI messages API filters HTTP headers correctly.""" + all_headers = self.get_all_beta_headers() + unsupported = self.get_unsupported_headers("vertex_ai") + + with patch( + "litellm.llms.custom_httpx.http_handler.get_async_httpx_client" + ) as mock_client_factory: + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.json.return_value = { + "id": "msg_123", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "Hello"}], + "model": "claude-3-5-sonnet-20241022", + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 20}, + } + mock_response.headers = {} + + mock_client = AsyncMock() + mock_client.post = AsyncMock(return_value=mock_response) + mock_client_factory.return_value = mock_client + + with patch( + "litellm.llms.vertex_ai.vertex_llm_base.VertexBase._ensure_access_token" + ) as mock_token: + mock_token.return_value = ("test-token", "test-project") + + try: + await litellm.acompletion( + model="vertex_ai/claude-3-5-sonnet-20241022", + messages=[{"role": "user", "content": "Hi"}], + vertex_project="test-project", + vertex_location="us-central1", + extra_headers={"anthropic-beta": ",".join(all_headers)}, + mock_response="Hello", + ) + except Exception: + pass + + if mock_client.post.called: + call_kwargs = mock_client.post.call_args.kwargs + headers = call_kwargs.get("headers", {}) + beta_header = headers.get("anthropic-beta", "") + + if beta_header: + beta_values = [b.strip() for b in beta_header.split(",")] + for unsupported_header in unsupported: + assert ( + unsupported_header not in beta_values + ), f"Unsupported header '{unsupported_header}' should not be in HTTP headers for Vertex AI" + + def test_header_mapping_correctness(self): + """Test that headers are mapped correctly for providers with transformations.""" + test_cases = [ + { + "provider": "bedrock", + "input": "advanced-tool-use-2025-11-20", + "expected": "tool-search-tool-2025-10-19", + }, + { + "provider": "vertex_ai", + "input": "advanced-tool-use-2025-11-20", + "expected": "tool-search-tool-2025-10-19", + }, + { + "provider": "anthropic", + "input": "advanced-tool-use-2025-11-20", + "expected": "advanced-tool-use-2025-11-20", + }, + { + "provider": "bedrock_converse", + "input": "computer-use-2025-01-24", + "expected": "computer-use-2025-01-24", + }, + { + "provider": "azure_ai", + "input": "advanced-tool-use-2025-11-20", + "expected": "advanced-tool-use-2025-11-20", + }, + ] + + for test_case in test_cases: + filtered = filter_and_transform_beta_headers( + beta_headers=[test_case["input"]], provider=test_case["provider"] + ) + + assert ( + test_case["expected"] in filtered + ), f"Header '{test_case['input']}' should be mapped to '{test_case['expected']}' for {test_case['provider']}, but got: {filtered}" + + def test_null_value_headers_filtered(self): + """Test that headers with null values are always filtered out.""" + for provider in ["anthropic", "azure_ai", "bedrock_converse", "bedrock", "vertex_ai"]: + unsupported = self.get_unsupported_headers(provider) + + if unsupported: + filtered = filter_and_transform_beta_headers( + beta_headers=unsupported, provider=provider + ) + + assert ( + len(filtered) == 0 + ), f"All null-value headers should be filtered out for {provider}, but got: {filtered}" + + def test_empty_headers_list(self): + """Test that empty headers list returns empty result.""" + for provider in ["anthropic", "azure_ai", "bedrock_converse", "bedrock", "vertex_ai"]: + filtered = filter_and_transform_beta_headers( + beta_headers=[], provider=provider + ) + + assert ( + len(filtered) == 0 + ), f"Empty headers list should return empty result for {provider}" + + def test_mixed_supported_and_unsupported_headers(self): + """Test filtering with a mix of supported, unsupported, and unknown headers.""" + for provider in ["anthropic", "azure_ai", "bedrock_converse", "bedrock", "vertex_ai"]: + supported = self.get_supported_headers(provider) + unsupported = self.get_unsupported_headers(provider) + mapped_headers = self.get_mapped_headers(provider) + + if not supported or not unsupported: + continue + + test_headers = ( + [supported[0]] + + [unsupported[0]] + + ["unknown-header-123"] + ) + + filtered = filter_and_transform_beta_headers( + beta_headers=test_headers, provider=provider + ) + + expected_mapped = mapped_headers[supported[0]] + assert ( + expected_mapped in filtered + ), f"Supported header should be in result for {provider}" + assert ( + unsupported[0] not in filtered + ), f"Unsupported header should not be in result for {provider}" + assert ( + "unknown-header-123" not in filtered + ), f"Unknown header should not be in result for {provider}" diff --git a/tests/test_litellm/test_anthropic_beta_headers_manager.py b/tests/test_litellm/test_anthropic_beta_headers_manager.py deleted file mode 100644 index d161426c22e..00000000000 --- a/tests/test_litellm/test_anthropic_beta_headers_manager.py +++ /dev/null @@ -1,306 +0,0 @@ -""" -Tests for the centralized Anthropic beta headers manager. - -Design: JSON config lists UNSUPPORTED headers for each provider. -Headers not in the unsupported list are passed through. -Header transformations (e.g., advanced-tool-use -> tool-search-tool) happen in code, not in JSON. -""" - -import pytest - -from litellm.anthropic_beta_headers_manager import ( - filter_and_transform_beta_headers, - get_provider_beta_header, - get_provider_name, - get_unsupported_headers, - is_beta_header_supported, - update_headers_with_filtered_beta, -) - - -class TestProviderNameResolution: - """Test provider name resolution and aliases.""" - - def test_get_provider_name_direct(self): - """Test direct provider names.""" - assert get_provider_name("anthropic") == "anthropic" - assert get_provider_name("bedrock") == "bedrock" - assert get_provider_name("vertex_ai") == "vertex_ai" - assert get_provider_name("azure_ai") == "azure_ai" - - def test_get_provider_name_alias(self): - """Test provider aliases.""" - # Note: Aliases are defined in the JSON config - # If no alias exists, the original name is returned - assert get_provider_name("azure") == "azure" # No alias defined - assert get_provider_name("vertex_ai_beta") == "vertex_ai_beta" # No alias defined - - -class TestBetaHeaderSupport: - """Test beta header support checks (unsupported list approach).""" - - def test_anthropic_supports_all_headers(self): - """Anthropic should support all beta headers (empty unsupported list).""" - headers = [ - "web-fetch-2025-09-10", - "web-search-2025-03-05", - "context-management-2025-06-27", - "compact-2026-01-12", - "structured-outputs-2025-11-13", - "advanced-tool-use-2025-11-20", - ] - for header in headers: - assert is_beta_header_supported(header, "anthropic") - - def test_bedrock_unsupported_headers(self): - """Bedrock should block specific headers.""" - # Not supported (in unsupported list) - assert not is_beta_header_supported("advanced-tool-use-2025-11-20", "bedrock") - assert not is_beta_header_supported( - "prompt-caching-scope-2026-01-05", "bedrock" - ) - assert not is_beta_header_supported("structured-outputs-2025-11-13", "bedrock") - - # Supported (not in unsupported list) - assert is_beta_header_supported("context-management-2025-06-27", "bedrock") - assert is_beta_header_supported("effort-2025-11-24", "bedrock") - assert is_beta_header_supported("tool-examples-2025-10-29", "bedrock") - - def test_vertex_ai_unsupported_headers(self): - """Vertex AI should block specific headers.""" - # Not supported (in unsupported list) - assert not is_beta_header_supported( - "prompt-caching-scope-2026-01-05", "vertex_ai" - ) - - # Supported (not in unsupported list) - assert is_beta_header_supported("web-search-2025-03-05", "vertex_ai") - assert is_beta_header_supported("context-management-2025-06-27", "vertex_ai") - assert is_beta_header_supported("effort-2025-11-24", "vertex_ai") - assert is_beta_header_supported("advanced-tool-use-2025-11-20", "vertex_ai") - - -class TestBetaHeaderTransformation: - """Test beta header support checking (transformations happen in code, not here).""" - - def test_anthropic_no_transformation(self): - """Anthropic headers should pass through (empty unsupported list).""" - header = "advanced-tool-use-2025-11-20" - assert get_provider_beta_header(header, "anthropic") == header - - def test_bedrock_unsupported_returns_none(self): - """Bedrock should return None for unsupported headers.""" - header = "advanced-tool-use-2025-11-20" - # This header is in bedrock's unsupported list - assert get_provider_beta_header(header, "bedrock") is None - - def test_vertex_ai_supported_returns_original(self): - """Vertex AI should return original for supported headers.""" - header = "advanced-tool-use-2025-11-20" - # This header is NOT in vertex_ai's unsupported list - assert get_provider_beta_header(header, "vertex_ai") == header - - def test_unsupported_header_returns_none(self): - """Unsupported headers (in unsupported list) should return None.""" - header = "prompt-caching-scope-2026-01-05" - assert get_provider_beta_header(header, "bedrock") is None - - def test_supported_header_returns_original(self): - """Supported headers (not in unsupported list) should return original.""" - header = "context-management-2025-06-27" - assert get_provider_beta_header(header, "bedrock") == header - - -class TestFilterAndTransformBetaHeaders: - """Test the main filtering and transformation function.""" - - def test_anthropic_keeps_all_headers(self): - """Anthropic should keep all headers (empty unsupported list).""" - headers = [ - "web-fetch-2025-09-10", - "context-management-2025-06-27", - "structured-outputs-2025-11-13", - "some-new-future-header-2026-01-01", # Even unknown headers pass through - ] - result = filter_and_transform_beta_headers(headers, "anthropic") - assert set(result) == set(headers) - - def test_bedrock_filters_unsupported(self): - """Bedrock should filter out headers in unsupported list.""" - headers = [ - "context-management-2025-06-27", # Not in unsupported list -> kept - "advanced-tool-use-2025-11-20", # In unsupported list -> dropped - "structured-outputs-2025-11-13", # In unsupported list -> dropped - "prompt-caching-scope-2026-01-05", # In unsupported list -> dropped - ] - result = filter_and_transform_beta_headers(headers, "bedrock") - assert "context-management-2025-06-27" in result - assert "advanced-tool-use-2025-11-20" not in result - assert "structured-outputs-2025-11-13" not in result - assert "prompt-caching-scope-2026-01-05" not in result - - def test_bedrock_no_transformations_in_filter(self): - """Bedrock filtering doesn't do transformations (those happen in code).""" - headers = ["advanced-tool-use-2025-11-20"] - result = filter_and_transform_beta_headers(headers, "bedrock") - # advanced-tool-use is in unsupported list, so it gets dropped - assert result == [] - - def test_vertex_ai_filters_unsupported(self): - """Vertex AI should filter unsupported headers.""" - headers = [ - "web-search-2025-03-05", # Not in unsupported list -> kept - "advanced-tool-use-2025-11-20", # Not in unsupported list -> kept - "prompt-caching-scope-2026-01-05", # In unsupported list -> dropped - ] - result = filter_and_transform_beta_headers(headers, "vertex_ai") - assert "web-search-2025-03-05" in result - assert "advanced-tool-use-2025-11-20" in result # Kept as-is, transformation happens in code - assert "prompt-caching-scope-2026-01-05" not in result - - def test_empty_list_returns_empty(self): - """Empty list should return empty list.""" - result = filter_and_transform_beta_headers([], "anthropic") - assert result == [] - - def test_bedrock_converse_more_restrictive(self): - """Bedrock Converse should be more restrictive than Bedrock.""" - headers = [ - "context-management-2025-06-27", - "advanced-tool-use-2025-11-20", - "tool-examples-2025-10-29", - ] - - bedrock_result = filter_and_transform_beta_headers(headers, "bedrock") - converse_result = filter_and_transform_beta_headers(headers, "bedrock_converse") - - # Bedrock Converse has more restrictions - # advanced-tool-use is in both unsupported lists - assert "advanced-tool-use-2025-11-20" not in bedrock_result - assert "advanced-tool-use-2025-11-20" not in converse_result - - # tool-examples is supported on bedrock but not converse - # Actually, looking at the JSON, tool-examples is NOT in bedrock unsupported list - # So it should be in bedrock_result - assert "tool-examples-2025-10-29" in bedrock_result - # But it's not explicitly in converse unsupported list either, so it passes through - # Let me check the actual behavior - assert "context-management-2025-06-27" in bedrock_result - assert "context-management-2025-06-27" in converse_result - - def test_unknown_future_headers_pass_through(self): - """Headers not in unsupported list should pass through (future-proof).""" - headers = ["some-new-beta-2026-05-01", "another-feature-2026-06-01"] - result = filter_and_transform_beta_headers(headers, "anthropic") - assert set(result) == set(headers) - - -class TestUpdateHeadersWithFilteredBeta: - """Test the headers update function.""" - - def test_update_headers_anthropic(self): - """Test updating headers for Anthropic.""" - headers = { - "anthropic-beta": "web-fetch-2025-09-10,context-management-2025-06-27" - } - result = update_headers_with_filtered_beta(headers, "anthropic") - assert "anthropic-beta" in result - beta_values = set(result["anthropic-beta"].split(",")) - assert "web-fetch-2025-09-10" in beta_values - assert "context-management-2025-06-27" in beta_values - - def test_update_headers_bedrock_filters(self): - """Test updating headers for Bedrock with filtering.""" - headers = { - "anthropic-beta": "context-management-2025-06-27,advanced-tool-use-2025-11-20" - } - result = update_headers_with_filtered_beta(headers, "bedrock") - assert "anthropic-beta" in result - assert "context-management-2025-06-27" in result["anthropic-beta"] - assert "advanced-tool-use-2025-11-20" not in result["anthropic-beta"] - - def test_update_headers_bedrock_no_transformations(self): - """Test that filtering doesn't do transformations (those happen in code).""" - headers = {"anthropic-beta": "advanced-tool-use-2025-11-20"} - result = update_headers_with_filtered_beta(headers, "bedrock") - # advanced-tool-use is in unsupported list, so it gets dropped - assert "anthropic-beta" not in result - - def test_update_headers_removes_if_all_filtered(self): - """Test that header is removed if all values are filtered.""" - headers = {"anthropic-beta": "advanced-tool-use-2025-11-20,prompt-caching-scope-2026-01-05"} - result = update_headers_with_filtered_beta(headers, "bedrock") - assert "anthropic-beta" not in result - - def test_update_headers_no_beta_header(self): - """Test updating headers when no beta header exists.""" - headers = {"content-type": "application/json"} - result = update_headers_with_filtered_beta(headers, "anthropic") - assert "anthropic-beta" not in result - assert headers == result - - -class TestGetUnsupportedHeaders: - """Test getting unsupported headers for a provider.""" - - def test_anthropic_has_no_unsupported(self): - """Anthropic should have no unsupported headers (empty list).""" - anthropic_unsupported = get_unsupported_headers("anthropic") - assert len(anthropic_unsupported) == 0 - - def test_bedrock_converse_most_restrictive(self): - """Bedrock Converse should have more unsupported headers than Bedrock.""" - bedrock_unsupported = get_unsupported_headers("bedrock") - converse_unsupported = get_unsupported_headers("bedrock_converse") - # Converse has more restrictions - assert len(converse_unsupported) >= len(bedrock_unsupported) - - def test_all_providers_have_config(self): - """All providers should have a configuration entry.""" - providers = ["anthropic", "azure_ai", "bedrock", "bedrock_converse", "vertex_ai"] - for provider in providers: - unsupported = get_unsupported_headers(provider) - # Should return a list (even if empty) - assert isinstance(unsupported, list), f"Provider {provider} should return a list" - - -class TestEdgeCases: - """Test edge cases and error handling.""" - - def test_unknown_provider(self): - """Unknown provider with no config should pass through all headers.""" - result = filter_and_transform_beta_headers( - ["context-management-2025-06-27"], "unknown_provider" - ) - # Unknown providers have no unsupported list, so headers pass through - assert "context-management-2025-06-27" in result - - def test_whitespace_handling(self): - """Headers with whitespace should be handled correctly.""" - headers = [ - " context-management-2025-06-27 ", - " web-search-2025-03-05 ", - ] - result = filter_and_transform_beta_headers(headers, "anthropic") - assert len(result) == 2 - - def test_duplicate_headers(self): - """Duplicate headers should be deduplicated.""" - headers = [ - "context-management-2025-06-27", - "context-management-2025-06-27", - ] - result = filter_and_transform_beta_headers(headers, "anthropic") - assert len(result) == 1 - - def test_case_sensitivity(self): - """Headers should be case-sensitive.""" - # Correct case - should pass through for anthropic (no unsupported list) - headers = ["context-management-2025-06-27"] - result = filter_and_transform_beta_headers(headers, "anthropic") - assert len(result) == 1 - - # Wrong case - should still pass through (not in unsupported list) - headers = ["Context-Management-2025-06-27"] - result = filter_and_transform_beta_headers(headers, "anthropic") - assert len(result) == 1 # Passes through because anthropic has empty unsupported list diff --git a/tests/test_litellm/test_deepseek_model_metadata.py b/tests/test_litellm/test_deepseek_model_metadata.py new file mode 100644 index 00000000000..4900af5d97d --- /dev/null +++ b/tests/test_litellm/test_deepseek_model_metadata.py @@ -0,0 +1,180 @@ +""" +Regression tests for #20885 – ``supports_response_schema`` (and related +capability flags) must be consistent between the bare model-name entry +(e.g. ``deepseek-chat``) and the provider-prefixed entry +(e.g. ``deepseek/deepseek-chat``) in the model-cost map. + +The bug caused ``supports_response_schema("deepseek/deepseek-chat")`` to +return ``False`` even though the canonical ``deepseek-chat`` entry has the +field set to ``True``. +""" + +import json +import os +import sys + +sys.path.insert( + 0, os.path.abspath("../..") +) # Adds the parent directory to the system path + +import litellm +from litellm.utils import ( + _supports_factory, + supports_response_schema, +) + + +# --------------------------------------------------------------------------- +# Data-level tests – verify the JSON files are in sync +# --------------------------------------------------------------------------- + + +def _load_backup_json() -> dict: + """Load the backup JSON directly from disk.""" + backup_path = os.path.join( + os.path.dirname(litellm.__file__), + "model_prices_and_context_window_backup.json", + ) + with open(backup_path, encoding="utf-8") as f: + return json.load(f) + + +class TestDeepSeekModelCostEntries: + """Verify that provider-prefixed DeepSeek entries contain the same + capability flags as their bare-name counterparts in the JSON files.""" + + def test_deepseek_chat_supports_response_schema_in_backup(self): + data = _load_backup_json() + entry = data.get("deepseek/deepseek-chat", {}) + assert entry.get("supports_response_schema") is True + + def test_deepseek_reasoner_supports_response_schema_in_backup(self): + data = _load_backup_json() + entry = data.get("deepseek/deepseek-reasoner", {}) + assert entry.get("supports_response_schema") is True + + def test_deepseek_chat_supports_system_messages_in_backup(self): + data = _load_backup_json() + entry = data.get("deepseek/deepseek-chat", {}) + assert entry.get("supports_system_messages") is True + + def test_deepseek_reasoner_supports_system_messages_in_backup(self): + data = _load_backup_json() + entry = data.get("deepseek/deepseek-reasoner", {}) + assert entry.get("supports_system_messages") is True + + def test_deepseek_chat_max_input_tokens_matches_bare_in_backup(self): + data = _load_backup_json() + bare = data.get("deepseek-chat", {}) + prefixed = data.get("deepseek/deepseek-chat", {}) + assert prefixed.get("max_input_tokens") == bare.get("max_input_tokens") + + def test_deepseek_reasoner_max_output_tokens_matches_bare_in_backup(self): + data = _load_backup_json() + bare = data.get("deepseek-reasoner", {}) + prefixed = data.get("deepseek/deepseek-reasoner", {}) + assert prefixed.get("max_output_tokens") == bare.get("max_output_tokens") + + def test_main_json_deepseek_chat_supports_response_schema(self): + main_path = os.path.join( + os.path.dirname(os.path.dirname(litellm.__file__)), + "model_prices_and_context_window.json", + ) + with open(main_path, encoding="utf-8") as f: + data = json.load(f) + entry = data.get("deepseek/deepseek-chat", {}) + assert entry.get("supports_response_schema") is True + + def test_main_json_deepseek_reasoner_supports_response_schema(self): + main_path = os.path.join( + os.path.dirname(os.path.dirname(litellm.__file__)), + "model_prices_and_context_window.json", + ) + with open(main_path, encoding="utf-8") as f: + data = json.load(f) + entry = data.get("deepseek/deepseek-reasoner", {}) + assert entry.get("supports_response_schema") is True + + +# --------------------------------------------------------------------------- +# API-level tests – verify supports_response_schema returns True +# --------------------------------------------------------------------------- + + +class TestSupportsResponseSchemaDeepSeek: + """All calling conventions for DeepSeek should return True for + ``supports_response_schema``.""" + + def test_provider_slash_model(self): + assert supports_response_schema(model="deepseek/deepseek-chat") is True + + def test_explicit_provider(self): + assert ( + supports_response_schema( + model="deepseek-chat", custom_llm_provider="deepseek" + ) + is True + ) + + def test_reasoner_provider_slash_model(self): + assert supports_response_schema(model="deepseek/deepseek-reasoner") is True + + def test_reasoner_explicit_provider(self): + assert ( + supports_response_schema( + model="deepseek-reasoner", custom_llm_provider="deepseek" + ) + is True + ) + + +# --------------------------------------------------------------------------- +# Fallback-logic test – bare model entry used when prefixed is incomplete +# --------------------------------------------------------------------------- + + +class TestBareModelFallback: + """When a provider-prefixed entry is missing a capability flag, the + ``_supports_factory`` fallback should consult the bare model-name + entry in ``litellm.model_cost``.""" + + def test_fallback_uses_bare_entry(self): + """Temporarily remove ``supports_response_schema`` from the prefixed + entry and verify the fallback still returns True.""" + key = "deepseek/deepseek-chat" + original = litellm.model_cost.get(key, {}).get("supports_response_schema") + try: + # Simulate the pre-fix state: field missing from prefixed entry + if key in litellm.model_cost: + litellm.model_cost[key].pop("supports_response_schema", None) + result = _supports_factory( + model="deepseek-chat", + custom_llm_provider="deepseek", + key="supports_response_schema", + ) + assert result is True + finally: + # Restore + if key in litellm.model_cost and original is not None: + litellm.model_cost[key]["supports_response_schema"] = original + + def test_no_fallback_when_explicitly_false(self): + """If the prefixed entry explicitly sets a capability to ``False``, + the fallback must NOT override it.""" + key = "deepseek/deepseek-reasoner" + # After the data fix, deepseek/deepseek-reasoner has + # supports_function_calling=false (matching the bare entry). + # Explicitly set it to False to test the guard. + original = litellm.model_cost.get(key, {}).get("supports_function_calling") + try: + if key in litellm.model_cost: + litellm.model_cost[key]["supports_function_calling"] = False + result = _supports_factory( + model="deepseek-reasoner", + custom_llm_provider="deepseek", + key="supports_function_calling", + ) + assert result is False + finally: + if key in litellm.model_cost and original is not None: + litellm.model_cost[key]["supports_function_calling"] = original diff --git a/tests/test_litellm/test_exception_exports.py b/tests/test_litellm/test_exception_exports.py new file mode 100644 index 00000000000..cde26295bad --- /dev/null +++ b/tests/test_litellm/test_exception_exports.py @@ -0,0 +1,31 @@ +""" +Test that all standard HTTP error exceptions are exported from litellm.__init__. +""" + +import litellm + + +def test_permission_denied_error_is_exported(): + """PermissionDeniedError (403) should be accessible as litellm.PermissionDeniedError.""" + assert hasattr(litellm, "PermissionDeniedError") + assert litellm.PermissionDeniedError is not None + + +def test_all_http_error_exceptions_exported(): + """All standard HTTP error exceptions should be accessible at module level.""" + expected_exceptions = [ + "BadRequestError", # 400 + "AuthenticationError", # 401 + "PermissionDeniedError", # 403 + "NotFoundError", # 404 + "Timeout", # 408 + "UnprocessableEntityError", # 422 + "RateLimitError", # 429 + "InternalServerError", # 500 + "BadGatewayError", # 502 + "ServiceUnavailableError", # 503 + ] + for exc_name in expected_exceptions: + assert hasattr(litellm, exc_name), ( + f"litellm.{exc_name} is not exported from litellm.__init__" + ) diff --git a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx index 2b7399c4565..74b2619f7f3 100644 --- a/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx +++ b/ui/litellm-dashboard/src/components/ModelSelect/ModelSelect.tsx @@ -200,6 +200,7 @@ export const ModelSelect = (props: ModelSelectProps) => { }, ]} mode="multiple" + placeholder="Select Models" allowClear maxTagCount="responsive" maxTagPlaceholder={(omittedValues) => ( diff --git a/ui/litellm-dashboard/src/components/OldTeams.tsx b/ui/litellm-dashboard/src/components/OldTeams.tsx index ecc6a624be0..7b906505759 100644 --- a/ui/litellm-dashboard/src/components/OldTeams.tsx +++ b/ui/litellm-dashboard/src/components/OldTeams.tsx @@ -52,7 +52,7 @@ import type { KeyResponse, Team } from "./key_team_helpers/key_list"; import MCPServerSelector from "./mcp_server_management/MCPServerSelector"; import MCPToolPermissions from "./mcp_server_management/MCPToolPermissions"; import NotificationsManager from "./molecules/notifications_manager"; -import { Organization, fetchMCPAccessGroups, getGuardrailsList, teamDeleteCall } from "./networking"; +import { Organization, fetchMCPAccessGroups, getGuardrailsList, getPoliciesList, teamDeleteCall } from "./networking"; import NumericalInput from "./shared/numerical_input"; import VectorStoreSelector from "./vector_store_management/VectorStoreSelector"; @@ -223,6 +223,7 @@ const Teams: React.FC = ({ const [isTeamDeleting, setIsTeamDeleting] = useState(false); // Add this state near the other useState declarations const [guardrailsList, setGuardrailsList] = useState([]); + const [policiesList, setPoliciesList] = useState([]); const [expandedAccordions, setExpandedAccordions] = useState>({}); const [loggingSettings, setLoggingSettings] = useState([]); const [mcpAccessGroups, setMcpAccessGroups] = useState([]); @@ -273,7 +274,22 @@ const Teams: React.FC = ({ } }; + const fetchPolicies = async () => { + try { + if (accessToken == null) { + return; + } + + const response = await getPoliciesList(accessToken); + const policyNames = response.policies.map((p: { policy_name: string }) => p.policy_name); + setPoliciesList(policyNames); + } catch (error) { + console.error("Failed to fetch policies:", error); + } + }; + fetchGuardrails(); + fetchPolicies(); }, [accessToken]); const fetchMcpAccessGroups = async () => { @@ -1330,6 +1346,36 @@ const Teams: React.FC = ({ } /> + + Policies{" "} + + e.stopPropagation()} + > + + + + + } + name="policies" + className="mt-8" + help="Select existing policies or enter new ones" + > + ({ + value: name, + label: name, + }))} + /> + diff --git a/ui/litellm-dashboard/src/components/UsageIndicator.test.tsx b/ui/litellm-dashboard/src/components/UsageIndicator.test.tsx index 71a37263980..8c7c15bc5a5 100644 --- a/ui/litellm-dashboard/src/components/UsageIndicator.test.tsx +++ b/ui/litellm-dashboard/src/components/UsageIndicator.test.tsx @@ -6,6 +6,7 @@ import UsageIndicator from "./UsageIndicator"; vi.mock("./networking", () => ({ getRemainingUsers: vi.fn(), + getLicenseInfo: vi.fn().mockResolvedValue(null), })); vi.mock("@/app/(dashboard)/hooks/useDisableUsageIndicator", () => ({ diff --git a/ui/litellm-dashboard/src/components/UsageIndicator.tsx b/ui/litellm-dashboard/src/components/UsageIndicator.tsx index 3976e4d3d6f..47da9f78bdc 100644 --- a/ui/litellm-dashboard/src/components/UsageIndicator.tsx +++ b/ui/litellm-dashboard/src/components/UsageIndicator.tsx @@ -1,8 +1,8 @@ import { useDisableUsageIndicator } from "@/app/(dashboard)/hooks/useDisableUsageIndicator"; import { Badge } from "@tremor/react"; -import { AlertTriangle, ChevronDown, ChevronUp, Loader2, Minus, TrendingUp, UserCheck, Users } from "lucide-react"; +import { AlertTriangle, Calendar, ChevronDown, ChevronUp, Loader2, Minus, TrendingUp, UserCheck, Users } from "lucide-react"; import { useEffect, useState } from "react"; -import { getRemainingUsers } from "./networking"; +import { getRemainingUsers, getLicenseInfo, LicenseInfo } from "./networking"; // Simple utility function to combine class names const cn = (...classes: (string | boolean | undefined)[]) => { @@ -23,11 +23,35 @@ interface UsageData { total_teams_remaining: number | null; } +// Calculate days until expiration +const getDaysUntilExpiration = (expirationDate: string | null): number | null => { + if (!expirationDate) return null; + const expDate = new Date(expirationDate + 'T00:00:00Z'); // Force UTC midnight + const now = new Date(); + now.setHours(0, 0, 0, 0); // Normalize to local midnight + const diffTime = expDate.getTime() - now.getTime(); + const diffDays = Math.ceil(diffTime / (1000 * 60 * 60 * 24)); + return diffDays; +}; + +// Format expiration for display +const formatExpirationDisplay = (daysRemaining: number | null): string => { + if (daysRemaining === null) return "No expiration"; + if (daysRemaining < 0) return "Expired"; + if (daysRemaining === 0) return "Expires today"; + if (daysRemaining === 1) return "1 day remaining"; + if (daysRemaining < 30) return `${daysRemaining} days remaining`; + if (daysRemaining < 60) return "1 month remaining"; + const months = Math.floor(daysRemaining / 30); + return `${months} months remaining`; +}; + export default function UsageIndicator({ accessToken, width = 220 }: UsageIndicatorProps) { const disableUsageIndicator = useDisableUsageIndicator(); const [isExpanded, setIsExpanded] = useState(false); const [isMinimized, setIsMinimized] = useState(false); const [data, setData] = useState(null); + const [licenseInfo, setLicenseInfo] = useState(null); const [isLoading, setIsLoading] = useState(false); const [error, setError] = useState(null); @@ -39,8 +63,12 @@ export default function UsageIndicator({ accessToken, width = 220 }: UsageIndica setError(null); try { - const result = await getRemainingUsers(accessToken); - setData(result); + const [usageResult, licenseResult] = await Promise.all([ + getRemainingUsers(accessToken), + getLicenseInfo(accessToken).catch(() => null), // Don't fail if license endpoint unavailable + ]); + setData(usageResult); + setLicenseInfo(licenseResult); } catch (err) { console.error("Failed to fetch usage data:", err); setError("Failed to load usage data"); @@ -52,6 +80,13 @@ export default function UsageIndicator({ accessToken, width = 220 }: UsageIndica fetchData(); }, [accessToken]); + // Calculate license expiration metrics + const daysUntilExpiration = licenseInfo?.expiration_date + ? getDaysUntilExpiration(licenseInfo.expiration_date) + : null; + const isLicenseExpired = daysUntilExpiration !== null && daysUntilExpiration < 0; + const isLicenseExpiringSoon = daysUntilExpiration !== null && daysUntilExpiration >= 0 && daysUntilExpiration < 30; + // Calculate derived values from data const getUsageMetrics = (data: UsageData | null) => { if (!data) { @@ -106,35 +141,38 @@ export default function UsageIndicator({ accessToken, width = 220 }: UsageIndica const { isOverLimit, isNearLimit, usagePercentage, userMetrics, teamMetrics } = getUsageMetrics(data); + // Include license status in overall status + const hasAnyIssue = isOverLimit || isNearLimit || isLicenseExpired || isLicenseExpiringSoon; + const hasError = isOverLimit || isLicenseExpired; + const hasWarning = (isNearLimit || isLicenseExpiringSoon) && !hasError; + const getStatusColor = () => { - if (isOverLimit) return "red"; - if (isNearLimit) return "yellow"; + if (hasError) return "red"; + if (hasWarning) return "yellow"; return "green"; }; const getStatusIcon = () => { - if (isOverLimit) return ; - if (isNearLimit) return ; + if (hasError) return ; + if (hasWarning) return ; return null; }; // Minimized view - just a small restore button const MinimizedView = () => { - const hasIssues = isOverLimit || isNearLimit; - return (
@@ -198,13 +245,13 @@ export default function UsageIndicator({ accessToken, width = 220 }: UsageIndica onClick={() => setIsExpanded(!isExpanded)} className={cn( "flex items-center gap-3 text-left hover:bg-gray-50 rounded-md px-0 py-1 transition-colors flex-1 min-w-0", - isOverLimit && "text-red-600", - isNearLimit && "text-yellow-600", + hasError && "text-red-600", + hasWarning && "text-yellow-600", )} > Usage Status - {(isOverLimit || isNearLimit) && ( + {hasAnyIssue && ( {getStatusIcon()} @@ -229,6 +276,28 @@ export default function UsageIndicator({ accessToken, width = 220 }: UsageIndica {/* Expanded details - simple and compact */} {isExpanded && (
+ {/* License expiration section */} + {licenseInfo?.has_license && licenseInfo.expiration_date && ( +
+
+ + License +
+
+ {isLicenseExpired ? ( + + ) : isLicenseExpiringSoon ? ( + + ) : null} + {formatExpirationDisplay(daysUntilExpiration)} +
+
+ )} + {/* Users section */} {data.total_users !== null && (
@@ -323,7 +392,6 @@ export default function UsageIndicator({ accessToken, width = 220 }: UsageIndica // Optimized CardStyleView for 220px width const CardStyleView = () => { if (isMinimized) { - const hasIssues = isOverLimit || isNearLimit; return ( @@ -416,6 +496,50 @@ export default function UsageIndicator({ accessToken, width = 220 }: UsageIndica {/* Compact stats optimized for 220px */}
+ {/* License expiration section */} + {licenseInfo?.has_license && licenseInfo.expiration_date && ( +
+
+ + License + + {isLicenseExpired ? "Expired" : isLicenseExpiringSoon ? "Expiring soon" : "OK"} + +
+
+ Status: + + {formatExpirationDisplay(daysUntilExpiration)} + +
+ {licenseInfo.license_type && ( +
+ Type: + {licenseInfo.license_type} +
+ )} +
+ )} + {/* Users section */} {data.total_users !== null && (
{ + it("should render", () => { + render(); + + expect(screen.getByText("Routes Configuration")).toBeInTheDocument(); + }); + + it("should display Add Route button", () => { + render(); + + expect(screen.getByRole("button", { name: /add route/i })).toBeInTheDocument(); + }); + + it("should show empty state when no routes are configured", () => { + render(); + + expect(screen.getByText(/no routes configured/i)).toBeInTheDocument(); + }); + + it("should add a route when Add Route is clicked", async () => { + const user = userEvent.setup(); + render(); + + await user.click(screen.getByRole("button", { name: /add route/i })); + + expect(screen.getByText("Route 1: Unnamed")).toBeInTheDocument(); + }); + + it("should call onChange when a route is added", async () => { + const user = userEvent.setup(); + const onChange = vi.fn(); + render(); + + await user.click(screen.getByRole("button", { name: /add route/i })); + + expect(onChange).toHaveBeenCalledWith({ + routes: [ + expect.objectContaining({ + name: "", + utterances: [], + description: "", + score_threshold: 0.5, + }), + ], + }); + }); + + it("should initialize routes from value prop", async () => { + const value = { + routes: [ + { + name: "gpt-4", + utterances: ["hello", "hi"], + description: "For greetings", + score_threshold: 0.7, + }, + ], + }; + render(); + + await waitFor(() => { + expect(screen.getByText("Route 1: gpt-4")).toBeInTheDocument(); + }); + }); + + it("should support both name and model fields in value prop", async () => { + const value = { + routes: [{ model: "gpt-3.5-turbo", utterances: [], description: "", score_threshold: 0.5 }], + }; + render(); + + await waitFor(() => { + expect(screen.getByText("Route 1: gpt-3.5-turbo")).toBeInTheDocument(); + }); + }); + + it("should remove a route when delete button is clicked", async () => { + const user = userEvent.setup(); + const value = { + routes: [ + { + name: "gpt-4", + utterances: [], + description: "", + score_threshold: 0.5, + }, + ], + }; + render(); + + await waitFor(() => { + expect(screen.getByText("Route 1: gpt-4")).toBeInTheDocument(); + }); + + const deleteButton = screen.getByRole("button", { name: "delete" }); + await user.click(deleteButton); + + await waitFor(() => { + expect(screen.queryByText("Route 1: gpt-4")).not.toBeInTheDocument(); + expect(screen.getByText(/no routes configured/i)).toBeInTheDocument(); + }); + }); + + it("should call onChange when route is removed", async () => { + const user = userEvent.setup(); + const onChange = vi.fn(); + const value = { + routes: [ + { + name: "gpt-4", + utterances: [], + description: "", + score_threshold: 0.5, + }, + ], + }; + render(); + + await waitFor(() => { + expect(screen.getByText("Route 1: gpt-4")).toBeInTheDocument(); + }); + + const deleteButton = screen.getByRole("button", { name: "delete" }); + await user.click(deleteButton); + + await waitFor(() => { + expect(onChange).toHaveBeenCalledWith({ routes: [] }); + }); + }); + + + + + it("should update route when description is changed", async () => { + const user = userEvent.setup(); + const onChange = vi.fn(); + const value = { + routes: [ + { + name: "gpt-4", + utterances: [], + description: "", + score_threshold: 0.5, + }, + ], + }; + render(); + + await waitFor(() => { + expect(screen.getByText("Route 1: gpt-4")).toBeInTheDocument(); + }); + + const descriptionInput = screen.getByPlaceholderText("Describe when this route should be used..."); + await user.type(descriptionInput, "For code generation"); + + await waitFor(() => { + const lastCall = onChange.mock.calls[onChange.mock.calls.length - 1]; + expect(lastCall[0].routes[0].description).toBe("For code generation"); + }); + }); + + it("should update route when score threshold is changed", async () => { + const onChange = vi.fn(); + const value = { + routes: [ + { + name: "gpt-4", + utterances: [], + description: "", + score_threshold: 0.5, + }, + ], + }; + render(); + + await waitFor(() => { + expect(screen.getByText("Route 1: gpt-4")).toBeInTheDocument(); + }); + + const scoreInput = screen.getByRole("spinbutton"); + fireEvent.change(scoreInput, { target: { value: "0.9" } }); + + await waitFor(() => { + const lastCall = onChange.mock.calls[onChange.mock.calls.length - 1]; + expect(lastCall[0].routes[0].score_threshold).toBe(0.9); + }); + }); + + it("should add multiple routes", async () => { + const user = userEvent.setup(); + render(); + + await user.click(screen.getByRole("button", { name: /add route/i })); + await user.click(screen.getByRole("button", { name: /add route/i })); + + expect(screen.getByText("Route 1: Unnamed")).toBeInTheDocument(); + expect(screen.getByText("Route 2: Unnamed")).toBeInTheDocument(); + }); + + it("should toggle JSON preview visibility", async () => { + const user = userEvent.setup(); + const { container } = render(); + + expect(screen.getByText("JSON Preview")).toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Show" })).toBeInTheDocument(); + expect(container.querySelector("pre")).not.toBeInTheDocument(); + + await user.click(screen.getByRole("button", { name: "Show" })); + + expect(screen.getByRole("button", { name: "Hide" })).toBeInTheDocument(); + expect(container.querySelector("pre")).toBeInTheDocument(); + + await user.click(screen.getByRole("button", { name: "Hide" })); + + expect(screen.getByRole("button", { name: "Show" })).toBeInTheDocument(); + expect(container.querySelector("pre")).not.toBeInTheDocument(); + }); + + it("should display JSON preview with route data when routes exist", async () => { + const user = userEvent.setup(); + const { container } = render( + , + ); + + await waitFor(() => { + expect(screen.getByText("Route 1: gpt-4")).toBeInTheDocument(); + }); + + await user.click(screen.getByRole("button", { name: "Show" })); + + const preElement = container.querySelector("pre"); + expect(preElement).toBeInTheDocument(); + expect(preElement?.textContent).toContain("gpt-4"); + expect(preElement?.textContent).toContain("hello"); + expect(preElement?.textContent).toContain("0.8"); + }); + + it("should display model selector with options from modelInfo", async () => { + const value = { + routes: [ + { name: "", utterances: [], description: "", score_threshold: 0.5 }, + ], + }; + render(); + + await waitFor(() => { + expect(screen.getByText("Route 1: Unnamed")).toBeInTheDocument(); + }); + + expect(screen.getByText("Model")).toBeInTheDocument(); + const comboboxes = screen.getAllByRole("combobox"); + expect(comboboxes.length).toBeGreaterThan(0); + }); + + it("should clear routes when value prop changes to empty", async () => { + const value = { + routes: [ + { + name: "gpt-4", + utterances: [], + description: "", + score_threshold: 0.5, + }, + ], + }; + const { rerender } = render(); + + await waitFor(() => { + expect(screen.getByText("Route 1: gpt-4")).toBeInTheDocument(); + }); + + rerender(); + + await waitFor(() => { + expect(screen.getByText(/no routes configured/i)).toBeInTheDocument(); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/RouterConfigBuilder.tsx b/ui/litellm-dashboard/src/components/add_model/RouterConfigBuilder.tsx new file mode 100644 index 00000000000..9b93287b36d --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/RouterConfigBuilder.tsx @@ -0,0 +1,277 @@ +import { DeleteOutlined, InfoCircleOutlined, PlusOutlined } from "@ant-design/icons"; +import { Select as AntdSelect, Button, Card, Collapse, Divider, Empty, Flex, Input, InputNumber, Space, Tooltip, Typography } from "antd"; +import React, { useEffect, useState } from "react"; +import { ModelGroup } from "../playground/llm_calls/fetch_models"; + +const { Text } = Typography; + +const { TextArea } = Input; + +interface Route { + id: string; + model: string; + utterances: string[]; + description: string; + score_threshold: number; +} + +interface SavedRoute { + id?: string; + name?: string; + model?: string; + utterances?: string[]; + description?: string; + score_threshold?: number; +} + +interface RouterConfig { + routes?: SavedRoute[]; +} + +interface RouterConfigBuilderProps { + modelInfo: ModelGroup[]; + value?: RouterConfig; + onChange?: (config: any) => void; +} + +const RouterConfigBuilder: React.FC = ({ modelInfo, value, onChange }) => { + const [routes, setRoutes] = useState([]); + const [showJsonPreview, setShowJsonPreview] = useState(false); + const [expandedRoutes, setExpandedRoutes] = useState([]); + + // Initialize routes from value prop - preserve existing route IDs to avoid focus loss when parent re-renders + useEffect(() => { + const routesFromValue = value?.routes; + if (routesFromValue) { + const routeIds: string[] = []; + setRoutes((prevRoutes) => { + const initializedRoutes = routesFromValue.map((route: SavedRoute, index: number) => { + const existingRoute = prevRoutes[index]; + const id = existingRoute?.id || route.id || `route-${index}-${Date.now()}`; + routeIds.push(id); + return { + id, + model: route.name || route.model || "", // handle both 'name' and 'model' fields + utterances: route.utterances || [], + description: route.description || "", + score_threshold: route.score_threshold ?? 0.5, + }; + }); + return initializedRoutes; + }); + setExpandedRoutes(routeIds); + } else { + setRoutes([]); + setExpandedRoutes([]); + } + }, [value]); + + // Handle adding a new route + const addRoute = () => { + const newRouteId = `route-${Date.now()}`; + const newRoute: Route = { + id: newRouteId, + model: "", + utterances: [], + description: "", + score_threshold: 0.5, + }; + const updatedRoutes = [...routes, newRoute]; + setRoutes(updatedRoutes); + updateConfig(updatedRoutes); + // Automatically expand the new route + setExpandedRoutes((prev) => [...prev, newRouteId]); + }; + + // Handle removing a route + const removeRoute = (routeId: string) => { + const updatedRoutes = routes.filter((route) => route.id !== routeId); + setRoutes(updatedRoutes); + updateConfig(updatedRoutes); + // Remove from expanded routes as well + setExpandedRoutes((prev) => prev.filter((id) => id !== routeId)); + }; + + // Handle updating a route + const updateRoute = (routeId: string, field: keyof Route, value: any) => { + const updatedRoutes = routes.map((route) => (route.id === routeId ? { ...route, [field]: value } : route)); + setRoutes(updatedRoutes); + updateConfig(updatedRoutes); + }; + + // Update the overall configuration + const updateConfig = (updatedRoutes: Route[]) => { + const config = { + routes: updatedRoutes.map((route) => ({ + name: route.model, + utterances: route.utterances, + description: route.description, + score_threshold: route.score_threshold, + })), + }; + onChange?.(config); + }; + + // Handle utterances change (convert textarea string to array) + const handleUtterancesChange = (routeId: string, utterancesText: string) => { + const utterancesArray = utterancesText + .split("\n") + .map((line) => line.trim()) // Only trims leading/trailing whitespace, preserves internal spaces + .filter((line) => line.length > 0); + updateRoute(routeId, "utterances", utterancesArray); + }; + + // Prepare model options for dropdowns + const modelOptions = modelInfo.map((model) => ({ + value: model.model_group, + label: model.model_group, + })); + + const generateConfig = () => { + return { + routes: routes.map((route) => ({ + name: route.model, + utterances: route.utterances, + description: route.description, + score_threshold: route.score_threshold, + })), + }; + }; + + return ( +
+ + + Routes Configuration + + + + + + + + {/* Routes */} + {routes.length === 0 ? ( + + + + ) : ( + setExpandedRoutes(Array.isArray(keys) ? keys : [keys].filter(Boolean))} + style={{ width: "100%" }} + items={routes.map((route, index) => ({ + key: route.id, + label: ( + + Route {index + 1}: {route.model || "Unnamed"} + + ), + extra: ( +