Merge pull request #20042 from BerriAI/main

merge main in passthrough
This commit is contained in:
Sameer Kankute
2026-01-30 08:44:00 +05:30
committed by GitHub
63 changed files with 4261 additions and 701 deletions
+112
View File
@@ -3407,6 +3407,110 @@ jobs:
- store_test_results:
path: test-results
proxy_e2e_anthropic_messages_tests:
machine:
image: ubuntu-2204:2023.10.1
resource_class: xlarge
working_directory: ~/project
steps:
- checkout
- setup_google_dns
- run:
name: Install Docker CLI (In case it's not already installed)
command: |
curl -fsSL https://get.docker.com | sh
sudo usermod -aG docker $USER
docker version
- run:
name: Install Python 3.10
command: |
curl https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh --output miniconda.sh
bash miniconda.sh -b -p $HOME/miniconda
export PATH="$HOME/miniconda/bin:$PATH"
conda init bash
source ~/.bashrc
conda create -n myenv python=3.10 -y
conda activate myenv
python --version
- run:
name: Install Dependencies
command: |
export PATH="$HOME/miniconda/bin:$PATH"
source $HOME/miniconda/etc/profile.d/conda.sh
conda activate myenv
pip install "pytest==7.3.1"
pip install "pytest-asyncio==0.21.1"
pip install "boto3==1.36.0"
pip install "httpx==0.27.0"
pip install "claude-agent-sdk"
pip install -r requirements.txt
- run:
name: Install dockerize
command: |
wget https://github.com/jwilder/dockerize/releases/download/v0.6.1/dockerize-linux-amd64-v0.6.1.tar.gz
sudo tar -C /usr/local/bin -xzvf dockerize-linux-amd64-v0.6.1.tar.gz
rm dockerize-linux-amd64-v0.6.1.tar.gz
- run:
name: Start PostgreSQL Database
command: |
docker run -d \
--name postgres-db \
-e POSTGRES_USER=postgres \
-e POSTGRES_PASSWORD=postgres \
-e POSTGRES_DB=circle_test \
-p 5432:5432 \
postgres:14
- run:
name: Wait for PostgreSQL to be ready
command: dockerize -wait tcp://localhost:5432 -timeout 1m
- attach_workspace:
at: ~/project
- run:
name: Load Docker Database Image
command: |
gunzip -c litellm-docker-database.tar.gz | docker load
docker images | grep litellm-docker-database
- run:
name: Run Docker container with test config
command: |
docker run -d \
-p 4000:4000 \
-e DATABASE_URL=postgresql://postgres:postgres@host.docker.internal:5432/circle_test \
-e LITELLM_MASTER_KEY="sk-1234" \
-e AWS_ACCESS_KEY_ID=$AWS_ACCESS_KEY_ID \
-e AWS_SECRET_ACCESS_KEY=$AWS_SECRET_ACCESS_KEY \
-e AWS_REGION_NAME="us-east-1" \
--add-host host.docker.internal:host-gateway \
--name my-app \
-v $(pwd)/tests/proxy_e2e_anthropic_messages_tests/test_config.yaml:/app/config.yaml \
litellm-docker-database:ci \
--config /app/config.yaml \
--port 4000 \
--detailed_debug
- run:
name: Start outputting logs
command: docker logs -f my-app
background: true
- run:
name: Wait for app to be ready
command: dockerize -wait http://localhost:4000 -timeout 5m
- run:
name: Run Claude Agent SDK E2E Tests
command: |
export PATH="$HOME/miniconda/bin:$PATH"
source $HOME/miniconda/etc/profile.d/conda.sh
conda activate myenv
export LITELLM_PROXY_URL="http://localhost:4000"
export LITELLM_API_KEY="sk-1234"
pwd
ls
python -m pytest -vv tests/proxy_e2e_anthropic_messages_tests/ -x -s --junitxml=test-results/junit.xml --durations=5
no_output_timeout: 120m
# Store test results
- store_test_results:
path: test-results
upload-coverage:
docker:
- image: cimg/python:3.9
@@ -4075,6 +4179,14 @@ workflows:
only:
- main
- /litellm_.*/
- proxy_e2e_anthropic_messages_tests:
requires:
- build_docker_database_image
filters:
branches:
only:
- main
- /litellm_.*/
- llm_translation_testing:
filters:
branches:
+1 -1
View File
@@ -73,4 +73,4 @@ jobs:
- name: Check import safety
run: |
poetry run python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)
poetry run python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)
+15
View File
@@ -0,0 +1,15 @@
name: Validate model_prices_and_context_window.json
on:
pull_request:
branches: [ main ]
jobs:
validate-model-prices-json:
runs-on: ubuntu-latest
steps:
- uses: actions/checkout@v4
- name: Validate model_prices_and_context_window.json
run: |
jq empty model_prices_and_context_window.json
+117
View File
@@ -0,0 +1,117 @@
# Claude Agent SDK with LiteLLM Gateway
A simple example showing how to use Claude's Agent SDK with LiteLLM as a proxy. This lets you use any LLM provider (OpenAI, Bedrock, Azure, etc.) through the Agent SDK.
## Quick Start
### 1. Install dependencies
```bash
pip install anthropic claude-agent-sdk litellm
```
### 2. Start LiteLLM proxy
```bash
# Simple start with Claude
litellm --model claude-sonnet-4-20250514
# Or with a config file
litellm --config config.yaml
```
### 3. Run the chat
```bash
python main.py
```
That's it! You can now chat with the agent in your terminal.
### Chat Commands
While chatting, you can use these commands:
- `models` - List all available models (fetched from your LiteLLM proxy)
- `model` - Switch to a different model
- `clear` - Start a new conversation
- `quit` or `exit` - End the chat
The chat automatically fetches available models from your LiteLLM proxy's `/models` endpoint, so you'll always see what's currently configured.
## Configuration
Set these environment variables if needed:
```bash
export LITELLM_PROXY_URL="http://localhost:4000"
export LITELLM_API_KEY="sk-1234"
export LITELLM_MODEL="claude-sonnet-4-20250514"
```
Or just use the defaults - it'll connect to `http://localhost:4000` by default.
## Example Config File
If you want to use multiple models, create a `config.yaml` (see `config.example.yaml`):
```yaml
model_list:
- model_name: bedrock-claude-sonnet-4
litellm_params:
model: "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-sonnet-4.5
litellm_params:
model: "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0"
aws_region_name: "us-east-1"
```
Then start LiteLLM with: `litellm --config config.yaml`
## How It Works
The key is pointing the Agent SDK to LiteLLM instead of directly to Anthropic:
```python
# Point to LiteLLM gateway (not Anthropic)
os.environ["ANTHROPIC_BASE_URL"] = "http://localhost:4000"
os.environ["ANTHROPIC_API_KEY"] = "sk-1234" # Your LiteLLM key
# Use any model configured in LiteLLM
options = ClaudeAgentOptions(
model="bedrock-claude-sonnet-4", # or gpt-4, or anything else
system_prompt="You are a helpful assistant.",
max_turns=50,
)
```
Note: Don't add `/anthropic` to the base URL - LiteLLM handles the routing automatically.
## Why Use This?
- **Switch providers easily**: Use the same code with OpenAI, Bedrock, Azure, etc.
- **Cost tracking**: LiteLLM tracks spending across all your agent conversations
- **Rate limiting**: Set budgets and limits on your agent usage
- **Load balancing**: Distribute requests across multiple API keys or regions
- **Fallbacks**: Automatically retry with a different model if one fails
## Troubleshooting
**Connection errors?**
- Make sure LiteLLM is running: `litellm --model your-model`
- Check the URL is correct (default: `http://localhost:4000`)
**Authentication errors?**
- Verify your LiteLLM API key is correct
- Make sure the model is configured in your LiteLLM setup
**Model not found?**
- Check the model name matches what's in your LiteLLM config
- Run `litellm --model your-model` to test it works
## Learn More
- [LiteLLM Docs](https://docs.litellm.ai/)
- [Claude Agent SDK](https://github.com/anthropics/anthropic-agent-sdk)
- [LiteLLM Proxy Guide](https://docs.litellm.ai/docs/proxy/quick_start)
@@ -0,0 +1,25 @@
model_list:
- model_name: bedrock-claude-sonnet-3.5
litellm_params:
model: "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-sonnet-4
litellm_params:
model: "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-sonnet-4.5
litellm_params:
model: "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-opus-4.5
litellm_params:
model: "bedrock/us.anthropic.claude-opus-4-5-20251101-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-nova-premier
litellm_params:
model: "bedrock/amazon.nova-premier-v1:0"
aws_region_name: "us-east-1"
+196
View File
@@ -0,0 +1,196 @@
"""
Simple Interactive Claude Agent SDK CLI using LiteLLM Gateway
This example demonstrates an interactive CLI chat with the Anthropic Agent SDK using LiteLLM as a proxy.
LiteLLM acts as a unified interface, allowing you to use any LLM provider (OpenAI, Azure, Bedrock, etc.)
through the Claude Agent SDK by pointing it to the LiteLLM gateway.
"""
import os
import asyncio
import httpx
from claude_agent_sdk import ClaudeSDKClient, ClaudeAgentOptions
class Config:
"""Configuration for LiteLLM Gateway connection"""
# LiteLLM proxy URL (default to local instance)
LITELLM_PROXY_URL = os.getenv("LITELLM_PROXY_URL", "http://localhost:4000")
# LiteLLM API key (master key or virtual key)
LITELLM_API_KEY = os.getenv("LITELLM_API_KEY", "sk-1234")
# Model name as configured in LiteLLM (e.g., "bedrock-claude-sonnet-4", "gpt-4", etc.)
LITELLM_MODEL = os.getenv("LITELLM_MODEL", "bedrock-claude-sonnet-4.5")
async def fetch_available_models(base_url: str, api_key: str) -> list[str]:
"""
Fetch available models from LiteLLM proxy /models endpoint
"""
try:
async with httpx.AsyncClient() as client:
response = await client.get(
f"{base_url}/models",
headers={"Authorization": f"Bearer {api_key}"},
timeout=10.0
)
response.raise_for_status()
data = response.json()
return [model["id"] for model in data.get("data", [])]
except Exception as e:
print(f"⚠️ Warning: Could not fetch models from proxy: {e}")
print("Using default model list...")
# Fallback to default models
return [
"bedrock-claude-sonnet-3.5",
"bedrock-claude-sonnet-4",
"bedrock-claude-sonnet-4.5",
"bedrock-claude-opus-4.5",
"bedrock-nova-premier",
]
async def interactive_chat():
"""
Interactive CLI chat with the agent
"""
config = Config()
# Configure Anthropic SDK to point to LiteLLM gateway
# Note: We don't add /anthropic to the base URL - LiteLLM handles routing
litellm_base_url = config.LITELLM_PROXY_URL.rstrip('/')
os.environ["ANTHROPIC_BASE_URL"] = litellm_base_url
os.environ["ANTHROPIC_API_KEY"] = config.LITELLM_API_KEY
# Fetch available models from proxy
available_models = await fetch_available_models(litellm_base_url, config.LITELLM_API_KEY)
current_model = config.LITELLM_MODEL
print("=" * 70)
print("🤖 Claude Agent SDK with LiteLLM Gateway - Interactive Chat")
print("=" * 70)
print(f"🚀 Connected to: {litellm_base_url}")
print(f"📦 Current model: {current_model}")
print("\nType your messages below. Commands:")
print(" - 'quit' or 'exit' to end the conversation")
print(" - 'clear' to start a new conversation")
print(" - 'model' to switch models")
print(" - 'models' to list available models")
print("=" * 70)
print()
while True:
# Configure agent options for each conversation
options = ClaudeAgentOptions(
system_prompt="You are a helpful AI assistant. Be concise, accurate, and friendly.",
model=current_model,
max_turns=50,
)
# Create agent client
async with ClaudeSDKClient(options=options) as client:
conversation_active = True
while conversation_active:
# Get user input
try:
user_input = input("\n👤 You: ").strip()
except (EOFError, KeyboardInterrupt):
print("\n\n👋 Goodbye!")
return
# Handle commands
if user_input.lower() in ['quit', 'exit']:
print("\n👋 Goodbye!")
return
if user_input.lower() == 'clear':
print("\n🔄 Starting new conversation...\n")
conversation_active = False
continue
if user_input.lower() == 'models':
print("\n📋 Available models:")
for i, model in enumerate(available_models, 1):
marker = "" if model == current_model else " "
print(f" {marker} {i}. {model}")
continue
if user_input.lower() == 'model':
print("\n📋 Select a model:")
for i, model in enumerate(available_models, 1):
marker = "" if model == current_model else " "
print(f" {marker} {i}. {model}")
try:
choice = input("\nEnter number (or press Enter to cancel): ").strip()
if choice:
idx = int(choice) - 1
if 0 <= idx < len(available_models):
current_model = available_models[idx]
print(f"\n✅ Switched to: {current_model}")
print("🔄 Starting new conversation with new model...\n")
conversation_active = False
else:
print("❌ Invalid choice")
except (ValueError, IndexError):
print("❌ Invalid input")
continue
if not user_input:
continue
# Send query to agent with loading indicator
print("\n🤖 Assistant: ", end='', flush=True)
try:
await client.query(user_input)
# Show loading indicator
print("⏳ thinking...", end='', flush=True)
# Stream the response
first_chunk = True
async for msg in client.receive_response():
# Clear loading indicator on first message
if first_chunk:
print("\r🤖 Assistant: ", end='', flush=True)
first_chunk = False
# Handle different message types
if hasattr(msg, 'type'):
if msg.type == 'content_block_delta':
# Streaming text delta
if hasattr(msg, 'delta') and hasattr(msg.delta, 'text'):
print(msg.delta.text, end='', flush=True)
elif msg.type == 'content_block_start':
# Start of content block
if hasattr(msg, 'content_block') and hasattr(msg.content_block, 'text'):
print(msg.content_block.text, end='', flush=True)
# Fallback to original content handling
if hasattr(msg, 'content'):
for content_block in msg.content:
if hasattr(content_block, 'text'):
print(content_block.text, end='', flush=True)
print() # New line after response
except Exception as e:
print(f"\r\n❌ Error: {e}")
print("Please check your LiteLLM gateway is running and configured correctly.")
def main():
"""Run interactive chat"""
try:
asyncio.run(interactive_chat())
except KeyboardInterrupt:
print("\n\n👋 Goodbye!")
if __name__ == "__main__":
main()
@@ -0,0 +1,2 @@
claude-agent-sdk
httpx>=0.27.0
+4 -2
View File
@@ -170,12 +170,14 @@ RUN sed -i 's/\r$//' docker/entrypoint.sh && \
[ -n "$LITELLM_PROXY_EXTRAS_PATH" ] && chmod -R g+w $LITELLM_PROXY_EXTRAS_PATH || true && \
chmod -R g+rX $PRISMA_PATH && \
chmod -R g+rX /app/.cache && \
mkdir -p /tmp/.npm /nonexistent /.npm && \
prisma generate
mkdir -p /tmp/.npm /nonexistent /.npm
# Switch to non-root user for runtime
USER nobody
# Generate Prisma client as nobody user to ensure correct file ownership
RUN prisma generate
# Prisma runtime knobs for offline containers
ENV PRISMA_SKIP_POSTINSTALL_GENERATE=1 \
PRISMA_HIDE_UPDATE_MESSAGE=1 \
@@ -0,0 +1,115 @@
import Tabs from '@theme/Tabs';
import TabItem from '@theme/TabItem';
# Claude Agent SDK with LiteLLM
Use Anthropic's Claude Agent SDK with any LLM provider through LiteLLM Proxy.
The Claude Agent SDK provides a high-level interface for building AI agents. By pointing it to LiteLLM, you can use the same agent code with OpenAI, Bedrock, Azure, Vertex AI, or any other provider.
## Quick Start
### 1. Install Dependencies
```bash
pip install claude-agent-sdk
```
### 2. Start LiteLLM Proxy
```yaml title="config.yaml" showLineNumbers
model_list:
- model_name: bedrock-claude-sonnet-3.5
litellm_params:
model: "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-sonnet-4
litellm_params:
model: "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-sonnet-4.5
litellm_params:
model: "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-opus-4.5
litellm_params:
model: "bedrock/us.anthropic.claude-opus-4-5-20251101-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-nova-premier
litellm_params:
model: "bedrock/amazon.nova-premier-v1:0"
aws_region_name: "us-east-1"
```
```bash
litellm --config config.yaml
```
### 3. Point Agent SDK to LiteLLM
| Environment Variable | Value | Description |
|---------------------|-------|-------------|
| `ANTHROPIC_BASE_URL` | `http://localhost:4000` | LiteLLM proxy URL |
| `ANTHROPIC_API_KEY` | `sk-1234` | Your LiteLLM API key (not Anthropic key) |
```python title="agent.py" showLineNumbers
import os
from claude_agent_sdk import ClaudeSDKClient, ClaudeAgentOptions
# Point to LiteLLM proxy (not Anthropic)
os.environ["ANTHROPIC_BASE_URL"] = "http://localhost:4000"
os.environ["ANTHROPIC_API_KEY"] = "sk-1234" # Your LiteLLM key
# Configure agent with any model from your config
options = ClaudeAgentOptions(
system_prompt="You are a helpful AI assistant.",
model="bedrock-claude-sonnet-4", # Use any model from config.yaml
max_turns=20,
)
async with ClaudeSDKClient(options=options) as client:
await client.query("What is LiteLLM?")
async for msg in client.receive_response():
if hasattr(msg, 'content'):
for content_block in msg.content:
if hasattr(content_block, 'text'):
print(content_block.text, end='', flush=True)
```
## Why Use LiteLLM with Agent SDK?
| Feature | Benefit |
|---------|---------|
| **Multi-Provider** | Use the same agent code with OpenAI, Bedrock, Azure, Vertex AI, etc. |
| **Cost Tracking** | Track spending across all agent conversations |
| **Rate Limiting** | Set budgets and limits on agent usage |
| **Load Balancing** | Distribute requests across multiple API keys or regions |
| **Fallbacks** | Automatically retry with different models if one fails |
## Complete Example
See our [cookbook example](https://github.com/BerriAI/litellm/tree/main/cookbook/anthropic_agent_sdk) for a complete interactive CLI agent that:
- Streams responses in real-time
- Switches between models dynamically
- Fetches available models from the proxy
```bash
# Clone and run the example
git clone https://github.com/BerriAI/litellm.git
cd litellm/cookbook/anthropic_agent_sdk
pip install -r requirements.txt
python main.py
```
## Related Resources
- [Claude Agent SDK Documentation](https://github.com/anthropics/anthropic-agent-sdk)
- [LiteLLM Proxy Quick Start](../proxy/quick_start)
- [Complete Cookbook Example](https://github.com/BerriAI/litellm/tree/main/cookbook/anthropic_agent_sdk)
+14 -1
View File
@@ -139,6 +139,20 @@ const sidebars = {
"tutorials/openai_codex"
]
},
{
type: "category",
label: "Agent SDKs",
link: {
type: "generated-index",
title: "Agent SDKs",
description: "Use LiteLLM with agent frameworks and SDKs",
slug: "/agent_sdks"
},
items: [
"tutorials/claude_agent_sdk",
"tutorials/google_adk",
]
},
],
// But you can create a sidebar manually
@@ -931,7 +945,6 @@ const sidebars = {
type: "category",
label: "LiteLLM Python SDK Tutorials",
items: [
'tutorials/google_adk',
'tutorials/azure_openai',
'tutorials/instructor',
"tutorials/gradio_integration",
@@ -1,12 +1,12 @@
-- DropIndex
DROP INDEX "LiteLLM_PromptTable_prompt_id_key";
DROP INDEX IF EXISTS "LiteLLM_PromptTable_prompt_id_key";
-- AlterTable
ALTER TABLE "LiteLLM_PromptTable" ADD COLUMN "version" INTEGER NOT NULL DEFAULT 1;
ALTER TABLE "LiteLLM_PromptTable"
ADD COLUMN "version" INTEGER NOT NULL DEFAULT 1;
-- CreateIndex
CREATE INDEX "LiteLLM_PromptTable_prompt_id_idx" ON "LiteLLM_PromptTable"("prompt_id");
CREATE INDEX "LiteLLM_PromptTable_prompt_id_idx" ON "LiteLLM_PromptTable" ("prompt_id");
-- CreateIndex
CREATE UNIQUE INDEX "LiteLLM_PromptTable_prompt_id_version_key" ON "LiteLLM_PromptTable"("prompt_id", "version");
CREATE UNIQUE INDEX "LiteLLM_PromptTable_prompt_id_version_key" ON "LiteLLM_PromptTable" ("prompt_id", "version");
+1
View File
@@ -980,6 +980,7 @@ BEDROCK_CONVERSE_MODELS = [
"meta.llama3-2-90b-instruct-v1:0",
"amazon.nova-lite-v1:0",
"amazon.nova-2-lite-v1:0",
"amazon.nova-2-pro-preview-20251202-v1:0",
"amazon.nova-pro-v1:0",
"writer.palmyra-x4-v1:0",
"writer.palmyra-x5-v1:0",
@@ -4752,7 +4752,14 @@ class StandardLoggingPayloadSetup:
) -> StandardLoggingPayloadErrorInformation:
from litellm.constants import MAXIMUM_TRACEBACK_LINES_TO_LOG
error_status: str = str(getattr(original_exception, "status_code", ""))
# Check for 'code' first (used by ProxyException), then fall back to 'status_code' (used by LiteLLM exceptions)
# Ensure error_code is always a string for Prisma Python JSON field compatibility
error_code_attr = getattr(original_exception, "code", None)
if error_code_attr is not None and str(error_code_attr) not in ("", "None"):
error_status: str = str(error_code_attr)
else:
status_code_attr = getattr(original_exception, "status_code", None)
error_status = str(status_code_attr) if status_code_attr is not None else ""
error_class: str = (
str(original_exception.__class__.__name__) if original_exception else ""
)
+13 -4
View File
@@ -290,10 +290,19 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
elif tool_choice == "none":
_tool_choice = AnthropicMessagesToolChoice(type="none")
elif isinstance(tool_choice, dict):
_tool_name = tool_choice.get("function", {}).get("name")
_tool_choice = AnthropicMessagesToolChoice(type="tool")
if _tool_name is not None:
_tool_choice["name"] = _tool_name
if "type" in tool_choice and "function" not in tool_choice:
tool_type = tool_choice.get("type")
if tool_type == "auto":
_tool_choice = AnthropicMessagesToolChoice(type="auto")
elif tool_type == "required" or tool_type == "any":
_tool_choice = AnthropicMessagesToolChoice(type="any")
elif tool_type == "none":
_tool_choice = AnthropicMessagesToolChoice(type="none")
else:
_tool_name = tool_choice.get("function", {}).get("name")
if _tool_name is not None:
_tool_choice = AnthropicMessagesToolChoice(type="tool")
_tool_choice["name"] = _tool_name
if parallel_tool_use is not None:
# Anthropic uses 'disable_parallel_tool_use' flag to determine if parallel tool use is allowed
@@ -22,7 +22,8 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config):
Accepts both explicit gpt-5 model names and the ``gpt5_series/`` prefix
used for manual routing.
"""
return "gpt-5" in model or "gpt5_series" in model
# gpt-5-chat* is a chat model and shouldn't go through GPT-5 reasoning restrictions.
return ("gpt-5" in model and "gpt-5-chat" not in model) or "gpt5_series" in model
def get_supported_openai_params(self, model: str) -> List[str]:
"""Get supported parameters for Azure OpenAI GPT-5 models.
@@ -37,6 +38,11 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config):
"""
params = OpenAIGPT5Config.get_supported_openai_params(self, model=model)
# Azure supports tool_choice for GPT-5 deployments, but the base GPT-5 config
# can drop it when the deployment name isn't in the OpenAI model registry.
if "tool_choice" not in params:
params.append("tool_choice")
# Only gpt-5.2 has been verified to support logprobs on Azure
if self.is_model_gpt_5_2_model(model):
azure_supported_params = ["logprobs", "top_logprobs"]
@@ -76,6 +76,13 @@ BEDROCK_COMPUTER_USE_TOOLS = [
"text_editor_",
]
# Beta header patterns that are not supported by Bedrock Converse API
# These will be filtered out to prevent errors
UNSUPPORTED_BEDROCK_CONVERSE_BETA_PATTERNS = [
"advanced-tool-use", # Bedrock Converse doesn't support advanced-tool-use beta headers
"prompt-caching", # Prompt caching not supported in Converse API
]
class AmazonConverseConfig(BaseConfig):
"""
@@ -610,6 +617,37 @@ class AmazonConverseConfig(BaseConfig):
return transformed_tools
def _filter_unsupported_beta_headers_for_bedrock(
self, model: str, beta_list: list
) -> list:
"""
Remove beta headers that are not supported on Bedrock Converse API for the given model.
Extended thinking beta headers are only supported on specific Claude 4+ models.
Some beta headers are universally unsupported on Bedrock Converse API.
Args:
model: The model name
beta_list: The list of beta headers to filter
Returns:
Filtered list of beta headers
"""
filtered_betas = []
# 1. Filter out beta headers that are universally unsupported on Bedrock Converse
for beta in beta_list:
should_keep = True
for unsupported_pattern in UNSUPPORTED_BEDROCK_CONVERSE_BETA_PATTERNS:
if unsupported_pattern in beta.lower():
should_keep = False
break
if should_keep:
filtered_betas.append(beta)
return filtered_betas
def _separate_computer_use_tools(
self, tools: List[OpenAIChatCompletionToolParam], model: str
) -> Tuple[
@@ -1088,7 +1126,14 @@ class AmazonConverseConfig(BaseConfig):
if beta not in seen:
unique_betas.append(beta)
seen.add(beta)
additional_request_params["anthropic_beta"] = unique_betas
# Filter out unsupported beta headers for Bedrock Converse API
filtered_betas = self._filter_unsupported_beta_headers_for_bedrock(
model=model,
beta_list=unique_betas,
)
additional_request_params["anthropic_beta"] = filtered_betas
return bedrock_tools, anthropic_beta_list
@@ -53,13 +53,26 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
model: str,
drop_params: bool,
) -> dict:
return AnthropicConfig.map_openai_params(
# Force tool-based structured outputs for Bedrock Invoke
# (similar to VertexAI fix in #19201)
# Bedrock Invoke doesn't support output_format parameter
original_model = model
if "response_format" in non_default_params:
# Use a model name that forces tool-based approach
model = "claude-3-sonnet-20240229"
optional_params = AnthropicConfig.map_openai_params(
self,
non_default_params,
optional_params,
model,
drop_params,
)
# Restore original model name
model = original_model
return optional_params
def transform_request(
@@ -90,6 +103,8 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
_anthropic_request.pop("model", None)
_anthropic_request.pop("stream", None)
# Bedrock Invoke doesn't support output_format parameter
_anthropic_request.pop("output_format", None)
if "anthropic_version" not in _anthropic_request:
_anthropic_request["anthropic_version"] = self.anthropic_version
@@ -117,6 +132,26 @@ class AmazonAnthropicClaudeConfig(AmazonInvokeConfig, AnthropicConfig):
if "opus-4" in model.lower() or "opus_4" in model.lower():
beta_set.add("tool-search-tool-2025-10-19")
# Filter out beta headers that Bedrock Invoke doesn't support
# AWS Bedrock only supports a specific whitelist of beta flags
# Reference: https://docs.aws.amazon.com/bedrock/latest/userguide/model-parameters-anthropic-claude-messages-request-response.html
BEDROCK_SUPPORTED_BETAS = {
"computer-use-2024-10-22", # Legacy computer use
"computer-use-2025-01-24", # Current computer use (Claude 3.7 Sonnet)
"token-efficient-tools-2025-02-19", # Tool use (Claude 3.7+ and Claude 4+)
"interleaved-thinking-2025-05-14", # Interleaved thinking (Claude 4+)
"output-128k-2025-02-19", # 128K output tokens (Claude 3.7 Sonnet)
"dev-full-thinking-2025-05-14", # Developer mode for raw thinking (Claude 4+)
"context-1m-2025-08-07", # 1 million tokens (Claude Sonnet 4)
"context-management-2025-06-27", # Context management (Claude Sonnet/Haiku 4.5)
"effort-2025-11-24", # Effort parameter (Claude Opus 4.5)
"tool-search-tool-2025-10-19", # Tool search (Claude Opus 4.5)
"tool-examples-2025-10-29", # Tool use examples (Claude Opus 4.5)
}
# Only keep beta headers that Bedrock supports
beta_set = {beta for beta in beta_set if beta in BEDROCK_SUPPORTED_BETAS}
if beta_set:
_anthropic_request["anthropic_beta"] = list(beta_set)
@@ -54,6 +54,7 @@ class AmazonAnthropicClaudeMessagesConfig(
# These will be filtered out to prevent 400 "invalid beta flag" errors
UNSUPPORTED_BEDROCK_INVOKE_BETA_PATTERNS = [
"advanced-tool-use", # Bedrock Invoke doesn't support advanced-tool-use beta headers
"prompt-caching-scope"
]
def __init__(self, **kwargs):
+1 -1
View File
@@ -92,7 +92,7 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
"parallel_tool_calls",
"web_search_options",
]
if supports_reasoning(model):
if supports_reasoning(model, custom_llm_provider="gemini"):
supported_params.append("reasoning_effort")
supported_params.append("thinking")
if self.is_model_gemini_audio_model(model):
@@ -19,7 +19,9 @@ class OpenAIGPT5Config(OpenAIGPTConfig):
@classmethod
def is_model_gpt_5_model(cls, model: str) -> bool:
return "gpt-5" in model
# gpt-5-chat* behaves like a regular chat model (supports temperature, etc.)
# Don't route it through GPT-5 reasoning-specific parameter restrictions.
return "gpt-5" in model and "gpt-5-chat" not in model
@classmethod
def is_model_gpt_5_codex_model(cls, model: str) -> bool:
@@ -1657,7 +1657,17 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig):
## This is necessary because promptTokensDetails includes both cached and non-cached tokens
## See: https://github.com/BerriAI/litellm/issues/18750
if cached_text_tokens is not None and prompt_text_tokens is not None:
# Explicit caching: subtract cached tokens per modality from cacheTokensDetails
prompt_text_tokens = prompt_text_tokens - cached_text_tokens
elif (
cached_tokens is not None
and prompt_text_tokens is not None
and cached_text_tokens is None
):
# Implicit caching: only cachedContentTokenCount is provided (no cacheTokensDetails)
# Subtract from text tokens since implicit caching is primarily for text content
# See: https://github.com/BerriAI/litellm/issues/16341
prompt_text_tokens = prompt_text_tokens - cached_tokens
if cached_audio_tokens is not None and prompt_audio_tokens is not None:
prompt_audio_tokens = prompt_audio_tokens - cached_audio_tokens
if cached_image_tokens is not None and prompt_image_tokens is not None:
+5 -2
View File
@@ -7280,8 +7280,11 @@ def _get_encoding():
def __getattr__(name: str) -> Any:
"""Lazy import handler for main module"""
if name == "encoding":
# Lazy load encoding to avoid heavy tiktoken import at module load time
_encoding = tiktoken.get_encoding("cl100k_base")
# Use _get_default_encoding which properly sets TIKTOKEN_CACHE_DIR
# before loading tiktoken, ensuring the local cache is used
# instead of downloading from the internet
from litellm._lazy_imports import _get_default_encoding
_encoding = _get_default_encoding()
# Cache it in the module's __dict__ for subsequent accesses
import sys
@@ -354,6 +354,25 @@
"supports_video_input": true,
"supports_vision": true
},
"amazon.nova-2-pro-preview-20251202-v1:0": {
"cache_read_input_token_cost": 5.46875e-07,
"input_cost_per_token": 2.1875e-06,
"input_cost_per_image_token": 2.1875e-06,
"input_cost_per_audio_token": 2.1875e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.75e-05,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_video_input": true,
"supports_vision": true
},
"apac.amazon.nova-2-lite-v1:0": {
"cache_read_input_token_cost": 8.25e-08,
"input_cost_per_token": 3.3e-07,
@@ -371,6 +390,25 @@
"supports_video_input": true,
"supports_vision": true
},
"apac.amazon.nova-2-pro-preview-20251202-v1:0": {
"cache_read_input_token_cost": 5.46875e-07,
"input_cost_per_token": 2.1875e-06,
"input_cost_per_image_token": 2.1875e-06,
"input_cost_per_audio_token": 2.1875e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.75e-05,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_video_input": true,
"supports_vision": true
},
"eu.amazon.nova-2-lite-v1:0": {
"cache_read_input_token_cost": 8.25e-08,
"input_cost_per_token": 3.3e-07,
@@ -388,6 +426,25 @@
"supports_video_input": true,
"supports_vision": true
},
"eu.amazon.nova-2-pro-preview-20251202-v1:0": {
"cache_read_input_token_cost": 5.46875e-07,
"input_cost_per_token": 2.1875e-06,
"input_cost_per_image_token": 2.1875e-06,
"input_cost_per_audio_token": 2.1875e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.75e-05,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_video_input": true,
"supports_vision": true
},
"us.amazon.nova-2-lite-v1:0": {
"cache_read_input_token_cost": 8.25e-08,
"input_cost_per_token": 3.3e-07,
@@ -405,6 +462,25 @@
"supports_video_input": true,
"supports_vision": true
},
"us.amazon.nova-2-pro-preview-20251202-v1:0": {
"cache_read_input_token_cost": 5.46875e-07,
"input_cost_per_token": 2.1875e-06,
"input_cost_per_image_token": 2.1875e-06,
"input_cost_per_audio_token": 2.1875e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.75e-05,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_video_input": true,
"supports_vision": true
},
"amazon.nova-2-multimodal-embeddings-v1:0": {
"litellm_provider": "bedrock",
"max_input_tokens": 8172,
@@ -3130,7 +3206,7 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": false,
"supports_tool_choice": true,
"supports_vision": true
},
"azure/gpt-5-chat-latest": {
@@ -3162,7 +3238,7 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": false,
"supports_tool_choice": true,
"supports_vision": true
},
"azure/gpt-5-codex": {
@@ -13525,7 +13601,7 @@
"cache_read_input_token_cost": 0,
"input_cost_per_token": 3e-07,
"input_cost_per_audio_token": 1e-06,
"litellm_provider": "gemini",
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 1048576,
"max_output_tokens": 65535,
"max_tokens": 65535,
+2
View File
@@ -227,6 +227,7 @@ class KeyManagementRoutes(str, enum.Enum):
KEY_REGENERATE_WITH_PATH_PARAM = "/key/{key_id}/regenerate"
KEY_BLOCK = "/key/block"
KEY_UNBLOCK = "/key/unblock"
KEY_BULK_UPDATE = "/key/bulk_update"
# info and health routes
KEY_INFO = "/key/info"
@@ -498,6 +499,7 @@ class LiteLLMRoutes(enum.Enum):
KeyManagementRoutes.KEY_LIST.value,
KeyManagementRoutes.KEY_BLOCK.value,
KeyManagementRoutes.KEY_UNBLOCK.value,
KeyManagementRoutes.KEY_BULK_UPDATE.value,
]
management_routes = [
+40 -32
View File
@@ -53,6 +53,7 @@ from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler
from litellm.proxy.auth.oauth2_check import Oauth2Handler
from litellm.proxy.auth.oauth2_proxy_hook import handle_oauth2_proxy_request
from litellm.proxy.auth.route_checks import RouteChecks
from litellm.proxy.common_utils.cache_coordinator import EventDrivenCacheCoordinator
from litellm.proxy.common_utils.http_parsing_utils import (
_read_request_body,
_safe_get_request_headers,
@@ -211,6 +212,33 @@ def update_valid_token_with_end_user_params(
return valid_token
# Reusable coordinator for global spend to prevent cache stampede
_global_spend_coordinator = EventDrivenCacheCoordinator(log_prefix="[GLOBAL SPEND]")
async def _fetch_global_spend_with_event_coordination(
cache_key: str,
user_api_key_cache: DualCache,
prisma_client: PrismaClient,
) -> Optional[float]:
"""
Fetch global spend with event-driven coordination to prevent cache stampede.
Uses EventDrivenCacheCoordinator: first request queries DB and signals others when done.
"""
async def _load_global_spend() -> Optional[float]:
sql_query = """SELECT SUM(spend) AS total_spend FROM "MonthlyGlobalSpend";"""
response = await prisma_client.db.query_raw(query=sql_query)
val = response[0]["total_spend"]
return float(val) if val is not None else None
return await _global_spend_coordinator.get_or_load(
cache_key=cache_key,
cache=user_api_key_cache,
load_fn=_load_global_spend,
)
async def get_global_proxy_spend(
litellm_proxy_admin_name: str,
user_api_key_cache: DualCache,
@@ -219,25 +247,14 @@ async def get_global_proxy_spend(
proxy_logging_obj: ProxyLogging,
) -> Optional[float]:
global_proxy_spend = None
if litellm.max_budget > 0: # user set proxy max budget
# check cache
global_proxy_spend = await user_api_key_cache.async_get_cache(
key="{}:spend".format(litellm_proxy_admin_name)
if litellm.max_budget > 0 and prisma_client is not None: # user set proxy max budget
# Use event-driven coordination to prevent cache stampede
cache_key = "{}:spend".format(litellm_proxy_admin_name)
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
cache_key=cache_key,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
)
if global_proxy_spend is None and prisma_client is not None:
# get from db
sql_query = (
"""SELECT SUM(spend) as total_spend FROM "MonthlyGlobalSpend";"""
)
response = await prisma_client.db.query_raw(query=sql_query)
global_proxy_spend = response[0]["total_spend"]
await user_api_key_cache.async_set_cache(
key="{}:spend".format(litellm_proxy_admin_name),
value=global_proxy_spend,
)
if global_proxy_spend is not None:
user_info = CallInfo(
user_id=litellm_proxy_admin_name,
@@ -1120,21 +1137,12 @@ async def _user_api_key_auth_builder( # noqa: PLR0915
if (
litellm.max_budget > 0 and prisma_client is not None
): # user set proxy max budget
# check cache
global_proxy_spend = await user_api_key_cache.async_get_cache(
key="{}:spend".format(litellm_proxy_admin_name)
cache_key = "{}:spend".format(litellm_proxy_admin_name)
global_proxy_spend = await _fetch_global_spend_with_event_coordination(
cache_key=cache_key,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
)
if global_proxy_spend is None:
# get from db
sql_query = """SELECT SUM(spend) as total_spend FROM "MonthlyGlobalSpend";"""
response = await prisma_client.db.query_raw(query=sql_query)
global_proxy_spend = response[0]["total_spend"]
await user_api_key_cache.async_set_cache(
key="{}:spend".format(litellm_proxy_admin_name),
value=global_proxy_spend,
)
if global_proxy_spend is not None:
call_info = CallInfo(
+8 -4
View File
@@ -650,11 +650,15 @@ class ProxyBaseLLMRequestProcessing:
)
tasks = []
# Start the moderation check (during_call_hook) as early as possible
# This gives it a head start to mask/validate input while the proxy handles routing
tasks.append(
proxy_logging_obj.during_call_hook(
data=self.data,
user_api_key_dict=user_api_key_dict,
call_type=route_type, # type: ignore
asyncio.create_task(
proxy_logging_obj.during_call_hook(
data=self.data,
user_api_key_dict=user_api_key_dict,
call_type=route_type, # type: ignore
)
)
)
@@ -0,0 +1,191 @@
"""
Event-driven cache coordinator to prevent cache stampede.
Use this when many requests can miss the same cache key at once (e.g. after
expiry or restart). Without coordination, they would all run the expensive
load (DB query, API call) in parallel and overload the backend.
This module ensures only one request performs the load; the rest wait for a
signal and then read the freshly cached value. Reuse it for any cache-aside
pattern: global spend, feature flags, config, or other shared read-through data.
"""
import asyncio
import time
from typing import Any, Awaitable, Callable, Optional, Protocol, TypeVar
from litellm._logging import verbose_proxy_logger
T = TypeVar("T")
class AsyncCacheProtocol(Protocol):
"""Protocol for cache backends used by EventDrivenCacheCoordinator."""
async def async_get_cache(self, key: str, **kwargs: Any) -> Any:
...
async def async_set_cache(self, key: str, value: Any, **kwargs: Any) -> Any:
...
class EventDrivenCacheCoordinator:
"""
Coordinates a single in-flight load per logical resource to prevent cache stampede.
Pattern:
- First request: loads data (e.g. DB query), caches it, then signals waiters.
- Other requests: wait for the signal, then read from cache.
Create one instance per resource (e.g. one for global spend, one for feature flags).
"""
def __init__(self, log_prefix: str = "[CACHE]"):
self._lock = asyncio.Lock()
self._event: Optional[asyncio.Event] = None
self._query_in_progress = False
self._log_prefix = log_prefix
async def _get_cached(
self, cache_key: str, cache: AsyncCacheProtocol
) -> Optional[T]:
"""Return value from cache if present, else None."""
return await cache.async_get_cache(key=cache_key)
def _log_cache_hit(self, value: T) -> None:
if self._log_prefix:
verbose_proxy_logger.debug(
"%s Cache hit, value: %s", self._log_prefix, value
)
def _log_cache_miss(self) -> None:
if self._log_prefix:
verbose_proxy_logger.debug("%s Cache miss", self._log_prefix)
async def _claim_role(self) -> Optional[asyncio.Event]:
"""
Under lock: return event to wait on if load is in progress, else set us as loader and return None.
"""
async with self._lock:
if self._query_in_progress and self._event is not None:
if self._log_prefix:
verbose_proxy_logger.debug(
"%s Load in flight, waiting for signal", self._log_prefix
)
return self._event
self._query_in_progress = True
self._event = asyncio.Event()
if self._log_prefix:
verbose_proxy_logger.debug(
"%s Starting load (will signal others when done)",
self._log_prefix,
)
return None
async def _wait_for_signal_and_get(
self,
event: asyncio.Event,
cache_key: str,
cache: AsyncCacheProtocol,
) -> Optional[T]:
"""Wait for loader to finish, then read from cache."""
await event.wait()
if self._log_prefix:
verbose_proxy_logger.debug(
"%s Signal received, reading from cache", self._log_prefix
)
value = await cache.async_get_cache(key=cache_key)
if value is not None and self._log_prefix:
verbose_proxy_logger.debug(
"%s Cache filled by other request, value: %s",
self._log_prefix,
value,
)
elif value is None and self._log_prefix:
verbose_proxy_logger.debug(
"%s Signal received but cache still empty", self._log_prefix
)
return value
async def _load_and_cache(
self,
cache_key: str,
cache: AsyncCacheProtocol,
load_fn: Callable[[], Awaitable[T]],
) -> Optional[T]:
"""Double-check cache, run load_fn, set cache, return value. Caller must call _signal_done in finally."""
value = await cache.async_get_cache(key=cache_key)
if value is not None:
if self._log_prefix:
verbose_proxy_logger.debug(
"%s Cache filled while acquiring lock, value: %s",
self._log_prefix,
value,
)
return value
if self._log_prefix:
verbose_proxy_logger.debug("%s Running load", self._log_prefix)
start = time.perf_counter()
value = await load_fn()
elapsed_ms = (time.perf_counter() - start) * 1000
if self._log_prefix:
verbose_proxy_logger.debug(
"%s Load completed in %.2fms, result: %s",
self._log_prefix,
elapsed_ms,
value,
)
await cache.async_set_cache(key=cache_key, value=value)
if self._log_prefix:
verbose_proxy_logger.debug("%s Result cached", self._log_prefix)
return value
async def _signal_done(self) -> None:
"""Reset loader state and signal all waiters."""
async with self._lock:
self._query_in_progress = False
if self._event is not None:
if self._log_prefix:
verbose_proxy_logger.debug(
"%s Signaling all waiting requests", self._log_prefix
)
self._event.set()
self._event = None
async def get_or_load(
self,
cache_key: str,
cache: AsyncCacheProtocol,
load_fn: Callable[[], Awaitable[T]],
) -> Optional[T]:
"""
Return cached value or load it once and signal waiters.
- cache_key: Key to read/write in the cache.
- cache: Object with async_get_cache(key) and async_set_cache(key, value).
- load_fn: Async callable that performs the load (e.g. DB query). No args.
Return value is cached and returned. If it raises, waiters are
still signaled so they can retry or handle empty cache.
Returns the value from cache or from load_fn, or None if load failed or
cache was still empty after waiting.
"""
value = await self._get_cached(cache_key, cache)
if value is not None:
self._log_cache_hit(value)
return value
self._log_cache_miss()
event_to_wait = await self._claim_role()
if event_to_wait is not None:
return await self._wait_for_signal_and_get(
event_to_wait, cache_key, cache
)
try:
return await self._load_and_cache(cache_key, cache, load_fn)
finally:
await self._signal_done()
@@ -198,6 +198,15 @@ class ContentFilterGuardrail(CustomGuardrail):
for pattern_config in normalized_patterns:
self._add_pattern(pattern_config)
# Warn if using during_call with MASK action (unstable)
if self.event_hook == GuardrailEventHooks.during_call and any(
p["action"] == ContentFilterAction.MASK for p in self.compiled_patterns
):
verbose_proxy_logger.warning(
f"ContentFilterGuardrail '{self.guardrail_name}': 'during_call' mode with 'MASK' action is unstable due to race conditions. "
"Use 'pre_call' mode for reliable request masking."
)
# Load blocked words - always initialize as dict
self.blocked_words: Dict[str, Tuple[ContentFilterAction, Optional[str]]] = {}
for word in normalized_blocked_words:
@@ -905,11 +914,15 @@ class ContentFilterGuardrail(CustomGuardrail):
elif isinstance(e.detail, str):
e.detail = e.detail + " (Image description): " + description
else:
e.detail = "Content blocked: Image description detected" + description
e.detail = (
"Content blocked: Image description detected" + description
)
raise e
def _count_masked_entities(
self, detections: List[ContentFilterDetection], masked_entity_count: Dict[str, int]
self,
detections: List[ContentFilterDetection],
masked_entity_count: Dict[str, int],
) -> None:
"""
Count masked entities by type from detections.
@@ -964,9 +977,11 @@ class ContentFilterGuardrail(CustomGuardrail):
dict(detection) for detection in detections
]
if status != "success":
guardrail_json_response = exception_str if exception_str else [
dict(detection) for detection in detections
]
guardrail_json_response = (
exception_str
if exception_str
else [dict(detection) for detection in detections]
)
self.add_standard_logging_guardrail_information_to_request_data(
guardrail_provider=self.guardrail_provider,
@@ -1066,99 +1081,84 @@ class ContentFilterGuardrail(CustomGuardrail):
Process streaming response chunks and check for blocked content.
For BLOCK action: Raises HTTPException immediately when blocked content is detected.
For MASK action: Content passes through (masking streaming responses is not supported).
For MASK action: Content is buffered to handle patterns split across chunks.
"""
accumulated_full_text = ""
yielded_masked_text_len = 0
buffer_size = 50 # Increased buffer to catch patterns split across many chunks
# Accumulate content as we iterate through chunks
accumulated_content = ""
verbose_proxy_logger.info(
f"ContentFilterGuardrail: Starting robust streaming masking for model {request_data.get('model')}"
)
async for item in response:
# Accumulate content from this chunk before checking
if isinstance(item, ModelResponseStream) and item.choices:
delta_content = ""
is_final = False
for choice in item.choices:
if hasattr(choice, "delta") and choice.delta:
content = getattr(choice.delta, "content", None)
if content and isinstance(content, str):
accumulated_content += content
delta_content += content
if getattr(choice, "finish_reason", None):
is_final = True
# Check accumulated content for blocked patterns/keywords after processing all choices
# Only check for BLOCK actions, not MASK (masking streaming is not supported)
if accumulated_content:
try:
# Check patterns
pattern_match = self._check_patterns(accumulated_content)
if pattern_match:
matched_text, pattern_name, action = pattern_match
if action == ContentFilterAction.BLOCK:
error_msg = (
f"Content blocked: {pattern_name} pattern detected"
)
verbose_proxy_logger.warning(error_msg)
raise HTTPException(
status_code=403,
detail={
"error": error_msg,
"pattern": pattern_name,
},
)
accumulated_full_text += delta_content
# Check blocked words
blocked_word_match = self._check_blocked_words(
accumulated_content
)
if blocked_word_match:
keyword, action, description = blocked_word_match
if action == ContentFilterAction.BLOCK:
error_msg = (
f"Content blocked: keyword '{keyword}' detected"
)
if description:
error_msg += f" ({description})"
verbose_proxy_logger.warning(error_msg)
raise HTTPException(
status_code=403,
detail={
"error": error_msg,
"keyword": keyword,
"description": description,
},
)
# Check for blocking or apply masking
# Add a space at the end if it's the final chunk to trigger word boundaries (\b)
text_to_check = accumulated_full_text
if is_final:
text_to_check += " "
# Check category keywords
all_exceptions = []
for category in self.loaded_categories.values():
all_exceptions.extend(category.exceptions)
category_match = self._check_category_keywords(
accumulated_content, all_exceptions
)
if category_match:
keyword, category_name, severity, action = category_match
if action == ContentFilterAction.BLOCK:
error_msg = (
f"Content blocked: {category_name} category keyword '{keyword}' detected "
f"(severity: {severity})"
)
verbose_proxy_logger.warning(error_msg)
raise HTTPException(
status_code=403,
detail={
"error": error_msg,
"category": category_name,
"keyword": keyword,
"severity": severity,
},
)
except HTTPException:
# Re-raise HTTPException (blocked content detected)
raise
except Exception as e:
# Log other exceptions but don't block the stream
verbose_proxy_logger.warning(
f"Error checking content filter in streaming: {e}"
)
try:
masked_text = self._filter_single_text(text_to_check)
if is_final and masked_text.endswith(" "):
masked_text = masked_text[:-1]
except HTTPException:
raise
except Exception as e:
verbose_proxy_logger.error(
f"ContentFilterGuardrail: Error in masking: {e}"
)
masked_text = text_to_check # Fallback to current text
# Yield the chunk (only if no exception was raised above)
yield item
# Determine how much can be safely yielded
if is_final:
safe_to_yield_len = len(masked_text)
else:
safe_to_yield_len = max(0, len(masked_text) - buffer_size)
if safe_to_yield_len > yielded_masked_text_len:
new_masked_content = masked_text[
yielded_masked_text_len:safe_to_yield_len
]
# Modify the chunk to contain only the new masked content
if (
item.choices
and hasattr(item.choices[0], "delta")
and item.choices[0].delta
):
item.choices[0].delta.content = new_masked_content
yielded_masked_text_len = safe_to_yield_len
yield item
else:
# Hold content by yielding empty content chunk (keeps metadata/structure)
if (
item.choices
and hasattr(item.choices[0], "delta")
and item.choices[0].delta
):
item.choices[0].delta.content = ""
yield item
else:
# Not a ModelResponseStream or no choices - yield as is
yield item
# Any remaining content (should have been handled by is_final, but just in case)
if yielded_masked_text_len < len(accumulated_full_text):
# We already reached the end of the generator
pass
@staticmethod
def get_config_model():
@@ -108,7 +108,7 @@
{
"name": "ipv6",
"display_name": "IP Address (IPv6)",
"pattern": "\\b(?:[0-9a-fA-F]{1,4}:){7}[0-9a-fA-F]{1,4}\\b",
"pattern": "(?<![0-9a-fA-F:])(?:(?:[0-9a-fA-F]{1,4}:){7}[0-9a-fA-F]{1,4}|(?:[0-9a-fA-F]{1,4}:){1,7}:|:(?::[0-9a-fA-F]{1,4}){1,7}|::|(?:[0-9a-fA-F]{1,4}:){1,6}:[0-9a-fA-F]{1,4}|(?:[0-9a-fA-F]{1,4}:){1,5}(?::[0-9a-fA-F]{1,4}){1,2}|(?:[0-9a-fA-F]{1,4}:){1,4}(?::[0-9a-fA-F]{1,4}){1,3}|(?:[0-9a-fA-F]{1,4}:){1,3}(?::[0-9a-fA-F]{1,4}){1,4}|(?:[0-9a-fA-F]{1,4}:){1,2}(?::[0-9a-fA-F]{1,4}){1,5}|[0-9a-fA-F]{1,4}:(?::[0-9a-fA-F]{1,4}){1,6})(?![0-9a-fA-F:])",
"category": "Network Patterns",
"description": "Detects IPv6 addresses"
},
@@ -122,9 +122,9 @@
{
"name": "passport_us",
"display_name": "Passport (US)",
"pattern": "\\b[0-9]{9}\\b",
"pattern": "\\b([A-Z][0-9]{8}|[0-9]{9})\\b",
"category": "PII Patterns",
"description": "US passport numbers (9 digits)"
"description": "US passport numbers (9 digits or alphanumeric letter + 8 digits)"
},
{
"name": "passport_uk",
@@ -157,9 +157,9 @@
{
"name": "passport_canada",
"display_name": "Passport (Canada)",
"pattern": "\\b[A-Z]{2}[0-9]{6}\\b",
"pattern": "\\b([A-Z]{2}[0-9]{6}|[A-Z][0-9]{6}[A-Z]{2})\\b",
"category": "PII Patterns",
"description": "Canadian passport numbers (2 letters + 6 digits)"
"description": "Canadian passport numbers (old: 2 letters + 6 digits; new: 1 letter + 6 digits + 2 letters)"
},
{
"name": "passport_india",
@@ -369,4 +369,4 @@
"description": "Detects Brazilian RG identity card numbers (common pattern for SP, RJ, MG states)"
}
]
}
}
@@ -37,6 +37,13 @@ from litellm.proxy._experimental.mcp_server.db import (
)
from litellm.proxy._types import *
from litellm.proxy._types import LiteLLM_VerificationToken
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateKeyRequest,
BulkUpdateKeyRequestItem,
BulkUpdateKeyResponse,
FailedKeyUpdate,
SuccessfulKeyUpdate,
)
from litellm.proxy.auth.auth_checks import (
_cache_key_object,
_delete_cache_key_object,
@@ -1438,6 +1445,211 @@ def is_different_team(
return data.team_id != existing_key_row.team_id
def _validate_max_budget(max_budget: Optional[float]) -> None:
"""
Validate that max_budget is not negative.
Args:
max_budget: The max_budget value to validate
Raises:
HTTPException: If max_budget is negative
"""
if max_budget is not None and max_budget < 0:
raise HTTPException(
status_code=400,
detail={
"error": f"max_budget cannot be negative. Received: {max_budget}"
},
)
async def _get_and_validate_existing_key(
token: str, prisma_client: Optional[PrismaClient]
) -> LiteLLM_VerificationToken:
"""
Get existing key from database and validate it exists.
Args:
token: The key token to look up
prisma_client: Prisma client instance
Returns:
LiteLLM_VerificationToken: The existing key row
Raises:
HTTPException: If key is not found
"""
if prisma_client is None:
raise HTTPException(
status_code=500,
detail={"error": "Database not connected"},
)
existing_key_row = await prisma_client.get_data(
token=token,
table_name="key",
query_type="find_unique",
)
if existing_key_row is None:
raise HTTPException(
status_code=404,
detail={"error": f"Key not found: {token}"},
)
return existing_key_row
async def _process_single_key_update(
key_update_item: BulkUpdateKeyRequestItem,
user_api_key_dict: UserAPIKeyAuth,
litellm_changed_by: Optional[str],
prisma_client: Optional[PrismaClient],
user_api_key_cache: DualCache,
proxy_logging_obj: Any,
llm_router: Optional[Router],
) -> Dict[str, Any]:
"""
Process a single key update with all validations and checks.
This function encapsulates all the logic for updating a single key,
including validation, permission checks, team checks, and database updates.
Args:
key_update_item: The key update request item
user_api_key_dict: The authenticated user's API key info
litellm_changed_by: Optional header for tracking who made the change
prisma_client: Prisma client instance
user_api_key_cache: User API key cache
proxy_logging_obj: Proxy logging object
llm_router: LLM router instance
Returns:
Dict containing the updated key information
Raises:
HTTPException: For various validation and permission errors
"""
# Validate max_budget
_validate_max_budget(key_update_item.max_budget)
# Get and validate existing key
existing_key_row = await _get_and_validate_existing_key(
token=key_update_item.key,
prisma_client=prisma_client,
)
# Check team member permissions
await TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint(
user_api_key_dict=user_api_key_dict,
route=KeyManagementRoutes.KEY_UPDATE,
prisma_client=prisma_client,
existing_key_row=existing_key_row,
user_api_key_cache=user_api_key_cache,
)
# Create UpdateKeyRequest from BulkUpdateKeyRequestItem
update_key_request = UpdateKeyRequest(
key=key_update_item.key,
budget_id=key_update_item.budget_id,
max_budget=key_update_item.max_budget,
team_id=key_update_item.team_id,
tags=key_update_item.tags,
)
# Get team object and check team limits if team_id is provided
team_obj: Optional[LiteLLM_TeamTableCachedObj] = None
if update_key_request.team_id is not None:
team_obj = await get_team_object(
team_id=update_key_request.team_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
check_db_only=True,
)
if team_obj is not None and prisma_client is not None:
await _check_team_key_limits(
team_table=team_obj,
data=update_key_request,
prisma_client=prisma_client,
)
# Validate team change if team is being changed
if is_different_team(
data=update_key_request, existing_key_row=existing_key_row
):
if llm_router is None:
raise HTTPException(
status_code=400,
detail={
"error": "LLM router not found. Please set it up by passing in a valid config.yaml or adding models via the UI."
},
)
if team_obj is None:
raise HTTPException(
status_code=500,
detail={
"error": "Team object not found for team change validation"
},
)
validate_key_team_change(
key=existing_key_row,
team=team_obj,
change_initiated_by=user_api_key_dict,
llm_router=llm_router,
)
# Prepare update data
non_default_values = await prepare_key_update_data(
data=update_key_request, existing_key_row=existing_key_row
)
# Update key in database
if prisma_client is None:
raise HTTPException(
status_code=500,
detail={"error": "Database not connected"},
)
_data = {**non_default_values, "token": key_update_item.key}
response = await prisma_client.update_data(
token=key_update_item.key, data=_data
)
# Delete cache
await _delete_cache_key_object(
hashed_token=hash_token(key_update_item.key),
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
)
# Trigger async hook
asyncio.create_task(
KeyManagementEventHooks.async_key_updated_hook(
data=update_key_request,
existing_key_row=existing_key_row,
response=response,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
)
)
if response is None:
raise ValueError("Failed to update key got response = None")
# Extract and format updated key info
updated_key_info = response.get("data", {})
if hasattr(updated_key_info, "model_dump"):
updated_key_info = updated_key_info.model_dump()
elif hasattr(updated_key_info, "dict"):
updated_key_info = updated_key_info.dict()
updated_key_info.pop("token", None)
return updated_key_info
@router.post(
"/key/update", tags=["key management"], dependencies=[Depends(user_api_key_auth)]
)
@@ -1684,6 +1896,167 @@ async def update_key_fn(
)
@router.post(
"/key/bulk_update",
tags=["key management"],
dependencies=[Depends(user_api_key_auth)],
response_model=BulkUpdateKeyResponse,
)
@management_endpoint_wrapper
async def bulk_update_keys(
data: BulkUpdateKeyRequest,
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
litellm_changed_by: Optional[str] = Header(
None,
description="The litellm-changed-by header enables tracking of actions performed by authorized users on behalf of other users, providing an audit trail for accountability",
),
):
"""
Bulk update multiple keys at once.
This endpoint allows updating multiple keys in a single request. Each key update
is processed independently - if some updates fail, others will still succeed.
Parameters:
- keys: List[BulkUpdateKeyRequestItem] - List of key update requests, each containing:
- key: str - The key identifier (token) to update
- budget_id: Optional[str] - Budget ID associated with the key
- max_budget: Optional[float] - Max budget for key
- team_id: Optional[str] - Team ID associated with key
- tags: Optional[List[str]] - Tags for organizing keys
Returns:
- total_requested: int - Total number of keys requested for update
- successful_updates: List[SuccessfulKeyUpdate] - List of successfully updated keys with their updated info
- failed_updates: List[FailedKeyUpdate] - List of failed updates with key_info and failed_reason
Example request:
```bash
curl --location 'http://0.0.0.0:4000/key/bulk_update' \
--header 'Authorization: Bearer sk-1234' \
--header 'Content-Type: application/json' \
--data '{
"keys": [
{
"key": "sk-1234",
"max_budget": 100.0,
"team_id": "team-123",
"tags": ["production", "api"]
},
{
"key": "sk-5678",
"budget_id": "budget-456",
"tags": ["staging"]
}
]
}'
```
"""
from litellm.proxy.proxy_server import (
llm_router,
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN.value:
raise HTTPException(
status_code=403,
detail={
"error": "Only proxy admins can perform bulk key updates"
},
)
if prisma_client is None:
raise HTTPException(
status_code=500,
detail={"error": "Database not connected"},
)
if not data.keys:
raise HTTPException(
status_code=400,
detail={"error": "No keys provided for update"},
)
MAX_BATCH_SIZE = 500
if len(data.keys) > MAX_BATCH_SIZE:
raise HTTPException(
status_code=400,
detail={
"error": f"Maximum {MAX_BATCH_SIZE} keys can be updated at once. Found {len(data.keys)} keys."
},
)
successful_updates: List[SuccessfulKeyUpdate] = []
failed_updates: List[FailedKeyUpdate] = []
for key_update_item in data.keys:
try:
# Process single key update using reusable function
updated_key_info = await _process_single_key_update(
key_update_item=key_update_item,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=litellm_changed_by,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
llm_router=llm_router,
)
successful_updates.append(
SuccessfulKeyUpdate(
key=key_update_item.key,
key_info=updated_key_info,
)
)
except Exception as e:
verbose_proxy_logger.exception(
f"Failed to update key {key_update_item.key}: {e}"
)
if isinstance(e, HTTPException):
error_detail = e.detail
if isinstance(error_detail, dict):
error_message = error_detail.get("error", str(e))
else:
error_message = str(error_detail)
else:
error_message = str(e)
key_info = None
try:
existing_key_row = await prisma_client.get_data(
token=key_update_item.key,
table_name="key",
query_type="find_unique",
)
if existing_key_row is not None:
if hasattr(existing_key_row, "model_dump"):
key_info = existing_key_row.model_dump()
elif hasattr(existing_key_row, "dict"):
key_info = existing_key_row.dict()
if key_info:
key_info.pop("token", None)
except Exception:
pass
failed_updates.append(
FailedKeyUpdate(
key=key_update_item.key,
key_info=key_info,
failed_reason=error_message,
)
)
return BulkUpdateKeyResponse(
total_requested=len(data.keys),
successful_updates=successful_updates,
failed_updates=failed_updates,
)
def validate_key_team_change(
key: LiteLLM_VerificationToken,
team: LiteLLM_TeamTable,
+19 -95
View File
@@ -1,101 +1,25 @@
model_list:
- model_name: "*"
- model_name: bedrock-claude-sonnet-3.5
litellm_params:
model: "*"
- model_name: "gpt-4"
model: "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-sonnet-4
litellm_params:
model: "gpt-4"
api_key: os.environ/OPENAI_API_KEY
- model_name: "gpt-3.5-turbo"
model: "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-sonnet-4.5
litellm_params:
model: "gpt-3.5-turbo"
api_key: os.environ/OPENAI_API_KEY
model: "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0"
aws_region_name: "us-east-1"
general_settings:
master_key: sk-1234
- model_name: bedrock-claude-opus-4.5
litellm_params:
model: "bedrock/converse/us.anthropic.claude-opus-4-5-20251101-v1:0"
aws_region_name: "us-east-1"
# ───────────────────────────────────────────────
# POLICIES - Define WHAT guardrails to apply
# ───────────────────────────────────────────────
#
# Policies define guardrails with:
# - inherit: Inherit guardrails from another policy
# - description: Human-readable description
# - guardrails.add: Add guardrails (on top of inherited)
# - guardrails.remove: Remove guardrails (from inherited)
# - condition.model: Model pattern (exact or regex) for when guardrails apply
#
policies:
# Global baseline policy
global-baseline:
description: "Base guardrails for all requests"
guardrails:
add:
- pii_blocker
# Healthcare policy - inherits from global-baseline
healthcare-compliance:
inherit: global-baseline
description: "HIPAA compliance for healthcare teams"
guardrails:
add:
- hipaa_audit
# Dev policy - inherits but removes PII blocker for testing
internal-dev:
inherit: global-baseline
description: "Relaxed policy for internal development"
guardrails:
add:
- toxicity_filter
remove:
- pii_blocker
# Policy with model condition (regex pattern)
gpt4-safety:
description: "Extra safety for GPT-4 models"
guardrails:
add:
- toxicity_filter
condition:
model: "gpt-4.*" # regex: matches gpt-4, gpt-4-turbo, gpt-4o, etc.
# Policy with model condition (exact match list)
bedrock-compliance:
description: "Compliance for Bedrock models"
guardrails:
add:
- strict_pii_blocker
condition:
model: ["bedrock/claude-3", "bedrock/claude-2"] # exact matches
# ───────────────────────────────────────────────
# POLICY ATTACHMENTS - Define WHERE policies apply
# ───────────────────────────────────────────────
#
# Attachments are REQUIRED to make policies active.
# A policy without an attachment will not be applied.
#
policy_attachments:
# Global attachment - applies to all requests
- policy: global-baseline
scope: "*"
# Team-specific attachment
- policy: healthcare-compliance
teams:
- healthcare-team
- medical-research
# Key pattern attachment
- policy: internal-dev
keys:
- "dev-key-*"
- "test-key-*"
# Model-specific policies (attached globally, condition filters by model)
- policy: gpt4-safety
scope: "*"
- policy: bedrock-compliance
scope: "*"
- model_name: bedrock-nova-premier
litellm_params:
model: "bedrock/us.amazon.nova-premier-v1:0"
aws_region_name: "us-east-1"
+95 -18
View File
@@ -11,7 +11,7 @@ import sys
import time
import traceback
import warnings
from datetime import datetime, timedelta
from datetime import datetime, timedelta, timezone
from typing import (
TYPE_CHECKING,
Any,
@@ -226,6 +226,7 @@ from litellm.proxy.auth.model_checks import (
get_team_models,
)
from litellm.proxy.auth.user_api_key_auth import (
_fetch_global_spend_with_event_coordination,
user_api_key_auth,
user_api_key_auth_websocket,
)
@@ -768,6 +769,13 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
ProxyStartupEvent._add_proxy_budget_to_db(
litellm_proxy_budget_name=litellm_proxy_admin_name
)
asyncio.create_task(
ProxyStartupEvent._warm_global_spend_cache(
litellm_proxy_admin_name=litellm_proxy_admin_name,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
)
)
### START BATCH WRITING DB + CHECKING NEW MODELS###
if prisma_client is not None:
@@ -4763,6 +4771,26 @@ class ProxyStartupEvent:
)
)
@classmethod
async def _warm_global_spend_cache(
cls,
litellm_proxy_admin_name: str,
user_api_key_cache: DualCache,
prisma_client: PrismaClient,
) -> None:
"""Warm global spend cache once at startup to reduce impact of first wave of requests."""
try:
cache_key = "{}:spend".format(litellm_proxy_admin_name)
await _fetch_global_spend_with_event_coordination(
cache_key=cache_key,
user_api_key_cache=user_api_key_cache,
prisma_client=prisma_client,
)
except Exception as e:
verbose_proxy_logger.debug(
"Global spend cache warm-up at startup skipped or failed: %s", e
)
@classmethod
async def _update_default_team_member_budget(cls):
"""Update the default team member budget"""
@@ -5201,7 +5229,20 @@ class ProxyStartupEvent:
except Exception as e:
raise e
await prisma_client.connect()
try:
await prisma_client.connect()
except Exception as e:
if "P3018" in str(e) or "P3009" in str(e):
verbose_proxy_logger.debug(
"CRITICAL: DATABASE MIGRATION FAILED"
)
verbose_proxy_logger.debug(
"Your database is in a 'dirty' state."
)
verbose_proxy_logger.debug(
"FIX: Run 'prisma migrate resolve --applied <migration_name>'"
)
raise e
## Start RDS IAM token refresh background task if enabled ##
# This proactively refreshes IAM tokens before they expire,
@@ -8042,6 +8083,48 @@ async def _apply_search_filter_to_models(
return filtered_models, search_total_count
def _normalize_datetime_for_sorting(dt: Any) -> Optional[datetime]:
"""
Normalize a datetime value to a timezone-aware UTC datetime for sorting.
This function handles:
- None values: returns None
- String values: parses ISO format strings and converts to UTC-aware datetime
- Datetime objects: converts naive datetimes to UTC-aware, and aware datetimes to UTC
Args:
dt: Datetime value (None, str, or datetime object)
Returns:
UTC-aware datetime object, or None if input is None or cannot be parsed
"""
if dt is None:
return None
if isinstance(dt, str):
try:
# Handle ISO format strings, including 'Z' suffix
dt_str = dt.replace("Z", "+00:00") if dt.endswith("Z") else dt
parsed_dt = datetime.fromisoformat(dt_str)
# Ensure it's UTC-aware
if parsed_dt.tzinfo is None:
parsed_dt = parsed_dt.replace(tzinfo=timezone.utc)
else:
parsed_dt = parsed_dt.astimezone(timezone.utc)
return parsed_dt
except (ValueError, AttributeError):
return None
if isinstance(dt, datetime):
# If naive, assume UTC and make it aware
if dt.tzinfo is None:
return dt.replace(tzinfo=timezone.utc)
# If aware, convert to UTC
return dt.astimezone(timezone.utc)
return None
def _sort_models(
all_models: List[Dict[str, Any]],
sort_by: Optional[str],
@@ -8071,26 +8154,18 @@ def _sort_models(
elif sort_by == "created_at":
created_at = model_info.get("created_at")
if created_at is None:
normalized_dt = _normalize_datetime_for_sorting(created_at)
if normalized_dt is None:
# Put None values at the end for asc, at the start for desc
return (datetime.max if not reverse else datetime.min)
if isinstance(created_at, str):
try:
return datetime.fromisoformat(created_at.replace("Z", "+00:00"))
except (ValueError, AttributeError):
return datetime.min if not reverse else datetime.max
return created_at
return (datetime.max.replace(tzinfo=timezone.utc) if not reverse else datetime.min.replace(tzinfo=timezone.utc))
return normalized_dt
elif sort_by == "updated_at":
updated_at = model_info.get("updated_at")
if updated_at is None:
return (datetime.max if not reverse else datetime.min)
if isinstance(updated_at, str):
try:
return datetime.fromisoformat(updated_at.replace("Z", "+00:00"))
except (ValueError, AttributeError):
return datetime.min if not reverse else datetime.max
return updated_at
normalized_dt = _normalize_datetime_for_sorting(updated_at)
if normalized_dt is None:
return (datetime.max.replace(tzinfo=timezone.utc) if not reverse else datetime.min.replace(tzinfo=timezone.utc))
return normalized_dt
elif sort_by == "costs":
input_cost = model_info.get("input_cost_per_token", 0) or 0
@@ -10569,6 +10644,8 @@ async def get_config_list(
"max_request_size_mb": {"type": "Integer"},
"max_response_size_mb": {"type": "Integer"},
"pass_through_endpoints": {"type": "PydanticModel"},
"store_prompts_in_spend_logs": {"type": "Boolean"},
"maximum_spend_logs_retention_period": {"type": "String"},
}
return_val = []
+1
View File
@@ -5,6 +5,7 @@ datasource client {
generator client {
provider = "prisma-client-py"
binaryTargets = ["native", "debian-openssl-1.1.x", "debian-openssl-3.0.x", "linux-musl", "linux-musl-openssl-3.0.x"]
}
// Budget / Rate Limits for an org
+12 -1
View File
@@ -1901,7 +1901,18 @@ class ProxyLogging:
) or _callback.should_run_guardrail(
data=request_data, event_type=GuardrailEventHooks.post_call
):
if "apply_guardrail" in type(callback).__dict__:
if (
"async_post_call_streaming_iterator_hook"
in type(callback).__dict__
):
current_response = (
_callback.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=current_response,
request_data=request_data,
)
)
elif "apply_guardrail" in type(callback).__dict__:
request_data["guardrail_to_apply"] = callback
current_response = (
unified_guardrail.async_post_call_streaming_iterator_hook(
+5 -5
View File
@@ -8729,11 +8729,6 @@ class Router:
if allowed_fails_policy is None:
return None
if (
isinstance(exception, litellm.BadRequestError)
and allowed_fails_policy.BadRequestErrorAllowedFails is not None
):
return allowed_fails_policy.BadRequestErrorAllowedFails
if (
isinstance(exception, litellm.AuthenticationError)
and allowed_fails_policy.AuthenticationErrorAllowedFails is not None
@@ -8754,6 +8749,11 @@ class Router:
and allowed_fails_policy.ContentPolicyViolationErrorAllowedFails is not None
):
return allowed_fails_policy.ContentPolicyViolationErrorAllowedFails
if (
isinstance(exception, litellm.BadRequestError)
and allowed_fails_policy.BadRequestErrorAllowedFails is not None
):
return allowed_fails_policy.BadRequestErrorAllowedFails
def _initialize_alerting(self):
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
@@ -43,11 +43,6 @@ def get_num_retries_from_retry_policy(
if isinstance(retry_policy, dict):
retry_policy = RetryPolicy(**retry_policy)
if (
isinstance(exception, BadRequestError)
and retry_policy.BadRequestErrorRetries is not None
):
return retry_policy.BadRequestErrorRetries
if (
isinstance(exception, AuthenticationError)
and retry_policy.AuthenticationErrorRetries is not None
@@ -65,6 +60,11 @@ def get_num_retries_from_retry_policy(
and retry_policy.ContentPolicyViolationErrorRetries is not None
):
return retry_policy.ContentPolicyViolationErrorRetries
if (
isinstance(exception, BadRequestError)
and retry_policy.BadRequestErrorRetries is not None
):
return retry_policy.BadRequestErrorRetries
def reset_retry_policy() -> RetryPolicy:
+22 -3
View File
@@ -20,6 +20,9 @@ from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import (
from litellm.types.proxy.guardrails.guardrail_hooks.qualifire import (
QualifireGuardrailConfigModel,
)
from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import (
ContentFilterCategoryConfig,
)
"""
Pydantic object defining how to set guardrails on litellm proxy
@@ -547,9 +550,27 @@ class ContentFilterConfigModel(BaseModel):
blocked_words_file: Optional[str] = Field(
default=None, description="Path to YAML file containing blocked_words list"
)
categories: Optional[List[ContentFilterCategoryConfig]] = Field(
default=None,
description="List of prebuilt categories to enable (harmful_*, bias_*)",
)
severity_threshold: Optional[str] = Field(
default=None,
description="Minimum severity to block (high, medium, low)",
)
pattern_redaction_format: Optional[str] = Field(
default=None,
description="Format string for pattern redaction (use {pattern_name} placeholder)",
)
keyword_redaction_tag: Optional[str] = Field(
default=None,
description="Tag to use for keyword redaction",
)
class BaseLitellmParams(BaseModel): # works for new and patch update guardrails
class BaseLitellmParams(
ContentFilterConfigModel
): # works for new and patch update guardrails
api_key: Optional[str] = Field(
default=None, description="API key for the guardrail service"
)
@@ -630,7 +651,6 @@ class BaseLitellmParams(BaseModel): # works for new and patch update guardrails
description="Whether to fail the request if Model Armor encounters an error",
)
# Generic Guardrail API params
additional_provider_specific_params: Optional[Dict[str, Any]] = Field(
default=None,
description="Additional provider-specific parameters for generic guardrail APIs",
@@ -657,7 +677,6 @@ class LitellmParams(
ToolPermissionGuardrailConfigModel,
ZscalerAIGuardConfigModel,
JavelinGuardrailConfigModel,
ContentFilterConfigModel,
BaseLitellmParams,
EnkryptAIGuardrailConfigs,
IBMGuardrailsBaseConfigModel,
@@ -0,0 +1,42 @@
from typing import Any, Dict, List, Optional
from pydantic import BaseModel
class BulkUpdateKeyRequestItem(BaseModel):
"""Individual key update request item"""
key: str # Key identifier (token)
budget_id: Optional[str] = None # Budget ID associated with the key
max_budget: Optional[float] = None # Max budget for key
team_id: Optional[str] = None # Team ID associated with key
tags: Optional[List[str]] = None # Tags for organizing keys
class BulkUpdateKeyRequest(BaseModel):
"""Request for bulk key updates"""
keys: List[BulkUpdateKeyRequestItem]
class SuccessfulKeyUpdate(BaseModel):
"""Successfully updated key with its updated information"""
key: str
key_info: Dict[str, Any]
class FailedKeyUpdate(BaseModel):
"""Failed key update with reason"""
key: str
key_info: Optional[Dict[str, Any]] = None
failed_reason: str
class BulkUpdateKeyResponse(BaseModel):
"""Response for bulk key update operations"""
total_requested: int
successful_updates: List[SuccessfulKeyUpdate]
failed_updates: List[FailedKeyUpdate]
+79 -3
View File
@@ -354,6 +354,25 @@
"supports_video_input": true,
"supports_vision": true
},
"amazon.nova-2-pro-preview-20251202-v1:0": {
"cache_read_input_token_cost": 5.46875e-07,
"input_cost_per_token": 2.1875e-06,
"input_cost_per_image_token": 2.1875e-06,
"input_cost_per_audio_token": 2.1875e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.75e-05,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_video_input": true,
"supports_vision": true
},
"apac.amazon.nova-2-lite-v1:0": {
"cache_read_input_token_cost": 8.25e-08,
"input_cost_per_token": 3.3e-07,
@@ -371,6 +390,25 @@
"supports_video_input": true,
"supports_vision": true
},
"apac.amazon.nova-2-pro-preview-20251202-v1:0": {
"cache_read_input_token_cost": 5.46875e-07,
"input_cost_per_token": 2.1875e-06,
"input_cost_per_image_token": 2.1875e-06,
"input_cost_per_audio_token": 2.1875e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.75e-05,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_video_input": true,
"supports_vision": true
},
"eu.amazon.nova-2-lite-v1:0": {
"cache_read_input_token_cost": 8.25e-08,
"input_cost_per_token": 3.3e-07,
@@ -388,6 +426,25 @@
"supports_video_input": true,
"supports_vision": true
},
"eu.amazon.nova-2-pro-preview-20251202-v1:0": {
"cache_read_input_token_cost": 5.46875e-07,
"input_cost_per_token": 2.1875e-06,
"input_cost_per_image_token": 2.1875e-06,
"input_cost_per_audio_token": 2.1875e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.75e-05,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_video_input": true,
"supports_vision": true
},
"us.amazon.nova-2-lite-v1:0": {
"cache_read_input_token_cost": 8.25e-08,
"input_cost_per_token": 3.3e-07,
@@ -405,6 +462,25 @@
"supports_video_input": true,
"supports_vision": true
},
"us.amazon.nova-2-pro-preview-20251202-v1:0": {
"cache_read_input_token_cost": 5.46875e-07,
"input_cost_per_token": 2.1875e-06,
"input_cost_per_image_token": 2.1875e-06,
"input_cost_per_audio_token": 2.1875e-06,
"litellm_provider": "bedrock_converse",
"max_input_tokens": 1000000,
"max_output_tokens": 64000,
"max_tokens": 64000,
"mode": "chat",
"output_cost_per_token": 1.75e-05,
"supports_function_calling": true,
"supports_pdf_input": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_video_input": true,
"supports_vision": true
},
"amazon.nova-2-multimodal-embeddings-v1:0": {
"litellm_provider": "bedrock",
"max_input_tokens": 8172,
@@ -3130,7 +3206,7 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": false,
"supports_tool_choice": true,
"supports_vision": true
},
"azure/gpt-5-chat-latest": {
@@ -3162,7 +3238,7 @@
"supports_reasoning": true,
"supports_response_schema": true,
"supports_system_messages": true,
"supports_tool_choice": false,
"supports_tool_choice": true,
"supports_vision": true
},
"azure/gpt-5-codex": {
@@ -13525,7 +13601,7 @@
"cache_read_input_token_cost": 0,
"input_cost_per_token": 3e-07,
"input_cost_per_audio_token": 1e-06,
"litellm_provider": "gemini",
"litellm_provider": "vertex_ai-language-models",
"max_input_tokens": 1048576,
"max_output_tokens": 65535,
"max_tokens": 65535,
+20
View File
@@ -236,6 +236,26 @@
"rag_query": true
}
},
"s3_vectors": {
"display_name": "AWS S3 Vectors (`s3_vectors`)",
"url": "https://docs.litellm.ai/docs/providers/s3_vectors",
"endpoints": {
"chat_completions": false,
"messages": false,
"responses": false,
"embeddings": false,
"image_generations": false,
"audio_transcriptions": false,
"audio_speech": false,
"moderations": false,
"batches": false,
"rerank": false,
"a2a": false,
"interactions": false,
"vector_stores_create": true,
"vector_stores_search": true
}
},
"sagemaker": {
"display_name": "AWS - Sagemaker (`sagemaker`)",
"url": "https://docs.litellm.ai/docs/providers/aws_sagemaker",
+19
View File
@@ -356,6 +356,25 @@ def test_get_internal_user_header_from_mapping_no_internal_returns_none():
"/openai/deployments/my-deployment/chat/completions",
"my-deployment"
),
# Custom model_name with slashes (e.g., gcp/google/gemini-2.5-flash)
# This is the NVIDIA P0 bug fix - regex should capture full model name including slashes
(
{},
"/vertex_ai/v1/projects/my-project/locations/us-central1/publishers/google/models/gcp/google/gemini-2.5-flash:generateContent",
"gcp/google/gemini-2.5-flash"
),
# Another custom model_name with slashes
(
{},
"/vertex_ai/v1/projects/my-project/locations/global/publishers/google/models/gcp/google/gemini-3-flash-preview:generateContent",
"gcp/google/gemini-3-flash-preview"
),
# Model name with single slash
(
{},
"/vertex_ai/v1/projects/my-project/locations/us-central1/publishers/google/models/custom/model:generateContent",
"custom/model"
),
],
)
def test_get_model_from_request_vertex_ai_passthrough(request_data, route, expected_model):
@@ -60,6 +60,7 @@ async def test_completion_with_retry_policy(sync_mode):
retry_number = 1
retry_policy = RetryPolicy(
BadRequestErrorRetries=10,
ContentPolicyViolationErrorRetries=retry_number, # run 3 retries for ContentPolicyViolationErrors
AuthenticationErrorRetries=0, # run 0 retries for AuthenticationErrorRetries
)
@@ -0,0 +1,120 @@
"""
E2E tests for Claude Agent SDK with LiteLLM Proxy using Bedrock models.
Tests streaming messages across different Bedrock models:
- Regular Bedrock Claude Sonnet 4.5
- Bedrock Converse Claude Sonnet 4.5
- AWS Nova Premier
"""
import os
import pytest
import asyncio
from claude_agent_sdk import ClaudeSDKClient, ClaudeAgentOptions
# Test models from proxy_config.yaml
TEST_MODELS = [
("bedrock-claude-sonnet-4.5", "Bedrock Invoke API"),
("bedrock-converse-claude-sonnet-4.5", "Bedrock Converse API"),
("bedrock-nova-premier", "AWS Nova Premier"),
]
@pytest.fixture(scope="module")
def litellm_proxy_config():
"""Configure connection to LiteLLM proxy"""
proxy_url = os.getenv("LITELLM_PROXY_URL", "http://localhost:4000")
api_key = os.getenv("LITELLM_API_KEY", "sk-1234")
# Set environment variables for Claude Agent SDK
os.environ["ANTHROPIC_BASE_URL"] = proxy_url.rstrip('/')
os.environ["ANTHROPIC_API_KEY"] = api_key
return {
"proxy_url": proxy_url,
"api_key": api_key,
}
@pytest.mark.asyncio
@pytest.mark.parametrize("model_name,model_description", TEST_MODELS)
async def test_claude_agent_sdk_streaming(litellm_proxy_config, model_name, model_description):
"""
Test streaming messages with Claude Agent SDK through LiteLLM proxy.
This validates:
1. Claude Agent SDK can connect to LiteLLM proxy
2. Streaming works correctly
3. Different Bedrock models (Invoke, Converse, Nova) work end-to-end
"""
print(f"\n{'='*60}")
print(f"Testing: {model_name} ({model_description})")
print(f"{'='*60}")
# Configure agent options
options = ClaudeAgentOptions(
system_prompt="You are a helpful AI assistant. Be concise.",
model=model_name,
max_turns=5,
)
# Test query
test_query = "Say 'Hello from LiteLLM!' and nothing else."
# Track streaming
received_chunks = []
full_response = ""
try:
async with ClaudeSDKClient(options=options) as client:
await client.query(test_query)
# Collect streaming response
async for msg in client.receive_response():
# Handle different message types
if hasattr(msg, 'type'):
if msg.type == 'content_block_delta':
# Streaming text delta
if hasattr(msg, 'delta') and hasattr(msg.delta, 'text'):
chunk_text = msg.delta.text
received_chunks.append(chunk_text)
full_response += chunk_text
elif msg.type == 'content_block_start':
# Start of content block
if hasattr(msg, 'content_block') and hasattr(msg.content_block, 'text'):
chunk_text = msg.content_block.text
received_chunks.append(chunk_text)
full_response += chunk_text
# Fallback to content handling
if hasattr(msg, 'content'):
for content_block in msg.content:
if hasattr(content_block, 'text'):
chunk_text = content_block.text
received_chunks.append(chunk_text)
full_response += chunk_text
# Assertions
print(f"\n✅ Received {len(received_chunks)} chunks")
print(f"📝 Full response: {full_response[:100]}...")
# Verify we got a response
assert len(full_response) > 0, f"No response received from {model_name}"
# Verify streaming (should have multiple chunks for most responses)
# Note: Very short responses might come in 1 chunk, so we just verify we got content
assert len(received_chunks) > 0, f"No chunks received from {model_name}"
# Verify response contains expected content (case insensitive)
assert "hello" in full_response.lower(), f"Response doesn't contain expected greeting: {full_response}"
print(f"✅ Test passed for {model_name}")
except Exception as e:
pytest.fail(f"Test failed for {model_name} ({model_description}): {str(e)}")
if __name__ == "__main__":
# Run tests
pytest.main([__file__, "-v", "-s"])
@@ -0,0 +1,31 @@
model_list:
- model_name: bedrock-claude-sonnet-3.5
litellm_params:
model: "bedrock/us.anthropic.claude-3-5-sonnet-20240620-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-sonnet-4
litellm_params:
model: "bedrock/us.anthropic.claude-sonnet-4-20250514-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-sonnet-4.5
litellm_params:
model: "bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-claude-opus-4.5
litellm_params:
model: "bedrock/us.anthropic.claude-opus-4-5-20251101-v1:0"
aws_region_name: "us-east-1"
- model_name: bedrock-nova-premier
litellm_params:
model: "bedrock/amazon.nova-premier-v1:0"
aws_region_name: "us-east-1"
# Converse API models
- model_name: bedrock-converse-claude-sonnet-4.5
litellm_params:
model: "bedrock_converse/us.anthropic.claude-sonnet-4-5-20250929-v1:0"
aws_region_name: "us-east-1"
@@ -1060,3 +1060,81 @@ def test_append_system_prompt_messages():
kwargs=None, messages=messages
)
assert result == messages
def test_get_error_information_error_code_priority():
"""
Test get_error_information prioritizes 'code' attribute over 'status_code' attribute
and handles edge cases like empty strings and "None" string values.
"""
from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup
# Test case 1: Exception with 'code' attribute (ProxyException style)
class ProxyException(Exception):
def __init__(self, code, message):
self.code = code
self.message = message
super().__init__(message)
proxy_exception = ProxyException(code="500", message="Internal Server Error")
result = StandardLoggingPayloadSetup.get_error_information(proxy_exception)
assert result["error_code"] == "500"
assert result["error_class"] == "ProxyException"
# Test case 2: Exception with 'status_code' attribute (LiteLLM style)
class LiteLLMException(Exception):
def __init__(self, status_code, message):
self.status_code = status_code
self.message = message
super().__init__(message)
litellm_exception = LiteLLMException(status_code=429, message="Rate limit exceeded")
result = StandardLoggingPayloadSetup.get_error_information(litellm_exception)
assert result["error_code"] == "429"
assert result["error_class"] == "LiteLLMException"
# Test case 3: Exception with both 'code' and 'status_code' - should prefer 'code'
class BothAttributesException(Exception):
def __init__(self, code, status_code, message):
self.code = code
self.status_code = status_code
self.message = message
super().__init__(message)
both_exception = BothAttributesException(
code="400", status_code=500, message="Bad Request"
)
result = StandardLoggingPayloadSetup.get_error_information(both_exception)
assert result["error_code"] == "400" # Should prefer 'code' over 'status_code'
# Test case 4: Exception with 'code' as empty string - should fall back to 'status_code'
empty_code_exception = BothAttributesException(
code="", status_code=404, message="Not Found"
)
result = StandardLoggingPayloadSetup.get_error_information(empty_code_exception)
assert result["error_code"] == "404" # Should fall back to status_code
# Test case 5: Exception with 'code' as "None" string - should fall back to 'status_code'
none_string_exception = BothAttributesException(
code="None", status_code=503, message="Service Unavailable"
)
result = StandardLoggingPayloadSetup.get_error_information(none_string_exception)
assert result["error_code"] == "503" # Should fall back to status_code
# Test case 6: Exception with 'code' as None - should fall back to 'status_code'
none_code_exception = BothAttributesException(
code=None, status_code=401, message="Unauthorized"
)
result = StandardLoggingPayloadSetup.get_error_information(none_code_exception)
assert result["error_code"] == "401" # Should fall back to status_code
# Test case 7: Exception with neither 'code' nor 'status_code' - should return empty string
class NoCodeException(Exception):
def __init__(self, message):
self.message = message
super().__init__(message)
no_code_exception = NoCodeException(message="Generic error")
result = StandardLoggingPayloadSetup.get_error_information(no_code_exception)
assert result["error_code"] == ""
assert result["error_class"] == "NoCodeException"
@@ -548,6 +548,59 @@ def test_map_tool_choice_dict_type_function_with_name():
assert result["name"] == "my_tool"
def test_map_tool_choice_dict_type_auto():
"""
Test that dict {"type": "auto"} maps to Anthropic type='auto'.
This handles Cursor's format for tool_choice.
"""
config = AnthropicConfig()
result = config._map_tool_choice(
tool_choice={"type": "auto"},
parallel_tool_use=None,
)
assert result is not None
assert result["type"] == "auto"
def test_map_tool_choice_dict_type_required():
"""
Test that dict {"type": "required"} maps to Anthropic type='any'.
"""
config = AnthropicConfig()
result = config._map_tool_choice(
tool_choice={"type": "required"},
parallel_tool_use=None,
)
assert result is not None
assert result["type"] == "any"
def test_map_tool_choice_dict_type_none():
"""
Test that dict {"type": "none"} maps to Anthropic type='none'.
"""
config = AnthropicConfig()
result = config._map_tool_choice(
tool_choice={"type": "none"},
parallel_tool_use=None,
)
assert result is not None
assert result["type"] == "none"
def test_map_tool_choice_dict_type_function_without_name():
"""
Test that dict {"type": "function"} without name is handled gracefully.
Should return None since there's no valid tool name.
"""
config = AnthropicConfig()
result = config._map_tool_choice(
tool_choice={"type": "function"},
parallel_tool_use=None,
)
assert result is None
def test_transform_response_with_prefix_prompt():
import httpx
@@ -16,6 +16,17 @@ def test_azure_gpt5_supports_reasoning_effort(config: AzureOpenAIGPT5Config):
)
def test_azure_gpt5_allows_tool_choice_for_deployment_names():
supported_params = litellm.get_supported_openai_params(
model="gpt-5-chat-2025-08-07", custom_llm_provider="azure"
)
assert supported_params is not None
assert "tool_choice" in supported_params
# gpt-5-chat* should not be treated as a GPT-5 reasoning model
assert "reasoning_effort" not in supported_params
assert "temperature" in supported_params
def test_azure_gpt5_maps_max_tokens(config: AzureOpenAIGPT5Config):
params = config.map_openai_params(
non_default_params={"max_tokens": 5},
@@ -464,3 +464,108 @@ def test_opus_4_5_model_detection():
for model in non_opus_4_5_models:
assert not config._is_claude_opus_4_5(model), \
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.
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()
messages = [{"role": "user", "content": "test"}]
# 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,
)
# 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"}
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}"
# 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,
)
# 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():
"""
Test that output_format parameter is removed from Bedrock Invoke requests.
Bedrock Invoke API doesn't support the output_format parameter (only supported
in Anthropic Messages API). This test ensures it's removed to prevent errors.
"""
config = AmazonAnthropicClaudeConfig()
messages = [{"role": "user", "content": "test"}]
# Create a request with output_format via map_openai_params
non_default_params = {
"response_format": {"type": "json_object"}
}
optional_params = {}
# This should trigger tool-based structured outputs
optional_params = config.map_openai_params(
non_default_params=non_default_params,
optional_params=optional_params,
model="anthropic.claude-4-0-sonnet-20250514-v1:0",
drop_params=False,
)
result = config.transform_request(
model="anthropic.claude-4-0-sonnet-20250514-v1:0",
messages=messages,
optional_params=optional_params,
litellm_params={},
headers={},
)
# Verify output_format is not in the request
assert "output_format" not in result, \
f"output_format should be removed for Bedrock Invoke, got keys: {result.keys()}"
@@ -20,6 +20,23 @@ def test_gpt5_supports_reasoning_effort(config: OpenAIConfig):
assert "reasoning_effort" in config.get_supported_openai_params(model="gpt-5-mini")
def test_gpt5_chat_does_not_support_reasoning_effort(config: OpenAIConfig):
assert (
"reasoning_effort"
not in config.get_supported_openai_params(model="gpt-5-chat-latest")
)
def test_gpt5_chat_supports_temperature(config: OpenAIConfig):
params = config.map_openai_params(
non_default_params={"temperature": 0.3},
optional_params={},
model="gpt-5-chat-latest",
drop_params=False,
)
assert params["temperature"] == 0.3
def test_gpt5_maps_max_tokens(config: OpenAIConfig):
params = config.map_openai_params(
non_default_params={"max_tokens": 10},
@@ -4,7 +4,7 @@ Tests for the Content Filter Guardrail
import os
import sys
from unittest.mock import AsyncMock, MagicMock, patch
from unittest.mock import MagicMock
import pytest
@@ -385,20 +385,12 @@ class TestContentFilterGuardrail:
assert result is not None
assert result[1] == "aws_access_key"
@pytest.mark.skip(
reason="Masking in streaming responses is no longer supported after unified_guardrail.py changes. Only blocking/rejecting is supported for responses."
)
@pytest.mark.asyncio
async def test_streaming_hook_mask(self):
"""
Test streaming hook with MASK action
Note: After changes to unified_guardrail.py, masking responses to users
is no longer supported. This test is skipped as the feature is deprecated.
Only BLOCK actions (test_streaming_hook_block) are supported for streaming responses.
Test streaming hook with MASK action.
This now works with the 50-char sliding window buffer.
"""
from unittest.mock import AsyncMock
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
patterns = [
@@ -415,51 +407,54 @@ class TestContentFilterGuardrail:
event_hook=GuardrailEventHooks.during_call,
)
# Create mock streaming chunks
# Create mock streaming chunks that split an email
async def mock_stream():
# Chunk 1: contains email
chunk1 = ModelResponseStream(
# Chunk 1: starts email
yield ModelResponseStream(
id="chunk1",
choices=[
StreamingChoices(
delta=Delta(content="Contact me at test@example.com"), index=0
delta=Delta(content="Contact me at test@ex"), index=0
)
],
model="gpt-4",
)
yield chunk1
# Chunk 2: normal content
chunk2 = ModelResponseStream(
# Chunk 2: ends email
yield ModelResponseStream(
id="chunk2",
choices=[
StreamingChoices(delta=Delta(content=" for more info"), index=0)
StreamingChoices(
delta=Delta(content="ample.com for info"),
index=0,
finish_reason="stop",
)
],
model="gpt-4",
)
yield chunk2
user_api_key_dict = MagicMock()
request_data = {}
# Process streaming response - no masking expected
result_chunks = []
# Process streaming response - masking IS expected now
full_content = ""
async for chunk in guardrail.async_post_call_streaming_iterator_hook(
user_api_key_dict=user_api_key_dict,
response=mock_stream(),
request_data=request_data,
):
result_chunks.append(chunk)
if chunk.choices[0].delta.content:
full_content += chunk.choices[0].delta.content
# Chunks should pass through unchanged since masking is no longer supported
assert len(result_chunks) == 2
# The email should be redacted even though it was split
assert "test@example.com" not in full_content
assert "[EMAIL_REDACTED]" in full_content
assert "Contact me at [EMAIL_REDACTED] for info" in full_content
@pytest.mark.asyncio
async def test_streaming_hook_block(self):
"""
Test streaming hook with BLOCK action
"""
from unittest.mock import AsyncMock
from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices
@@ -715,7 +710,10 @@ class TestContentFilterGuardrail:
assert result is not None
assert len(result) == 1
# All matches should be redacted
assert result[0] == "[CUSTOM_KEY_REDACTED] [CUSTOM_KEY_REDACTED] [CUSTOM_KEY_REDACTED]"
assert (
result[0]
== "[CUSTOM_KEY_REDACTED] [CUSTOM_KEY_REDACTED] [CUSTOM_KEY_REDACTED]"
)
assert "Key1" not in result[0]
assert "Key2" not in result[0]
@@ -797,7 +795,7 @@ class TestContentFilterGuardrail:
# Apply guardrail with content that triggers detections
# Email will be masked, blocked word will be masked
result = await guardrail.apply_guardrail(
await guardrail.apply_guardrail(
inputs={"texts": ["Contact me at test@example.com for confidential info"]},
request_data=request_data,
input_type="request",
@@ -807,7 +805,9 @@ class TestContentFilterGuardrail:
assert "metadata" in request_data
assert "standard_logging_guardrail_information" in request_data["metadata"]
guardrail_info_list = request_data["metadata"]["standard_logging_guardrail_information"]
guardrail_info_list = request_data["metadata"][
"standard_logging_guardrail_information"
]
assert isinstance(guardrail_info_list, list)
assert len(guardrail_info_list) == 1
@@ -820,8 +820,8 @@ class TestContentFilterGuardrail:
assert "start_time" in guardrail_info
assert "end_time" in guardrail_info
assert "duration" in guardrail_info
assert guardrail_info["duration"] > 0
assert guardrail_info["start_time"] < guardrail_info["end_time"]
assert guardrail_info["duration"] >= 0
assert guardrail_info["start_time"] <= guardrail_info["end_time"]
# Verify detections are logged
assert "guardrail_response" in guardrail_info
@@ -839,15 +839,21 @@ class TestContentFilterGuardrail:
assert "action" in detection
assert detection["action"] == "MASK"
# Verify sensitive content (matched_text) is NOT included
assert "matched_text" not in detection, "Sensitive content should not be logged"
assert (
"matched_text" not in detection
), "Sensitive content should not be logged"
# Verify blocked word detection structure
blocked_word_detections = [d for d in detections if d.get("type") == "blocked_word"]
blocked_word_detections = [
d for d in detections if d.get("type") == "blocked_word"
]
assert len(blocked_word_detections) > 0
for detection in blocked_word_detections:
assert detection["type"] == "blocked_word"
assert "keyword" in detection
assert detection["keyword"] == "confidential" # Config keyword, not user content
assert (
detection["keyword"] == "confidential"
) # Config keyword, not user content
assert "action" in detection
assert detection["action"] == "MASK"
assert "description" in detection
@@ -896,7 +902,9 @@ class TestContentFilterGuardrail:
assert "metadata" in request_data
assert "standard_logging_guardrail_information" in request_data["metadata"]
guardrail_info_list = request_data["metadata"]["standard_logging_guardrail_information"]
guardrail_info_list = request_data["metadata"][
"standard_logging_guardrail_information"
]
assert len(guardrail_info_list) == 1
guardrail_info = guardrail_info_list[0]
@@ -909,4 +917,6 @@ class TestContentFilterGuardrail:
# If detections are logged, verify they don't contain sensitive content
for detection in detections:
if detection.get("type") == "pattern":
assert "matched_text" not in detection, "Sensitive content should not be logged"
assert (
"matched_text" not in detection
), "Sensitive content should not be logged"
@@ -30,10 +30,13 @@ from litellm.proxy.management_endpoints.key_management_endpoints import (
_check_org_key_limits,
_check_team_key_limits,
_common_key_generation_helper,
_get_and_validate_existing_key,
_list_key_helper,
_persist_deleted_verification_tokens,
_process_single_key_update,
_save_deleted_verification_token_records,
_transform_verification_tokens_to_deleted_records,
_validate_max_budget,
can_modify_verification_token,
check_org_key_model_specific_limits,
check_team_key_model_specific_limits,
@@ -4223,3 +4226,467 @@ async def test_update_key_with_router_settings(monkeypatch):
# Verify router_settings can be deserialized and matches input
deserialized_settings = json.loads(result["router_settings"])
assert deserialized_settings == router_settings_data
@pytest.mark.asyncio
async def test_validate_max_budget():
"""
Test _validate_max_budget helper function.
Tests:
1. Positive max_budget should pass
2. Zero max_budget should pass
3. Negative max_budget should raise HTTPException
4. None max_budget should pass
"""
from fastapi import HTTPException
# Test Case 1: Positive max_budget should pass
try:
_validate_max_budget(100.0)
_validate_max_budget(0.0)
except HTTPException:
pytest.fail("_validate_max_budget raised HTTPException for valid values")
# Test Case 2: None max_budget should pass
try:
_validate_max_budget(None)
except HTTPException:
pytest.fail("_validate_max_budget raised HTTPException for None")
# Test Case 3: Negative max_budget should raise HTTPException
with pytest.raises(HTTPException) as exc_info:
_validate_max_budget(-10.0)
assert exc_info.value.status_code == 400
assert "max_budget cannot be negative" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_get_and_validate_existing_key():
"""
Test _get_and_validate_existing_key helper function.
Tests:
1. Successfully retrieve existing key
2. Key not found raises HTTPException
3. Database not connected raises HTTPException
"""
from fastapi import HTTPException
# Test Case 1: Successfully retrieve existing key
mock_prisma_client = AsyncMock()
mock_key = LiteLLM_VerificationToken(
token="test-key-123",
user_id="user-123",
models=["gpt-4"],
team_id=None,
)
mock_prisma_client.get_data = AsyncMock(return_value=mock_key)
result = await _get_and_validate_existing_key(
token="test-key-123",
prisma_client=mock_prisma_client,
)
assert result == mock_key
mock_prisma_client.get_data.assert_called_once_with(
token="test-key-123",
table_name="key",
query_type="find_unique",
)
# Test Case 2: Key not found raises HTTPException
mock_prisma_client.get_data = AsyncMock(return_value=None)
with pytest.raises(HTTPException) as exc_info:
await _get_and_validate_existing_key(
token="non-existent-key",
prisma_client=mock_prisma_client,
)
assert exc_info.value.status_code == 404
assert "Key not found" in str(exc_info.value.detail)
# Test Case 3: Database not connected raises HTTPException
with pytest.raises(HTTPException) as exc_info:
await _get_and_validate_existing_key(
token="test-key-123",
prisma_client=None,
)
assert exc_info.value.status_code == 500
assert "Database not connected" in str(exc_info.value.detail)
@pytest.mark.asyncio
async def test_process_single_key_update():
"""
Test _process_single_key_update helper function.
Tests successful key update with all validations passing.
"""
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateKeyRequestItem,
)
# Setup mocks
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
mock_llm_router = MagicMock()
# Mock existing key
existing_key = LiteLLM_VerificationToken(
token="test-key-123",
user_id="user-123",
models=["gpt-4"],
team_id=None,
max_budget=None,
tags=None,
)
# Mock updated key response
updated_key_data = {
"user_id": "user-123",
"models": ["gpt-4"],
"team_id": None,
"max_budget": 100.0,
"tags": ["production"],
}
mock_prisma_client.get_data = AsyncMock(return_value=existing_key)
mock_updated_key_obj = MagicMock()
mock_updated_key_obj.model_dump.return_value = updated_key_data
mock_prisma_client.update_data = AsyncMock(
return_value={"data": mock_updated_key_obj}
)
# Mock prepare_key_update_data
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
) as mock_prepare:
mock_prepare.return_value = {"max_budget": 100.0, "tags": ["production"]}
# Mock TeamMemberPermissionChecks
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
) as mock_permission_check:
mock_permission_check.return_value = None
# Mock _delete_cache_key_object
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
) as mock_delete_cache:
mock_delete_cache.return_value = None
# Mock hash_token (imported from litellm.proxy._types)
with patch(
"litellm.proxy._types.hash_token"
) as mock_hash:
mock_hash.return_value = "hashed-test-key-123"
# Mock KeyManagementEventHooks
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
):
# Create update request
key_update_item = BulkUpdateKeyRequestItem(
key="test-key-123",
max_budget=100.0,
tags=["production"],
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call the function
result = await _process_single_key_update(
key_update_item=key_update_item,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
prisma_client=mock_prisma_client,
user_api_key_cache=mock_user_api_key_cache,
proxy_logging_obj=mock_proxy_logging_obj,
llm_router=mock_llm_router,
)
# Verify results
assert result is not None
assert "token" not in result # Token should be removed
assert result.get("max_budget") == 100.0
assert result.get("tags") == ["production"]
# Verify mocks were called
mock_prisma_client.get_data.assert_called_once()
mock_prisma_client.update_data.assert_called_once()
mock_delete_cache.assert_called_once()
@pytest.mark.asyncio
async def test_bulk_update_keys_success(monkeypatch):
"""
Test /key/bulk_update endpoint with successful updates.
Tests:
1. Multiple keys updated successfully
2. Response contains correct counts and data
"""
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateKeyRequest,
BulkUpdateKeyRequestItem,
)
from litellm.proxy.management_endpoints.key_management_endpoints import (
bulk_update_keys,
)
from litellm.proxy.proxy_server import (
llm_router,
prisma_client,
proxy_logging_obj,
user_api_key_cache,
)
# Setup mocks
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
mock_llm_router = MagicMock()
# Mock existing keys
existing_key_1 = LiteLLM_VerificationToken(
token="test-key-1",
user_id="user-123",
models=["gpt-4"],
team_id=None,
max_budget=None,
)
existing_key_2 = LiteLLM_VerificationToken(
token="test-key-2",
user_id="user-123",
models=["gpt-3.5-turbo"],
team_id=None,
max_budget=50.0,
)
# Mock updated key responses
updated_key_1_data = {
"user_id": "user-123",
"models": ["gpt-4"],
"max_budget": 100.0,
"tags": ["production"],
}
updated_key_2_data = {
"user_id": "user-123",
"models": ["gpt-3.5-turbo"],
"max_budget": 200.0,
"tags": ["staging"],
}
mock_prisma_client.get_data = AsyncMock(
side_effect=[existing_key_1, existing_key_2]
)
mock_updated_key_1_obj = MagicMock()
mock_updated_key_1_obj.model_dump.return_value = updated_key_1_data
mock_updated_key_2_obj = MagicMock()
mock_updated_key_2_obj.model_dump.return_value = updated_key_2_data
mock_prisma_client.update_data = AsyncMock(
side_effect=[
{"data": mock_updated_key_1_obj},
{"data": mock_updated_key_2_obj},
]
)
# Patch dependencies
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_llm_router)
# Mock helper functions
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
) as mock_prepare:
mock_prepare.side_effect = [
{"max_budget": 100.0, "tags": ["production"]},
{"max_budget": 200.0, "tags": ["staging"]},
]
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
):
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
):
with patch(
"litellm.proxy._types.hash_token"
) as mock_hash:
mock_hash.side_effect = ["hashed-key-1", "hashed-key-2"]
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
):
# Create request
request_data = BulkUpdateKeyRequest(
keys=[
BulkUpdateKeyRequestItem(
key="test-key-1",
max_budget=100.0,
tags=["production"],
),
BulkUpdateKeyRequestItem(
key="test-key-2",
max_budget=200.0,
tags=["staging"],
),
]
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call endpoint
response = await bulk_update_keys(
data=request_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify response
assert response.total_requested == 2
assert len(response.successful_updates) == 2
assert len(response.failed_updates) == 0
assert response.successful_updates[0].key == "test-key-1"
assert response.successful_updates[1].key == "test-key-2"
@pytest.mark.asyncio
async def test_bulk_update_keys_partial_failures(monkeypatch):
"""
Test /key/bulk_update endpoint with partial failures.
Tests:
1. Some keys update successfully, others fail
2. Response contains both successful and failed updates
3. Failed updates include error messages
"""
from litellm.types.proxy.management_endpoints.key_management_endpoints import (
BulkUpdateKeyRequest,
BulkUpdateKeyRequestItem,
)
from litellm.proxy.management_endpoints.key_management_endpoints import (
bulk_update_keys,
)
# Setup mocks
mock_prisma_client = AsyncMock()
mock_user_api_key_cache = MagicMock()
mock_proxy_logging_obj = MagicMock()
mock_llm_router = MagicMock()
# Mock existing keys
existing_key_1 = LiteLLM_VerificationToken(
token="test-key-1",
user_id="user-123",
models=["gpt-4"],
team_id=None,
max_budget=None,
)
# Mock updated key response for successful update
updated_key_1_data = {
"user_id": "user-123",
"models": ["gpt-4"],
"max_budget": 100.0,
"tags": ["production"],
}
# First key exists, second key doesn't exist
mock_prisma_client.get_data = AsyncMock(
side_effect=[existing_key_1, None] # Second key not found
)
mock_updated_key_1_obj = MagicMock()
mock_updated_key_1_obj.model_dump.return_value = updated_key_1_data
mock_prisma_client.update_data = AsyncMock(
return_value={"data": mock_updated_key_1_obj}
)
# Patch dependencies
monkeypatch.setattr(
"litellm.proxy.proxy_server.prisma_client", mock_prisma_client
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.user_api_key_cache", mock_user_api_key_cache
)
monkeypatch.setattr(
"litellm.proxy.proxy_server.proxy_logging_obj", mock_proxy_logging_obj
)
monkeypatch.setattr("litellm.proxy.proxy_server.llm_router", mock_llm_router)
# Mock helper functions
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.prepare_key_update_data"
) as mock_prepare:
mock_prepare.return_value = {"max_budget": 100.0, "tags": ["production"]}
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.TeamMemberPermissionChecks.can_team_member_execute_key_management_endpoint"
):
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints._delete_cache_key_object"
):
with patch(
"litellm.proxy._types.hash_token"
) as mock_hash:
mock_hash.return_value = "hashed-key-1"
with patch(
"litellm.proxy.management_endpoints.key_management_endpoints.KeyManagementEventHooks.async_key_updated_hook"
):
# Create request with one valid and one invalid key
request_data = BulkUpdateKeyRequest(
keys=[
BulkUpdateKeyRequestItem(
key="test-key-1",
max_budget=100.0,
tags=["production"],
),
BulkUpdateKeyRequestItem(
key="non-existent-key",
max_budget=200.0,
tags=["staging"],
),
]
)
user_api_key_dict = UserAPIKeyAuth(
user_role=LitellmUserRoles.PROXY_ADMIN,
api_key="sk-admin",
user_id="admin-user",
)
# Call endpoint
response = await bulk_update_keys(
data=request_data,
user_api_key_dict=user_api_key_dict,
litellm_changed_by=None,
)
# Verify response
assert response.total_requested == 2
assert len(response.successful_updates) == 1
assert len(response.failed_updates) == 1
assert response.successful_updates[0].key == "test-key-1"
assert response.failed_updates[0].key == "non-existent-key"
assert "Key not found" in response.failed_updates[0].failed_reason
+77 -1
View File
@@ -5,7 +5,7 @@ import os
import socket
import subprocess
import sys
from datetime import datetime
from datetime import datetime, timezone
from pathlib import Path
from unittest import mock
from unittest.mock import AsyncMock, MagicMock, mock_open, patch
@@ -1049,6 +1049,82 @@ async def test_get_config_from_file(tmp_path, monkeypatch):
assert result == test_config
def test_normalize_datetime_for_sorting():
"""
Test the _normalize_datetime_for_sorting function.
Tests various scenarios: None values, ISO format strings, datetime objects (naive and aware).
"""
from litellm.proxy.proxy_server import _normalize_datetime_for_sorting
# Test Case 1: None value
assert _normalize_datetime_for_sorting(None) is None
# Test Case 2: ISO format string with 'Z' suffix
dt_str_z = "2024-01-15T10:30:00Z"
result = _normalize_datetime_for_sorting(dt_str_z)
assert result is not None
assert isinstance(result, datetime)
assert result.tzinfo == timezone.utc
assert result.year == 2024
assert result.month == 1
assert result.day == 15
assert result.hour == 10
assert result.minute == 30
# Test Case 3: ISO format string without 'Z' suffix (naive)
dt_str_naive = "2024-01-15T10:30:00"
result = _normalize_datetime_for_sorting(dt_str_naive)
assert result is not None
assert isinstance(result, datetime)
assert result.tzinfo == timezone.utc
# Test Case 4: ISO format string with timezone offset
dt_str_tz = "2024-01-15T10:30:00+05:00"
result = _normalize_datetime_for_sorting(dt_str_tz)
assert result is not None
assert isinstance(result, datetime)
assert result.tzinfo == timezone.utc
# Should convert from +05:00 to UTC (subtract 5 hours)
assert result.hour == 5 # 10:30 - 5 hours = 5:30 UTC
# Test Case 5: Naive datetime object
naive_dt = datetime(2024, 1, 15, 10, 30, 0)
result = _normalize_datetime_for_sorting(naive_dt)
assert result is not None
assert isinstance(result, datetime)
assert result.tzinfo == timezone.utc
assert result.year == 2024
assert result.month == 1
assert result.day == 15
# Test Case 6: Timezone-aware datetime object (non-UTC)
from datetime import timedelta
aware_dt = datetime(2024, 1, 15, 10, 30, 0, tzinfo=timezone(timedelta(hours=5)))
result = _normalize_datetime_for_sorting(aware_dt)
assert result is not None
assert isinstance(result, datetime)
assert result.tzinfo == timezone.utc
# Should convert from +05:00 to UTC
assert result.hour == 5
# Test Case 7: UTC-aware datetime object
utc_dt = datetime(2024, 1, 15, 10, 30, 0, tzinfo=timezone.utc)
result = _normalize_datetime_for_sorting(utc_dt)
assert result is not None
assert isinstance(result, datetime)
assert result.tzinfo == timezone.utc
assert result == utc_dt
# Test Case 8: Invalid string format
invalid_str = "not-a-date"
result = _normalize_datetime_for_sorting(invalid_str)
assert result is None
# Test Case 9: Invalid type (should return None)
result = _normalize_datetime_for_sorting(12345)
assert result is None
@pytest.mark.asyncio
async def test_add_proxy_budget_to_db_only_creates_user_no_keys():
"""
+67 -43
View File
@@ -1695,60 +1695,84 @@ def test_gemini_without_cache_tokens_details():
print("✅ Gemini without cacheTokensDetails works correctly")
def test_generic_provider_cached_token_cost():
def test_gemini_implicit_caching_cost_calculation():
"""
Test that the generic cost calculator correctly handles cached tokens
for providers like z.ai/deepseek that are not explicitly handled.
"""
from litellm.cost_calculator import completion_cost
from litellm.types.utils import ModelResponse, PromptTokensDetailsWrapper, Usage
Test for Issue #16341: Gemini implicit cached tokens not counted in spend log
# Setup model cost for a generic provider
# We use a name that will bypass complex provider mapping logic
model_name = "custom-cached-model"
litellm.model_cost[model_name] = {
"input_cost_per_token": 0.0000006,
"output_cost_per_token": 0.0000006,
"cache_read_input_token_cost": 0.0000001,
"litellm_provider": "openai",
When Gemini uses implicit caching, it returns cachedContentTokenCount but NOT
cacheTokensDetails. In this case, we should subtract cachedContentTokenCount
from text_tokens to correctly calculate costs.
See: https://github.com/BerriAI/litellm/issues/16341
"""
from litellm import completion_cost
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
VertexGeminiConfig,
)
from litellm.types.utils import Choices, Message, ModelResponse
# Simulate Gemini response with implicit caching (cachedContentTokenCount only)
completion_response = {
"usageMetadata": {
"promptTokenCount": 10000,
"candidatesTokenCount": 5,
"totalTokenCount": 10005,
"cachedContentTokenCount": 8000, # Implicit caching - no cacheTokensDetails
"promptTokensDetails": [{"modality": "TEXT", "tokenCount": 10000}],
"candidatesTokensDetails": [{"modality": "TEXT", "tokenCount": 5}],
}
}
# Case 1: Standard nested cached tokens (prompt_tokens_details.cached_tokens)
usage = Usage(
prompt_tokens=10000,
completion_tokens=0,
prompt_tokens_details=PromptTokensDetailsWrapper(cached_tokens=9000),
usage = VertexGeminiConfig._calculate_usage(completion_response)
# Verify parsing
assert (
usage.cache_read_input_tokens == 8000
), f"cache_read_input_tokens should be 8000, got {usage.cache_read_input_tokens}"
assert (
usage.prompt_tokens_details.cached_tokens == 8000
), f"cached_tokens should be 8000, got {usage.prompt_tokens_details.cached_tokens}"
# CRITICAL: text_tokens should be (10000 - 8000) = 2000, NOT 10000
# This is the fix for issue #16341
assert (
usage.prompt_tokens_details.text_tokens == 2000
), f"text_tokens should be 2000 (10000 - 8000), got {usage.prompt_tokens_details.text_tokens}"
# Verify cost calculation uses cached token pricing
response = ModelResponse(
id="mock-id",
model="gemini-2.0-flash",
choices=[
Choices(
index=0,
message=Message(role="assistant", content="Hello!"),
finish_reason="stop",
)
],
usage=usage,
)
response = ModelResponse(usage=usage, model=model_name)
cost = completion_cost(
completion_response=response,
model=model_name,
custom_llm_provider="openai", # Explicitly set provider to trigger generic path
model="gemini-2.0-flash",
custom_llm_provider="gemini",
)
# Expected: (1000 * 0.0000006) + (9000 * 0.0000001) = 0.0006 + 0.0009 = 0.0015
expected_cost = 0.0015
assert (
abs(cost - expected_cost) < 1e-9
), f"Nested cache cost failed. Got {cost}, expected {expected_cost}"
# Get model pricing for verification
import litellm
# Case 2: Top-level cached tokens (cache_read_input_tokens)
usage_top = Usage(
prompt_tokens=10000,
completion_tokens=0,
cache_read_input_tokens=9000,
)
response_top = ModelResponse(usage=usage_top, model=model_name)
model_info = litellm.get_model_info("gemini/gemini-2.0-flash")
input_cost = model_info.get("input_cost_per_token", 0)
cache_read_cost = model_info.get("cache_read_input_token_cost", input_cost)
output_cost = model_info.get("output_cost_per_token", 0)
cost_top = completion_cost(
completion_response=response_top,
model=model_name,
custom_llm_provider="openai",
# Expected cost: (2000 * input) + (8000 * cache_read) + (5 * output)
expected_cost = (2000 * input_cost) + (8000 * cache_read_cost) + (5 * output_cost)
assert abs(cost - expected_cost) < 1e-9, (
f"Cost calculation is wrong. Got ${cost:.6f}, expected ${expected_cost:.6f}. "
f"Cached tokens may not be using reduced pricing."
)
assert (
abs(cost_top - expected_cost) < 1e-9
), f"Top-level cache cost failed. Got {cost_top}, expected {expected_cost}"
print("✅ Generic provider cached token cost verified")
print("✅ Issue #16341 fix verified: Gemini implicit caching cost calculated correctly")
@@ -78,6 +78,35 @@ def test_lazy_loading_default():
assert len(tokens) > 0, "Encoding should work"
def test_tiktoken_cache_dir_set_on_lazy_load():
"""Test that TIKTOKEN_CACHE_DIR is set when encoding is lazy loaded.
This ensures the local tiktoken cache is used instead of downloading
from the internet. Regression test for issue #19768.
"""
# Remove environment variables to ensure clean state
if "LITELLM_DISABLE_LAZY_LOADING" in os.environ:
del os.environ["LITELLM_DISABLE_LAZY_LOADING"]
if "TIKTOKEN_CACHE_DIR" in os.environ:
del os.environ["TIKTOKEN_CACHE_DIR"]
# Clear any cached modules
modules_to_clear = [k for k in sys.modules.keys() if k.startswith("litellm")]
for module in modules_to_clear:
del sys.modules[module]
# Import litellm fresh
import litellm
# Access encoding (triggers lazy load)
_ = litellm.encoding
# Verify TIKTOKEN_CACHE_DIR is now set and points to local tokenizers
assert "TIKTOKEN_CACHE_DIR" in os.environ, "TIKTOKEN_CACHE_DIR should be set after lazy loading encoding"
cache_dir = os.environ["TIKTOKEN_CACHE_DIR"]
assert "tokenizers" in cache_dir, f"TIKTOKEN_CACHE_DIR should point to tokenizers directory, got: {cache_dir}"
@pytest.fixture(autouse=True)
def cleanup_env():
"""Clean up environment variable after each test"""
@@ -0,0 +1,554 @@
import { describe, it, expect, vi, beforeEach, afterEach } from "vitest";
import { renderHook, waitFor } from "@testing-library/react";
import { QueryClient, QueryClientProvider } from "@tanstack/react-query";
import React, { ReactNode } from "react";
import {
useProxyConfig,
useDeleteProxyConfigField,
getProxyConfigCall,
deleteProxyConfigFieldCall,
ConfigType,
GeneralSettingsFieldName,
type ProxyConfigResponse,
type DeleteProxyConfigFieldRequest,
type DeleteProxyConfigFieldResponse,
} from "./useProxyConfig";
const {
mockProxyBaseUrl,
mockAccessToken,
mockHeaderName,
mockProxyConfigResponse,
mockDeleteResponse,
mockUseAuthorized,
mockGetGlobalLitellmHeaderName,
mockDeriveErrorMessage,
mockHandleError,
} = vi.hoisted(() => {
const mockProxyBaseUrl = "https://proxy.example.com";
const mockAccessToken = "test-access-token";
const mockHeaderName = "X-LiteLLM-API-Key";
const mockProxyConfigResponse: ProxyConfigResponse = [
{
field_name: "maximum_spend_logs_retention_period",
field_type: "int",
field_description: "Maximum retention period for spend logs",
field_value: 30,
stored_in_db: true,
field_default_value: 7,
premium_field: false,
nested_fields: null,
},
{
field_name: "another_field",
field_type: "string",
field_description: "Another config field",
field_value: "test-value",
stored_in_db: false,
field_default_value: "default-value",
premium_field: true,
nested_fields: [
{
field_name: "nested_field",
field_type: "string",
field_description: "Nested field description",
field_default_value: "nested-default",
stored_in_db: true,
},
],
},
];
const mockDeleteResponse: DeleteProxyConfigFieldResponse = {
message: "Field deleted successfully",
};
const mockUseAuthorized = vi.fn();
const mockGetGlobalLitellmHeaderName = vi.fn(() => mockHeaderName);
const mockDeriveErrorMessage = vi.fn((errorData: any) => {
if (typeof errorData === "string") return errorData;
return errorData?.message || errorData?.error || "An error occurred";
});
const mockHandleError = vi.fn();
return {
mockProxyBaseUrl,
mockAccessToken,
mockHeaderName,
mockProxyConfigResponse,
mockDeleteResponse,
mockUseAuthorized,
mockGetGlobalLitellmHeaderName,
mockDeriveErrorMessage,
mockHandleError,
};
});
vi.mock("../useAuthorized", () => ({
default: () => mockUseAuthorized(),
}));
vi.mock("@/components/networking", () => ({
proxyBaseUrl: mockProxyBaseUrl,
getGlobalLitellmHeaderName: mockGetGlobalLitellmHeaderName,
deriveErrorMessage: mockDeriveErrorMessage,
handleError: mockHandleError,
}));
vi.mock("../common/queryKeysFactory", () => ({
createQueryKeys: vi.fn((resource: string) => ({
all: [resource],
lists: () => [resource, "list"],
list: (params?: any) => [resource, "list", { params }],
details: () => [resource, "detail"],
detail: (uid: string) => [resource, "detail", uid],
})),
}));
describe("useProxyConfig", () => {
let queryClient: QueryClient;
let fetchSpy: ReturnType<typeof vi.fn>;
beforeEach(() => {
queryClient = new QueryClient({
defaultOptions: {
queries: {
retry: false,
},
},
});
vi.clearAllMocks();
mockUseAuthorized.mockReturnValue({
accessToken: mockAccessToken,
userId: "test-user-id",
userRole: "Admin",
token: "test-token",
userEmail: "test@example.com",
premiumUser: false,
disabledPersonalKeyCreation: null,
showSSOBanner: false,
});
fetchSpy = vi.fn();
global.fetch = fetchSpy;
});
afterEach(() => {
vi.restoreAllMocks();
});
const wrapper = ({ children }: { children: ReactNode }) =>
React.createElement(QueryClientProvider, { client: queryClient }, children);
it("should render successfully", () => {
(fetchSpy as any).mockResolvedValue({
ok: true,
json: async () => mockProxyConfigResponse,
});
const { result } = renderHook(() => useProxyConfig(ConfigType.GENERAL_SETTINGS), { wrapper });
expect(result.current).toBeDefined();
expect(result.current.isLoading).toBe(true);
});
it("should return proxy config data when query is successful", async () => {
(fetchSpy as any).mockResolvedValue({
ok: true,
json: async () => mockProxyConfigResponse,
});
const { result } = renderHook(() => useProxyConfig(ConfigType.GENERAL_SETTINGS), { wrapper });
expect(result.current.isLoading).toBe(true);
expect(result.current.data).toBeUndefined();
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
expect(result.current.isSuccess).toBe(true);
});
expect(result.current.data).toEqual(mockProxyConfigResponse);
expect(result.current.error).toBeNull();
expect(fetchSpy).toHaveBeenCalledWith(
`${mockProxyBaseUrl}/config/list?config_type=${ConfigType.GENERAL_SETTINGS}`,
{
method: "GET",
headers: {
[mockHeaderName]: `Bearer ${mockAccessToken}`,
"Content-Type": "application/json",
},
},
);
expect(fetchSpy).toHaveBeenCalledTimes(1);
});
it("should handle error when API call fails", async () => {
const errorMessage = "Failed to fetch proxy config";
const errorResponse = { message: errorMessage };
(fetchSpy as any).mockResolvedValue({
ok: false,
json: async () => errorResponse,
});
const { result } = renderHook(() => useProxyConfig(ConfigType.GENERAL_SETTINGS), { wrapper });
expect(result.current.isLoading).toBe(true);
await waitFor(() => {
expect(result.current.isLoading).toBe(false);
expect(result.current.isError).toBe(true);
});
expect(result.current.error).toBeDefined();
expect(result.current.data).toBeUndefined();
});
it("should not execute query when accessToken is missing", async () => {
mockUseAuthorized.mockReturnValue({
accessToken: null,
userId: "test-user-id",
userRole: "Admin",
token: null,
userEmail: "test@example.com",
premiumUser: false,
disabledPersonalKeyCreation: null,
showSSOBanner: false,
});
const { result } = renderHook(() => useProxyConfig(ConfigType.GENERAL_SETTINGS), { wrapper });
expect(result.current.isLoading).toBe(false);
expect(result.current.data).toBeUndefined();
expect(result.current.isFetched).toBe(false);
expect(fetchSpy).not.toHaveBeenCalled();
});
it("should use correct query key with config type filter", async () => {
(fetchSpy as any).mockResolvedValue({
ok: true,
json: async () => mockProxyConfigResponse,
});
const { result } = renderHook(() => useProxyConfig(ConfigType.GENERAL_SETTINGS), { wrapper });
await waitFor(() => {
expect(result.current.isSuccess).toBe(true);
});
expect(result.current.data).toEqual(mockProxyConfigResponse);
});
it("should handle network errors", async () => {
const networkError = new Error("Network error");
(fetchSpy as any).mockRejectedValue(networkError);
const { result } = renderHook(() => useProxyConfig(ConfigType.GENERAL_SETTINGS), { wrapper });
await waitFor(() => {
expect(result.current.isError).toBe(true);
});
expect(result.current.error).toBeDefined();
expect(result.current.data).toBeUndefined();
});
it("should handle empty config response", async () => {
(fetchSpy as any).mockResolvedValue({
ok: true,
json: async () => [],
});
const { result } = renderHook(() => useProxyConfig(ConfigType.GENERAL_SETTINGS), { wrapper });
await waitFor(() => {
expect(result.current.isSuccess).toBe(true);
});
expect(result.current.data).toEqual([]);
});
});
describe("useDeleteProxyConfigField", () => {
let queryClient: QueryClient;
let fetchSpy: ReturnType<typeof vi.fn>;
beforeEach(() => {
queryClient = new QueryClient({
defaultOptions: {
queries: {
retry: false,
},
mutations: {
retry: false,
},
},
});
vi.clearAllMocks();
mockUseAuthorized.mockReturnValue({
accessToken: mockAccessToken,
userId: "test-user-id",
userRole: "Admin",
token: "test-token",
userEmail: "test@example.com",
premiumUser: false,
disabledPersonalKeyCreation: null,
showSSOBanner: false,
});
fetchSpy = vi.fn();
global.fetch = fetchSpy;
});
afterEach(() => {
vi.restoreAllMocks();
});
const wrapper = ({ children }: { children: ReactNode }) =>
React.createElement(QueryClientProvider, { client: queryClient }, children);
it("should render successfully", () => {
const { result } = renderHook(() => useDeleteProxyConfigField(), { wrapper });
expect(result.current).toBeDefined();
expect(result.current.isIdle).toBe(true);
});
it("should successfully delete a proxy config field", async () => {
(fetchSpy as any).mockResolvedValue({
ok: true,
json: async () => mockDeleteResponse,
});
const { result } = renderHook(() => useDeleteProxyConfigField(), { wrapper });
const deleteRequest: DeleteProxyConfigFieldRequest = {
config_type: ConfigType.GENERAL_SETTINGS,
field_name: GeneralSettingsFieldName.MAXIMUM_SPEND_LOGS_RETENTION_PERIOD,
};
result.current.mutate(deleteRequest);
await waitFor(() => {
expect(result.current.isSuccess).toBe(true);
});
expect(result.current.data).toEqual(mockDeleteResponse);
expect(fetchSpy).toHaveBeenCalledWith(`${mockProxyBaseUrl}/config/field/delete`, {
method: "POST",
headers: {
[mockHeaderName]: `Bearer ${mockAccessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify(deleteRequest),
});
expect(fetchSpy).toHaveBeenCalledTimes(1);
});
it("should handle error when delete request fails", async () => {
const errorMessage = "Failed to delete field";
const errorResponse = { message: errorMessage };
(fetchSpy as any).mockResolvedValue({
ok: false,
json: async () => errorResponse,
});
const { result } = renderHook(() => useDeleteProxyConfigField(), { wrapper });
const deleteRequest: DeleteProxyConfigFieldRequest = {
config_type: ConfigType.GENERAL_SETTINGS,
field_name: GeneralSettingsFieldName.MAXIMUM_SPEND_LOGS_RETENTION_PERIOD,
};
result.current.mutate(deleteRequest);
await waitFor(() => {
expect(result.current.isError).toBe(true);
});
expect(result.current.error).toBeDefined();
expect(result.current.data).toBeUndefined();
});
it("should throw error when accessToken is missing", async () => {
mockUseAuthorized.mockReturnValue({
accessToken: null,
userId: "test-user-id",
userRole: "Admin",
token: null,
userEmail: "test@example.com",
premiumUser: false,
disabledPersonalKeyCreation: null,
showSSOBanner: false,
});
const { result } = renderHook(() => useDeleteProxyConfigField(), { wrapper });
const deleteRequest: DeleteProxyConfigFieldRequest = {
config_type: ConfigType.GENERAL_SETTINGS,
field_name: GeneralSettingsFieldName.MAXIMUM_SPEND_LOGS_RETENTION_PERIOD,
};
result.current.mutate(deleteRequest);
await waitFor(() => {
expect(result.current.isError).toBe(true);
});
expect(result.current.error?.message).toBe("Access token is required");
expect(fetchSpy).not.toHaveBeenCalled();
});
it("should handle network errors during delete", async () => {
const networkError = new Error("Network error");
(fetchSpy as any).mockRejectedValue(networkError);
const { result } = renderHook(() => useDeleteProxyConfigField(), { wrapper });
const deleteRequest: DeleteProxyConfigFieldRequest = {
config_type: ConfigType.GENERAL_SETTINGS,
field_name: GeneralSettingsFieldName.MAXIMUM_SPEND_LOGS_RETENTION_PERIOD,
};
result.current.mutate(deleteRequest);
await waitFor(() => {
expect(result.current.isError).toBe(true);
});
expect(result.current.error).toBeDefined();
});
});
describe("getProxyConfigCall", () => {
let fetchSpy: ReturnType<typeof vi.fn>;
let consoleErrorSpy: ReturnType<typeof vi.spyOn>;
beforeEach(() => {
fetchSpy = vi.fn();
global.fetch = fetchSpy;
consoleErrorSpy = vi.spyOn(console, "error").mockImplementation(() => {});
});
afterEach(() => {
vi.restoreAllMocks();
});
it("should successfully fetch proxy config", async () => {
(fetchSpy as any).mockResolvedValue({
ok: true,
json: async () => mockProxyConfigResponse,
});
const result = await getProxyConfigCall(mockAccessToken, ConfigType.GENERAL_SETTINGS);
expect(result).toEqual(mockProxyConfigResponse);
expect(fetchSpy).toHaveBeenCalledWith(
`${mockProxyBaseUrl}/config/list?config_type=${ConfigType.GENERAL_SETTINGS}`,
{
method: "GET",
headers: {
[mockHeaderName]: `Bearer ${mockAccessToken}`,
"Content-Type": "application/json",
},
},
);
});
it("should throw error when API returns error response", async () => {
const errorMessage = "Failed to fetch config";
const errorResponse = { message: errorMessage };
(fetchSpy as any).mockResolvedValue({
ok: false,
json: async () => errorResponse,
});
await expect(getProxyConfigCall(mockAccessToken, ConfigType.GENERAL_SETTINGS)).rejects.toThrow(errorMessage);
});
it("should handle network errors", async () => {
const networkError = new Error("Network error");
(fetchSpy as any).mockRejectedValue(networkError);
await expect(getProxyConfigCall(mockAccessToken, ConfigType.GENERAL_SETTINGS)).rejects.toThrow("Network error");
expect(consoleErrorSpy).toHaveBeenCalled();
});
});
describe("deleteProxyConfigFieldCall", () => {
let fetchSpy: ReturnType<typeof vi.fn>;
let consoleErrorSpy: ReturnType<typeof vi.spyOn>;
beforeEach(() => {
fetchSpy = vi.fn();
global.fetch = fetchSpy;
consoleErrorSpy = vi.spyOn(console, "error").mockImplementation(() => {});
});
afterEach(() => {
vi.restoreAllMocks();
});
it("should successfully delete proxy config field", async () => {
(fetchSpy as any).mockResolvedValue({
ok: true,
json: async () => mockDeleteResponse,
});
const deleteRequest: DeleteProxyConfigFieldRequest = {
config_type: ConfigType.GENERAL_SETTINGS,
field_name: GeneralSettingsFieldName.MAXIMUM_SPEND_LOGS_RETENTION_PERIOD,
};
const result = await deleteProxyConfigFieldCall(mockAccessToken, deleteRequest);
expect(result).toEqual(mockDeleteResponse);
expect(fetchSpy).toHaveBeenCalledWith(`${mockProxyBaseUrl}/config/field/delete`, {
method: "POST",
headers: {
[mockHeaderName]: `Bearer ${mockAccessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify(deleteRequest),
});
});
it("should throw error when API returns error response", async () => {
const errorMessage = "Failed to delete field";
const errorResponse = { message: errorMessage };
(fetchSpy as any).mockResolvedValue({
ok: false,
json: async () => errorResponse,
});
const deleteRequest: DeleteProxyConfigFieldRequest = {
config_type: ConfigType.GENERAL_SETTINGS,
field_name: GeneralSettingsFieldName.MAXIMUM_SPEND_LOGS_RETENTION_PERIOD,
};
await expect(deleteProxyConfigFieldCall(mockAccessToken, deleteRequest)).rejects.toThrow(errorMessage);
});
it("should handle network errors", async () => {
const networkError = new Error("Network error");
(fetchSpy as any).mockRejectedValue(networkError);
const deleteRequest: DeleteProxyConfigFieldRequest = {
config_type: ConfigType.GENERAL_SETTINGS,
field_name: GeneralSettingsFieldName.MAXIMUM_SPEND_LOGS_RETENTION_PERIOD,
};
await expect(deleteProxyConfigFieldCall(mockAccessToken, deleteRequest)).rejects.toThrow("Network error");
expect(consoleErrorSpy).toHaveBeenCalled();
});
});
@@ -0,0 +1,180 @@
import { useQuery, useMutation, UseMutationResult } from "@tanstack/react-query";
import { createQueryKeys } from "../common/queryKeysFactory";
import useAuthorized from "../useAuthorized";
import { proxyBaseUrl, getGlobalLitellmHeaderName, deriveErrorMessage, handleError } from "@/components/networking";
/**
* Enum for config types that can be fetched from the proxy config endpoint.
* Currently supports general_settings, but can be extended as more config types are added.
*/
export enum ConfigType {
GENERAL_SETTINGS = "general_settings",
}
/**
* Enum for supported field names that can be deleted from general_settings.
* This should match the fields available in ConfigGeneralSettings.
*/
export enum GeneralSettingsFieldName {
MAXIMUM_SPEND_LOGS_RETENTION_PERIOD = "maximum_spend_logs_retention_period",
// Add more field names here as needed
}
/**
* Field detail for nested fields within a config field
*/
export interface FieldDetail {
field_name: string;
field_type: string;
field_description: string;
field_default_value: any;
stored_in_db: boolean | null;
}
/**
* Configuration list item returned from /config/list endpoint
*/
export interface ConfigListItem {
field_name: string;
field_type: string;
field_description: string;
field_value: any;
stored_in_db: boolean | null;
field_default_value: any;
premium_field?: boolean;
nested_fields?: FieldDetail[] | null;
}
/**
* Response type for /config/list endpoint
*/
export type ProxyConfigResponse = ConfigListItem[];
/**
* Request body for /config/field/delete endpoint
*/
export interface DeleteProxyConfigFieldRequest {
config_type: ConfigType;
field_name: string;
}
/**
* Response type for /config/field/delete endpoint
*/
export interface DeleteProxyConfigFieldResponse {
message?: string;
[key: string]: any;
}
/**
* Network call function to fetch proxy config by config type
* @param accessToken - The access token for authentication
* @param configType - The type of config to fetch (from ConfigType enum)
* @returns Promise resolving to the config list response
*/
export const getProxyConfigCall = async (accessToken: string, configType: ConfigType): Promise<ProxyConfigResponse> => {
try {
const url = proxyBaseUrl
? `${proxyBaseUrl}/config/list?config_type=${configType}`
: `/config/list?config_type=${configType}`;
const response = await fetch(url, {
method: "GET",
headers: {
[getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
});
if (!response.ok) {
const errorData = await response.json();
const errorMessage = deriveErrorMessage(errorData);
handleError(errorMessage);
throw new Error(errorMessage);
}
const data = await response.json();
return data;
} catch (error) {
console.error(`Failed to get proxy config for ${configType}:`, error);
throw error;
}
};
const proxyConfigKeys = createQueryKeys("proxyConfig");
/**
* Network call function to delete a proxy config field
* @param accessToken - The access token for authentication
* @param request - The delete request containing config_type and field_name
* @returns Promise resolving to the delete response
*/
export const deleteProxyConfigFieldCall = async (
accessToken: string,
request: DeleteProxyConfigFieldRequest,
): Promise<DeleteProxyConfigFieldResponse> => {
try {
const url = proxyBaseUrl ? `${proxyBaseUrl}/config/field/delete` : `/config/field/delete`;
const response = await fetch(url, {
method: "POST",
headers: {
[getGlobalLitellmHeaderName()]: `Bearer ${accessToken}`,
"Content-Type": "application/json",
},
body: JSON.stringify(request),
});
if (!response.ok) {
const errorData = await response.json();
const errorMessage = deriveErrorMessage(errorData);
handleError(errorMessage);
throw new Error(errorMessage);
}
const data = await response.json();
return data;
} catch (error) {
console.error(`Failed to delete proxy config field ${request.field_name}:`, error);
throw error;
}
};
/**
* React Query hook to fetch proxy config by config type
* @param configType - The type of config to fetch (from ConfigType enum)
* @returns React Query result with the config list data
*/
export const useProxyConfig = (configType: ConfigType) => {
const { accessToken } = useAuthorized();
return useQuery<ProxyConfigResponse>({
queryKey: proxyConfigKeys.list({
filters: {
configType,
},
}),
queryFn: async () => await getProxyConfigCall(accessToken!, configType),
enabled: Boolean(accessToken),
});
};
/**
* React Query hook to delete a proxy config field
* @returns React Query mutation result for deleting config fields
*/
export const useDeleteProxyConfigField = (): UseMutationResult<
DeleteProxyConfigFieldResponse,
Error,
DeleteProxyConfigFieldRequest
> => {
const { accessToken } = useAuthorized();
return useMutation<DeleteProxyConfigFieldResponse, Error, DeleteProxyConfigFieldRequest>({
mutationFn: async (request: DeleteProxyConfigFieldRequest) => {
if (!accessToken) {
throw new Error("Access token is required");
}
return await deleteProxyConfigFieldCall(accessToken, request);
},
});
};
@@ -49,4 +49,39 @@ describe("NewBadge", () => {
expect(container.firstChild).toBeNull();
});
it("should render badge with dot when dot prop is true", () => {
mockUseDisableShowNewBadge.mockReturnValue(false);
render(<NewBadge dot={true}>Test Content</NewBadge>);
expect(screen.queryByText("New")).not.toBeInTheDocument();
expect(screen.getByText("Test Content")).toBeInTheDocument();
});
it("should render badge with 'New' text when dot prop is false", () => {
mockUseDisableShowNewBadge.mockReturnValue(false);
render(<NewBadge dot={false}>Test Content</NewBadge>);
expect(screen.getByText("New")).toBeInTheDocument();
expect(screen.getByText("Test Content")).toBeInTheDocument();
});
it("should render badge with 'New' text when dot prop is not provided (defaults to false)", () => {
mockUseDisableShowNewBadge.mockReturnValue(false);
render(<NewBadge>Test Content</NewBadge>);
expect(screen.getByText("New")).toBeInTheDocument();
expect(screen.getByText("Test Content")).toBeInTheDocument();
});
it("should render badge with dot when dot is true and no children", () => {
mockUseDisableShowNewBadge.mockReturnValue(false);
render(<NewBadge dot={true} />);
expect(screen.queryByText("New")).not.toBeInTheDocument();
});
});
@@ -1,7 +1,13 @@
import { Badge } from "antd";
import { useDisableShowNewBadge } from "@/app/(dashboard)/hooks/useDisableShowNewBadge";
export default function NewBadge({ children }: { children?: React.ReactNode }) {
export default function NewBadge({
children,
dot = false
}: {
children?: React.ReactNode;
dot?: boolean;
}) {
const disableShowNewBadge = useDisableShowNewBadge();
if (disableShowNewBadge) {
@@ -9,10 +15,10 @@ export default function NewBadge({ children }: { children?: React.ReactNode }) {
}
return children ? (
<Badge color="blue" count="New">
<Badge color="blue" count={dot ? undefined : "New"} dot={dot}>
{children}
</Badge>
) : (
<Badge color="blue" count="New" />
<Badge color="blue" count={dot ? undefined : "New"} dot={dot} />
);
}
+261 -261
View File
@@ -62,260 +62,260 @@ interface MenuGroup {
// Menu groups organized by category - defined outside component for export
const menuGroups: MenuGroup[] = [
{
groupLabel: "AI GATEWAY",
items: [
{
key: "api-keys",
page: "api-keys",
label: "Virtual Keys",
icon: <KeyOutlined />,
},
{
key: "llm-playground",
page: "llm-playground",
label: "Playground",
icon: <PlayCircleOutlined />,
roles: rolesWithWriteAccess,
},
{
key: "models",
page: "models",
label: "Models + Endpoints",
icon: <BlockOutlined />,
roles: rolesWithWriteAccess,
},
{
key: "agents",
page: "agents",
label: "Agents",
icon: <RobotOutlined />,
roles: rolesWithWriteAccess,
},
{
key: "mcp-servers",
page: "mcp-servers",
label: "MCP Servers",
icon: <ToolOutlined />,
},
{
key: "guardrails",
page: "guardrails",
label: "Guardrails",
icon: <SafetyOutlined />,
roles: all_admin_roles,
},
{
key: "policies",
page: "policies",
label: (
<span className="flex items-center gap-4">
Policies <NewBadge />
</span>
),
icon: <AuditOutlined />,
roles: all_admin_roles,
},
{
key: "tools",
page: "tools",
label: "Tools",
icon: <ToolOutlined />,
children: [
{
key: "search-tools",
page: "search-tools",
label: "Search Tools",
icon: <SearchOutlined />,
},
{
key: "vector-stores",
page: "vector-stores",
label: "Vector Stores",
icon: <DatabaseOutlined />,
},
],
},
],
},
{
groupLabel: "OBSERVABILITY",
items: [
{
key: "new_usage",
page: "new_usage",
icon: <BarChartOutlined />,
roles: [...all_admin_roles, ...internalUserRoles],
label: "Usage",
},
{
key: "logs",
page: "logs",
label: (
<span className="flex items-center gap-4">
Logs <NewBadge />
</span>
),
icon: <LineChartOutlined />,
},
],
},
{
groupLabel: "ACCESS CONTROL",
items: [
{
key: "users",
page: "users",
label: "Internal Users",
icon: <UserOutlined />,
roles: all_admin_roles,
},
{
key: "teams",
page: "teams",
label: "Teams",
icon: <TeamOutlined />,
},
{
key: "organizations",
page: "organizations",
label: "Organizations",
icon: <BankOutlined />,
roles: all_admin_roles,
},
{
key: "budgets",
page: "budgets",
label: "Budgets",
icon: <CreditCardOutlined />,
roles: all_admin_roles,
},
],
},
{
groupLabel: "DEVELOPER TOOLS",
items: [
{
key: "api_ref",
page: "api_ref",
label: "API Reference",
icon: <ApiOutlined />,
},
{
key: "model-hub-table",
page: "model-hub-table",
label: "AI Hub",
icon: <AppstoreOutlined />,
},
{
key: "learning-resources",
page: "learning-resources",
label: "Learning Resources",
icon: <BookOutlined />,
external_url: "https://models.litellm.ai/cookbook",
},
{
key: "experimental",
page: "experimental",
label: "Experimental",
icon: <ExperimentOutlined />,
children: [
{
key: "caching",
page: "caching",
label: "Caching",
icon: <DatabaseOutlined />,
roles: all_admin_roles,
},
{
key: "prompts",
page: "prompts",
label: "Prompts",
icon: <FileTextOutlined />,
roles: all_admin_roles,
},
{
key: "transform-request",
page: "transform-request",
label: "API Playground",
icon: <ApiOutlined />,
roles: [...all_admin_roles, ...internalUserRoles],
},
{
key: "tag-management",
page: "tag-management",
label: "Tag Management",
icon: <TagsOutlined />,
roles: all_admin_roles,
},
{
key: "claude-code-plugins",
page: "claude-code-plugins",
label: "Claude Code Plugins",
icon: <ToolOutlined />,
roles: all_admin_roles,
},
{
key: "4",
page: "usage",
label: "Old Usage",
icon: <BarChartOutlined />,
}
],
},
],
},
{
groupLabel: "SETTINGS",
roles: all_admin_roles,
items: [
{
key: "settings",
page: "settings",
label: <span className="flex items-center gap-4">Settings</span>,
icon: <SettingOutlined />,
roles: all_admin_roles,
children: [
{
key: "router-settings",
page: "router-settings",
label: "Router Settings",
icon: <SettingOutlined />,
roles: all_admin_roles,
},
{
key: "logging-and-alerts",
page: "logging-and-alerts",
label: "Logging & Alerts",
icon: <SettingOutlined />,
roles: all_admin_roles,
},
{
key: "admin-panel",
page: "admin-panel",
label: "Admin Settings",
icon: <SettingOutlined />,
roles: all_admin_roles,
},
{
key: "cost-tracking",
page: "cost-tracking",
label: "Cost Tracking",
icon: <BarChartOutlined />,
roles: all_admin_roles,
},
{
key: "ui-theme",
page: "ui-theme",
label: "UI Theme",
icon: <BgColorsOutlined />,
roles: all_admin_roles,
},
],
},
],
},
];
{
groupLabel: "AI GATEWAY",
items: [
{
key: "api-keys",
page: "api-keys",
label: "Virtual Keys",
icon: <KeyOutlined />,
},
{
key: "llm-playground",
page: "llm-playground",
label: "Playground",
icon: <PlayCircleOutlined />,
roles: rolesWithWriteAccess,
},
{
key: "models",
page: "models",
label: "Models + Endpoints",
icon: <BlockOutlined />,
roles: rolesWithWriteAccess,
},
{
key: "agents",
page: "agents",
label: "Agents",
icon: <RobotOutlined />,
roles: rolesWithWriteAccess,
},
{
key: "mcp-servers",
page: "mcp-servers",
label: "MCP Servers",
icon: <ToolOutlined />,
},
{
key: "guardrails",
page: "guardrails",
label: "Guardrails",
icon: <SafetyOutlined />,
roles: all_admin_roles,
},
{
key: "policies",
page: "policies",
label: (
<span className="flex items-center gap-4">
Policies
</span>
),
icon: <AuditOutlined />,
roles: all_admin_roles,
},
{
key: "tools",
page: "tools",
label: "Tools",
icon: <ToolOutlined />,
children: [
{
key: "search-tools",
page: "search-tools",
label: "Search Tools",
icon: <SearchOutlined />,
},
{
key: "vector-stores",
page: "vector-stores",
label: "Vector Stores",
icon: <DatabaseOutlined />,
},
],
},
],
},
{
groupLabel: "OBSERVABILITY",
items: [
{
key: "new_usage",
page: "new_usage",
icon: <BarChartOutlined />,
roles: [...all_admin_roles, ...internalUserRoles],
label: "Usage",
},
{
key: "logs",
page: "logs",
label: (
<span className="flex items-center gap-4">
Logs <NewBadge />
</span>
),
icon: <LineChartOutlined />,
},
],
},
{
groupLabel: "ACCESS CONTROL",
items: [
{
key: "users",
page: "users",
label: "Internal Users",
icon: <UserOutlined />,
roles: all_admin_roles,
},
{
key: "teams",
page: "teams",
label: "Teams",
icon: <TeamOutlined />,
},
{
key: "organizations",
page: "organizations",
label: "Organizations",
icon: <BankOutlined />,
roles: all_admin_roles,
},
{
key: "budgets",
page: "budgets",
label: "Budgets",
icon: <CreditCardOutlined />,
roles: all_admin_roles,
},
],
},
{
groupLabel: "DEVELOPER TOOLS",
items: [
{
key: "api_ref",
page: "api_ref",
label: "API Reference",
icon: <ApiOutlined />,
},
{
key: "model-hub-table",
page: "model-hub-table",
label: "AI Hub",
icon: <AppstoreOutlined />,
},
{
key: "learning-resources",
page: "learning-resources",
label: "Learning Resources",
icon: <BookOutlined />,
external_url: "https://models.litellm.ai/cookbook",
},
{
key: "experimental",
page: "experimental",
label: "Experimental",
icon: <ExperimentOutlined />,
children: [
{
key: "caching",
page: "caching",
label: "Caching",
icon: <DatabaseOutlined />,
roles: all_admin_roles,
},
{
key: "prompts",
page: "prompts",
label: "Prompts",
icon: <FileTextOutlined />,
roles: all_admin_roles,
},
{
key: "transform-request",
page: "transform-request",
label: "API Playground",
icon: <ApiOutlined />,
roles: [...all_admin_roles, ...internalUserRoles],
},
{
key: "tag-management",
page: "tag-management",
label: "Tag Management",
icon: <TagsOutlined />,
roles: all_admin_roles,
},
{
key: "claude-code-plugins",
page: "claude-code-plugins",
label: "Claude Code Plugins",
icon: <ToolOutlined />,
roles: all_admin_roles,
},
{
key: "4",
page: "usage",
label: "Old Usage",
icon: <BarChartOutlined />,
}
],
},
],
},
{
groupLabel: "SETTINGS",
roles: all_admin_roles,
items: [
{
key: "settings",
page: "settings",
label: <span className="flex items-center gap-4">Settings</span>,
icon: <SettingOutlined />,
roles: all_admin_roles,
children: [
{
key: "router-settings",
page: "router-settings",
label: "Router Settings",
icon: <SettingOutlined />,
roles: all_admin_roles,
},
{
key: "logging-and-alerts",
page: "logging-and-alerts",
label: "Logging & Alerts",
icon: <SettingOutlined />,
roles: all_admin_roles,
},
{
key: "admin-panel",
page: "admin-panel",
label: "Admin Settings",
icon: <SettingOutlined />,
roles: all_admin_roles,
},
{
key: "cost-tracking",
page: "cost-tracking",
label: "Cost Tracking",
icon: <BarChartOutlined />,
roles: all_admin_roles,
},
{
key: "ui-theme",
page: "ui-theme",
label: "UI Theme",
icon: <BgColorsOutlined />,
roles: all_admin_roles,
},
],
},
],
},
];
const Sidebar: React.FC<SidebarProps> = ({ setPage, defaultSelectedKey, collapsed = false, enabledPagesInternalUsers }) => {
const { userId, accessToken, userRole } = useAuthorized();
@@ -377,7 +377,7 @@ const Sidebar: React.FC<SidebarProps> = ({ setPage, defaultSelectedKey, collapse
if (!isAdmin && enabledPagesInternalUsers !== null && enabledPagesInternalUsers !== undefined) {
// If item has children, check if any children are visible
if (item.children && item.children.length > 0) {
const hasVisibleChildren = item.children.some((child) =>
const hasVisibleChildren = item.children.some((child) =>
enabledPagesInternalUsers.includes(child.page)
);
if (hasVisibleChildren) {
@@ -385,7 +385,7 @@ const Sidebar: React.FC<SidebarProps> = ({ setPage, defaultSelectedKey, collapse
return true;
}
}
const isIncluded = enabledPagesInternalUsers.includes(item.page);
console.log(`[LeftNav] Page "${item.page}" (${item.key}): ${isIncluded ? "VISIBLE" : "HIDDEN"}`);
return isIncluded;
@@ -444,12 +444,12 @@ const Sidebar: React.FC<SidebarProps> = ({ setPage, defaultSelectedKey, collapse
})),
onClick: !item.children
? () => {
if (item.external_url) {
window.open(item.external_url, "_blank");
} else {
navigateToPage(item.page);
}
if (item.external_url) {
window.open(item.external_url, "_blank");
} else {
navigateToPage(item.page);
}
}
: undefined,
})),
});
@@ -2,9 +2,10 @@ import React from "react";
interface ConfigInfoMessageProps {
show: boolean;
onOpenSettings?: () => void;
}
export const ConfigInfoMessage: React.FC<ConfigInfoMessageProps> = ({ show }) => {
export const ConfigInfoMessage: React.FC<ConfigInfoMessageProps> = ({ show, onOpenSettings }) => {
if (!show) return null;
return (
@@ -30,7 +31,18 @@ export const ConfigInfoMessage: React.FC<ConfigInfoMessageProps> = ({ show }) =>
<h4 className="text-sm font-medium text-blue-800">Request/Response Data Not Available</h4>
<p className="text-sm text-blue-700 mt-1">
To view request and response details, enable prompt storage in your LiteLLM configuration by adding the
following to your <code className="bg-blue-100 px-1 py-0.5 rounded">proxy_config.yaml</code> file:
following to your <code className="bg-blue-100 px-1 py-0.5 rounded">proxy_config.yaml</code> file
{onOpenSettings && (
<> or{" "}
<button
onClick={onOpenSettings}
className="text-blue-600 hover:text-blue-800 underline font-medium"
>
open the settings
</button>
{" "}to configure this directly.
</>
)}
</p>
<pre className="mt-2 bg-white p-3 rounded border border-blue-200 text-xs font-mono overflow-auto">
{`general_settings:
@@ -1,3 +1,4 @@
import { useDeleteProxyConfigField, useProxyConfig } from "@/app/(dashboard)/hooks/proxyConfig/useProxyConfig";
import { useStoreRequestInSpendLogs } from "@/app/(dashboard)/hooks/storeRequestInSpendLogs/useStoreRequestInSpendLogs";
import NotificationsManager from "@/components/molecules/notifications_manager";
import { parseErrorMessage } from "@/components/shared/errorUtils";
@@ -8,6 +9,7 @@ import { renderWithProviders } from "../../../../tests/test-utils";
import SpendLogsSettingsModal from "./SpendLogsSettingsModal";
vi.mock("@/app/(dashboard)/hooks/storeRequestInSpendLogs/useStoreRequestInSpendLogs");
vi.mock("@/app/(dashboard)/hooks/proxyConfig/useProxyConfig");
vi.mock("@/components/molecules/notifications_manager", () => ({
default: {
success: vi.fn(),
@@ -19,6 +21,8 @@ vi.mock("@/components/shared/errorUtils", () => ({
}));
const mockUseStoreRequestInSpendLogs = vi.mocked(useStoreRequestInSpendLogs);
const mockUseProxyConfig = vi.mocked(useProxyConfig);
const mockUseDeleteProxyConfigField = vi.mocked(useDeleteProxyConfigField);
const mockNotificationsManager = vi.mocked(NotificationsManager);
const mockParseErrorMessage = vi.mocked(parseErrorMessage);
@@ -26,6 +30,8 @@ describe("SpendLogsSettingsModal", () => {
const mockOnCancel = vi.fn();
const mockOnSuccess = vi.fn();
const mockMutateAsync = vi.fn();
const mockDeleteField = vi.fn();
const mockRefetch = vi.fn();
const defaultProps = {
isVisible: true,
@@ -39,6 +45,15 @@ describe("SpendLogsSettingsModal", () => {
mutateAsync: mockMutateAsync,
isPending: false,
} as any);
mockUseDeleteProxyConfigField.mockReturnValue({
mutateAsync: mockDeleteField,
isPending: false,
} as any);
mockUseProxyConfig.mockReturnValue({
data: [],
isLoading: false,
refetch: mockRefetch,
} as any);
mockParseErrorMessage.mockImplementation((error: any) => error?.message || String(error));
});
@@ -127,6 +142,7 @@ describe("SpendLogsSettingsModal", () => {
await user.click(saveButton);
await waitFor(() => {
expect(mockDeleteField).not.toHaveBeenCalled();
expect(mockMutateAsync).toHaveBeenCalledWith(
{
store_prompts_in_spend_logs: true,
@@ -139,6 +155,7 @@ describe("SpendLogsSettingsModal", () => {
it("should submit form with store prompts disabled and no retention period", async () => {
const user = userEvent.setup();
mockDeleteField.mockResolvedValue({ message: "Field deleted successfully" });
mockMutateAsync.mockImplementation(async (params, options) => {
await Promise.resolve();
options?.onSuccess?.();
@@ -151,10 +168,10 @@ describe("SpendLogsSettingsModal", () => {
await user.click(saveButton);
await waitFor(() => {
expect(mockDeleteField).toHaveBeenCalled();
expect(mockMutateAsync).toHaveBeenCalledWith(
{
store_prompts_in_spend_logs: false,
maximum_spend_logs_retention_period: undefined,
},
expect.any(Object)
);
@@ -163,6 +180,7 @@ describe("SpendLogsSettingsModal", () => {
it("should show success notification and call onSuccess on successful submission", async () => {
const user = userEvent.setup();
mockDeleteField.mockResolvedValue({ message: "Field deleted successfully" });
mockMutateAsync.mockImplementation(async (params, options) => {
await Promise.resolve();
options?.onSuccess?.();
@@ -176,6 +194,7 @@ describe("SpendLogsSettingsModal", () => {
await waitFor(() => {
expect(mockNotificationsManager.success).toHaveBeenCalledWith("Spend logs settings updated successfully");
expect(mockRefetch).toHaveBeenCalled();
expect(mockOnSuccess).toHaveBeenCalledTimes(1);
});
});
@@ -227,6 +246,31 @@ describe("SpendLogsSettingsModal", () => {
expect(cancelButton).toBeDisabled();
});
it("should disable cancel button when deleting field", () => {
mockUseDeleteProxyConfigField.mockReturnValue({
mutateAsync: mockDeleteField,
isPending: true,
} as any);
renderWithProviders(<SpendLogsSettingsModal {...defaultProps} />);
const cancelButton = screen.getByRole("button", { name: "Cancel" });
expect(cancelButton).toBeDisabled();
});
it("should disable cancel button when loading config", () => {
mockUseProxyConfig.mockReturnValue({
data: undefined,
isLoading: true,
refetch: mockRefetch,
} as any);
renderWithProviders(<SpendLogsSettingsModal {...defaultProps} />);
const cancelButton = screen.getByRole("button", { name: "Cancel" });
expect(cancelButton).toBeDisabled();
});
it("should show loading state on save button when pending", () => {
mockUseStoreRequestInSpendLogs.mockReturnValue({
mutateAsync: mockMutateAsync,
@@ -240,6 +284,19 @@ describe("SpendLogsSettingsModal", () => {
expect(saveButton.className).toContain("ant-btn-loading");
});
it("should show loading state on save button when deleting field", () => {
mockUseDeleteProxyConfigField.mockReturnValue({
mutateAsync: mockDeleteField,
isPending: true,
} as any);
renderWithProviders(<SpendLogsSettingsModal {...defaultProps} />);
const saveButton = screen.getByRole("button", { name: /Saving/i });
expect(saveButton).toBeInTheDocument();
expect(saveButton.className).toContain("ant-btn-loading");
});
it("should call onCancel when cancel button is clicked after modifying form", async () => {
const user = userEvent.setup();
renderWithProviders(<SpendLogsSettingsModal {...defaultProps} />);
@@ -259,15 +316,16 @@ describe("SpendLogsSettingsModal", () => {
expect(mockOnCancel).toHaveBeenCalledTimes(1);
});
it("should reset form fields after successful submission", async () => {
it("should call refetch after successful submission", async () => {
const user = userEvent.setup();
mockDeleteField.mockResolvedValue({ message: "Field deleted successfully" });
mockMutateAsync.mockImplementation(async (params, options) => {
await Promise.resolve();
options?.onSuccess?.();
return { message: "Success" };
});
const { rerender } = renderWithProviders(<SpendLogsSettingsModal {...defaultProps} />);
renderWithProviders(<SpendLogsSettingsModal {...defaultProps} />);
const switchElement = screen.getByRole("switch");
await user.click(switchElement);
@@ -283,20 +341,13 @@ describe("SpendLogsSettingsModal", () => {
await waitFor(() => {
expect(mockNotificationsManager.success).toHaveBeenCalled();
});
rerender(<SpendLogsSettingsModal {...defaultProps} />);
await waitFor(() => {
const updatedSwitchElement = screen.getByRole("switch");
const updatedRetentionInput = screen.getByPlaceholderText("e.g., 7d, 30d");
expect(updatedSwitchElement).not.toBeChecked();
expect(updatedRetentionInput).toHaveValue("");
expect(mockRefetch).toHaveBeenCalled();
});
});
it("should not call onSuccess when it is not provided", async () => {
const user = userEvent.setup();
mockDeleteField.mockResolvedValue({ message: "Field deleted successfully" });
mockMutateAsync.mockImplementation(async (params, options) => {
await Promise.resolve();
options?.onSuccess?.();
@@ -319,8 +370,93 @@ describe("SpendLogsSettingsModal", () => {
expect(screen.queryByRole("dialog")).not.toBeInTheDocument();
});
it("should call refetch when modal opens", () => {
renderWithProviders(<SpendLogsSettingsModal {...defaultProps} />);
expect(mockRefetch).toHaveBeenCalledTimes(1);
});
it("should render form with initial values from config data", () => {
mockUseProxyConfig.mockReturnValue({
data: [
{
field_name: "store_prompts_in_spend_logs",
field_type: "bool",
field_description: "Store prompts in spend logs",
field_value: true,
stored_in_db: true,
field_default_value: false,
},
{
field_name: "maximum_spend_logs_retention_period",
field_type: "string",
field_description: "Maximum retention period",
field_value: "30d",
stored_in_db: true,
field_default_value: undefined,
},
],
isLoading: false,
refetch: mockRefetch,
} as any);
renderWithProviders(<SpendLogsSettingsModal {...defaultProps} />);
const switchElement = screen.getByRole("switch");
const retentionInput = screen.getByPlaceholderText("e.g., 7d, 30d");
expect(switchElement).toBeChecked();
expect(retentionInput).toHaveValue("30d");
});
it("should show skeleton loaders when config is loading", () => {
mockUseProxyConfig.mockReturnValue({
data: undefined,
isLoading: true,
refetch: mockRefetch,
} as any);
renderWithProviders(<SpendLogsSettingsModal {...defaultProps} />);
// Check that switch and input are not present when loading (skeletons are shown instead)
expect(screen.queryByRole("switch")).not.toBeInTheDocument();
expect(screen.queryByPlaceholderText("e.g., 7d, 30d")).not.toBeInTheDocument();
// Check for skeleton elements (Ant Design Skeleton.Input renders with ant-skeleton class)
const skeletons = document.querySelectorAll(".ant-skeleton");
expect(skeletons.length).toBeGreaterThan(0);
});
it("should continue with update even if deleteField fails", async () => {
const user = userEvent.setup();
const deleteError = new Error("Field does not exist");
mockDeleteField.mockRejectedValue(deleteError);
mockMutateAsync.mockImplementation(async (params, options) => {
await Promise.resolve();
options?.onSuccess?.();
return { message: "Success" };
});
renderWithProviders(<SpendLogsSettingsModal {...defaultProps} />);
const saveButton = screen.getByRole("button", { name: "Save Settings" });
await user.click(saveButton);
await waitFor(() => {
expect(mockDeleteField).toHaveBeenCalled();
expect(mockMutateAsync).toHaveBeenCalledWith(
{
store_prompts_in_spend_logs: false,
},
expect.any(Object)
);
expect(mockNotificationsManager.success).toHaveBeenCalled();
});
});
it("should submit form with only store prompts enabled and no retention period", async () => {
const user = userEvent.setup();
mockDeleteField.mockResolvedValue({ message: "Field deleted successfully" });
mockMutateAsync.mockImplementation(async (params, options) => {
await Promise.resolve();
options?.onSuccess?.();
@@ -336,10 +472,10 @@ describe("SpendLogsSettingsModal", () => {
await user.click(saveButton);
await waitFor(() => {
expect(mockDeleteField).toHaveBeenCalled();
expect(mockMutateAsync).toHaveBeenCalledWith(
{
store_prompts_in_spend_logs: true,
maximum_spend_logs_retention_period: undefined,
},
expect.any(Object)
);
@@ -1,11 +1,13 @@
"use client";
import { ConfigType, GeneralSettingsFieldName, useDeleteProxyConfigField, useProxyConfig } from "@/app/(dashboard)/hooks/proxyConfig/useProxyConfig";
import { StoreRequestInSpendLogsParams, useStoreRequestInSpendLogs } from "@/app/(dashboard)/hooks/storeRequestInSpendLogs/useStoreRequestInSpendLogs";
import NewBadge from "@/components/common_components/NewBadge";
import NotificationsManager from "@/components/molecules/notifications_manager";
import { parseErrorMessage } from "@/components/shared/errorUtils";
import { ClockCircleOutlined } from "@ant-design/icons";
import { Button, Form, Input, Modal, Space, Switch } from "antd";
import React from "react";
import { Button, Form, Input, Modal, Skeleton, Space, Switch, Typography } from "antd";
import React, { useEffect, useMemo } from "react";
interface SpendLogsSettingsModalProps {
isVisible: boolean;
@@ -16,14 +18,69 @@ interface SpendLogsSettingsModalProps {
const SpendLogsSettingsModal: React.FC<SpendLogsSettingsModalProps> = ({ isVisible, onCancel, onSuccess }) => {
const [form] = Form.useForm();
const { mutateAsync, isPending } = useStoreRequestInSpendLogs();
const { mutateAsync: deleteField, isPending: isDeletingField } = useDeleteProxyConfigField();
const { data: proxyConfigData, isLoading: isLoadingConfig, refetch } = useProxyConfig(ConfigType.GENERAL_SETTINGS);
const storePromptsValue = Form.useWatch('store_prompts_in_spend_logs', form);
// Refetch config when modal opens to ensure we have the latest values
useEffect(() => {
if (isVisible) {
refetch();
}
}, [isVisible, refetch]);
// Compute initial values from fetched config data
const initialValues = useMemo(() => {
if (!proxyConfigData) {
return {
store_prompts_in_spend_logs: false,
maximum_spend_logs_retention_period: undefined,
};
}
const storePromptsField = proxyConfigData.find(field => field.field_name === 'store_prompts_in_spend_logs');
const retentionPeriodField = proxyConfigData.find(field => field.field_name === 'maximum_spend_logs_retention_period');
return {
store_prompts_in_spend_logs: storePromptsField?.field_value ?? false,
maximum_spend_logs_retention_period: retentionPeriodField?.field_value ?? undefined,
};
}, [proxyConfigData]);
const handleFormSubmit = async (formValues: StoreRequestInSpendLogsParams) => {
try {
await mutateAsync(formValues, {
// If maximum_spend_logs_retention_period is empty/null, delete the field first
const retentionPeriodValue = formValues.maximum_spend_logs_retention_period;
const shouldDeleteRetentionPeriod =
!retentionPeriodValue ||
(typeof retentionPeriodValue === "string" && retentionPeriodValue.trim() === "");
if (shouldDeleteRetentionPeriod) {
try {
await deleteField({
config_type: ConfigType.GENERAL_SETTINGS,
field_name: GeneralSettingsFieldName.MAXIMUM_SPEND_LOGS_RETENTION_PERIOD,
});
} catch (deleteError) {
// If field doesn't exist, that's okay - continue with update
console.warn("Failed to delete retention period field (may not exist):", deleteError);
}
}
// Update the settings (excluding maximum_spend_logs_retention_period if it's empty)
const updateParams: StoreRequestInSpendLogsParams = {
store_prompts_in_spend_logs: formValues.store_prompts_in_spend_logs,
...(retentionPeriodValue &&
typeof retentionPeriodValue === "string" &&
retentionPeriodValue.trim() !== "" && {
maximum_spend_logs_retention_period: retentionPeriodValue,
}),
};
await mutateAsync(updateParams, {
onSuccess: () => {
NotificationsManager.success("Spend logs settings updated successfully");
form.resetFields();
refetch(); // Refetch config to get updated values
onSuccess?.();
},
onError: (error) => {
@@ -42,55 +99,55 @@ const SpendLogsSettingsModal: React.FC<SpendLogsSettingsModalProps> = ({ isVisib
return (
<Modal
title="Spend Logs Settings"
title={<span className="flex gap-2"><Typography.Title level={5}>Spend Logs Settings</Typography.Title><NewBadge /></span>}
open={isVisible}
width={600}
footer={
<Space>
<Button onClick={handleCancel} disabled={isPending}>
<Button onClick={handleCancel} disabled={isPending || isDeletingField || isLoadingConfig}>
Cancel
</Button>
<Button type="primary" loading={isPending} onClick={() => form.submit()}>
{isPending ? "Saving..." : "Save Settings"}
<Button type="primary" loading={isPending || isDeletingField} disabled={isLoadingConfig} onClick={() => form.submit()}>
{isPending || isDeletingField ? "Saving..." : "Save Settings"}
</Button>
</Space>
}
onCancel={handleCancel}
>
<Form
key={proxyConfigData ? JSON.stringify(initialValues) : 'loading'}
form={form}
layout="horizontal"
labelCol={{ flex: "auto", style: { textAlign: "left" } }}
wrapperCol={{ flex: "auto", style: { textAlign: "right" } }}
onFinish={handleFormSubmit}
initialValues={{
store_prompts_in_spend_logs: false,
maximum_spend_logs_retention_period: undefined,
}}
initialValues={initialValues}
>
<Form.Item
label="Store Prompts in Spend Logs"
name="store_prompts_in_spend_logs"
tooltip="When enabled, prompts will be stored in spend logs for tracking and analysis purposes."
tooltip={
proxyConfigData?.find(f => f.field_name === 'store_prompts_in_spend_logs')?.field_description ||
"When enabled, prompts will be stored in spend logs for tracking and analysis purposes."
}
valuePropName="checked"
>
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center' }}>
<span>Store Prompts in Spend Logs</span>
<Switch checked={storePromptsValue ?? false} onChange={(checked) => form.setFieldValue('store_prompts_in_spend_logs', checked)} />
<div>
{isLoadingConfig ? <Skeleton.Input active block /> : <Switch checked={storePromptsValue ?? false} onChange={(checked) => form.setFieldValue('store_prompts_in_spend_logs', checked)} />}
</div>
</Form.Item>
<Form.Item
label="Maximum Spend Logs Retention Period (Optional)"
name="maximum_spend_logs_retention_period"
tooltip="Set the maximum retention period for spend logs (e.g., '7d' for 7 days, '30d' for 30 days). Leave empty for no limit."
labelCol={{ flex: "auto", style: { textAlign: "left" } }}
wrapperCol={{ flex: "0 0 25%", style: { textAlign: "right" } }}
tooltip={
proxyConfigData?.find(f => f.field_name === 'maximum_spend_logs_retention_period')?.field_description ||
"Set the maximum retention period for spend logs (e.g., '7d' for 7 days, '30d' for 30 days). Leave empty for no limit."
}
>
<Input
{isLoadingConfig ? <Skeleton.Input active block /> : <Input
placeholder="e.g., 7d, 30d"
prefix={<ClockCircleOutlined />}
/>
/>}
</Form.Item>
</Form>
</Modal>
@@ -1,38 +1,36 @@
import { useQuery, useQueryClient } from "@tanstack/react-query";
import moment from "moment";
import { useQuery } from "@tanstack/react-query";
import { useState, useRef, useEffect, useCallback } from "react";
import { useQueryClient } from "@tanstack/react-query";
import { useCallback, useEffect, useRef, useState } from "react";
import { uiSpendLogsCall, keyInfoV1Call, sessionSpendLogsCall, keyListCall, allEndUsersCall } from "../networking";
import { DataTable } from "./table";
import { columns, LogEntry } from "./columns";
import { Row } from "@tanstack/react-table";
import { prefetchLogDetails } from "./prefetch";
import { RequestResponsePanel } from "./RequestResponsePanel";
import { ErrorViewer } from "./ErrorViewer";
import { internalUserRoles } from "../../utils/roles";
import { ConfigInfoMessage } from "./ConfigInfoMessage";
import { Button, Tooltip } from "antd";
import { KeyResponse, Team } from "../key_team_helpers/key_list";
import KeyInfoView from "../templates/key_info_view";
import { SessionView } from "./SessionView";
import { VectorStoreViewer } from "./VectorStoreViewer";
import GuardrailViewer from "@/components/view_logs/GuardrailViewer/GuardrailViewer";
import { CostBreakdownViewer } from "./CostBreakdownViewer";
import FilterComponent from "../molecules/filter";
import { FilterOption } from "../molecules/filter";
import { useLogFilterLogic } from "./log_filter_logic";
import { fetchAllKeyAliases } from "../key_team_helpers/filter_helpers";
import { Tab, TabGroup, TabList, TabPanels, TabPanel, Switch } from "@tremor/react";
import AuditLogs from "./audit_logs";
import { getTimeRangeDisplay } from "./logs_utils";
import { formatNumberWithCommas } from "@/utils/dataUtils";
import { truncateString } from "@/utils/textUtils";
import { SettingOutlined } from "@ant-design/icons";
import { Row } from "@tanstack/react-table";
import { Switch, Tab, TabGroup, TabList, TabPanel, TabPanels } from "@tremor/react";
import { Button, Tooltip } from "antd";
import { internalUserRoles } from "../../utils/roles";
import DeletedKeysPage from "../DeletedKeysPage/DeletedKeysPage";
import DeletedTeamsPage from "../DeletedTeamsPage/DeletedTeamsPage";
import NewBadge from "../common_components/NewBadge";
import { fetchAllKeyAliases } from "../key_team_helpers/filter_helpers";
import { KeyResponse, Team } from "../key_team_helpers/key_list";
import FilterComponent, { FilterOption } from "../molecules/filter";
import { allEndUsersCall, keyInfoV1Call, keyListCall, sessionSpendLogsCall, uiSpendLogsCall } from "../networking";
import KeyInfoView from "../templates/key_info_view";
import AuditLogs from "./audit_logs";
import { columns, LogEntry } from "./columns";
import { ConfigInfoMessage } from "./ConfigInfoMessage";
import { CostBreakdownViewer } from "./CostBreakdownViewer";
import { ErrorViewer } from "./ErrorViewer";
import { useLogFilterLogic } from "./log_filter_logic";
import { getTimeRangeDisplay } from "./logs_utils";
import { prefetchLogDetails } from "./prefetch";
import { RequestResponsePanel } from "./RequestResponsePanel";
import { SessionView } from "./SessionView";
import SpendLogsSettingsModal from "./SpendLogsSettingsModal/SpendLogsSettingsModal";
import { SettingOutlined } from "@ant-design/icons";
import { DataTable } from "./table";
import { VectorStoreViewer } from "./VectorStoreViewer";
import NewBadge from "../common_components/NewBadge";
interface SpendLogsTableProps {
accessToken: string | null;
@@ -513,8 +511,8 @@ export default function SpendLogsTable({
<TabList>
<Tab>Request Logs</Tab>
<Tab>Audit Logs</Tab>
<Tab><>Deleted Keys <NewBadge /></></Tab>
<Tab><>Deleted Teams <NewBadge /></></Tab>
<Tab>Deleted Keys</Tab>
<Tab>Deleted Teams</Tab>
</TabList>
<TabPanels>
<TabPanel>
@@ -535,11 +533,12 @@ export default function SpendLogsTable({
)}
</h1>
{!selectedSessionId && (
<Button
<NewBadge dot><Button
icon={<SettingOutlined />}
onClick={() => setIsSpendLogsSettingsModalVisible(true)}
title="Spend Logs Settings"
/>
/></NewBadge>
)}
</div>
{selectedKeyInfo && selectedKeyIdInfoView && selectedKeyInfo.api_key === selectedKeyIdInfoView ? (
@@ -555,7 +554,7 @@ export default function SpendLogsTable({
<DataTable
columns={columns}
data={sessionData}
renderSubComponent={RequestViewer}
renderSubComponent={({ row }) => <RequestViewer row={row} onOpenSettings={() => setIsSpendLogsSettingsModalVisible(true)} />}
getRowCanExpand={() => true}
// Optionally: add session-specific row expansion state
/>
@@ -754,7 +753,7 @@ export default function SpendLogsTable({
<DataTable
columns={columns}
data={filteredData}
renderSubComponent={RequestViewer}
renderSubComponent={({ row }) => <RequestViewer row={row} onOpenSettings={() => setIsSpendLogsSettingsModalVisible(true)} />}
getRowCanExpand={() => true}
/>
</div>
@@ -780,7 +779,7 @@ export default function SpendLogsTable({
);
}
export function RequestViewer({ row }: { row: Row<LogEntry> }) {
export function RequestViewer({ row, onOpenSettings }: { row: Row<LogEntry>; onOpenSettings?: () => void }) {
// Helper function to clean metadata by removing specific fields
const formatData = (input: any) => {
if (typeof input === "string") {
@@ -991,7 +990,7 @@ export function RequestViewer({ row }: { row: Row<LogEntry> }) {
<CostBreakdownViewer costBreakdown={row.original.metadata?.cost_breakdown} totalSpend={row.original.spend || 0} />
{/* Configuration Info Message - Show when data is missing */}
<ConfigInfoMessage show={missingData} />
<ConfigInfoMessage show={missingData} onOpenSettings={onOpenSettings} />
{/* Request/Response Panel */}
<div className="w-full max-w-full overflow-hidden">