diff --git a/.circleci/config.yml b/.circleci/config.yml index d99c485af94..8672561f654 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -1255,7 +1255,15 @@ jobs: ls # Add --timeout to kill hanging tests after 120s (2 min) # Add --durations=20 to show 20 slowest tests for debugging - python -m pytest -vv tests/llm_translation --cov=litellm --cov-report=xml -v --junitxml=test-results/junit.xml --durations=20 -n 4 --timeout=120 --timeout_method=thread + # Subdirectories with dedicated jobs (maintain this list as new jobs are added) + IGNORE_DIRS=( + "tests/llm_translation/realtime" + ) + IGNORE_ARGS="" + for dir in "${IGNORE_DIRS[@]}"; do + IGNORE_ARGS="$IGNORE_ARGS --ignore=$dir" + done + python -m pytest -vv tests/llm_translation $IGNORE_ARGS --cov=litellm --cov-report=xml -v --junitxml=test-results/junit.xml --durations=20 -n 4 --timeout=120 --timeout_method=thread no_output_timeout: 120m - run: name: Rename the coverage files @@ -1271,6 +1279,54 @@ jobs: paths: - llm_translation_coverage.xml - llm_translation_coverage + realtime_translation_testing: + docker: + - image: cimg/python:3.11 + auth: + username: ${DOCKERHUB_USERNAME} + password: ${DOCKERHUB_PASSWORD} + working_directory: ~/project + + steps: + - checkout + - setup_google_dns + - run: + name: Install Dependencies + command: | + python -m pip install --upgrade pip + python -m pip install -r requirements.txt + pip install "pytest==7.3.1" + pip install "pytest-retry==1.6.3" + pip install "pytest-cov==5.0.0" + pip install "pytest-asyncio==0.21.1" + pip install "respx==0.22.0" + pip install "pytest-xdist==3.6.1" + pip install "pytest-timeout==2.2.0" + pip install "websockets" + # Run pytest and generate JUnit XML report + - run: + name: Run realtime tests + command: | + pwd + ls + # Add --timeout to kill hanging tests after 120s (2 min) + # Add --durations=20 to show 20 slowest tests for debugging + python -m pytest -vv tests/llm_translation/realtime --cov=litellm --cov-report=xml -v --junitxml=test-results/junit.xml --durations=20 -n 4 --timeout=120 --timeout_method=thread + no_output_timeout: 120m + - run: + name: Rename the coverage files + command: | + mv coverage.xml realtime_translation_coverage.xml + mv .coverage realtime_translation_coverage + + # Store test results + - store_test_results: + path: test-results + - persist_to_workspace: + root: . + paths: + - realtime_translation_coverage.xml + - realtime_translation_coverage mcp_testing: docker: - image: cimg/python:3.11 @@ -3532,7 +3588,7 @@ jobs: python -m venv venv . venv/bin/activate pip install coverage - coverage combine llm_translation_coverage llm_responses_api_coverage ocr_coverage search_coverage mcp_coverage logging_coverage audio_coverage litellm_router_coverage litellm_router_unit_coverage local_testing_part1_coverage local_testing_part2_coverage litellm_assistants_api_coverage auth_ui_unit_tests_coverage langfuse_coverage caching_coverage litellm_proxy_unit_tests_part1_coverage litellm_proxy_unit_tests_part2_coverage image_gen_coverage pass_through_unit_tests_coverage batches_coverage litellm_security_tests_coverage guardrails_coverage litellm_mapped_tests_coverage + coverage combine llm_translation_coverage realtime_translation_coverage llm_responses_api_coverage ocr_coverage search_coverage mcp_coverage logging_coverage audio_coverage litellm_router_coverage litellm_router_unit_coverage local_testing_part1_coverage local_testing_part2_coverage litellm_assistants_api_coverage auth_ui_unit_tests_coverage langfuse_coverage caching_coverage litellm_proxy_unit_tests_part1_coverage litellm_proxy_unit_tests_part2_coverage image_gen_coverage pass_through_unit_tests_coverage batches_coverage litellm_security_tests_coverage guardrails_coverage litellm_mapped_tests_coverage coverage xml - codecov/upload: file: ./coverage.xml @@ -4196,6 +4252,12 @@ workflows: only: - main - /litellm_.*/ + - realtime_translation_testing: + filters: + branches: + only: + - main + - /litellm_.*/ - mcp_testing: filters: branches: @@ -4307,6 +4369,7 @@ workflows: - upload-coverage: requires: - llm_translation_testing + - realtime_translation_testing - mcp_testing - google_generate_content_endpoint_testing - guardrails_testing @@ -4384,6 +4447,7 @@ workflows: - e2e_openai_endpoints - test_bad_database_url - llm_translation_testing + - realtime_translation_testing - mcp_testing - google_generate_content_endpoint_testing - llm_responses_api_testing diff --git a/Makefile b/Makefile index 0da83c363cd..b867d7ea35e 100644 --- a/Makefile +++ b/Makefile @@ -1,7 +1,10 @@ # LiteLLM Makefile # Simple Makefile for running tests and basic development tasks -.PHONY: help test test-unit test-integration test-unit-helm lint format install-dev install-proxy-dev install-test-deps install-helm-unittest check-circular-imports check-import-safety +.PHONY: help test test-unit test-integration test-unit-helm \ + info lint lint-dev format \ + install-dev install-proxy-dev install-test-deps \ + install-helm-unittest check-circular-imports check-import-safety # Default target help: @@ -25,6 +28,13 @@ help: @echo " make test-integration - Run integration tests" @echo " make test-unit-helm - Run helm unit tests" +# Keep PIP simple for edge cases: +PIP := $(shell command -v pip > /dev/null 2>&1 && echo "pip" || echo "python3 -m pip") + +# Show info +info: + @echo "PIP: $(PIP)" + # Installation targets install-dev: poetry install --with dev @@ -34,19 +44,19 @@ install-proxy-dev: # CI-compatible installations (matches GitHub workflows exactly) install-dev-ci: - pip install openai==2.8.0 + $(PIP) install openai==2.8.0 poetry install --with dev - pip install openai==2.8.0 + $(PIP) install openai==2.8.0 install-proxy-dev-ci: poetry install --with dev,proxy-dev --extras proxy - pip install openai==2.8.0 + $(PIP) install openai==2.8.0 install-test-deps: install-proxy-dev - poetry run pip install "pytest-retry==1.6.3" - poetry run pip install pytest-xdist - poetry run pip install openapi-core - cd enterprise && poetry run pip install -e . && cd .. + poetry run $(PIP) install "pytest-retry==1.6.3" + poetry run $(PIP) install pytest-xdist + poetry run $(PIP) install openapi-core + cd enterprise && poetry run $(PIP) install -e . && cd .. install-helm-unittest: helm plugin install https://github.com/helm-unittest/helm-unittest --version v0.4.4 || echo "ignore error if plugin exists" @@ -62,8 +72,40 @@ format-check: install-dev lint-ruff: install-dev cd litellm && poetry run ruff check . && cd .. +# faster linter for developing ... +# inspiration from: +# https://github.com/astral-sh/ruff/discussions/10977 +# https://github.com/astral-sh/ruff/discussions/4049 +lint-format-changed: install-dev + @git diff origin/main --unified=0 --no-color -- '*.py' | \ + perl -ne '\ + if (/^diff --git a\/(.*) b\//) { $$file = $$1; } \ + if (/^@@ .* \+(\d+)(?:,(\d+))? @@/) { \ + $$start = $$1; $$count = $$2 || 1; $$end = $$start + $$count - 1; \ + print "$$file:$$start:1-$$end:999\n"; \ + }' | \ + while read range; do \ + file="$${range%%:*}"; \ + lines="$${range#*:}"; \ + echo "Formatting $$file (lines $$lines)"; \ + poetry run ruff format --range "$$lines" "$$file"; \ + done + +lint-ruff-dev: install-dev + @tmpfile=$$(mktemp /tmp/ruff-dev.XXXXXX) && \ + cd litellm && \ + (poetry run ruff check . --output-format=pylint || true) > "$$tmpfile" && \ + poetry run diff-quality --violations=pylint "$$tmpfile" --compare-branch=origin/main && \ + cd .. ; \ + rm -f "$$tmpfile" + +lint-ruff-FULL-dev: install-dev + @files=$$(git diff --name-only origin/main -- '*.py'); \ + if [ -n "$$files" ]; then echo "$$files" | xargs poetry run ruff check; \ + else echo "No changed .py files to check."; fi + lint-mypy: install-dev - poetry run pip install types-requests types-setuptools types-redis types-PyYAML + poetry run $(PIP) install types-requests types-setuptools types-redis types-PyYAML cd litellm && poetry run mypy . --ignore-missing-imports && cd .. lint-black: format-check @@ -72,11 +114,14 @@ check-circular-imports: install-dev cd litellm && poetry run python ../tests/documentation_tests/test_circular_imports.py && cd .. check-import-safety: install-dev - 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 *; print('[from litellm import *] OK! no issues!');" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1) # Combined linting (matches test-linting.yml workflow) lint: format-check lint-ruff lint-mypy check-circular-imports check-import-safety +# Faster linting for local development (only checks changed code) +lint-dev: lint-format-changed lint-mypy check-circular-imports check-import-safety + # Testing targets test: poetry run pytest tests/ diff --git a/cookbook/livekit_agent_sdk/README.md b/cookbook/livekit_agent_sdk/README.md new file mode 100644 index 00000000000..1c3f0bf9564 --- /dev/null +++ b/cookbook/livekit_agent_sdk/README.md @@ -0,0 +1,114 @@ +# LiveKit Voice Agent with LiteLLM Gateway + +Simple example showing how to use LiveKit's xAI realtime plugin with LiteLLM as a proxy. This lets you switch between xAI, OpenAI, and Azure realtime APIs without changing your code. + +## Quick Start + +### 1. Install dependencies + +```bash +pip install livekit-agents[xai] websockets +``` + +### 2. Start LiteLLM proxy + +```bash +# With xAI +export XAI_API_KEY="your-xai-key" +litellm --config config.yaml --port 4000 +``` + +### 3. Run the voice agent + +```bash +python main.py +``` + +Type your message and get a voice response from Grok! + +## Configuration + +Set these environment variables if needed: + +```bash +export LITELLM_PROXY_URL="http://localhost:4000" +export LITELLM_API_KEY="sk-1234" +export LITELLM_MODEL="grok-voice-agent" +``` + +Or use the defaults - connects to `http://localhost:4000` by default. + +## Example Config File + +Create a `config.yaml` with your realtime models: + +```yaml +model_list: + - model_name: grok-voice-agent + litellm_params: + model: xai/grok-2-vision-1212 + api_key: os.environ/XAI_API_KEY + model_info: + mode: realtime + + - model_name: openai-voice-agent + litellm_params: + model: gpt-4o-realtime-preview + api_key: os.environ/OPENAI_API_KEY + model_info: + mode: realtime + +general_settings: + master_key: sk-1234 +``` + +Then start: `litellm --config config.yaml --port 4000` + +## How It Works + +LiveKit's xAI plugin connects through LiteLLM proxy by setting `base_url`: + +```python +from livekit.plugins import xai + +model = xai.realtime.RealtimeModel( + voice="ara", + api_key="sk-1234", # LiteLLM proxy key + base_url="http://localhost:4000", # Point to LiteLLM +) +``` + +## Switching Providers + +Just change the model in your config - no code changes needed: + +**xAI Grok:** +```yaml +model: xai/grok-2-vision-1212 +``` + +**OpenAI:** +```yaml +model: gpt-4o-realtime-preview +``` + +**Azure OpenAI:** +```yaml +model: azure/gpt-4o-realtime-preview +api_base: https://your-endpoint.openai.azure.com/ +``` + +## Why Use LiteLLM? + +- āœ… **Switch providers** without changing agent code +- āœ… **Cost tracking** across all voice sessions +- āœ… **Rate limiting** and budgets +- āœ… **Load balancing** across multiple API keys +- āœ… **Fallbacks** to backup models + +## Learn More + +- [LiveKit xAI Realtime Tutorial](/docs/tutorials/livekit_xai_realtime) +- [xAI Realtime Docs](/docs/providers/xai_realtime) +- [LiveKit Agents Documentation](https://docs.livekit.io/agents/) +- [LiteLLM Realtime API](/docs/realtime) diff --git a/cookbook/livekit_agent_sdk/config.example.yaml b/cookbook/livekit_agent_sdk/config.example.yaml new file mode 100644 index 00000000000..1361f36af34 --- /dev/null +++ b/cookbook/livekit_agent_sdk/config.example.yaml @@ -0,0 +1,21 @@ +model_list: + - model_name: grok-voice-agent + litellm_params: + model: xai/grok-2-vision-1212 + api_key: os.environ/XAI_API_KEY + model_info: + mode: realtime + + - model_name: openai-voice-agent + litellm_params: + model: gpt-4o-realtime-preview + api_key: os.environ/OPENAI_API_KEY + model_info: + mode: realtime + +litellm_settings: + drop_params: True + telemetry: False + +general_settings: + master_key: sk-1234 # Change this to a secure key diff --git a/cookbook/livekit_agent_sdk/main.py b/cookbook/livekit_agent_sdk/main.py new file mode 100644 index 00000000000..0e2d7ebdfaf --- /dev/null +++ b/cookbook/livekit_agent_sdk/main.py @@ -0,0 +1,112 @@ +""" +Simple xAI Voice Agent using LiveKit SDK with LiteLLM Gateway + +This example shows how to use LiveKit's xAI realtime plugin through LiteLLM proxy. +LiteLLM acts as a unified interface, allowing you to switch between xAI, OpenAI, +and Azure realtime APIs without changing your agent code. +""" +import asyncio +import json +import os +import websockets + +# Configuration +PROXY_URL = os.getenv("LITELLM_PROXY_URL", "http://localhost:4000") +API_KEY = os.getenv("LITELLM_API_KEY", "sk-1234") +MODEL = os.getenv("LITELLM_MODEL", "grok-voice-agent") + + +async def run_voice_agent(): + """ + Simple voice agent that: + 1. Connects to xAI realtime API through LiteLLM proxy + 2. Sends a user message + 3. Streams back the response + """ + + url = f"ws://{PROXY_URL.replace('http://', '').replace('https://', '')}/v1/realtime?model={MODEL}" + headers = {"Authorization": f"Bearer {API_KEY}"} + + print(f"šŸŽ™ļø Connecting to voice agent...") + print(f" Model: {MODEL}") + print(f" Proxy: {PROXY_URL}") + print() + + async with websockets.connect(url, additional_headers=headers) as ws: + # Receive initial connection event + initial = json.loads(await ws.recv()) + print(f"āœ… Connected! Event: {initial['type']}\n") + + # Get user input + user_message = input("šŸ’¬ Your message: ").strip() + if not user_message: + user_message = "Tell me a fun fact about AI!" + + print(f"\nšŸ¤– Sending to {MODEL}...\n") + + # Send user message + await ws.send(json.dumps({ + "type": "conversation.item.create", + "item": { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": user_message}] + } + })) + + # Request response + await ws.send(json.dumps({ + "type": "response.create", + "response": {"modalities": ["text", "audio"]} + })) + + # Stream response + print("šŸŽ¤ Response: ", end='', flush=True) + transcript = [] + + try: + while True: + msg = await asyncio.wait_for(ws.recv(), timeout=15.0) + event = json.loads(msg) + + # Capture transcript deltas + if event['type'] == 'response.output_audio_transcript.delta': + delta = event.get('delta', '') + if delta: + print(delta, end='', flush=True) + transcript.append(delta) + + # Done when response completes + elif event['type'] == 'response.done': + break + + except asyncio.TimeoutError: + pass + + print("\n") + + if transcript: + print(f"āœ… Complete response: {''.join(transcript)}") + + await ws.close() + + +def main(): + """Run the voice agent""" + print("=" * 70) + print("LiveKit xAI Voice Agent via LiteLLM Proxy") + print("=" * 70) + print() + + try: + asyncio.run(run_voice_agent()) + except KeyboardInterrupt: + print("\n\nšŸ‘‹ Goodbye!") + except Exception as e: + print(f"\nāŒ Error: {e}") + print("\nMake sure LiteLLM proxy is running:") + print(f" litellm --config config.yaml --port 4000") + + +if __name__ == "__main__": + main() diff --git a/cookbook/livekit_agent_sdk/requirements.txt b/cookbook/livekit_agent_sdk/requirements.txt new file mode 100644 index 00000000000..9e3542fac27 --- /dev/null +++ b/cookbook/livekit_agent_sdk/requirements.txt @@ -0,0 +1,2 @@ +livekit-agents[xai]>=1.3.12 +websockets>=15.0.1 diff --git a/docs/my-website/docs/a2a.md b/docs/my-website/docs/a2a.md index a7e8b52d99a..b1166a7809c 100644 --- a/docs/my-website/docs/a2a.md +++ b/docs/my-website/docs/a2a.md @@ -68,116 +68,9 @@ Follow [this guide, to add your pydantic ai agent to LiteLLM Agent Gateway](./pr ## Invoking your Agents -Use the [A2A Python SDK](https://pypi.org/project/a2a-sdk) to invoke agents through LiteLLM. - -This example shows how to: -1. **List available agents** - Query `/v1/agents` to see which agents your key can access -2. **Select an agent** - Pick an agent from the list -3. **Invoke via A2A** - Use the A2A protocol to send messages to the agent - -```python showLineNumbers title="invoke_a2a_agent.py" -from uuid import uuid4 -import httpx -import asyncio -from a2a.client import A2ACardResolver, A2AClient -from a2a.types import MessageSendParams, SendMessageRequest - -# === CONFIGURE THESE === -LITELLM_BASE_URL = "http://localhost:4000" # Your LiteLLM proxy URL -LITELLM_VIRTUAL_KEY = "sk-1234" # Your LiteLLM Virtual Key -# ======================= - -async def main(): - headers = {"Authorization": f"Bearer {LITELLM_VIRTUAL_KEY}"} - - async with httpx.AsyncClient(headers=headers) as client: - # Step 1: List available agents - response = await client.get(f"{LITELLM_BASE_URL}/v1/agents") - agents = response.json() - - print("Available agents:") - for agent in agents: - print(f" - {agent['agent_name']} (ID: {agent['agent_id']})") - - if not agents: - print("No agents available for this key") - return - - # Step 2: Select an agent and invoke it - selected_agent = agents[0] - agent_id = selected_agent["agent_id"] - agent_name = selected_agent["agent_name"] - print(f"\nInvoking: {agent_name}") - - # Step 3: Use A2A protocol to invoke the agent - base_url = f"{LITELLM_BASE_URL}/a2a/{agent_id}" - resolver = A2ACardResolver(httpx_client=client, base_url=base_url) - agent_card = await resolver.get_agent_card() - a2a_client = A2AClient(httpx_client=client, agent_card=agent_card) - - request = SendMessageRequest( - id=str(uuid4()), - params=MessageSendParams( - message={ - "role": "user", - "parts": [{"kind": "text", "text": "Hello, what can you do?"}], - "messageId": uuid4().hex, - } - ), - ) - response = await a2a_client.send_message(request) - print(f"Response: {response.model_dump(mode='json', exclude_none=True, indent=4)}") - -if __name__ == "__main__": - asyncio.run(main()) -``` - -### Streaming Responses - -For streaming responses, use `send_message_streaming`: - -```python showLineNumbers title="invoke_a2a_agent_streaming.py" -from uuid import uuid4 -import httpx -import asyncio -from a2a.client import A2ACardResolver, A2AClient -from a2a.types import MessageSendParams, SendStreamingMessageRequest - -# === CONFIGURE THESE === -LITELLM_BASE_URL = "http://localhost:4000" # Your LiteLLM proxy URL -LITELLM_VIRTUAL_KEY = "sk-1234" # Your LiteLLM Virtual Key -LITELLM_AGENT_NAME = "ij-local" # Agent name registered in LiteLLM -# ======================= - -async def main(): - base_url = f"{LITELLM_BASE_URL}/a2a/{LITELLM_AGENT_NAME}" - headers = {"Authorization": f"Bearer {LITELLM_VIRTUAL_KEY}"} - - async with httpx.AsyncClient(headers=headers) as httpx_client: - # Resolve agent card and create client - resolver = A2ACardResolver(httpx_client=httpx_client, base_url=base_url) - agent_card = await resolver.get_agent_card() - client = A2AClient(httpx_client=httpx_client, agent_card=agent_card) - - # Send a streaming message - request = SendStreamingMessageRequest( - id=str(uuid4()), - params=MessageSendParams( - message={ - "role": "user", - "parts": [{"kind": "text", "text": "Hello, what can you do?"}], - "messageId": uuid4().hex, - } - ), - ) - - # Stream the response - async for chunk in client.send_message_streaming(request): - print(chunk.model_dump(mode="json", exclude_none=True)) - -if __name__ == "__main__": - asyncio.run(main()) -``` +See the [Invoking A2A Agents](./a2a_invoking_agents) guide to learn how to call your agents using: +- **A2A SDK** - Native A2A protocol with full support for tasks and artifacts +- **OpenAI SDK** - Familiar `/chat/completions` interface with `a2a/` model prefix ## Tracking Agent Logs diff --git a/docs/my-website/docs/a2a_invoking_agents.md b/docs/my-website/docs/a2a_invoking_agents.md new file mode 100644 index 00000000000..3bb248e4561 --- /dev/null +++ b/docs/my-website/docs/a2a_invoking_agents.md @@ -0,0 +1,280 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Invoking A2A Agents + +Learn how to invoke A2A agents through LiteLLM using different methods. + +:::tip Deploy Your Own A2A Agent + +Want to test with your own agent? Deploy this template A2A agent powered by Google Gemini: + +[**shin-bot-litellm/a2a-gemini-agent**](https://github.com/shin-bot-litellm/a2a-gemini-agent) - Simple deployable A2A agent with streaming support + +::: + +## A2A SDK + +Use the [A2A Python SDK](https://pypi.org/project/a2a-sdk) to invoke agents through LiteLLM using the A2A protocol. + +### Non-Streaming + +This example shows how to: +1. **List available agents** - Query `/v1/agents` to see which agents your key can access +2. **Select an agent** - Pick an agent from the list +3. **Invoke via A2A** - Use the A2A protocol to send messages to the agent + +```python showLineNumbers title="invoke_a2a_agent.py" +from uuid import uuid4 +import httpx +import asyncio +from a2a.client import A2ACardResolver, A2AClient +from a2a.types import MessageSendParams, SendMessageRequest + +# === CONFIGURE THESE === +LITELLM_BASE_URL = "http://localhost:4000" # Your LiteLLM proxy URL +LITELLM_VIRTUAL_KEY = "sk-1234" # Your LiteLLM Virtual Key +# ======================= + +async def main(): + headers = {"Authorization": f"Bearer {LITELLM_VIRTUAL_KEY}"} + + async with httpx.AsyncClient(headers=headers) as client: + # Step 1: List available agents + response = await client.get(f"{LITELLM_BASE_URL}/v1/agents") + agents = response.json() + + print("Available agents:") + for agent in agents: + print(f" - {agent['agent_name']} (ID: {agent['agent_id']})") + + if not agents: + print("No agents available for this key") + return + + # Step 2: Select an agent and invoke it + selected_agent = agents[0] + agent_id = selected_agent["agent_id"] + agent_name = selected_agent["agent_name"] + print(f"\nInvoking: {agent_name}") + + # Step 3: Use A2A protocol to invoke the agent + base_url = f"{LITELLM_BASE_URL}/a2a/{agent_id}" + resolver = A2ACardResolver(httpx_client=client, base_url=base_url) + agent_card = await resolver.get_agent_card() + a2a_client = A2AClient(httpx_client=client, agent_card=agent_card) + + request = SendMessageRequest( + id=str(uuid4()), + params=MessageSendParams( + message={ + "role": "user", + "parts": [{"kind": "text", "text": "Hello, what can you do?"}], + "messageId": uuid4().hex, + } + ), + ) + response = await a2a_client.send_message(request) + print(f"Response: {response.model_dump(mode='json', exclude_none=True, indent=4)}") + +if __name__ == "__main__": + asyncio.run(main()) +``` + +### Streaming + +For streaming responses, use `send_message_streaming`: + +```python showLineNumbers title="invoke_a2a_agent_streaming.py" +from uuid import uuid4 +import httpx +import asyncio +from a2a.client import A2ACardResolver, A2AClient +from a2a.types import MessageSendParams, SendStreamingMessageRequest + +# === CONFIGURE THESE === +LITELLM_BASE_URL = "http://localhost:4000" # Your LiteLLM proxy URL +LITELLM_VIRTUAL_KEY = "sk-1234" # Your LiteLLM Virtual Key +LITELLM_AGENT_NAME = "ij-local" # Agent name registered in LiteLLM +# ======================= + +async def main(): + base_url = f"{LITELLM_BASE_URL}/a2a/{LITELLM_AGENT_NAME}" + headers = {"Authorization": f"Bearer {LITELLM_VIRTUAL_KEY}"} + + async with httpx.AsyncClient(headers=headers) as httpx_client: + # Resolve agent card and create client + resolver = A2ACardResolver(httpx_client=httpx_client, base_url=base_url) + agent_card = await resolver.get_agent_card() + client = A2AClient(httpx_client=httpx_client, agent_card=agent_card) + + # Send a streaming message + request = SendStreamingMessageRequest( + id=str(uuid4()), + params=MessageSendParams( + message={ + "role": "user", + "parts": [{"kind": "text", "text": "Tell me a long story"}], + "messageId": uuid4().hex, + } + ), + ) + + # Stream the response + async for chunk in client.send_message_streaming(request): + print(chunk.model_dump(mode="json", exclude_none=True)) + +if __name__ == "__main__": + asyncio.run(main()) +``` + +## /chat/completions API (OpenAI SDK) + +You can also invoke A2A agents using the familiar OpenAI SDK by using the `a2a/` model prefix. + +### Non-Streaming + + + + +```python showLineNumbers title="openai_non_streaming.py" +import openai + +client = openai.OpenAI( + api_key="sk-1234", # Your LiteLLM Virtual Key + base_url="http://localhost:4000" # Your LiteLLM proxy URL +) + +response = client.chat.completions.create( + model="a2a/my-agent", # Use a2a/ prefix with your agent name + messages=[ + {"role": "user", "content": "Hello, what can you do?"} + ] +) + +print(response.choices[0].message.content) +``` + + + + +```typescript showLineNumbers title="openai_non_streaming.ts" +import OpenAI from 'openai'; + +const client = new OpenAI({ + apiKey: 'sk-1234', // Your LiteLLM Virtual Key + baseURL: 'http://localhost:4000' // Your LiteLLM proxy URL +}); + +const response = await client.chat.completions.create({ + model: 'a2a/my-agent', // Use a2a/ prefix with your agent name + messages: [ + { role: 'user', content: 'Hello, what can you do?' } + ] +}); + +console.log(response.choices[0].message.content); +``` + + + + +```bash showLineNumbers title="curl_non_streaming.sh" +curl -X POST http://localhost:4000/v1/chat/completions \ + -H "Authorization: Bearer sk-1234" \ + -H "Content-Type: application/json" \ + -d '{ + "model": "a2a/my-agent", + "messages": [ + {"role": "user", "content": "Hello, what can you do?"} + ] + }' +``` + + + + +### Streaming + + + + +```python showLineNumbers title="openai_streaming.py" +import openai + +client = openai.OpenAI( + api_key="sk-1234", # Your LiteLLM Virtual Key + base_url="http://localhost:4000" # Your LiteLLM proxy URL +) + +stream = client.chat.completions.create( + model="a2a/my-agent", # Use a2a/ prefix with your agent name + messages=[ + {"role": "user", "content": "Tell me a long story"} + ], + stream=True +) + +for chunk in stream: + if chunk.choices[0].delta.content: + print(chunk.choices[0].delta.content, end="", flush=True) +``` + + + + +```typescript showLineNumbers title="openai_streaming.ts" +import OpenAI from 'openai'; + +const client = new OpenAI({ + apiKey: 'sk-1234', // Your LiteLLM Virtual Key + baseURL: 'http://localhost:4000' // Your LiteLLM proxy URL +}); + +const stream = await client.chat.completions.create({ + model: 'a2a/my-agent', // Use a2a/ prefix with your agent name + messages: [ + { role: 'user', content: 'Tell me a long story' } + ], + stream: true +}); + +for await (const chunk of stream) { + const content = chunk.choices[0]?.delta?.content; + if (content) { + process.stdout.write(content); + } +} +``` + + + + +```bash showLineNumbers title="curl_streaming.sh" +curl -X POST http://localhost:4000/v1/chat/completions \ + -H "Authorization: Bearer sk-1234" \ + -H "Content-Type: application/json" \ + -d '{ + "model": "a2a/my-agent", + "messages": [ + {"role": "user", "content": "Tell me a long story"} + ], + "stream": true + }' +``` + + + + +## Key Differences + +| Method | Use Case | Advantages | +|--------|----------|------------| +| **A2A SDK** | Native A2A protocol integration | • Full A2A protocol support
• Access to task states and artifacts
• Context management | +| **OpenAI SDK** | Familiar OpenAI-style interface | • Drop-in replacement for OpenAI calls
• Easier migration from LLM to agent workflows
• Works with existing OpenAI tooling | + +:::tip Model Prefix + +When using the OpenAI SDK, always prefix your agent name with `a2a/` (e.g., `a2a/my-agent`) to route requests to the A2A agent instead of an LLM provider. + +::: diff --git a/docs/my-website/docs/adding_provider/simple_guardrail_tutorial.md b/docs/my-website/docs/adding_provider/simple_guardrail_tutorial.md index 9c654cd1560..884a7397bde 100644 --- a/docs/my-website/docs/adding_provider/simple_guardrail_tutorial.md +++ b/docs/my-website/docs/adding_provider/simple_guardrail_tutorial.md @@ -101,12 +101,11 @@ model_list: - model_name: gpt-4 litellm_params: model: gpt-4 - api_key: os.environ/OPENAI_API_KEY + api_key: os.environ/OPENAI_API_KEY -litellm_settings: - guardrails: +guardrails: - guardrail_name: my_guardrail - litellm_params: + litellm_params: guardrail: my_guardrail mode: during_call api_key: os.environ/MY_GUARDRAIL_API_KEY diff --git a/docs/my-website/docs/observability/langfuse_integration.md b/docs/my-website/docs/observability/langfuse_integration.md index a81336c5bc6..d3c5a44d481 100644 --- a/docs/my-website/docs/observability/langfuse_integration.md +++ b/docs/my-website/docs/observability/langfuse_integration.md @@ -215,6 +215,66 @@ The following parameters can be updated on a continuation of a trace by passing Any other key value pairs passed into the metadata not listed in the above spec for a `litellm` completion will be added as a metadata key value pair for the generation. +#### Multiple Langfuse Projects (Per-Request Credentials) + +You can send traces to different Langfuse projects per request by passing credentials directly to `completion()` or `acompletion()`. This works alongside (or instead of) the global env vars and is useful when different teams or business processes use different Langfuse projects. + +Pass **`langfuse_public_key`**, **`langfuse_secret_key`** (or **`langfuse_secret`**), and optionally **`langfuse_host`** as keyword arguments: + +```python +import litellm +from litellm import completion + +# Optional: set a default via env for requests that don't pass credentials +# os.environ["LANGFUSE_PUBLIC_KEY"] = "pk-default..." +# os.environ["LANGFUSE_SECRET_KEY"] = "sk-default..." + +litellm.success_callback = ["langfuse"] +litellm.failure_callback = ["langfuse"] + +# Request 1 → Langfuse Project A +response_a = completion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "Hello from team A"}], + langfuse_public_key="pk-lf-project-a...", + langfuse_secret_key="sk-lf-project-a...", + langfuse_host="https://us.cloud.langfuse.com", # optional +) + +# Request 2 → Langfuse Project B (different project) +response_b = completion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "Hello from team B"}], + langfuse_public_key="pk-lf-project-b...", + langfuse_secret_key="sk-lf-project-b...", + langfuse_host="https://eu.cloud.langfuse.com", # optional, can differ per project +) +``` + +Async usage with per-request credentials: + +```python +import litellm +from litellm import acompletion + +litellm.success_callback = ["langfuse"] +litellm.failure_callback = ["langfuse"] + +response = await acompletion( + model="gpt-3.5-turbo", + messages=[{"role": "user", "content": "Hi"}], + langfuse_public_key="pk-lf-...", + langfuse_secret_key="sk-lf-...", + langfuse_host="https://us.cloud.langfuse.com", # optional +) +``` + +- **`langfuse_public_key`** – Langfuse project public key (required for per-request override). +- **`langfuse_secret_key`** or **`langfuse_secret`** – Langfuse secret key (either name is accepted). +- **`langfuse_host`** – Langfuse host URL (e.g. `https://us.cloud.langfuse.com`); optional, defaults to env or Langfuse cloud. + +When these are passed, that request uses this project (and host) for the Langfuse callback; when omitted, the callback uses the global Langfuse client (from env vars if set). LiteLLM caches a Langfuse client per credential set to avoid creating a new client on every request. + #### Disable Logging - Specific Calls To disable logging for specific calls use the `no-log` flag. diff --git a/docs/my-website/docs/providers/github_copilot.md b/docs/my-website/docs/providers/github_copilot.md index 306c9f949ec..e9fd3444f5f 100644 --- a/docs/my-website/docs/providers/github_copilot.md +++ b/docs/my-website/docs/providers/github_copilot.md @@ -35,11 +35,10 @@ from litellm import completion response = completion( model="github_copilot/gpt-4", - messages=[{"role": "user", "content": "Write a Python function to calculate fibonacci numbers"}], - extra_headers={ - "editor-version": "vscode/1.85.1", - "Copilot-Integration-Id": "vscode-chat" - } + messages=[ + {"role": "system", "content": "You are a helpful coding assistant"}, + {"role": "user", "content": "Write a Python function to calculate fibonacci numbers"} + ] ) print(response) ``` @@ -50,11 +49,7 @@ from litellm import completion stream = completion( model="github_copilot/gpt-4", messages=[{"role": "user", "content": "Explain async/await in Python"}], - stream=True, - extra_headers={ - "editor-version": "vscode/1.85.1", - "Copilot-Integration-Id": "vscode-chat" - } + stream=True ) for chunk in stream: @@ -134,11 +129,7 @@ client = OpenAI( # Non-streaming response response = client.chat.completions.create( model="github_copilot/gpt-4", - messages=[{"role": "user", "content": "How do I optimize this SQL query?"}], - extra_headers={ - "editor-version": "vscode/1.85.1", - "Copilot-Integration-Id": "vscode-chat" - } + messages=[{"role": "user", "content": "How do I optimize this SQL query?"}] ) print(response.choices[0].message.content) @@ -156,11 +147,7 @@ response = litellm.completion( model="litellm_proxy/github_copilot/gpt-4", messages=[{"role": "user", "content": "Review this code for bugs"}], api_base="http://localhost:4000", - api_key="your-proxy-api-key", - extra_headers={ - "editor-version": "vscode/1.85.1", - "Copilot-Integration-Id": "vscode-chat" - } + api_key="your-proxy-api-key" ) print(response.choices[0].message.content) @@ -174,8 +161,6 @@ print(response.choices[0].message.content) curl http://localhost:4000/v1/chat/completions \ -H "Content-Type: application/json" \ -H "Authorization: Bearer your-proxy-api-key" \ - -H "editor-version: vscode/1.85.1" \ - -H "Copilot-Integration-Id: vscode-chat" \ -d '{ "model": "github_copilot/gpt-4", "messages": [{"role": "user", "content": "Explain this error message"}] @@ -211,9 +196,11 @@ export GITHUB_COPILOT_API_KEY_FILE="api-key.json" ### Headers -GitHub Copilot supports various editor-specific headers: +LiteLLM automatically injects the required GitHub Copilot headers (simulating VSCode). You don't need to specify them manually. -```python showLineNumbers title="Common Headers" +If you want to override the defaults (e.g., to simulate a different editor), you can use `extra_headers`: + +```python showLineNumbers title="Custom Headers (Optional)" extra_headers = { "editor-version": "vscode/1.85.1", # Editor version "editor-plugin-version": "copilot/1.155.0", # Plugin version diff --git a/docs/my-website/docs/providers/xai_realtime.md b/docs/my-website/docs/providers/xai_realtime.md new file mode 100644 index 00000000000..b36908c4686 --- /dev/null +++ b/docs/my-website/docs/providers/xai_realtime.md @@ -0,0 +1,308 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# xAI Voice Agent (Realtime API) + +xAI's Grok Voice Agent provides real-time voice conversation capabilities through WebSocket connections, enabling natural bidirectional audio interactions. + +| Feature | Description | Comments | +| --- | --- | --- | +| LiteLLM AI Gateway | āœ… | | +| LiteLLM Python SDK | āœ… | Full support via `litellm.realtime()` | + +## Quick Start + +### Supported Model + +| Model | Context | Features | +|-------|---------|----------| +| `xai/grok-4-1-fast-non-reasoning` | 2M tokens | Voice conversation, Function calling, Vision, Audio, Web search, Caching | + +**Note:** xAI Realtime API uses the non-reasoning variant for optimal real-time performance. + +## Python SDK Usage + +### Basic Realtime Connection + +```python +import asyncio +from litellm import realtime + +async def test_xai_realtime(): + """ + Test xAI Grok Voice Agent via LiteLLM SDK + """ + # Initialize realtime connection + ws = await realtime( + model="xai/grok-4-1-fast-non-reasoning", + api_key="your-xai-api-key", # or set XAI_API_KEY env var + ) + + # Connection established, xAI sends "conversation.created" event + print("Connected to xAI Grok Voice Agent") + + # Send a message + await ws.send_text(json.dumps({ + "type": "conversation.item.create", + "item": { + "type": "message", + "role": "user", + "content": [{ + "type": "input_text", + "text": "Hello! How are you?" + }] + } + })) + + # Request a response + await ws.send_text(json.dumps({ + "type": "response.create" + })) + + # Listen for responses + async for message in ws: + data = json.loads(message) + print(f"Received: {data['type']}") + + if data['type'] == 'response.done': + break + + await ws.close() + +# Run the async function +asyncio.run(test_xai_realtime()) +``` + +### With Audio Input/Output + +```python +import asyncio +import json +from litellm import realtime + +async def xai_voice_conversation(): + """ + Voice conversation with xAI Grok Voice Agent + """ + ws = await realtime( + model="xai/grok-4-1-fast-non-reasoning", + api_key="your-xai-api-key", + ) + + # Send audio data (base64 encoded PCM16 24kHz) + await ws.send_text(json.dumps({ + "type": "conversation.item.create", + "item": { + "type": "message", + "role": "user", + "content": [{ + "type": "input_audio", + "audio": "base64_encoded_audio_data_here" + }] + } + })) + + # Request response with audio + await ws.send_text(json.dumps({ + "type": "response.create", + "response": { + "modalities": ["text", "audio"], + "instructions": "Please respond in a friendly tone." + } + })) + + # Process streaming audio response + async for message in ws: + data = json.loads(message) + + if data['type'] == 'response.audio.delta': + # Handle audio chunks + audio_chunk = data['delta'] + # Process audio_chunk (play it, save it, etc.) + + elif data['type'] == 'response.done': + break + + await ws.close() + +asyncio.run(xai_voice_conversation()) +``` + +## LiteLLM Proxy (AI Gateway) Usage + +Load balance across multiple xAI deployments or combine with other providers. + +### 1. Add Model to Config + +```yaml +model_list: + - model_name: grok-voice-agent + litellm_params: + model: xai/grok-4-1-fast-non-reasoning + api_key: os.environ/XAI_API_KEY + model_info: + mode: realtime + + # Optional: Add fallback to OpenAI + - model_name: grok-voice-agent + litellm_params: + model: openai/gpt-4o-realtime-preview-2024-10-01 + api_key: os.environ/OPENAI_API_KEY + model_info: + mode: realtime +``` + +### 2. Start Proxy + +```bash +litellm --config /path/to/config.yaml + +# RUNNING on http://0.0.0.0:4000 +``` + +### 3. Test Connection + +#### Python Client + +```python +import asyncio +import websockets +import json + +async def test_proxy(): + url = "ws://0.0.0.0:4000/v1/realtime?model=grok-voice-agent" + + async with websockets.connect( + url, + extra_headers={ + "Authorization": "Bearer sk-1234", # Your LiteLLM proxy key + "OpenAI-Beta": "realtime=v1" + } + ) as ws: + # Wait for conversation.created event from xAI + message = await ws.recv() + print(f"Connected: {message}") + + # Send a message + await ws.send(json.dumps({ + "type": "conversation.item.create", + "item": { + "type": "message", + "role": "user", + "content": [{ + "type": "input_text", + "text": "Hello from LiteLLM proxy!" + }] + } + })) + + # Request response + await ws.send(json.dumps({ + "type": "response.create" + })) + + # Listen for response + async for message in ws: + data = json.loads(message) + print(f"Event: {data['type']}") + + if data['type'] == 'response.done': + break + +asyncio.run(test_proxy()) +``` + +#### Node.js Client + +```javascript +// test.js - Run with: node test.js +const WebSocket = require("ws"); + +const url = "ws://0.0.0.0:4000/v1/realtime?model=grok-voice-agent"; + +const ws = new WebSocket(url, { + headers: { + "Authorization": "Bearer sk-1234", + "OpenAI-Beta": "realtime=v1", + }, +}); + +ws.on("open", function open() { + console.log("Connected to xAI via LiteLLM proxy"); + + // Send a message + ws.send(JSON.stringify({ + type: "conversation.item.create", + item: { + type: "message", + role: "user", + content: [{ + type: "input_text", + text: "What's the weather like?" + }] + } + })); + + // Request response + ws.send(JSON.stringify({ + type: "response.create", + response: { + modalities: ["text"], + instructions: "Please assist the user." + } + })); +}); + +ws.on("message", function incoming(message) { + const data = JSON.parse(message.toString()); + console.log(`Event: ${data.type}`); + + if (data.type === 'response.done') { + ws.close(); + } +}); + +ws.on("error", function handleError(error) { + console.error("Error: ", error); +}); +``` + +## Key Differences from OpenAI + +xAI's Grok Voice Agent has some differences from OpenAI's Realtime API: + +| Feature | xAI | OpenAI | LiteLLM Handling | +|---------|-----|--------|------------------| +| Initial Event | `conversation.created` | `session.created` | āš ļø Passed through as-is | +| WebSocket URL | `wss://api.x.ai/v1/realtime` | `wss://api.openai.com/v1/realtime` | āœ… Auto-configured | +| Model | `grok-4-1-fast-non-reasoning` | `gpt-4o-realtime-preview` | āœ… Via model prefix | +| Audio Format | PCM16 24kHz mono | PCM16 24kHz mono | āœ… Compatible | +| Context Window | 2M tokens | 128K tokens | N/A | + +**What LiteLLM Handles:** +- āœ… Automatic URL routing to correct provider +- āœ… Authentication headers (no `OpenAI-Beta` header for xAI) +- āœ… WebSocket connection management +- āœ… All other event types are compatible + +**What You Need to Handle:** +- āš ļø Initial event type difference (`conversation.created` vs `session.created`) + +**Tip:** Make your client compatible with both event types: +```python +# Handle both providers +if event['type'] in ['session.created', 'conversation.created']: + print("Connection established") +``` + +## Related Documentation + +- [xAI Chat/Text Models](/docs/providers/xai) +- [LiteLLM Realtime API Overview](/docs/realtime) +- [xAI Official Documentation](https://docs.x.ai/docs) + +## Support + +For issues or questions: +- [LiteLLM GitHub Issues](https://github.com/BerriAI/litellm/issues) +- [xAI Documentation](https://docs.x.ai/docs) diff --git a/docs/my-website/docs/proxy/admin_ui_sso.md b/docs/my-website/docs/proxy/admin_ui_sso.md index 7b299429db7..37e45b50284 100644 --- a/docs/my-website/docs/proxy/admin_ui_sso.md +++ b/docs/my-website/docs/proxy/admin_ui_sso.md @@ -23,26 +23,75 @@ From v1.76.0, SSO is now Free for up to 5 users. -1. Add Okta credentials to your .env +#### Step 1: Create an OIDC Application in Okta + +In your Okta Admin Console, create a new **OIDC Web Application**. See [Okta's guide on creating OIDC app integrations](https://help.okta.com/en-us/content/topics/apps/apps_app_integration_wizard_oidc.htm) for detailed instructions. + +When configuring the application: +- **Sign-in redirect URI**: `https:///sso/callback` +- **Sign-out redirect URI** (optional): `https://` + + + +After creating the app, copy your **Client ID** and **Client Secret** from the application's General tab: + + + +#### Step 2: Assign Users to the Application + +Ensure users are assigned to the app in the **Assignments** tab. If Federation Broker Mode is enabled, you may need to disable it to assign users manually. + +#### Step 3: Configure Authorization Server Access Policy + +:::warning Important +This step is required. Without an Access Policy for your app, users will get a `no_matching_policy` error when attempting to log in. +::: + +1. Go to **Security** → **API** + + + +2. Select the **default** authorization server (or your custom one) + + + +3. Click on **Access Policies** tab, create a new policy assigned to your LiteLLM app +4. Add a rule that allows the **Authorization Code** grant type + + + +See [Okta's Access Policy documentation](https://help.okta.com/en-us/content/topics/security/api-access-management/access-policies.htm) for more details. + +#### Step 4: Configure LiteLLM Environment Variables ```bash -GENERIC_CLIENT_ID = "" -GENERIC_CLIENT_SECRET = "" -GENERIC_AUTHORIZATION_ENDPOINT = "/authorize" # https://dev-2kqkcd6lx6kdkuzt.us.auth0.com/authorize -GENERIC_TOKEN_ENDPOINT = "/token" # https://dev-2kqkcd6lx6kdkuzt.us.auth0.com/oauth/token -GENERIC_USERINFO_ENDPOINT = "/userinfo" # https://dev-2kqkcd6lx6kdkuzt.us.auth0.com/userinfo -GENERIC_CLIENT_STATE = "random-string" # [OPTIONAL] REQUIRED BY OKTA, if not set random state value is generated -GENERIC_SSO_HEADERS = "Content-Type=application/json, X-Custom-Header=custom-value" # [OPTIONAL] Comma-separated list of additional headers to add to the request - e.g. Content-Type=application/json, etc. +GENERIC_CLIENT_ID="" +GENERIC_CLIENT_SECRET="" +GENERIC_AUTHORIZATION_ENDPOINT="https:///oauth2/default/v1/authorize" +GENERIC_TOKEN_ENDPOINT="https:///oauth2/default/v1/token" +GENERIC_USERINFO_ENDPOINT="https:///oauth2/default/v1/userinfo" +GENERIC_CLIENT_STATE="random-string" +PROXY_BASE_URL="https://" ``` -You can get your domain specific auth/token/userinfo endpoints at `/.well-known/openid-configuration` +:::tip +You can find all OAuth endpoints at `https:///.well-known/openid-configuration` +::: -2. Add proxy url as callback_url on Okta +#### Step 5: Test the SSO Flow -On Okta, add the 'callback_url' as `/sso/callback` +1. Start your LiteLLM proxy +2. Navigate to `https:///ui` +3. Click the SSO login button +4. Authenticate with Okta and verify you're redirected back to LiteLLM +#### Troubleshooting - +| Error | Cause | Solution | +|-------|-------|----------| +| `redirect_uri` error | Redirect URI not configured | Add `/sso/callback` to Sign-in redirect URIs in Okta | +| `access_denied` | User not assigned to app | Assign the user in the Assignments tab | +| `no_matching_policy` | Missing Access Policy | Create an Access Policy in the Authorization Server (see Step 3) | diff --git a/docs/my-website/docs/proxy/cli.md b/docs/my-website/docs/proxy/cli.md index 9244f75b756..d3624000a32 100644 --- a/docs/my-website/docs/proxy/cli.md +++ b/docs/my-website/docs/proxy/cli.md @@ -1,7 +1,10 @@ # CLI Arguments -Cli arguments, --host, --port, --num_workers -## --host +This page documents all command-line interface (CLI) arguments available for the LiteLLM proxy server. + +## Server Configuration + +### --host - **Default:** `'0.0.0.0'` - The host for the server to listen on. - **Usage:** @@ -14,7 +17,7 @@ Cli arguments, --host, --port, --num_workers litellm ``` -## --port +### --port - **Default:** `4000` - The port to bind the server to. - **Usage:** @@ -27,9 +30,9 @@ Cli arguments, --host, --port, --num_workers litellm ``` -## --num_workers - - **Default:** `1` - - The number of uvicorn workers to spin up. +### --num_workers + - **Default:** Number of logical CPUs in the system, or `4` if that cannot be determined + - The number of uvicorn / gunicorn workers to spin up. - **Usage:** ```shell litellm --num_workers 4 @@ -40,55 +43,273 @@ Cli arguments, --host, --port, --num_workers litellm ``` -## --api_base +### --config + - **Short form:** `-c` - **Default:** `None` - - The API base for the model litellm should call. + - Path to the proxy configuration file (e.g., config.yaml). + - **Usage:** + ```shell + litellm --config path/to/config.yaml + ``` + +### --log_config + - **Default:** `None` + - **Type:** `str` + - Path to the logging configuration file for uvicorn. + - **Usage:** + ```shell + litellm --log_config path/to/log_config.conf + ``` + +### --keepalive_timeout + - **Default:** `None` + - **Type:** `int` + - Set the uvicorn keepalive timeout in seconds (uvicorn timeout_keep_alive parameter). + - **Usage:** + ```shell + litellm --keepalive_timeout 30 + ``` + - **Usage - set Environment Variable:** `KEEPALIVE_TIMEOUT` + ```shell + export KEEPALIVE_TIMEOUT=30 + litellm + ``` + +### --max_requests_before_restart + - **Default:** `None` + - **Type:** `int` + - Restart worker after this many requests. This is useful for mitigating memory growth over time. + - For uvicorn: maps to `limit_max_requests` + - For gunicorn: maps to `max_requests` + - **Usage:** + ```shell + litellm --max_requests_before_restart 10000 + ``` + - **Usage - set Environment Variable:** `MAX_REQUESTS_BEFORE_RESTART` + ```shell + export MAX_REQUESTS_BEFORE_RESTART=10000 + litellm + ``` + +## Server Backend Options + +### --run_gunicorn + - **Default:** `False` + - **Type:** `bool` (Flag) + - Starts proxy via gunicorn instead of uvicorn. Better for managing multiple workers in production. + - **Usage:** + ```shell + litellm --run_gunicorn + ``` + +### --run_hypercorn + - **Default:** `False` + - **Type:** `bool` (Flag) + - Starts proxy via hypercorn instead of uvicorn. Supports HTTP/2. + - **Usage:** + ```shell + litellm --run_hypercorn + ``` + +### --skip_server_startup + - **Default:** `False` + - **Type:** `bool` (Flag) + - Skip starting the server after setup (useful for database migrations only). + - **Usage:** + ```shell + litellm --skip_server_startup + ``` + +## SSL/TLS Configuration + +### --ssl_keyfile_path + - **Default:** `None` + - **Type:** `str` + - Path to the SSL keyfile. Use this when you want to provide SSL certificate when starting proxy. + - **Usage:** + ```shell + litellm --ssl_keyfile_path /path/to/key.pem --ssl_certfile_path /path/to/cert.pem + ``` + - **Usage - set Environment Variable:** `SSL_KEYFILE_PATH` + ```shell + export SSL_KEYFILE_PATH=/path/to/key.pem + litellm + ``` + +### --ssl_certfile_path + - **Default:** `None` + - **Type:** `str` + - Path to the SSL certfile. Use this when you want to provide SSL certificate when starting proxy. + - **Usage:** + ```shell + litellm --ssl_certfile_path /path/to/cert.pem --ssl_keyfile_path /path/to/key.pem + ``` + - **Usage - set Environment Variable:** `SSL_CERTFILE_PATH` + ```shell + export SSL_CERTFILE_PATH=/path/to/cert.pem + litellm + ``` + +### --ciphers + - **Default:** `None` + - **Type:** `str` + - Ciphers to use for the SSL setup. Only used with `--run_hypercorn`. + - **Usage:** + ```shell + litellm --run_hypercorn --ssl_keyfile_path /path/to/key.pem --ssl_certfile_path /path/to/cert.pem --ciphers "ECDHE+AESGCM" + ``` + +## Model Configuration + +### --model or -m + - **Default:** `None` + - The model name to pass to LiteLLM. + - **Usage:** + ```shell + litellm --model gpt-3.5-turbo + ``` + +### --alias + - **Default:** `None` + - An alias for the model, for user-friendly reference. Use this to give a litellm model name (e.g., "huggingface/codellama/CodeLlama-7b-Instruct-hf") a more user-friendly name ("codellama"). + - **Usage:** + ```shell + litellm --alias my-gpt-model + ``` + +### --api_base + - **Default:** `None` + - The API base for the model LiteLLM should call. - **Usage:** ```shell litellm --model huggingface/tinyllama --api_base https://k58ory32yinf1ly0.us-east-1.aws.endpoints.huggingface.cloud ``` -## --api_version - - **Default:** `None` +### --api_version + - **Default:** `2024-07-01-preview` - For Azure services, specify the API version. - **Usage:** ```shell litellm --model azure/gpt-deployment --api_version 2023-08-01 --api_base https://" ``` -## --model or -m +### --headers - **Default:** `None` - - The model name to pass to Litellm. + - Headers for the API call (as JSON string). - **Usage:** ```shell - litellm --model gpt-3.5-turbo + litellm --model my-model --headers '{"Authorization": "Bearer token"}' ``` -## --test - - **Type:** `bool` (Flag) - - Proxy chat completions URL to make a test request. - - **Usage:** - ```shell - litellm --test - ``` - -## --health - - **Type:** `bool` (Flag) - - Runs a health check on all models in config.yaml - - **Usage:** - ```shell - litellm --health - ``` - -## --alias +### --add_key - **Default:** `None` - - An alias for the model, for user-friendly reference. + - Add a key to the model configuration. - **Usage:** ```shell - litellm --alias my-gpt-model + litellm --add_key my-api-key ``` -## --debug +### --save + - **Type:** `bool` (Flag) + - Save the model-specific config. + - **Usage:** + ```shell + litellm --model gpt-3.5-turbo --save + ``` + +## Model Parameters + +### --temperature + - **Default:** `None` + - **Type:** `float` + - Set the temperature for the model. + - **Usage:** + ```shell + litellm --temperature 0.7 + ``` + +### --max_tokens + - **Default:** `None` + - **Type:** `int` + - Set the maximum number of tokens for the model output. + - **Usage:** + ```shell + litellm --max_tokens 50 + ``` + +### --request_timeout + - **Default:** `None` + - **Type:** `int` + - Set the timeout in seconds for completion calls. + - **Usage:** + ```shell + litellm --request_timeout 300 + ``` + +### --max_budget + - **Default:** `None` + - **Type:** `float` + - Set max budget for API calls. Works for hosted models like OpenAI, TogetherAI, Anthropic, etc. + - **Usage:** + ```shell + litellm --max_budget 100.0 + ``` + +### --drop_params + - **Type:** `bool` (Flag) + - Drop any unmapped params. + - **Usage:** + ```shell + litellm --drop_params + ``` + +### --add_function_to_prompt + - **Type:** `bool` (Flag) + - If a function passed but unsupported, pass it as a part of the prompt. + - **Usage:** + ```shell + litellm --add_function_to_prompt + ``` + +## Database Configuration + +### --iam_token_db_auth + - **Default:** `False` + - **Type:** `bool` (Flag) + - Connects to an RDS database using IAM token authentication instead of a password. This is useful for AWS RDS instances that are configured to use IAM database authentication. + - When enabled, LiteLLM will generate an IAM authentication token to connect to the database. + - **Required Environment Variables:** + - `DATABASE_HOST` - The RDS database host + - `DATABASE_PORT` - The database port + - `DATABASE_USER` - The database user + - `DATABASE_NAME` - The database name + - `DATABASE_SCHEMA` (optional) - The database schema + - **Usage:** + ```shell + litellm --iam_token_db_auth + ``` + - **Usage - set Environment Variable:** `IAM_TOKEN_DB_AUTH` + ```shell + export IAM_TOKEN_DB_AUTH=True + export DATABASE_HOST=mydb.us-east-1.rds.amazonaws.com + export DATABASE_PORT=5432 + export DATABASE_USER=mydbuser + export DATABASE_NAME=mydb + litellm + ``` + +### --use_prisma_db_push + - **Default:** `False` + - **Type:** `bool` (Flag) + - Use `prisma db push` instead of `prisma migrate` for database schema updates. This is useful when you want to quickly sync your database schema without creating migration files. + - **Usage:** + ```shell + litellm --use_prisma_db_push + ``` + +## Debugging + +### --debug - **Default:** `False` - **Type:** `bool` (Flag) - Enable debugging mode for the input. @@ -102,10 +323,10 @@ Cli arguments, --host, --port, --num_workers litellm ``` -## --detailed_debug +### --detailed_debug - **Default:** `False` - **Type:** `bool` (Flag) - - Enable debugging mode for the input. + - Enable detailed debugging mode to view verbose debug logs. - **Usage:** ```shell litellm --detailed_debug @@ -116,80 +337,76 @@ Cli arguments, --host, --port, --num_workers litellm ``` -#### --temperature - - **Default:** `None` - - **Type:** `float` - - Set the temperature for the model. - - **Usage:** - ```shell - litellm --temperature 0.7 - ``` - -## --max_tokens - - **Default:** `None` - - **Type:** `int` - - Set the maximum number of tokens for the model output. - - **Usage:** - ```shell - litellm --max_tokens 50 - ``` - -## --request_timeout - - **Default:** `6000` - - **Type:** `int` - - Set the timeout in seconds for completion calls. - - **Usage:** - ```shell - litellm --request_timeout 300 - ``` - -## --drop_params +### --local + - **Default:** `False` - **Type:** `bool` (Flag) - - Drop any unmapped params. + - For local debugging purposes. - **Usage:** ```shell - litellm --drop_params + litellm --local ``` -## --add_function_to_prompt +## Testing & Health Checks + +### --test - **Type:** `bool` (Flag) - - If a function passed but unsupported, pass it as a part of the prompt. + - Proxy chat completions URL to make a test request to. - **Usage:** ```shell - litellm --add_function_to_prompt + litellm --test ``` -## --config - - Configure Litellm by providing a configuration file path. +### --test_async + - **Default:** `False` + - **Type:** `bool` (Flag) + - Calls async endpoints `/queue/requests` and `/queue/response`. - **Usage:** ```shell - litellm --config path/to/config.yaml + litellm --test_async ``` -## --telemetry +### --num_requests + - **Default:** `10` + - **Type:** `int` + - Number of requests to hit async endpoint with (used with `--test_async`). + - **Usage:** + ```shell + litellm --test_async --num_requests 100 + ``` + +### --health + - **Type:** `bool` (Flag) + - Runs a health check on all models in config.yaml. + - **Usage:** + ```shell + litellm --health + ``` + +## Other Options + +### --version + - **Short form:** `-v` + - **Type:** `bool` (Flag) + - Print LiteLLM version and exit. + - **Usage:** + ```shell + litellm --version + ``` + +### --telemetry - **Default:** `True` - **Type:** `bool` - - Help track usage of this feature. + - Help track usage of this feature. Turn off for privacy. - **Usage:** ```shell litellm --telemetry False ``` - -## --log_config - - **Default:** `None` - - **Type:** `str` - - Specify a log configuration file for uvicorn. - - **Usage:** - ```shell - litellm --log_config path/to/log_config.conf - ``` - -## --skip_server_startup +### --use_queue - **Default:** `False` - **Type:** `bool` (Flag) - - Skip starting the server after setup (useful for DB migrations only). + - To use celery workers for async endpoints. - **Usage:** ```shell - litellm --skip_server_startup - ``` \ No newline at end of file + litellm --use_queue + ``` diff --git a/docs/my-website/docs/proxy/config_settings.md b/docs/my-website/docs/proxy/config_settings.md index 385b4b0de32..5cdae51f448 100644 --- a/docs/my-website/docs/proxy/config_settings.md +++ b/docs/my-website/docs/proxy/config_settings.md @@ -94,7 +94,7 @@ litellm_settings: # /chat/completions, /completions, /embeddings, /audio/transcriptions mode: default_off # if default_off, you need to opt in to caching on a per call basis ttl: 600 # ttl for caching - disable_copilot_system_to_assistant: False # If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior. + disable_copilot_system_to_assistant: False # DEPRECATED - GitHub Copilot API supports system prompts. callback_settings: otel: @@ -197,7 +197,7 @@ router_settings: | disable_add_transform_inline_image_block | boolean | For Fireworks AI models - if true, turns off the auto-add of `#transform=inline` to the url of the image_url, if the model is not a vision model. | | disable_hf_tokenizer_download | boolean | If true, it defaults to using the openai tokenizer for all models (including huggingface models). | | enable_json_schema_validation | boolean | If true, enables json schema validation for all requests. | -| disable_copilot_system_to_assistant | boolean | If false (default), converts all 'system' role messages to 'assistant' for GitHub Copilot compatibility. Set to true to disable this behavior. Useful for tools (like Claude Code) that send system messages, which Copilot does not support. | +| disable_copilot_system_to_assistant | boolean | **DEPRECATED** - GitHub Copilot API supports system prompts. | ### general_settings - Reference diff --git a/docs/my-website/docs/proxy/guardrails/custom_code_guardrail.md b/docs/my-website/docs/proxy/guardrails/custom_code_guardrail.md new file mode 100644 index 00000000000..cb246144497 --- /dev/null +++ b/docs/my-website/docs/proxy/guardrails/custom_code_guardrail.md @@ -0,0 +1,278 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# Custom Code Guardrail + +Write custom guardrail logic using Python-like code that runs in a sandboxed environment. + +## Quick Start + +### 1. Define the guardrail in config + +```yaml +model_list: + - model_name: gpt-4 + litellm_params: + model: gpt-4 + api_key: os.environ/OPENAI_API_KEY + +guardrails: + - guardrail_name: block-ssn + litellm_params: + guardrail: custom_code + mode: pre_call + custom_code: | + def apply_guardrail(inputs, request_data, input_type): + for text in inputs["texts"]: + if regex_match(text, r"\d{3}-\d{2}-\d{4}"): + return block("SSN detected") + return allow() +``` + +### 2. Start proxy + +```bash +litellm --config config.yaml +``` + +### 3. Test + +```bash +curl -X POST http://localhost:4000/chat/completions \ + -H "Authorization: Bearer sk-1234" \ + -H "Content-Type: application/json" \ + -d '{ + "model": "gpt-4", + "messages": [{"role": "user", "content": "My SSN is 123-45-6789"}], + "guardrails": ["block-ssn"] + }' +``` + +## Configuration + +| Parameter | Type | Required | Description | +|-----------|------|----------|-------------| +| `guardrail` | string | āœ… | Must be `custom_code` | +| `mode` | string | āœ… | When to run: `pre_call`, `post_call`, `during_call` | +| `custom_code` | string | āœ… | Python-like code with `apply_guardrail` function | +| `default_on` | bool | āŒ | Run on all requests (default: `false`) | + +## Writing Custom Code + +### Function Signature + +Your code must define an `apply_guardrail` function: + +```python +def apply_guardrail(inputs, request_data, input_type): + # inputs: see table below + # request_data: {"model": "...", "user_id": "...", "team_id": "...", "metadata": {...}} + # input_type: "request" or "response" + + return allow() # or block() or modify() +``` + +### `inputs` Parameter + +| Field | Type | Description | +|-------|------|-------------| +| `texts` | `List[str]` | Extracted text from the request/response | +| `images` | `List[str]` | Extracted images (for image guardrails) | +| `tools` | `List[dict]` | Tools sent to the LLM | +| `tool_calls` | `List[dict]` | Tool calls returned from the LLM | +| `structured_messages` | `List[dict]` | Full messages with role info (system/user/assistant) | +| `model` | `str` | The model being used | + +### `request_data` Parameter + +| Field | Type | Description | +|-------|------|-------------| +| `model` | `str` | Model name | +| `user_id` | `str` | User ID from API key | +| `team_id` | `str` | Team ID from API key | +| `end_user_id` | `str` | End user ID | +| `metadata` | `dict` | Request metadata | + +### Return Values + +| Function | Description | +|----------|-------------| +| `allow()` | Let request/response through | +| `block(reason)` | Reject with message | +| `modify(texts=[], images=[], tool_calls=[])` | Transform content | + +## Built-in Primitives + +### Regex + +| Function | Description | +|----------|-------------| +| `regex_match(text, pattern)` | Returns `True` if pattern found | +| `regex_replace(text, pattern, replacement)` | Replace all matches | +| `regex_find_all(text, pattern)` | Return list of matches | + +### JSON + +| Function | Description | +|----------|-------------| +| `json_parse(text)` | Parse JSON string, returns `None` on error | +| `json_stringify(obj)` | Convert to JSON string | +| `json_schema_valid(obj, schema)` | Validate against JSON schema | + +### URL + +| Function | Description | +|----------|-------------| +| `extract_urls(text)` | Extract all URLs from text | +| `is_valid_url(url)` | Check if URL is valid | +| `all_urls_valid(text)` | Check all URLs in text are valid | + +### Code Detection + +| Function | Description | +|----------|-------------| +| `detect_code(text)` | Returns `True` if code detected | +| `detect_code_languages(text)` | Returns list of detected languages | +| `contains_code_language(text, ["sql", "python"])` | Check for specific languages | + +### Text Utilities + +| Function | Description | +|----------|-------------| +| `contains(text, substring)` | Check if substring exists | +| `contains_any(text, [substr1, substr2])` | Check if any substring exists | +| `word_count(text)` | Count words | +| `char_count(text)` | Count characters | +| `lower(text)` / `upper(text)` / `trim(text)` | String transforms | + +## Examples + +### Block PII (SSN) + +```python +def apply_guardrail(inputs, request_data, input_type): + for text in inputs["texts"]: + if regex_match(text, r"\d{3}-\d{2}-\d{4}"): + return block("SSN detected") + return allow() +``` + +### Redact Email Addresses + +```python +def apply_guardrail(inputs, request_data, input_type): + pattern = r"[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}" + modified = [] + for text in inputs["texts"]: + modified.append(regex_replace(text, pattern, "[EMAIL REDACTED]")) + return modify(texts=modified) +``` + +### Block SQL Injection + +```python +def apply_guardrail(inputs, request_data, input_type): + if input_type != "request": + return allow() + for text in inputs["texts"]: + if contains_code_language(text, ["sql"]): + return block("SQL code not allowed") + return allow() +``` + +### Validate JSON Response + +```python +def apply_guardrail(inputs, request_data, input_type): + if input_type != "response": + return allow() + + schema = { + "type": "object", + "required": ["name", "value"] + } + + for text in inputs["texts"]: + obj = json_parse(text) + if obj is None: + return block("Invalid JSON response") + if not json_schema_valid(obj, schema): + return block("Response missing required fields") + return allow() +``` + +### Check URLs in Response + +```python +def apply_guardrail(inputs, request_data, input_type): + if input_type != "response": + return allow() + for text in inputs["texts"]: + if not all_urls_valid(text): + return block("Response contains invalid URLs") + return allow() +``` + +### Combine Multiple Checks + +```python +def apply_guardrail(inputs, request_data, input_type): + modified = [] + + for text in inputs["texts"]: + # Redact SSN + text = regex_replace(text, r"\d{3}-\d{2}-\d{4}", "[SSN]") + # Redact credit cards + text = regex_replace(text, r"\d{16}", "[CARD]") + modified.append(text) + + # Block SQL in requests + if input_type == "request": + for text in inputs["texts"]: + if contains_code_language(text, ["sql"]): + return block("SQL injection blocked") + + return modify(texts=modified) +``` + +## Sandbox Restrictions + +Custom code runs in a restricted environment: + +- āŒ No `import` statements +- āŒ No file I/O +- āŒ No network access +- āŒ No `exec()` or `eval()` +- āœ… Only LiteLLM-provided primitives available + +## Per-Request Usage + +Enable guardrail per request: + +```bash +curl -X POST http://localhost:4000/chat/completions \ + -H "Authorization: Bearer sk-1234" \ + -H "Content-Type: application/json" \ + -d '{ + "model": "gpt-4", + "messages": [{"role": "user", "content": "Hello"}], + "guardrails": ["block-ssn"] + }' +``` + +## Default On + +Run guardrail on all requests: + +```yaml +litellm_settings: + guardrails: + - guardrail_name: block-ssn + litellm_params: + guardrail: custom_code + mode: pre_call + default_on: true + custom_code: | + def apply_guardrail(inputs, request_data, input_type): + ... +``` diff --git a/docs/my-website/docs/proxy/guardrails/grayswan.md b/docs/my-website/docs/proxy/guardrails/grayswan.md index d6efaf15504..6c0ccbc293d 100644 --- a/docs/my-website/docs/proxy/guardrails/grayswan.md +++ b/docs/my-website/docs/proxy/guardrails/grayswan.md @@ -13,20 +13,26 @@ Cygnal returns a `violation` score between `0` and `1` (higher means more likely ### 1. Obtain Credentials -1. Create a Gray Swan account and generate a Cygnal API key. +1. Log in to our Gray Swan platform and generate a Cygnal API key. + + For existing customers, you should already have access to our [platform](https://platform.grayswan.ai). + + For new users, please register at this [page](https://hubs.ly/Q03-sX1J0) and we are more than happy to give you an onboarding! + + 2. Configure environment variables for the LiteLLM proxy host: -```bash -export GRAYSWAN_API_KEY="your-grayswan-key" -export GRAYSWAN_API_BASE="https://api.grayswan.ai" -``` + ```bash + export GRAYSWAN_API_KEY="your-grayswan-key" + export GRAYSWAN_API_BASE="https://api.grayswan.ai" + ``` ### 2. Configure `config.yaml` -Add a guardrail entry that references the Gray Swan integration. Below is a balanced example that monitors both input and output but only blocks once the violation score reaches the configured threshold. +Add a guardrail entry that references the Gray Swan integration. Below is our recommmended settings. ```yaml -model_list: +model_list: # this part is a standard litellm configuration for reference - model_name: openai/gpt-4.1-mini litellm_params: model: openai/gpt-4.1-mini @@ -40,13 +46,14 @@ guardrails: api_key: os.environ/GRAYSWAN_API_KEY api_base: os.environ/GRAYSWAN_API_BASE # optional optional_params: - on_flagged_action: monitor # or "block" + on_flagged_action: passthrough # or "block" or "monitor" violation_threshold: 0.5 # score >= threshold is flagged reasoning_mode: hybrid # off | hybrid | thinking - categories: - safety: "Detect jailbreaks and policy violations" - policy_id: "your-cygnal-policy-id" + policy_id: "your-cygnal-policy-id" # Optional: Your Cygnal policy ID. Defaults to a content safety policy if empty. + streaming_end_of_stream_only: true # For streaming API, only send the assembled message to Cygnal (post_call only). Defaults to false. default_on: true + guardrail_timeout: 30 # Defaults to 30 seconds. Change accordingly. + fail_open: true # Defaults to true; set to false to propagate guardrail errors. general_settings: master_key: "your-litellm-master-key" @@ -65,13 +72,13 @@ litellm --config config.yaml --port 4000 ## Choosing Guardrail Modes -Gray Swan can run during `pre_call`, `during_call`, and `post_call` stages. Combine modes based on your latency and coverage requirements. +Gray Swan can run during `pre_call`, `during_call`, and `post_call` stages. Combine modes based on your latency and coverage requirements. | Mode | When it Runs | Protects | Typical Use Case | |--------------|-------------------|-----------------------|------------------| | `pre_call` | Before LLM call | User input only | Block prompt injection before it reaches the model | | `during_call`| Parallel to call | User input only | Low-latency monitoring without blocking | -| `post_call` | After response | Full conversation | Scan output for policy violations, leaked secrets, or IPI | +| `post_call` | After response | Model Outputs | Scan output for policy violations, leaked secrets, or IPI | When using `during_call` with `on_flagged_action: block` or `on_flagged_action: passthrough`: @@ -81,87 +88,110 @@ When using `during_call` with `on_flagged_action: block` or `on_flagged_action: - The guardrail exception prevents the response from reaching the user, but **does not cancel the running LLM task** - This means you pay full LLM costs while returning an error/passthrough message to the user -**Recommendation:** For cost-sensitive applications, use `pre_call` and `post_call` instead of `during_call` for blocking or passthrough modes. Reserve `during_call` for `monitor` mode where you want low-latency logging without impacting the user experience. +**Recommendation:** Use `pre_call` and `post_call` instead of `during_call` for `passthrough` (or `block`) `on_flagged_action` (see our recommended configuration above). Reserve `during_call` for `monitor` mode ONLY when you want low-latency logging without impacting the user experience. - - +--- -```yaml -guardrails: - - guardrail_name: "cygnal-monitor-only" - litellm_params: - guardrail: grayswan - mode: "during_call" - api_key: os.environ/GRAYSWAN_API_KEY - optional_params: - on_flagged_action: monitor - violation_threshold: 0.6 - default_on: true +## Work with Claude Code + +Follow the official litellm [guide](https://docs.litellm.ai/docs/tutorials/claude_responses_api) on setting up Claude Code with litellm, with the guardrail part mentioned above added to your litellm configuration. Cygnal natively supports coding agent policies defense. Define your own policy or use the provided coding policies on the platform. The example config we show above is also the recommended setup for Claude Code (with the `policy_id` replaced with an appropriate one). + +--- + +## Per-request overrides via `extra_body` + +You can override parts of the Gray Swan guardrail configuration on a per-request basis by passing `litellm_metadata.guardrails[*].grayswan.extra_body`. + +`extra_body` is merged into the Cygnal request body and takes precedence over specific fields from `config.yaml`, which are `policy_id`, `violation_threshold`, and `reasoning_mode`. + +If you include a `metadata` field inside `extra_body`, it is forwarded to the Cygnal API as-is under the request body's `metadata` field. + +Example: + +```bash +curl -X POST "http://0.0.0.0:4000/v1/messages?beta=true" \ + -H "Authorization: Bearer token" \ + -H "Content-Type: application/json" \ + -d '{ + "model": "openrouter/anthropic/claude-sonnet-4.5", + "messages": [{"role": "user", "content": "hello"}], + "litellm_metadata": { + "guardrails": [ + { + "cygnal-monitor": { + "extra_body": { + "policy_id": "specific policy id you want to use", + "metadata": { + "user": "health-check" + } + } + } + } + ] + } + }' ``` -Best for visibility without blocking. Alerts are logged via LiteLLM’s standard logging callbacks. +OpenAI client: - - +```python +from openai import OpenAI -```yaml -guardrails: - - guardrail_name: "cygnal-block-input" - litellm_params: - guardrail: grayswan - mode: "pre_call" - api_key: os.environ/GRAYSWAN_API_KEY - optional_params: - on_flagged_action: block - violation_threshold: 0.4 - categories: - pii: "Detect sensitive data" - default_on: true +client = OpenAI(api_key="anything", base_url="http://0.0.0.0:4000") + +resp = client.responses.create( + model="openrouter/anthropic/claude-sonnet-4.5", + input="hello", + extra_body={ + "litellm_metadata": { + "guardrails": [ + { + "cygnal-monitor": { + "extra_body": { + "policy_id": "69038214e5cdb6befc5e991e", + "metadata": {"trace_id": "trace-123"}, + } + } + } + ] + } + }, +) ``` -Stops malicious or sensitive prompts before any tokens are generated. +Anthropic client: - - +```python +from anthropic import Anthropic -```yaml -guardrails: - - guardrail_name: "cygnal-full-coverage" - litellm_params: - guardrail: grayswan - mode: [pre_call, post_call] - api_key: os.environ/GRAYSWAN_API_KEY - optional_params: - on_flagged_action: block - violation_threshold: 0.5 - reasoning_mode: thinking - policy_id: "policy-id-from-grayswan" - default_on: true +client = Anthropic(api_key="anything", base_url="http://0.0.0.0:4000") + +resp = client.messages.create( + model="openrouter/anthropic/claude-sonnet-4.5", + max_tokens=256, + messages=[{"role": "user", "content": "hello"}], + extra_body={ + "litellm_metadata": { + "guardrails": [ + { + "cygnal-monitor": { + "extra_body": { + "policy_id": "69038214e5cdb6befc5e991e", + "metadata": {"trace_id": "trace-123"}, + } + } + } + ] + } + }, +) ``` -Provides the strongest enforcement by inspecting both prompts and responses. +Notes: - - - -```yaml -guardrails: - - guardrail_name: "cygnal-passthrough" - litellm_params: - guardrail: grayswan - mode: [pre_call, post_call] - api_key: os.environ/GRAYSWAN_API_KEY - optional_params: - on_flagged_action: passthrough - violation_threshold: 0.5 - default_on: true -``` - -Allows requests to proceed without raising a 400 error when content is flagged. Instead of blocking, the model response content is replaced with a detailed violation message including violation score, violated rules, and detection flags (mutation, IPI). **Supported Response Formats:** OpenAI chat/text completions, Anthropic Messages API. Other response types (embeddings, images, etc.) will log a warning and return unchanged. - - - +- The guardrail name (for example, `cygnal-monitor`) must match the `guardrail_name` in `config.yaml`. +- Per-request guardrail overrides may require a premium license, depending on your proxy settings. --- @@ -170,9 +200,14 @@ Allows requests to proceed without raising a 400 error when content is flagged. | Parameter | Type | Description | |---------------------------------------|-----------------|-------------| | `api_key` | string | Gray Swan Cygnal API key. Reads from `GRAYSWAN_API_KEY` if omitted. | +| `api_base` | string | Override for the Gray Swan API base URL. Defaults to `https://api.grayswan.ai` or `GRAYSWAN_API_BASE`. | | `mode` | string or list | Guardrail stages (`pre_call`, `during_call`, `post_call`). | | `optional_params.on_flagged_action` | string | `monitor` (log only), `block` (raise `HTTPException`), or `passthrough` (replace response content with violation message, no 400 error). | -| `.optional_params.violation_threshold`| number (0-1) | Scores at or above this value are considered violations. | +| `optional_params.violation_threshold` | number (0-1) | Scores at or above this value are considered violations. | | `optional_params.reasoning_mode` | string | `off`, `hybrid`, or `thinking`. Enables Cygnal's reasoning capabilities. | | `optional_params.categories` | object | Map of custom category names to descriptions. | | `optional_params.policy_id` | string | Gray Swan policy identifier. | +| `guardrail_timeout` | number | Timeout in seconds for the Cygnal request. Defaults to 30. | +| `fail_open` | boolean | If true, errors contacting Cygnal are logged and the request proceeds; if false, errors propagate. Defaults to treu. | +| `streaming_end_of_stream_only` | boolean | For streaming `post_call`, only send the final assembled response to Cygnal. Defaults to false. | +| `default_on` | boolean | Run the guardrail on every request by default. | diff --git a/docs/my-website/docs/realtime.md b/docs/my-website/docs/realtime.md index f4627c78da3..b191c82c670 100644 --- a/docs/my-website/docs/realtime.md +++ b/docs/my-website/docs/realtime.md @@ -3,11 +3,12 @@ import TabItem from '@theme/TabItem'; # /realtime -Use this to loadbalance across Azure + OpenAI. +Use this to loadbalance across Azure + OpenAI + xAI and more. Supported Providers: - OpenAI - Azure +- xAI ([see full docs](/docs/providers/xai_realtime)) - Google AI Studio (Gemini) - Vertex AI - Bedrock @@ -46,6 +47,21 @@ model_list: api_key: os.environ/OPENAI_API_KEY ``` + + + +```yaml +model_list: + - model_name: grok-voice-agent + litellm_params: + model: xai/grok-4-1-fast-non-reasoning + api_key: os.environ/XAI_API_KEY + model_info: + mode: realtime +``` + +**[See full xAI Realtime documentation →](/docs/providers/xai_realtime)** + diff --git a/docs/my-website/docs/tutorials/copilotkit_sdk.md b/docs/my-website/docs/tutorials/copilotkit_sdk.md new file mode 100644 index 00000000000..fc4db8bfe3e --- /dev/null +++ b/docs/my-website/docs/tutorials/copilotkit_sdk.md @@ -0,0 +1,99 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# CopilotKit SDK with LiteLLM + +Use CopilotKit SDK with any LLM provider through LiteLLM Proxy. + +> **Note:** CopilotKit SDK integration with LiteLLM Proxy works with LiteLLM v1.81.7-nightly or higher. + + +## Quick Start + +### 1. Add Model to Config + +```yaml title="config.yaml" +model_list: + - model_name: claude-sonnet-4-5 + litellm_params: + model: "anthropic/claude-sonnet-4-5-20250514-v1:0" + api_key: "os.environ/ANTHROPIC_API_KEY" +``` + +### 2. Start LiteLLM Proxy + +```bash +litellm --config config.yaml +``` + +### 3. Use CopilotKit SDK + +```typescript +import OpenAI from "openai"; +import { + CopilotRuntime, + OpenAIAdapter, + copilotRuntimeNextJSAppRouterEndpoint, +} from "@copilotkit/runtime"; +import { NextRequest } from "next/server"; + +const model = "claude-sonnet-4-5"; + +const openai = new OpenAI({ + apiKey: process.env.OPENAI_API_KEY || "sk-12345", + baseURL: process.env.OPENAI_BASE_URL || "http://localhost:4000/v1", +}); + +const serviceAdapter = new OpenAIAdapter({ openai, model }); +const runtime = new CopilotRuntime(); + +export const POST = async (req: NextRequest) => { + const { handleRequest } = copilotRuntimeNextJSAppRouterEndpoint({ + runtime, + serviceAdapter, + endpoint: "/api/copilotkit", + }); + return handleRequest(req); +}; +``` + +### 4. Test + +```bash +curl -X POST http://localhost:3000/api/copilotkit \ + -H "Content-Type: application/json" \ + -d '{ + "method": "agent/run", + "params": { + "agentId": "default" + }, + "runId": "your_run_id", + "threadId": "your_thread_id", + "runId": ""your_run_id"", + "tools": [], + "context": [], + "forwardedProps": {}, + "state": {}, + "messages": [ + { + "id": "166e573e-f7c6-4c0f-8685-04dbefec18be", + "content": "Hi", + "role": "user" + } + ] + } +}' +``` + +## Environment Variables + +| Variable | Value | Description | +|----------|-------|-------------| +| `OPENAI_API_KEY` | `sk-12345` | Your LiteLLM API key | +| `OPENAI_BASE_URL` | `http://localhost:4000/v1` | LiteLLM proxy URL | + + +## Related Resources + +- [CopilotKit Documentation](https://docs.copilotkit.ai) +- [LiteLLM Proxy Quick Start](../proxy/quick_start) diff --git a/docs/my-website/docs/tutorials/livekit_xai_realtime.md b/docs/my-website/docs/tutorials/livekit_xai_realtime.md new file mode 100644 index 00000000000..1d70186382f --- /dev/null +++ b/docs/my-website/docs/tutorials/livekit_xai_realtime.md @@ -0,0 +1,190 @@ +import Tabs from '@theme/Tabs'; +import TabItem from '@theme/TabItem'; + +# LiveKit xAI Realtime Voice Agent + +Use LiveKit's xAI Grok Voice Agent plugin with LiteLLM Proxy to build low-latency voice AI agents. + +The LiveKit Agents framework provides tools for building real-time voice and video AI applications. By routing through LiteLLM Proxy, you get unified access to multiple realtime voice providers, cost tracking, rate limiting, and more. + +## Quick Start + +### 1. Install Dependencies + +```bash +pip install livekit-agents[xai] +``` + +### 2. Start LiteLLM Proxy + +Create a config file with your xAI realtime model: + +```yaml title="config.yaml" showLineNumbers +model_list: + - model_name: grok-voice-agent + litellm_params: + model: xai/grok-2-vision-1212 + api_key: os.environ/XAI_API_KEY + model_info: + mode: realtime + +litellm_settings: + drop_params: True + +general_settings: + master_key: sk-1234 # Change this to a secure key +``` + +Start the proxy: + +```bash +litellm --config config.yaml --port 4000 +``` + +### 3. Configure LiveKit xAI Plugin + +Point LiveKit's xAI plugin to your LiteLLM proxy: + +```python +from livekit.plugins import xai + +# Configure xAI to use LiteLLM proxy +model = xai.realtime.RealtimeModel( + voice="ara", # Voice option + api_key="sk-1234", # Your LiteLLM proxy master key + base_url="http://localhost:4000", # LiteLLM proxy URL +) +``` + +## Complete Example + +Here's a complete working example: + + + + +```python +#!/usr/bin/env python3 +""" +Simple xAI realtime voice agent through LiteLLM proxy. +""" +import asyncio +import json +import websockets + +PROXY_URL = "ws://localhost:4000/v1/realtime" +API_KEY = "sk-1234" +MODEL = "grok-voice-agent" + +async def run_voice_agent(): + """Connect to xAI realtime API through LiteLLM proxy""" + url = f"{PROXY_URL}?model={MODEL}" + headers = {"Authorization": f"Bearer {API_KEY}"} + + async with websockets.connect(url, extra_headers=headers) as ws: + # Wait for initial connection event + initial = json.loads(await ws.recv()) + print(f"āœ… Connected: {initial['type']}") + + # Send user message + await ws.send(json.dumps({ + "type": "conversation.item.create", + "item": { + "type": "message", + "role": "user", + "content": [{ + "type": "input_text", + "text": "Hello! Tell me a joke." + }] + } + })) + + # Request response + await ws.send(json.dumps({ + "type": "response.create", + "response": {"modalities": ["text", "audio"]} + })) + + # Collect response + transcript = [] + async for message in ws: + event = json.loads(message) + + # Capture text response + if event['type'] == 'response.output_audio_transcript.delta': + transcript.append(event['delta']) + print(event['delta'], end='', flush=True) + + # Done when response completes + elif event['type'] == 'response.done': + break + + print(f"\n\nāœ… Full response: {''.join(transcript)}") + +if __name__ == "__main__": + asyncio.run(run_voice_agent()) +``` + + + + + +```python +from livekit.agents import Agent, AgentSession, WorkerOptions, cli +from livekit.plugins import xai + +class VoiceAgent(Agent): + def __init__(self): + super().__init__( + instructions="You are a helpful voice assistant.", + llm=xai.realtime.RealtimeModel( + voice="ara", + api_key="sk-1234", + base_url="http://localhost:4000", + ), + ) + +if __name__ == "__main__": + cli.run_app( + WorkerOptions( + agent_factory=VoiceAgent, + ) + ) +``` + + + + +## Running the Example + +1. **Start LiteLLM Proxy** (if not already running): + ```bash + litellm --config config.yaml --port 4000 + ``` + +2. **Run the example**: + ```bash + python your_script.py + ``` + +## Expected Output + +``` +āœ… Connected: conversation.created +Hello! Here's a joke for you: Why don't scientists trust atoms? +Because they make up everything! + +āœ… Full response: Hello! Here's a joke for you: Why don't scientists trust atoms? Because they make up everything! +``` + + +## Complete Working Example + +**[LiveKit Agent SDK Cookbook](https://github.com/BerriAI/litellm/tree/main/cookbook/livekit_agent_sdk)** + + +## Learn More + +- [xAI Realtime API](/docs/providers/xai_realtime) +- [LiveKit xAI Plugin](https://docs.livekit.io/agents/models/realtime/plugins/xai/) +- [LiteLLM Realtime API](/docs/realtime) diff --git a/docs/my-website/img/okta_access_policies.png b/docs/my-website/img/okta_access_policies.png new file mode 100644 index 00000000000..e09adc2ce7f Binary files /dev/null and b/docs/my-website/img/okta_access_policies.png differ diff --git a/docs/my-website/img/okta_authorization_server.png b/docs/my-website/img/okta_authorization_server.png new file mode 100644 index 00000000000..bddb3e07a4a Binary files /dev/null and b/docs/my-website/img/okta_authorization_server.png differ diff --git a/docs/my-website/img/okta_client_credentials.png b/docs/my-website/img/okta_client_credentials.png new file mode 100644 index 00000000000..a00a9f4657e Binary files /dev/null and b/docs/my-website/img/okta_client_credentials.png differ diff --git a/docs/my-website/img/okta_redirect_uri.png b/docs/my-website/img/okta_redirect_uri.png new file mode 100644 index 00000000000..a1e58560c72 Binary files /dev/null and b/docs/my-website/img/okta_redirect_uri.png differ diff --git a/docs/my-website/img/okta_security_api.png b/docs/my-website/img/okta_security_api.png new file mode 100644 index 00000000000..7f9e218074c Binary files /dev/null and b/docs/my-website/img/okta_security_api.png differ diff --git a/docs/my-website/sidebars.js b/docs/my-website/sidebars.js index d932b6af250..fda0e3be4e4 100644 --- a/docs/my-website/sidebars.js +++ b/docs/my-website/sidebars.js @@ -79,6 +79,7 @@ const sidebars = { "proxy/guardrails/panw_prisma_airs", "proxy/guardrails/secret_detection", "proxy/guardrails/custom_guardrail", + "proxy/guardrails/custom_code_guardrail", "proxy/guardrails/prompt_injection", "proxy/guardrails/tool_permission", "proxy/guardrails/zscaler_ai_guard", @@ -150,7 +151,9 @@ const sidebars = { }, items: [ "tutorials/claude_agent_sdk", + "tutorials/copilotkit_sdk", "tutorials/google_adk", + "tutorials/livekit_xai_realtime", ] }, @@ -469,6 +472,7 @@ const sidebars = { label: "/a2a - A2A Agent Gateway", items: [ "a2a", + "a2a_invoking_agents", "a2a_cost_tracking", "a2a_agent_permissions" ], @@ -850,7 +854,14 @@ const sidebars = { "providers/watsonx/audio_transcription", ] }, - "providers/xai", + { + type: "category", + label: "xAI", + items: [ + "providers/xai", + "providers/xai_realtime", + ] + }, "providers/xiaomi_mimo", "providers/xinference", "providers/zai", diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index d4ee4042b1a..bb25e4f0626 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -53,7 +53,7 @@ class CheckBatchCost: jobs = await self.prisma_client.db.litellm_managedobjecttable.find_many( where={ - "status": "validating", + "status": {"in": ["validating", "in_progress", "finalizing"]}, "file_purpose": "batch", } ) diff --git a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py index 5ee3372cca7..569ea17f6d8 100644 --- a/enterprise/litellm_enterprise/proxy/hooks/managed_files.py +++ b/enterprise/litellm_enterprise/proxy/hooks/managed_files.py @@ -166,7 +166,11 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): "updated_by": user_api_key_dict.user_id, "status": file_object.status, }, - "update": {}, # don't do anything if it already exists + "update": { + "file_object": file_object.model_dump_json(), + "status": file_object.status, + "updated_by": user_api_key_dict.user_id, + }, # FIX: Update status and file_object on every operation to keep state in sync }, ) @@ -354,6 +358,31 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) return False + async def check_file_ids_access( + self, file_ids: List[str], user_api_key_dict: UserAPIKeyAuth + ) -> None: + """ + Check if the user has access to a list of file IDs. + Only checks managed (unified) file IDs. + + Args: + file_ids: List of file IDs to check access for + user_api_key_dict: User API key authentication details + + Raises: + HTTPException: If user doesn't have access to any of the files + """ + for file_id in file_ids: + is_unified_file_id = _is_base64_encoded_unified_file_id(file_id) + if is_unified_file_id: + if not await self.can_user_call_unified_file_id( + file_id, user_api_key_dict + ): + raise HTTPException( + status_code=403, + detail=f"User {user_api_key_dict.user_id} does not have access to the file {file_id}", + ) + async def async_pre_call_hook( # noqa: PLR0915 self, user_api_key_dict: UserAPIKeyAuth, @@ -387,6 +416,9 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): if messages: file_ids = self.get_file_ids_from_messages(messages) if file_ids: + # Check user has access to all managed files + await self.check_file_ids_access(file_ids, user_api_key_dict) + # Check if any files are stored in storage backends and need base64 conversion # This is needed for Vertex AI/Gemini which requires base64 content is_vertex_ai = model and ("vertex_ai" in model or "gemini" in model.lower()) @@ -402,15 +434,27 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): ) data["model_file_id_mapping"] = model_file_id_mapping elif call_type == CallTypes.aresponses.value or call_type == CallTypes.responses.value: - # Handle managed files in responses API input + # Handle managed files in responses API input and tools + file_ids = [] + + # Extract file IDs from input parameter input_data = data.get("input") if input_data: - file_ids = self.get_file_ids_from_responses_input(input_data) - if file_ids: - model_file_id_mapping = await self.get_model_file_id_mapping( - file_ids, user_api_key_dict.parent_otel_span - ) - data["model_file_id_mapping"] = model_file_id_mapping + file_ids.extend(self.get_file_ids_from_responses_input(input_data)) + + # Extract file IDs from tools parameter (e.g., code_interpreter container) + tools = data.get("tools") + if tools: + file_ids.extend(self.get_file_ids_from_responses_tools(tools)) + + if file_ids: + # Check user has access to all managed files + await self.check_file_ids_access(file_ids, user_api_key_dict) + + model_file_id_mapping = await self.get_model_file_id_mapping( + file_ids, user_api_key_dict.parent_otel_span + ) + data["model_file_id_mapping"] = model_file_id_mapping elif call_type == CallTypes.afile_content.value: retrieve_file_id = cast(Optional[str], data.get("file_id")) potential_file_id = ( @@ -460,8 +504,6 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): if retrieve_object_id else False ) - print(f"šŸ”„potential_llm_object_id: {potential_llm_object_id}") - print(f"šŸ”„retrieve_object_id: {retrieve_object_id}") if potential_llm_object_id and retrieve_object_id: ## VALIDATE USER HAS ACCESS TO THE OBJECT ## if not await self.can_user_call_unified_object_id( @@ -614,6 +656,41 @@ class _PROXY_LiteLLMManagedFiles(CustomLogger, BaseFileEndpoints): return file_ids + def get_file_ids_from_responses_tools( + self, tools: List[Dict[str, Any]] + ) -> List[str]: + """ + Gets file ids from responses API tools parameter. + + The tools can contain code_interpreter with container.file_ids: + [ + { + "type": "code_interpreter", + "container": {"type": "auto", "file_ids": ["file-123", "file-456"]} + } + ] + """ + file_ids: List[str] = [] + + if not isinstance(tools, list): + return file_ids + + for tool in tools: + if not isinstance(tool, dict): + continue + + # Check for code_interpreter with container file_ids + if tool.get("type") == "code_interpreter": + container = tool.get("container") + if isinstance(container, dict): + container_file_ids = container.get("file_ids") + if isinstance(container_file_ids, list): + for file_id in container_file_ids: + if isinstance(file_id, str): + file_ids.append(file_id) + + return file_ids + async def get_model_file_id_mapping( self, file_ids: List[str], litellm_parent_otel_span: Span ) -> dict: diff --git a/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30-py3-none-any.whl b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30-py3-none-any.whl new file mode 100644 index 00000000000..383f9b7b43f Binary files /dev/null and b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30-py3-none-any.whl differ diff --git a/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30.tar.gz b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30.tar.gz new file mode 100644 index 00000000000..484c28ba7b1 Binary files /dev/null and b/litellm-proxy-extras/dist/litellm_proxy_extras-0.4.30.tar.gz differ diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260205091235_allow_team_guardrail_config/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260205091235_allow_team_guardrail_config/migration.sql new file mode 100644 index 00000000000..000b96b3b87 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260205091235_allow_team_guardrail_config/migration.sql @@ -0,0 +1,6 @@ +-- AlterTable +ALTER TABLE "LiteLLM_DeletedTeamTable" ADD COLUMN "allow_team_guardrail_config" BOOLEAN NOT NULL DEFAULT false; + +-- AlterTable +ALTER TABLE "LiteLLM_TeamTable" ADD COLUMN "allow_team_guardrail_config" BOOLEAN NOT NULL DEFAULT false; + diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index b118400b620..dc49036cb15 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -129,6 +129,7 @@ model LiteLLM_TeamTable { team_member_permissions String[] @default([]) policies String[] @default([]) model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases + allow_team_guardrail_config Boolean @default(false) // if true, team admin can configure guardrails for this team litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id]) litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id]) object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id]) @@ -160,7 +161,8 @@ model LiteLLM_DeletedTeamTable { team_member_permissions String[] @default([]) policies String[] @default([]) model_id Int? // id for LiteLLM_ModelTable -> stores team-level model aliases - + allow_team_guardrail_config Boolean @default(false) + // Original timestamps from team creation/updates created_at DateTime? @map("created_at") updated_at DateTime? @map("updated_at") @@ -774,6 +776,7 @@ model LiteLLM_GuardrailsTable { guardrail_name String @unique litellm_params Json guardrail_info Json? + team_id String? created_at DateTime @default(now()) updated_at DateTime @updatedAt } diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index fb6996b71db..d43b591686c 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm-proxy-extras" -version = "0.4.29" +version = "0.4.30" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." authors = ["BerriAI"] readme = "README.md" @@ -22,7 +22,7 @@ requires = ["poetry-core"] build-backend = "poetry.core.masonry.api" [tool.commitizen] -version = "0.4.29" +version = "0.4.30" version_files = [ "pyproject.toml:version", "../requirements.txt:litellm-proxy-extras==", diff --git a/litellm/__init__.py b/litellm/__init__.py index f857e10eed3..8174b9d2655 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -261,6 +261,8 @@ extra_spend_tag_headers: Optional[List[str]] = None in_memory_llm_clients_cache: "LLMClientCache" safe_memory_mode: bool = False enable_azure_ad_token_refresh: Optional[bool] = False +# Proxy Authentication - auto-obtain/refresh OAuth2/JWT tokens for LiteLLM Proxy +proxy_auth: Optional[Any] = None ### DEFAULT AZURE API VERSION ### AZURE_DEFAULT_API_VERSION = "2025-02-01-preview" # this is updated to the latest ### DEFAULT WATSONX API VERSION ### diff --git a/litellm/completion_extras/litellm_responses_transformation/transformation.py b/litellm/completion_extras/litellm_responses_transformation/transformation.py index 57bd05124aa..753a94295b3 100644 --- a/litellm/completion_extras/litellm_responses_transformation/transformation.py +++ b/litellm/completion_extras/litellm_responses_transformation/transformation.py @@ -329,6 +329,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge): else: request_data[key] = value + if headers: + request_data["extra_headers"] = headers + return request_data @staticmethod diff --git a/litellm/constants.py b/litellm/constants.py index 6427c367924..872ad899f84 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -81,6 +81,11 @@ MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH = int( os.getenv("MAX_MCP_SEMANTIC_FILTER_TOOLS_HEADER_LENGTH", 150) ) +LITELLM_UI_ALLOW_HEADERS = [ + "x-litellm-semantic-filter", + "x-litellm-semantic-filter-tools", +] + # Gemini model-specific minimal thinking budget constants DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH = int( os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH", 1) @@ -99,6 +104,9 @@ DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET = int( os.getenv("DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET", 128) ) +# Provider-specific API base URLs +XAI_API_BASE = "https://api.x.ai/v1" + DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET = int( os.getenv("DEFAULT_REASONING_EFFORT_LOW_THINKING_BUDGET", 1024) ) diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index a5bb530fc56..1652ec2aa0c 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -475,11 +475,18 @@ class CustomGuardrail(CustomLogger): guardrail_config: DynamicGuardrailParams = DynamicGuardrailParams( **guardrail[self.guardrail_name] ) + extra_body = guardrail_config.get("extra_body", {}) if self._validate_premium_user() is not True: + if isinstance(extra_body, dict) and extra_body: + verbose_logger.warning( + "Guardrail %s: ignoring dynamic extra_body keys %s because premium_user is False", + self.guardrail_name, + list(extra_body.keys()), + ) return {} # Return the extra_body if it exists, otherwise empty dict - return guardrail_config.get("extra_body", {}) + return extra_body return {} diff --git a/litellm/integrations/langfuse/langfuse_otel.py b/litellm/integrations/langfuse/langfuse_otel.py index 08493a0e8ec..8955d3619f7 100644 --- a/litellm/integrations/langfuse/langfuse_otel.py +++ b/litellm/integrations/langfuse/langfuse_otel.py @@ -8,9 +8,8 @@ from litellm.integrations.arize import _utils from litellm.integrations.langfuse.langfuse_otel_attributes import ( LangfuseLLMObsOTELAttributes, ) -from litellm.integrations.opentelemetry import OpenTelemetry +from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig from litellm.types.integrations.langfuse_otel import ( - LangfuseOtelConfig, LangfuseSpanAttributes, ) from litellm.types.utils import StandardCallbackDynamicParams @@ -18,17 +17,8 @@ from litellm.types.utils import StandardCallbackDynamicParams if TYPE_CHECKING: from opentelemetry.trace import Span as _Span - from litellm.integrations.opentelemetry import ( - OpenTelemetryConfig as _OpenTelemetryConfig, - ) - from litellm.types.integrations.arize import Protocol as _Protocol - - Protocol = _Protocol - OpenTelemetryConfig = _OpenTelemetryConfig Span = Union[_Span, Any] else: - Protocol = Any - OpenTelemetryConfig = Any Span = Any @@ -37,8 +27,12 @@ LANGFUSE_CLOUD_US_ENDPOINT = "https://us.cloud.langfuse.com/api/public/otel" class LangfuseOtelLogger(OpenTelemetry): - def __init__(self, *args, **kwargs): - super().__init__(*args, **kwargs) + def __init__(self, config=None, *args, **kwargs): + # Prevent LangfuseOtelLogger from modifying global environment variables by constructing config manually + # and passing it to the parent OpenTelemetry class + if config is None: + config = self._create_open_telemetry_config_from_langfuse_env() + super().__init__(config=config, *args, **kwargs) @staticmethod def set_langfuse_otel_attributes(span: Span, kwargs, response_obj): @@ -114,6 +108,10 @@ class LangfuseOtelLogger(OpenTelemetry): for key, enum_attr in mapping.items(): if key in metadata and metadata[key] is not None: value = metadata[key] + if key == "trace_id" and isinstance(value, str): + # trace_id must be 32 hex char no dashes for langfuse : Litellm sends uuid with dashes (might be breaking at some point) + value = value.replace("-", "") + if isinstance(value, (list, dict)): try: value = json.dumps(value) @@ -265,8 +263,47 @@ class LangfuseOtelLogger(OpenTelemetry): """ return os.environ.get("LANGFUSE_OTEL_HOST") or os.environ.get("LANGFUSE_HOST") + def _create_open_telemetry_config_from_langfuse_env(self) -> OpenTelemetryConfig: + """ + Creates OpenTelemetryConfig from Langfuse environment variables. + Does NOT modify global environment variables. + """ + from litellm.integrations.opentelemetry import OpenTelemetryConfig + + public_key = os.environ.get("LANGFUSE_PUBLIC_KEY", None) + secret_key = os.environ.get("LANGFUSE_SECRET_KEY", None) + + if not public_key or not secret_key: + # If no keys, return default from env (likely logging to console or something else) + return OpenTelemetryConfig.from_env() + + # Determine endpoint - default to US cloud + langfuse_host = LangfuseOtelLogger._get_langfuse_otel_host() + + if langfuse_host: + # If LANGFUSE_HOST is provided, construct OTEL endpoint from it + if not langfuse_host.startswith("http"): + langfuse_host = "https://" + langfuse_host + endpoint = f"{langfuse_host.rstrip('/')}/api/public/otel" + verbose_logger.debug(f"Using Langfuse OTEL endpoint from host: {endpoint}") + else: + # Default to US cloud endpoint + endpoint = LANGFUSE_CLOUD_US_ENDPOINT + verbose_logger.debug(f"Using Langfuse US cloud endpoint: {endpoint}") + + auth_header = LangfuseOtelLogger._get_langfuse_authorization_header( + public_key=public_key, secret_key=secret_key + ) + otlp_auth_headers = f"Authorization={auth_header}" + + return OpenTelemetryConfig( + exporter="otlp_http", + endpoint=endpoint, + headers=otlp_auth_headers, + ) + @staticmethod - def get_langfuse_otel_config() -> LangfuseOtelConfig: + def get_langfuse_otel_config() -> "OpenTelemetryConfig": """ Retrieves the Langfuse OpenTelemetry configuration based on environment variables. @@ -276,7 +313,7 @@ class LangfuseOtelLogger(OpenTelemetry): LANGFUSE_HOST: Optional. Custom Langfuse host URL. Defaults to US cloud. Returns: - LangfuseOtelConfig: A Pydantic model containing Langfuse OTEL configuration. + OpenTelemetryConfig: A Pydantic model containing Langfuse OTEL configuration. Raises: ValueError: If required keys are missing. @@ -308,12 +345,14 @@ class LangfuseOtelLogger(OpenTelemetry): ) otlp_auth_headers = f"Authorization={auth_header}" - # Set standard OTEL environment variables - os.environ["OTEL_EXPORTER_OTLP_ENDPOINT"] = endpoint - os.environ["OTEL_EXPORTER_OTLP_HEADERS"] = otlp_auth_headers + # Prevent modification of global env vars which causes leakage + # os.environ["OTEL_EXPORTER_OTLP_ENDPOINT"] = endpoint + # os.environ["OTEL_EXPORTER_OTLP_HEADERS"] = otlp_auth_headers - return LangfuseOtelConfig( - otlp_auth_headers=otlp_auth_headers, protocol="otlp_http" + return OpenTelemetryConfig( + exporter="otlp_http", + endpoint=endpoint, + headers=otlp_auth_headers, ) @staticmethod diff --git a/litellm/integrations/opentelemetry.py b/litellm/integrations/opentelemetry.py index 18898be7dce..296a88f9a0b 100644 --- a/litellm/integrations/opentelemetry.py +++ b/litellm/integrations/opentelemetry.py @@ -599,9 +599,9 @@ class OpenTelemetry(CustomLogger): def _get_dynamic_otel_headers_from_kwargs(self, kwargs) -> Optional[dict]: """Extract dynamic headers from kwargs if available.""" - standard_callback_dynamic_params: Optional[StandardCallbackDynamicParams] = ( - kwargs.get("standard_callback_dynamic_params") - ) + standard_callback_dynamic_params: Optional[ + StandardCallbackDynamicParams + ] = kwargs.get("standard_callback_dynamic_params") if not standard_callback_dynamic_params: return None @@ -619,7 +619,9 @@ class OpenTelemetry(CustomLogger): # Prevents thread exhaustion by reusing providers for the same credential sets (e.g. per-team keys) cache_key = str(sorted(dynamic_headers.items())) if cache_key in self._tracer_provider_cache: - return self._tracer_provider_cache[cache_key].get_tracer(LITELLM_TRACER_NAME) + return self._tracer_provider_cache[cache_key].get_tracer( + LITELLM_TRACER_NAME + ) # Create a temporary tracer provider with dynamic headers temp_provider = TracerProvider(resource=self._get_litellm_resource(self.config)) @@ -674,7 +676,10 @@ class OpenTelemetry(CustomLogger): kwargs, response_obj, start_time, end_time, span ) # Ensure proxy-request parent span is annotated with the actual operation kind - if parent_span is not None and parent_span.name == LITELLM_PROXY_REQUEST_SPAN_NAME: + if ( + parent_span is not None + and parent_span.name == LITELLM_PROXY_REQUEST_SPAN_NAME + ): self.set_attributes(parent_span, kwargs, response_obj) else: # Do not create primary span (keep hierarchy shallow when parent exists) @@ -1003,14 +1008,11 @@ class OpenTelemetry(CustomLogger): # TODO: Refactor to use the proper OTEL Logs API instead of directly creating SDK LogRecords from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider + try: - from opentelemetry.sdk._logs import ( - LogRecord as SdkLogRecord, # type: ignore[attr-defined] # OTEL < 1.39.0 - ) + from opentelemetry.sdk._logs import LogRecord as SdkLogRecord # type: ignore[attr-defined] # OTEL < 1.39.0 except ImportError: - from opentelemetry.sdk._logs._internal import ( - LogRecord as SdkLogRecord, # OTEL >= 1.39.0 - ) + from opentelemetry.sdk._logs._internal import LogRecord as SdkLogRecord # type: ignore[attr-defined, no-redef] # OTEL >= 1.39.0 otel_logger = get_logger(LITELLM_LOGGER_NAME) @@ -1618,7 +1620,6 @@ class OpenTelemetry(CustomLogger): for idx, choice in enumerate(response_obj.get("choices")): if choice.get("finish_reason"): - message = choice.get("message") tool_calls = message.get("tool_calls") if tool_calls: @@ -1631,7 +1632,9 @@ class OpenTelemetry(CustomLogger): ) except Exception as e: - self.handle_callback_failure(callback_name=self.callback_name or "opentelemetry") + self.handle_callback_failure( + callback_name=self.callback_name or "opentelemetry" + ) verbose_logger.exception( "OpenTelemetry logging error in set_attributes %s", str(e) ) @@ -1722,6 +1725,7 @@ class OpenTelemetry(CustomLogger): def set_raw_request_attributes(self, span: Span, kwargs, response_obj): try: + self.set_attributes(span, kwargs, response_obj) kwargs.get("optional_params", {}) litellm_params = kwargs.get("litellm_params", {}) or {} custom_llm_provider = litellm_params.get("custom_llm_provider", "Unknown") diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 2c897cb0692..00c38eac188 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -1683,6 +1683,108 @@ class PrometheusLogger(CustomLogger): ) pass + def _safe_get(self, obj: Any, key: str, default: Any = None) -> Any: + """Get value from dict or Pydantic model.""" + if obj is None: + return default + if isinstance(obj, dict): + return obj.get(key, default) + return getattr(obj, key, default) + + def _extract_deployment_failure_label_values( + self, request_kwargs: dict + ) -> Dict[str, Optional[str]]: + """ + Extract label values for deployment failure metrics from all available + sources in request_kwargs. Falls back to litellm_params metadata and + user_api_key_auth when standard_logging_payload has None values. + """ + standard_logging_payload = ( + request_kwargs.get("standard_logging_object", {}) or {} + ) + _litellm_params = request_kwargs.get("litellm_params", {}) or {} + _metadata_raw = self._safe_get(standard_logging_payload, "metadata") or {} + if isinstance(_metadata_raw, dict): + _metadata = _metadata_raw + else: + _metadata = { + "user_api_key_alias": getattr( + _metadata_raw, "user_api_key_alias", None + ), + "user_api_key_team_id": getattr( + _metadata_raw, "user_api_key_team_id", None + ), + "user_api_key_team_alias": getattr( + _metadata_raw, "user_api_key_team_alias", None + ), + "user_api_key_hash": getattr(_metadata_raw, "user_api_key_hash", None), + "requester_ip_address": getattr( + _metadata_raw, "requester_ip_address", None + ), + "user_agent": getattr(_metadata_raw, "user_agent", None), + } + _litellm_params_metadata = _litellm_params.get("metadata", {}) or {} + + # Extract user_api_key_auth if present (proxy injects this, skipped in merge) + user_api_key_auth = _litellm_params_metadata.get("user_api_key_auth") + + def _get_api_key_alias() -> Optional[str]: + val = _metadata.get("user_api_key_alias") + if val is not None: + return val + val = _litellm_params_metadata.get("user_api_key_alias") + if val is not None: + return val + if user_api_key_auth is not None: + return getattr(user_api_key_auth, "key_alias", None) + return None + + def _get_team_id() -> Optional[str]: + val = _metadata.get("user_api_key_team_id") + if val is not None: + return val + val = _litellm_params_metadata.get("user_api_key_team_id") + if val is not None: + return val + if user_api_key_auth is not None: + return getattr(user_api_key_auth, "team_id", None) + return None + + def _get_team_alias() -> Optional[str]: + val = _metadata.get("user_api_key_team_alias") + if val is not None: + return val + val = _litellm_params_metadata.get("user_api_key_team_alias") + if val is not None: + return val + if user_api_key_auth is not None: + return getattr(user_api_key_auth, "team_alias", None) + return None + + def _get_hashed_api_key() -> Optional[str]: + val = _metadata.get("user_api_key_hash") + if val is not None: + return val + val = _litellm_params_metadata.get("user_api_key_hash") + if val is not None: + return val + if user_api_key_auth is not None: + return getattr(user_api_key_auth, "api_key", None) or getattr( + user_api_key_auth, "api_key_hash", None + ) + return None + + return { + "api_key_alias": _get_api_key_alias(), + "team": _get_team_id(), + "team_alias": _get_team_alias(), + "hashed_api_key": _get_hashed_api_key(), + "client_ip": _metadata.get("requester_ip_address") + or _litellm_params_metadata.get("requester_ip_address"), + "user_agent": _metadata.get("user_agent") + or _litellm_params_metadata.get("user_agent"), + } + def set_llm_deployment_failure_metrics(self, request_kwargs: dict): """ Sets Failure metrics when an LLM API call fails @@ -1707,6 +1809,21 @@ class PrometheusLogger(CustomLogger): model_id = standard_logging_payload.get("model_id", None) exception = request_kwargs.get("exception", None) + # Fallback: model_id from litellm_metadata.model_info + if model_id is None: + _model_info = ( + (_litellm_params.get("litellm_metadata") or {}).get("model_info") + or (_litellm_params.get("metadata") or {}).get("model_info") + or {} + ) + model_id = _model_info.get("id") + + # Fallback: model_group from litellm_metadata + if model_group is None: + model_group = (_litellm_params.get("litellm_metadata") or {}).get( + "model_group" + ) or (_litellm_params.get("metadata") or {}).get("model_group") + llm_provider = _litellm_params.get("custom_llm_provider", None) if self._should_skip_metrics_for_invalid_key( @@ -1714,9 +1831,37 @@ class PrometheusLogger(CustomLogger): standard_logging_payload=standard_logging_payload, ): return - hashed_api_key = standard_logging_payload.get("metadata", {}).get( + + # Extract context labels from all available sources (fix for None labels) + fallback_values = self._extract_deployment_failure_label_values( + request_kwargs + ) + _metadata = standard_logging_payload.get("metadata", {}) or {} + hashed_api_key = fallback_values.get("hashed_api_key") or _metadata.get( "user_api_key_hash" ) + api_key_alias = fallback_values.get("api_key_alias") or _metadata.get( + "user_api_key_alias" + ) + team = fallback_values.get("team") or _metadata.get("user_api_key_team_id") + team_alias = fallback_values.get("team_alias") or _metadata.get( + "user_api_key_team_alias" + ) + client_ip = fallback_values.get("client_ip") or _metadata.get( + "requester_ip_address" + ) + user_agent = fallback_values.get("user_agent") or _metadata.get( + "user_agent" + ) + + # exception_status: prefer status_code, fallback to exception class for known types + exception_status = None + if exception is not None: + exception_status = str(getattr(exception, "status_code", None)) + if exception_status == "None" or not exception_status: + code = getattr(exception, "code", None) + if code is not None: + exception_status = str(code) # Create enum_values for the label factory (always create for use in different metrics) enum_values = UserAPIKeyLabelValues( @@ -1724,26 +1869,18 @@ class PrometheusLogger(CustomLogger): model_id=model_id, api_base=api_base, api_provider=llm_provider, - exception_status=( - str(getattr(exception, "status_code", None)) if exception else None - ), + exception_status=exception_status, exception_class=( self._get_exception_class_name(exception) if exception else None ), - requested_model=model_group, + requested_model=model_group or litellm_model_name, hashed_api_key=hashed_api_key, - api_key_alias=standard_logging_payload["metadata"][ - "user_api_key_alias" - ], - team=standard_logging_payload["metadata"]["user_api_key_team_id"], - team_alias=standard_logging_payload["metadata"][ - "user_api_key_team_alias" - ], + api_key_alias=api_key_alias, + team=team, + team_alias=team_alias, tags=standard_logging_payload.get("request_tags", []), - client_ip=standard_logging_payload["metadata"].get( - "requester_ip_address" - ), - user_agent=standard_logging_payload["metadata"].get("user_agent"), + client_ip=client_ip, + user_agent=user_agent, ) """ diff --git a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py index c425319b4d4..ff521d47804 100644 --- a/litellm/litellm_core_utils/initialize_dynamic_callback_params.py +++ b/litellm/litellm_core_utils/initialize_dynamic_callback_params.py @@ -1,8 +1,35 @@ from typing import Dict, Optional - from litellm.secret_managers.main import get_secret_str from litellm.types.utils import StandardCallbackDynamicParams +# Hardcoded list of supported callback params to avoid runtime inspection issues with TypedDict +_supported_callback_params = [ + "langfuse_public_key", + "langfuse_secret", + "langfuse_secret_key", + "langfuse_host", + "langfuse_prompt_version", + "gcs_bucket_name", + "gcs_path_service_account", + "langsmith_api_key", + "langsmith_project", + "langsmith_base_url", + "langsmith_sampling_rate", + "langsmith_tenant_id", + "humanloop_api_key", + "arize_api_key", + "arize_space_key", + "arize_space_id", + "posthog_api_key", + "posthog_host", + "braintrust_api_key", + "braintrust_project", + "braintrust_host", + "slack_webhook_url", + "lunary_public_key", + "turn_off_message_logging", +] + def initialize_standard_callback_dynamic_params( kwargs: Optional[Dict] = None, @@ -15,13 +42,10 @@ def initialize_standard_callback_dynamic_params( standard_callback_dynamic_params = StandardCallbackDynamicParams() if kwargs: - _supported_callback_params = ( - StandardCallbackDynamicParams.__annotations__.keys() - ) - + # 1. Check top-level kwargs for param in _supported_callback_params: if param in kwargs: - _param_value = kwargs.pop(param) + _param_value = kwargs.get(param) if ( _param_value is not None and isinstance(_param_value, str) @@ -30,4 +54,22 @@ def initialize_standard_callback_dynamic_params( _param_value = get_secret_str(secret_name=_param_value) standard_callback_dynamic_params[param] = _param_value # type: ignore + # 2. Fallback: check "metadata" or "litellm_params" -> "metadata" + metadata = (kwargs.get("metadata") or {}).copy() + litellm_params = kwargs.get("litellm_params") or {} + if isinstance(litellm_params, dict): + metadata.update(litellm_params.get("metadata") or {}) + + if isinstance(metadata, dict): + for param in _supported_callback_params: + if param not in standard_callback_dynamic_params and param in metadata: + _param_value = metadata.get(param) + if ( + _param_value is not None + and isinstance(_param_value, str) + and "os.environ/" in _param_value + ): + _param_value = get_secret_str(secret_name=_param_value) + standard_callback_dynamic_params[param] = _param_value # type: ignore + return standard_callback_dynamic_params diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 1b3a687f1f3..14015225f38 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -3917,18 +3917,6 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 return langfuse_logger # type: ignore elif logging_integration == "langfuse_otel": from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger - from litellm.integrations.opentelemetry import ( - OpenTelemetry, - OpenTelemetryConfig, - ) - - langfuse_otel_config = LangfuseOtelLogger.get_langfuse_otel_config() - - # The endpoint and headers are now set as environment variables by get_langfuse_otel_config() - otel_config = OpenTelemetryConfig( - exporter=langfuse_otel_config.protocol, - headers=langfuse_otel_config.otlp_auth_headers, - ) for callback in _in_memory_loggers: if ( @@ -3936,8 +3924,10 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915 and callback.callback_name == "langfuse_otel" ): return callback # type: ignore + # Allow LangfuseOtelLogger to initialize its own config safely + # This prevents startup crashes if LANGFUSE keys are not in env (e.g. for dynamic usage) _otel_logger = LangfuseOtelLogger( - config=otel_config, callback_name="langfuse_otel" + config=None, callback_name="langfuse_otel" ) _in_memory_loggers.append(_otel_logger) return _otel_logger # type: ignore diff --git a/litellm/litellm_core_utils/logging_callback_manager.py b/litellm/litellm_core_utils/logging_callback_manager.py index 4f76a5bad03..435ae078a65 100644 --- a/litellm/litellm_core_utils/logging_callback_manager.py +++ b/litellm/litellm_core_utils/logging_callback_manager.py @@ -114,6 +114,27 @@ class LoggingCallbackManager: for c in remove_list: callback_list.remove(c) + def remove_callbacks_by_type(self, callback_list, callback_type): + """ + Remove all callbacks of a specific type from a callback list. + + Args: + callback_list: The list to remove callbacks from (e.g., litellm.callbacks) + callback_type: The class type to match (e.g., SemanticToolFilterHook) + + Example: + litellm.logging_callback_manager.remove_callbacks_by_type( + litellm.callbacks, SemanticToolFilterHook + ) + """ + if not isinstance(callback_list, list): + return + + remove_list = [c for c in callback_list if isinstance(c, callback_type)] + + for c in remove_list: + callback_list.remove(c) + def _add_string_callback_to_list( self, callback: str, parent_list: List[Union[CustomLogger, Callable, str]] ): diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 7790fb83361..b1c2d0a52f5 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -443,13 +443,21 @@ def update_messages_with_model_file_ids( def update_responses_input_with_model_file_ids( input: Any, + model_id: Optional[str] = None, + model_file_id_mapping: Optional[Dict[str, Dict[str, str]]] = None, ) -> Union[str, List[Dict[str, Any]]]: """ Updates responses API input with provider-specific file IDs. File IDs are always inside the content array, not as direct input_file items. - For managed files (unified file IDs), decodes the base64-encoded unified file ID - and extracts the llm_output_file_id directly. + For managed files (unified file IDs), uses model_file_id_mapping if provided, + otherwise decodes the base64-encoded unified file ID and extracts the llm_output_file_id directly. + + Args: + input: The responses API input parameter + model_id: The model ID to use for looking up provider-specific file IDs + model_file_id_mapping: Dictionary mapping litellm file IDs to provider file IDs + Format: {"litellm_file_id": {"model_id": "provider_file_id"}} """ from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, @@ -479,22 +487,35 @@ def update_responses_input_with_model_file_ids( ): file_id = content_item.get("file_id") if file_id: - # Check if this is a managed file ID (base64-encoded unified file ID) - is_unified_file_id = _is_base64_encoded_unified_file_id(file_id) - if is_unified_file_id: - unified_file_id = convert_b64_uid_to_unified_uid(file_id) - if "llm_output_file_id," in unified_file_id: - provider_file_id = unified_file_id.split( - "llm_output_file_id," - )[1].split(";")[0] - else: - # Fallback: keep original if we can't extract - provider_file_id = file_id + provider_file_id = file_id # Default to original + + # Check if we have a mapping for this file ID + if model_file_id_mapping and model_id and file_id in model_file_id_mapping: + # Use the model-specific file ID from mapping + provider_file_id = ( + model_file_id_mapping.get(file_id, {}).get(model_id) + or file_id + ) updated_content_item = content_item.copy() updated_content_item["file_id"] = provider_file_id updated_content.append(updated_content_item) else: - updated_content.append(content_item) + # Check if this is a base64-encoded unified file ID without mapping + is_unified_file_id = _is_base64_encoded_unified_file_id(file_id) + if is_unified_file_id: + # Fallback: decode unified file ID + unified_file_id = convert_b64_uid_to_unified_uid(file_id) + if "llm_output_file_id," in unified_file_id: + provider_file_id = unified_file_id.split( + "llm_output_file_id," + )[1].split(";")[0] + + updated_content_item = content_item.copy() + updated_content_item["file_id"] = provider_file_id + updated_content.append(updated_content_item) + else: + # Not a managed file, keep as-is + updated_content.append(content_item) else: updated_content.append(content_item) else: @@ -506,6 +527,68 @@ def update_responses_input_with_model_file_ids( return updated_input +def update_responses_tools_with_model_file_ids( + tools: Optional[List[Dict[str, Any]]], + model_id: Optional[str] = None, + model_file_id_mapping: Optional[Dict[str, Dict[str, str]]] = None, +) -> Optional[List[Dict[str, Any]]]: + """ + Updates responses API tools with provider-specific file IDs. + + Handles code_interpreter tools with container.file_ids. + + Args: + tools: The responses API tools parameter + model_id: The model ID to use for looking up provider-specific file IDs + model_file_id_mapping: Dictionary mapping litellm file IDs to provider file IDs + Format: {"litellm_file_id": {"model_id": "provider_file_id"}} + """ + if not tools or not isinstance(tools, list): + return tools + + if not model_file_id_mapping or not model_id: + return tools + + updated_tools = [] + for tool in tools: + if not isinstance(tool, dict): + updated_tools.append(tool) + continue + + updated_tool = tool.copy() + + # Handle code_interpreter with container file_ids + if tool.get("type") == "code_interpreter": + container = tool.get("container") + if isinstance(container, dict): + container_file_ids = container.get("file_ids") + if isinstance(container_file_ids, list): + updated_file_ids = [] + for file_id in container_file_ids: + if isinstance(file_id, str): + # Check if we have a mapping for this file ID + if file_id in model_file_id_mapping: + # Map to provider-specific file ID + provider_file_id = ( + model_file_id_mapping.get(file_id, {}).get(model_id) + or file_id + ) + updated_file_ids.append(provider_file_id) + else: + updated_file_ids.append(file_id) + else: + updated_file_ids.append(file_id) + + # Update the tool with new file IDs + updated_container = container.copy() + updated_container["file_ids"] = updated_file_ids + updated_tool["container"] = updated_container + + updated_tools.append(updated_tool) + + return updated_tools + + def extract_file_data(file_data: FileTypes) -> ExtractedFileData: """ Extracts and processes file data from various input formats. diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index c4c56a8d335..53d2ca2f23f 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -3987,10 +3987,12 @@ class BedrockConverseMessagesProcessor: assistant_parts=assistants_parts, ) elif element["type"] == "text": - assistants_part = BedrockContentBlock( - text=element["text"] - ) - assistants_parts.append(assistants_part) + # Skip completely empty strings to avoid blank content blocks + if element.get("text", "").strip(): + assistants_part = BedrockContentBlock( + text=element["text"] + ) + assistants_parts.append(assistants_part) elif element["type"] == "image_url": if isinstance(element["image_url"], dict): image_url = element["image_url"]["url"] @@ -4015,9 +4017,12 @@ class BedrockConverseMessagesProcessor: elif _assistant_content is not None and isinstance( _assistant_content, str ): - assistant_content.append( - BedrockContentBlock(text=_assistant_content) - ) + # Skip completely empty strings to avoid blank content blocks + if _assistant_content.strip(): + assistant_content.append( + BedrockContentBlock(text=_assistant_content) + ) + # If content is empty/whitespace, skip it (don't add a placeholder) # Add cache point block for assistant string content _cache_point_block = ( litellm.AmazonConverseConfig()._get_cache_point_block( @@ -4348,12 +4353,11 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 assistant_parts=assistants_parts, ) elif element["type"] == "text": - # AWS Bedrock doesn't allow empty or whitespace-only text content, so use placeholder for empty strings - text_content = ( - element["text"] if element["text"].strip() else "." - ) - assistants_part = BedrockContentBlock(text=text_content) - assistants_parts.append(assistants_part) + # AWS Bedrock doesn't allow empty or whitespace-only text content + # Skip completely empty strings to avoid blank content blocks + if element.get("text", "").strip(): + assistants_part = BedrockContentBlock(text=element["text"]) + assistants_parts.append(assistants_part) elif element["type"] == "image_url": if isinstance(element["image_url"], dict): image_url = element["image_url"]["url"] @@ -4376,9 +4380,9 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915 assistants_parts.append(_cache_point_block) assistant_content.extend(assistants_parts) elif _assistant_content is not None and isinstance(_assistant_content, str): - # AWS Bedrock doesn't allow empty or whitespace-only text content, so use placeholder for empty strings - text_content = _assistant_content if _assistant_content.strip() else "." - assistant_content.append(BedrockContentBlock(text=text_content)) + # Skip completely empty strings to avoid blank content blocks + if _assistant_content.strip(): + assistant_content.append(BedrockContentBlock(text=_assistant_content)) # Add cache point block for assistant string content _cache_point_block = ( litellm.AmazonConverseConfig()._get_cache_point_block( diff --git a/litellm/litellm_core_utils/redact_messages.py b/litellm/litellm_core_utils/redact_messages.py index 0effed3db70..aa763dc9899 100644 --- a/litellm/litellm_core_utils/redact_messages.py +++ b/litellm/litellm_core_utils/redact_messages.py @@ -130,6 +130,11 @@ def perform_redaction(model_call_details: dict, result): def should_redact_message_logging(model_call_details: dict) -> bool: """ Determine if message logging should be redacted. + + Priority order: + 1. Dynamic parameter (turn_off_message_logging in request) + 2. Headers (litellm-disable-message-redaction / litellm-enable-message-redaction) + 3. Global setting (litellm.turn_off_message_logging) """ litellm_params = model_call_details.get("litellm_params", {}) @@ -139,36 +144,36 @@ def should_redact_message_logging(model_call_details: dict) -> bool: # Get headers from the metadata request_headers = metadata.get("headers", {}) if isinstance(metadata, dict) else {} - possible_request_headers = [ + # Check for headers that explicitly control redaction + if request_headers and bool( + request_headers.get("litellm-disable-message-redaction", False) + ): + # User explicitly disabled redaction via header + return False + + possible_enable_headers = [ "litellm-enable-message-redaction", # old header. maintain backwards compatibility "x-litellm-enable-message-redaction", # new header ] is_redaction_enabled_via_header = False - for header in possible_request_headers: + for header in possible_enable_headers: if bool(request_headers.get(header, False)): is_redaction_enabled_via_header = True break - # check if user opted out of logging message/response to callbacks - if ( - litellm.turn_off_message_logging is not True - and is_redaction_enabled_via_header is not True - and _get_turn_off_message_logging_from_dynamic_params(model_call_details) - is not True - ): - return False - - if request_headers and bool( - request_headers.get("litellm-disable-message-redaction", False) - ): - return False - - # user has OPTED OUT of message redaction - if _get_turn_off_message_logging_from_dynamic_params(model_call_details) is False: - return False - - return True + # Priority 1: Check dynamic parameter first (if explicitly set) + dynamic_turn_off = _get_turn_off_message_logging_from_dynamic_params(model_call_details) + if dynamic_turn_off is not None: + # Dynamic parameter is explicitly set, use it + return dynamic_turn_off + + # Priority 2: Check if header explicitly enables redaction + if is_redaction_enabled_via_header: + return True + + # Priority 3: Fall back to global setting + return litellm.turn_off_message_logging is True def redact_message_input_output_from_logging( diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 84b6fffaa31..4b689414ddd 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -94,7 +94,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator): if state == "completed": return "stop" elif state == "failed": - return "error" + return "stop" # Map failed state to 'stop' (valid finish_reason) # Check for [DONE] marker if chunk.get("done") is True: diff --git a/litellm/llms/a2a/chat/transformation.py b/litellm/llms/a2a/chat/transformation.py index 243fba63719..163cd5ab22e 100644 --- a/litellm/llms/a2a/chat/transformation.py +++ b/litellm/llms/a2a/chat/transformation.py @@ -2,10 +2,9 @@ A2A Protocol Transformation for LiteLLM """ import uuid -from typing import Any, Dict, Iterator, List, Optional, Union, cast +from typing import Any, Dict, Iterator, List, Optional, Union import httpx -from pydantic import BaseModel from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException @@ -27,6 +26,68 @@ class A2AConfig(BaseConfig): Handles transformation between OpenAI and A2A JSON-RPC 2.0 formats. """ + @staticmethod + def resolve_agent_config_from_registry( + model: str, + api_base: Optional[str], + api_key: Optional[str], + headers: Optional[Dict[str, Any]], + optional_params: Dict[str, Any], + ) -> tuple[Optional[str], Optional[str], Optional[Dict[str, Any]]]: + """ + Resolve agent configuration from registry if model format is "a2a/". + + Extracts agent name from model string and looks up configuration in the + agent registry (if available in proxy context). + + Args: + model: Model string (e.g., "a2a/my-agent") + api_base: Explicit api_base (takes precedence over registry) + api_key: Explicit api_key (takes precedence over registry) + headers: Explicit headers (takes precedence over registry) + optional_params: Dict to merge additional litellm_params into + + Returns: + Tuple of (api_base, api_key, headers) with registry values filled in + """ + # Extract agent name from model (e.g., "a2a/my-agent" -> "my-agent") + agent_name = model.split("/", 1)[1] if "/" in model else None + + # Only lookup if agent name exists and some config is missing + if not agent_name or (api_base is not None and api_key is not None and headers is not None): + return api_base, api_key, headers + + # Try registry lookup (only available in proxy context) + try: + from litellm.proxy.agent_endpoints.agent_registry import ( + global_agent_registry, + ) + + agent = global_agent_registry.get_agent_by_name(agent_name) + if agent: + # Get api_base from agent card URL + if api_base is None and agent.agent_card_params: + api_base = agent.agent_card_params.get("url") + + # Get api_key, headers, and other params from litellm_params + if agent.litellm_params: + if api_key is None: + api_key = agent.litellm_params.get("api_key") + + if headers is None: + agent_headers = agent.litellm_params.get("headers") + if agent_headers: + headers = agent_headers + + # Merge other litellm_params (timeout, max_retries, etc.) + for key, value in agent.litellm_params.items(): + if key not in ["api_key", "api_base", "headers", "model"] and key not in optional_params: + optional_params[key] = value + except ImportError: + pass # Registry not available (not running in proxy context) + + return api_base, api_key, headers + def get_supported_openai_params(self, model: str) -> List[str]: """Return list of supported OpenAI parameters""" return [ @@ -46,9 +107,14 @@ class A2AConfig(BaseConfig): """ Map OpenAI parameters to A2A parameters. - For A2A protocol, we don't need to map most parameters since - they're handled in the transform_request method. + For A2A protocol, we need to map the stream parameter so + transform_request can determine which JSON-RPC method to use. """ + # Map stream parameter + for param, value in non_default_params.items(): + if param == "stream" and value is True: + optional_params["stream"] = value + return optional_params def validate_environment( @@ -160,8 +226,9 @@ class A2AConfig(BaseConfig): # Build JSON-RPC 2.0 request # For A2A protocol, the method is "message/send" for non-streaming - # and "message/stream" for streaming (handled by optional_params["stream"]) - method = "message/stream" if optional_params.get("stream") else "message/send" + # and "message/stream" for streaming + stream = optional_params.get("stream", False) + method = "message/stream" if stream else "message/send" request_data = { "jsonrpc": "2.0", diff --git a/litellm/llms/a2a/common_utils.py b/litellm/llms/a2a/common_utils.py index 4c7da78b42f..116e1205409 100644 --- a/litellm/llms/a2a/common_utils.py +++ b/litellm/llms/a2a/common_utils.py @@ -113,6 +113,8 @@ def extract_text_from_a2a_response( # 1. Direct message: {"result": {"kind": "message", "parts": [...]}} # 2. Nested message: {"result": {"message": {"parts": [...]}}} # 3. Task with artifacts: {"result": {"kind": "task", "artifacts": [{"parts": [...]}]}} + # 4. Task with status message: {"result": {"kind": "task", "status": {"message": {"parts": [...]}}}} + # 5. Streaming artifact-update: {"result": {"kind": "artifact-update", "artifact": {"parts": [...]}}} # Check if result itself has parts (direct message) if "parts" in result: @@ -123,7 +125,23 @@ def extract_text_from_a2a_response( if message: return extract_text_from_a2a_message(message, depth=0, max_depth=max_depth) - # Handle task result with artifacts + # Check for streaming artifact-update (singular artifact) + artifact = result.get("artifact") + if artifact and isinstance(artifact, dict): + return extract_text_from_a2a_message( + artifact, depth=0, max_depth=max_depth + ) + + # Check for task status message (common in Gemini A2A agents) + status = result.get("status", {}) + if isinstance(status, dict): + status_message = status.get("message") + if status_message: + return extract_text_from_a2a_message( + status_message, depth=0, max_depth=max_depth + ) + + # Handle task result with artifacts (plural, array) artifacts = result.get("artifacts", []) if artifacts and len(artifacts) > 0: first_artifact = artifacts[0] diff --git a/litellm/llms/anthropic/chat/guardrail_translation/handler.py b/litellm/llms/anthropic/chat/guardrail_translation/handler.py index 8e1016bd5bd..a14e7d118e8 100644 --- a/litellm/llms/anthropic/chat/guardrail_translation/handler.py +++ b/litellm/llms/anthropic/chat/guardrail_translation/handler.py @@ -34,6 +34,7 @@ from litellm.types.llms.openai import ( ) from litellm.types.utils import ( ChatCompletionMessageToolCall, + Choices, GenericGuardrailAPIInputs, ModelResponse, ) @@ -76,7 +77,8 @@ class AnthropicMessagesHandler(BaseTranslation): chat_completion_compatible_request, tool_name_mapping = ( LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai( - anthropic_message_request=cast(AnthropicMessagesRequest, data) + # Use a shallow copy to avoid mutating request data (pop on litellm_metadata). + anthropic_message_request=cast(AnthropicMessagesRequest, data.copy()) ) ) @@ -84,9 +86,9 @@ class AnthropicMessagesHandler(BaseTranslation): texts_to_check: List[str] = [] images_to_check: List[str] = [] - tools_to_check: List[ChatCompletionToolParam] = ( - chat_completion_compatible_request.get("tools", []) - ) + tools_to_check: List[ + ChatCompletionToolParam + ] = chat_completion_compatible_request.get("tools", []) task_mappings: List[Tuple[int, Optional[int]]] = [] # Track (message_index, content_index) for each text # content_index is None for string content, int for list content @@ -282,7 +284,10 @@ class AnthropicMessagesHandler(BaseTranslation): if hasattr(content_block, "model_dump"): block_dict = content_block.model_dump() else: - block_dict = {"type": block_type, "text": getattr(content_block, "text", None)} + block_dict = { + "type": block_type, + "text": getattr(content_block, "text", None), + } else: continue @@ -358,30 +363,40 @@ class AnthropicMessagesHandler(BaseTranslation): """ has_ended = self._check_streaming_has_ended(responses_so_far) if has_ended: - # build the model response from the responses_so_far - model_response = cast( - ModelResponse, - AnthropicPassthroughLoggingHandler._build_complete_streaming_response( - all_chunks=responses_so_far, - litellm_logging_obj=cast("LiteLLMLoggingObj", litellm_logging_obj), - model="", - ), + built_response = AnthropicPassthroughLoggingHandler._build_complete_streaming_response( + all_chunks=responses_so_far, + litellm_logging_obj=cast("LiteLLMLoggingObj", litellm_logging_obj), + model="", ) - tool_calls_list = cast(Optional[List[ChatCompletionMessageToolCall]], model_response.choices[0].message.tool_calls) # type: ignore - string_so_far = model_response.choices[0].message.content # type: ignore - guardrail_inputs = GenericGuardrailAPIInputs() - if string_so_far: - guardrail_inputs["texts"] = [string_so_far] - if tool_calls_list: - guardrail_inputs["tool_calls"] = tool_calls_list - _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid - inputs=guardrail_inputs, - request_data={}, - input_type="response", - logging_obj=litellm_logging_obj, - ) + # Check if model_response is valid and has choices before accessing + if ( + built_response is not None + and hasattr(built_response, "choices") + and built_response.choices + ): + model_response = cast(ModelResponse, built_response) + first_choice = cast(Choices, model_response.choices[0]) + tool_calls_list = cast( + Optional[List[ChatCompletionMessageToolCall]], + first_choice.message.tool_calls, + ) + string_so_far = first_choice.message.content + guardrail_inputs = GenericGuardrailAPIInputs() + if string_so_far: + guardrail_inputs["texts"] = [string_so_far] + if tool_calls_list: + guardrail_inputs["tool_calls"] = tool_calls_list + + _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid + inputs=guardrail_inputs, + request_data={}, + input_type="response", + logging_obj=litellm_logging_obj, + ) + else: + verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices") return responses_so_far string_so_far = self.get_streaming_string_so_far(responses_so_far) @@ -648,7 +663,10 @@ class AnthropicMessagesHandler(BaseTranslation): if isinstance(content_block, dict): if content_block.get("type") == "text": cast(Dict[str, Any], content_block)["text"] = guardrail_response - elif hasattr(content_block, "type") and getattr(content_block, "type", None) == "text": + elif ( + hasattr(content_block, "type") + and getattr(content_block, "type", None) == "text" + ): # Update Pydantic object's text attribute if hasattr(content_block, "text"): content_block.text = guardrail_response diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index 444f821c20a..169b138a5f7 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -939,22 +939,8 @@ class LiteLLMAnthropicMessagesAdapter: self, choices: List[Choices], tool_name_mapping: Optional[Dict[str, str]] = None, - ) -> List[ - Union[ - AnthropicResponseContentBlockText, - AnthropicResponseContentBlockToolUse, - AnthropicResponseContentBlockThinking, - AnthropicResponseContentBlockRedactedThinking, - ] - ]: - new_content: List[ - Union[ - AnthropicResponseContentBlockText, - AnthropicResponseContentBlockToolUse, - AnthropicResponseContentBlockThinking, - AnthropicResponseContentBlockRedactedThinking, - ] - ] = [] + ) -> List[Dict[str, Any]]: + new_content: List[Dict[str, Any]] = [] for choice in choices: # Handle thinking blocks first if ( @@ -978,7 +964,7 @@ class LiteLLMAnthropicMessagesAdapter: if signature_value is not None else None ), - ) + ).model_dump() ) elif thinking_block.get("type") == "redacted_thinking": data_value = thinking_block.get("data", "") @@ -986,7 +972,7 @@ class LiteLLMAnthropicMessagesAdapter: AnthropicResponseContentBlockRedactedThinking( type="redacted_thinking", data=str(data_value) if data_value is not None else "", - ) + ).model_dump() ) # Handle reasoning_content when thinking_blocks is not present elif ( @@ -998,7 +984,7 @@ class LiteLLMAnthropicMessagesAdapter: type="thinking", thinking=str(choice.message.reasoning_content), signature=None, - ) + ).model_dump() ) # Handle text content @@ -1006,7 +992,7 @@ class LiteLLMAnthropicMessagesAdapter: new_content.append( AnthropicResponseContentBlockText( type="text", text=choice.message.content - ) + ).model_dump() ) # Handle tool calls (in parallel to text content) if ( @@ -1044,7 +1030,7 @@ class LiteLLMAnthropicMessagesAdapter: tool_use_block.provider_specific_fields = ( provider_specific_fields ) - new_content.append(tool_use_block) + new_content.append(tool_use_block.model_dump()) return new_content diff --git a/litellm/llms/bedrock/chat/converse_transformation.py b/litellm/llms/bedrock/chat/converse_transformation.py index f6d7e128580..6591e152a14 100644 --- a/litellm/llms/bedrock/chat/converse_transformation.py +++ b/litellm/llms/bedrock/chat/converse_transformation.py @@ -30,6 +30,8 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( from litellm.llms.anthropic.chat.transformation import AnthropicConfig from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.types.llms.bedrock import * + +from ..common_utils import is_claude_4_5_on_bedrock from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionAssistantMessage, @@ -306,9 +308,7 @@ class AmazonConverseConfig(BaseConfig): return "nova-2-lite" in model_without_region def _map_web_search_options( - self, - web_search_options: dict, - model: str + self, web_search_options: dict, model: str ) -> Optional[BedrockToolBlock]: """ Map web_search_options to Nova grounding systemTool. @@ -634,7 +634,7 @@ class AmazonConverseConfig(BaseConfig): 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 @@ -642,10 +642,10 @@ class AmazonConverseConfig(BaseConfig): 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( @@ -808,11 +808,11 @@ class AmazonConverseConfig(BaseConfig): if param == "web_search_options" and isinstance(value, dict): # Note: we use `isinstance(value, dict)` instead of `value and isinstance(value, dict)` # because empty dict {} is falsy but is a valid way to enable Nova grounding - grounding_tool = self._map_web_search_options(value, model) - if grounding_tool is not None: - optional_params = self._add_tools_to_optional_params( - optional_params=optional_params, tools=[grounding_tool] - ) + grounding_tool = self._map_web_search_options(value, model) + if grounding_tool is not None: + optional_params = self._add_tools_to_optional_params( + optional_params=optional_params, tools=[grounding_tool] + ) # Only update thinking tokens for non-GPT-OSS models and non-Nova-Lite-2 models # Nova Lite 2 handles token budgeting differently through reasoningConfig @@ -926,6 +926,7 @@ class AmazonConverseConfig(BaseConfig): ChatCompletionAssistantMessage, ], block_type: Literal["system"], + model: Optional[str] = None, ) -> Optional[SystemContentBlock]: pass @@ -939,6 +940,7 @@ class AmazonConverseConfig(BaseConfig): ChatCompletionAssistantMessage, ], block_type: Literal["content_block"], + model: Optional[str] = None, ) -> Optional[ContentBlock]: pass @@ -951,16 +953,26 @@ class AmazonConverseConfig(BaseConfig): ChatCompletionAssistantMessage, ], block_type: Literal["system", "content_block"], + model: Optional[str] = None, ) -> Optional[Union[SystemContentBlock, ContentBlock]]: - if message_block.get("cache_control", None) is None: + cache_control = message_block.get("cache_control", None) + if cache_control is None: return None + + cache_point = CachePointBlock(type="default") + if isinstance(cache_control, dict) and "ttl" in cache_control: + ttl = cache_control["ttl"] + if ttl in ["5m", "1h"] and model is not None: + if is_claude_4_5_on_bedrock(model): + cache_point["ttl"] = ttl + if block_type == "system": - return SystemContentBlock(cachePoint=CachePointBlock(type="default")) + return SystemContentBlock(cachePoint=cache_point) else: - return ContentBlock(cachePoint=CachePointBlock(type="default")) + return ContentBlock(cachePoint=cache_point) def _transform_system_message( - self, messages: List[AllMessageValues] + self, messages: List[AllMessageValues], model: Optional[str] = None ) -> Tuple[List[AllMessageValues], List[SystemContentBlock]]: system_prompt_indices = [] system_content_blocks: List[SystemContentBlock] = [] @@ -972,7 +984,7 @@ class AmazonConverseConfig(BaseConfig): SystemContentBlock(text=message["content"]) ) cache_block = self._get_cache_point_block( - message, block_type="system" + message, block_type="system", model=model ) if cache_block: system_content_blocks.append(cache_block) @@ -983,7 +995,7 @@ class AmazonConverseConfig(BaseConfig): SystemContentBlock(text=m["text"]) ) cache_block = self._get_cache_point_block( - m, block_type="system" + m, block_type="system", model=model ) if cache_block: system_content_blocks.append(cache_block) @@ -1137,13 +1149,13 @@ class AmazonConverseConfig(BaseConfig): if beta not in seen: unique_betas.append(beta) seen.add(beta) - + # 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 @@ -1196,9 +1208,11 @@ class AmazonConverseConfig(BaseConfig): ) # Prepare and separate parameters - inference_params, additional_request_params, request_metadata = self._prepare_request_params( - optional_params, model - ) + ( + inference_params, + additional_request_params, + request_metadata, + ) = self._prepare_request_params(optional_params, model) original_tools = inference_params.pop("tools", []) @@ -1250,7 +1264,9 @@ class AmazonConverseConfig(BaseConfig): litellm_params: dict, headers: Optional[dict] = None, ) -> RequestObject: - messages, system_content_blocks = self._transform_system_message(messages) + messages, system_content_blocks = self._transform_system_message( + messages, model=model + ) # Convert last user message to guarded_text if guardrailConfig is present messages = self._convert_consecutive_user_messages_to_guarded_text( @@ -1306,7 +1322,9 @@ class AmazonConverseConfig(BaseConfig): litellm_params: dict, headers: Optional[dict] = None, ) -> RequestObject: - messages, system_content_blocks = self._transform_system_message(messages) + messages, system_content_blocks = self._transform_system_message( + messages, model=model + ) # Convert last user message to guarded_text if guardrailConfig is present messages = self._convert_consecutive_user_messages_to_guarded_text( @@ -1484,7 +1502,9 @@ class AmazonConverseConfig(BaseConfig): return message, returned_finish_reason - def _translate_message_content(self, content_blocks: List[ContentBlock]) -> Tuple[ + def _translate_message_content( + self, content_blocks: List[ContentBlock] + ) -> Tuple[ str, List[ChatCompletionToolCallChunk], Optional[List[BedrockConverseReasoningContentBlock]], @@ -1501,9 +1521,9 @@ class AmazonConverseConfig(BaseConfig): """ content_str = "" tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( - None - ) + reasoningContentBlocks: Optional[ + List[BedrockConverseReasoningContentBlock] + ] = None citationsContentBlocks: Optional[List[CitationsContentBlock]] = None for idx, content in enumerate(content_blocks): """ @@ -1557,7 +1577,7 @@ class AmazonConverseConfig(BaseConfig): return content_str, tools, reasoningContentBlocks, citationsContentBlocks - def _transform_response( # noqa: PLR0915 + def _transform_response( # noqa: PLR0915 self, model: str, response: httpx.Response, @@ -1630,9 +1650,9 @@ class AmazonConverseConfig(BaseConfig): chat_completion_message: ChatCompletionResponseMessage = {"role": "assistant"} content_str = "" tools: List[ChatCompletionToolCallChunk] = [] - reasoningContentBlocks: Optional[List[BedrockConverseReasoningContentBlock]] = ( - None - ) + reasoningContentBlocks: Optional[ + List[BedrockConverseReasoningContentBlock] + ] = None citationsContentBlocks: Optional[List[CitationsContentBlock]] = None if message is not None: @@ -1651,15 +1671,17 @@ class AmazonConverseConfig(BaseConfig): provider_specific_fields["citationsContent"] = citationsContentBlocks if provider_specific_fields: - chat_completion_message["provider_specific_fields"] = provider_specific_fields + chat_completion_message[ + "provider_specific_fields" + ] = provider_specific_fields if reasoningContentBlocks is not None: - chat_completion_message["reasoning_content"] = ( - self._transform_reasoning_content(reasoningContentBlocks) - ) - chat_completion_message["thinking_blocks"] = ( - self._transform_thinking_blocks(reasoningContentBlocks) - ) + chat_completion_message[ + "reasoning_content" + ] = self._transform_reasoning_content(reasoningContentBlocks) + chat_completion_message[ + "thinking_blocks" + ] = self._transform_thinking_blocks(reasoningContentBlocks) chat_completion_message["content"] = content_str if ( json_mode is True diff --git a/litellm/llms/bedrock/common_utils.py b/litellm/llms/bedrock/common_utils.py index 65d237bdbdf..4c87f6fa994 100644 --- a/litellm/llms/bedrock/common_utils.py +++ b/litellm/llms/bedrock/common_utils.py @@ -446,6 +446,29 @@ def get_bedrock_base_model(model: str) -> str: return model +def is_claude_4_5_on_bedrock(model: str) -> bool: + """ + Check if the model is a Claude 4.5 model on Bedrock. + Claude 4.5 models support prompt caching with '5m' and '1h' TTL on Bedrock. + """ + model_lower = model.lower() + claude_4_5_patterns = [ + "sonnet-4.5", + "sonnet_4.5", + "sonnet-4-5", + "sonnet_4_5", + "haiku-4.5", + "haiku_4.5", + "haiku-4-5", + "haiku_4_5", + "opus-4.5", + "opus_4.5", + "opus-4-5", + "opus_4_5", + ] + return any(pattern in model_lower for pattern in claude_4_5_patterns) + + # Import after standalone functions to avoid circular imports from litellm.llms.bedrock.count_tokens.bedrock_token_counter import BedrockTokenCounter @@ -815,21 +838,23 @@ def get_anthropic_beta_from_headers(headers: dict) -> List[str]: # If it's already a list, return it if isinstance(anthropic_beta_header, list): return anthropic_beta_header - + # Try to parse as JSON array first (e.g., '["interleaved-thinking-2025-05-14", "claude-code-20250219"]') if isinstance(anthropic_beta_header, str): anthropic_beta_header = anthropic_beta_header.strip() - if anthropic_beta_header.startswith("[") and anthropic_beta_header.endswith("]"): + if anthropic_beta_header.startswith("[") and anthropic_beta_header.endswith( + "]" + ): try: parsed = json.loads(anthropic_beta_header) if isinstance(parsed, list): return [str(beta).strip() for beta in parsed] except json.JSONDecodeError: pass # Fall through to comma-separated parsing - + # Fall back to comma-separated values return [beta.strip() for beta in anthropic_beta_header.split(",")] - + return [] diff --git a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py index b1c45ea83a2..90de67a822f 100644 --- a/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py +++ b/litellm/llms/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.py @@ -23,7 +23,10 @@ from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( AmazonInvokeConfig, ) -from litellm.llms.bedrock.common_utils import get_anthropic_beta_from_headers +from litellm.llms.bedrock.common_utils import ( + get_anthropic_beta_from_headers, + is_claude_4_5_on_bedrock, +) from litellm.types.llms.anthropic import ANTHROPIC_TOOL_SEARCH_BETA_HEADER from litellm.types.llms.openai import AllMessageValues from litellm.types.router import GenericLiteLLMParams @@ -54,7 +57,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" + "prompt-caching-scope", ] def __init__(self, **kwargs): @@ -116,15 +119,22 @@ class AmazonAnthropicClaudeMessagesConfig( ) def _remove_ttl_from_cache_control( - self, anthropic_messages_request: Dict + self, anthropic_messages_request: Dict, model: Optional[str] = None ) -> None: """ Remove `ttl` field from cache_control in messages. Bedrock doesn't support the ttl field in cache_control. + Update: Bedock supports `5m` and `1h` for Claude 4.5 models. + Args: anthropic_messages_request: The request dictionary to modify in-place + model: The model name to check if it supports ttl """ + is_claude_4_5 = False + if model: + is_claude_4_5 = self._is_claude_4_5_on_bedrock(model) + if "messages" in anthropic_messages_request: for message in anthropic_messages_request["messages"]: if isinstance(message, dict) and "content" in message: @@ -133,7 +143,14 @@ class AmazonAnthropicClaudeMessagesConfig( for item in content: if isinstance(item, dict) and "cache_control" in item: cache_control = item["cache_control"] - if isinstance(cache_control, dict) and "ttl" in cache_control: + if ( + isinstance(cache_control, dict) + and "ttl" in cache_control + ): + ttl = cache_control["ttl"] + if is_claude_4_5 and ttl in ["5m", "1h"]: + continue + cache_control.pop("ttl", None) def _supports_extended_thinking_on_bedrock(self, model: str) -> bool: @@ -155,10 +172,18 @@ class AmazonAnthropicClaudeMessagesConfig( # Supported models on Bedrock for extended thinking supported_patterns = [ - "opus-4.5", "opus_4.5", "opus-4-5", "opus_4_5", # Opus 4.5 - "opus-4.1", "opus_4.1", "opus-4-1", "opus_4_1", # Opus 4.1 - "opus-4", "opus_4", # Opus 4 - "sonnet-4", "sonnet_4", # Sonnet 4 + "opus-4.5", + "opus_4.5", + "opus-4-5", + "opus_4_5", # Opus 4.5 + "opus-4.1", + "opus_4.1", + "opus-4-1", + "opus_4_1", # Opus 4.1 + "opus-4", + "opus_4", # Opus 4 + "sonnet-4", + "sonnet_4", # Sonnet 4 ] return any(pattern in model_lower for pattern in supported_patterns) @@ -175,10 +200,27 @@ class AmazonAnthropicClaudeMessagesConfig( """ model_lower = model.lower() opus_4_5_patterns = [ - "opus-4.5", "opus_4.5", "opus-4-5", "opus_4_5", + "opus-4.5", + "opus_4.5", + "opus-4-5", + "opus_4_5", ] return any(pattern in model_lower for pattern in opus_4_5_patterns) + def _is_claude_4_5_on_bedrock(self, model: str) -> bool: + """ + Check if the model is Claude 4.5 on Bedrock. + + Claude Sonnet 4.5, Haiku 4.5, and Opus 4.5 support 1-hour prompt caching. + + Args: + model: The model name + + Returns: + True if the model is Claude 4.5 + """ + return is_claude_4_5_on_bedrock(model) + def _supports_tool_search_on_bedrock(self, model: str) -> bool: """ Check if the model supports tool search on Bedrock. @@ -199,9 +241,15 @@ class AmazonAnthropicClaudeMessagesConfig( # Supported models for tool search on Bedrock supported_patterns = [ # Opus 4.5 - "opus-4.5", "opus_4.5", "opus-4-5", "opus_4_5", + "opus-4.5", + "opus_4.5", + "opus-4-5", + "opus_4_5", # Sonnet 4.5 - "sonnet-4.5", "sonnet_4.5", "sonnet-4-5", "sonnet_4_5", + "sonnet-4.5", + "sonnet_4.5", + "sonnet-4-5", + "sonnet_4_5", ] return any(pattern in model_lower for pattern in supported_patterns) @@ -238,8 +286,7 @@ class AmazonAnthropicClaudeMessagesConfig( beta_headers_to_remove.add(beta) has_advanced_tool_use = True break - - + # 2. Filter out extended thinking headers for models that don't support them extended_thinking_patterns = [ "extended-thinking", @@ -263,7 +310,6 @@ class AmazonAnthropicClaudeMessagesConfig( beta_set.add("tool-search-tool-2025-10-19") beta_set.add("tool-examples-2025-10-29") - def _get_tool_search_beta_header_for_bedrock( self, model: str, @@ -290,7 +336,9 @@ class AmazonAnthropicClaudeMessagesConfig( input_examples_used: Whether input examples are used beta_set: The set of beta headers to modify in-place """ - if tool_search_used and not (programmatic_tool_calling_used or input_examples_used): + if tool_search_used and not ( + programmatic_tool_calling_used or input_examples_used + ): beta_set.discard(ANTHROPIC_TOOL_SEARCH_BETA_HEADER) if "opus-4" in model.lower() or "opus_4" in model.lower(): beta_set.add("tool-search-tool-2025-10-19") @@ -302,13 +350,13 @@ class AmazonAnthropicClaudeMessagesConfig( ) -> None: """ Convert Anthropic output_format to inline schema in message content. - + Bedrock Invoke doesn't support the output_format parameter, so we embed the schema directly into the user message content as text instructions. - + This approach adds the schema to the last user message, instructing the model to respond in the specified JSON format. - + Args: output_format: The output_format dict with 'type' and 'schema' anthropic_messages_request: The request dict to modify in-place @@ -321,35 +369,32 @@ class AmazonAnthropicClaudeMessagesConfig( schema = output_format.get("schema") if not schema: return - + # Get messages from the request messages = anthropic_messages_request.get("messages", []) if not messages: return - + # Find the last user message last_user_message_idx = None for idx in range(len(messages) - 1, -1, -1): if messages[idx].get("role") == "user": last_user_message_idx = idx break - + if last_user_message_idx is None: return - + last_user_message = messages[last_user_message_idx] content = last_user_message.get("content", []) - + # Ensure content is a list if isinstance(content, str): content = [{"type": "text", "text": content}] last_user_message["content"] = content - + # Add schema as text content to the message - schema_text = { - "type": "text", - "text": json.dumps(schema) - } + schema_text = {"type": "text", "text": json.dumps(schema)} content.append(schema_text) def transform_anthropic_messages_request( @@ -374,9 +419,9 @@ class AmazonAnthropicClaudeMessagesConfig( # 1. anthropic_version is required for all claude models if "anthropic_version" not in anthropic_messages_request: - anthropic_messages_request["anthropic_version"] = ( - self.DEFAULT_BEDROCK_ANTHROPIC_API_VERSION - ) + anthropic_messages_request[ + "anthropic_version" + ] = self.DEFAULT_BEDROCK_ANTHROPIC_API_VERSION # 2. `stream` is not allowed in request body for bedrock invoke if "stream" in anthropic_messages_request: @@ -386,8 +431,10 @@ class AmazonAnthropicClaudeMessagesConfig( if "model" in anthropic_messages_request: anthropic_messages_request.pop("model", None) - # 4. Remove `ttl` field from cache_control in messages (Bedrock doesn't support it) - self._remove_ttl_from_cache_control(anthropic_messages_request) + # 4. Remove `ttl` field from cache_control in messages (Bedrock doesn't support it for older models) + self._remove_ttl_from_cache_control( + anthropic_messages_request=anthropic_messages_request, model=model + ) # 5. Convert `output_format` to inline schema (Bedrock invoke doesn't support output_format) output_format = anthropic_messages_request.pop("output_format", None) @@ -396,14 +443,14 @@ class AmazonAnthropicClaudeMessagesConfig( output_format=output_format, anthropic_messages_request=anthropic_messages_request, ) - + # 6. AUTO-INJECT beta headers based on features used anthropic_model_info = AnthropicModelInfo() tools = anthropic_messages_optional_request_params.get("tools") messages_typed = cast(List[AllMessageValues], messages) tool_search_used = anthropic_model_info.is_tool_search_used(tools) - programmatic_tool_calling_used = anthropic_model_info.is_programmatic_tool_calling_used( - tools + programmatic_tool_calling_used = ( + anthropic_model_info.is_programmatic_tool_calling_used(tools) ) input_examples_used = anthropic_model_info.is_input_examples_used(tools) @@ -436,8 +483,7 @@ class AmazonAnthropicClaudeMessagesConfig( if beta_set: anthropic_messages_request["anthropic_beta"] = list(beta_set) - - + return anthropic_messages_request def get_async_streaming_response_iterator( @@ -455,7 +501,7 @@ class AmazonAnthropicClaudeMessagesConfig( ) # Convert decoded Bedrock events to Server-Sent Events expected by Anthropic clients. return self.bedrock_sse_wrapper( - completion_stream=completion_stream, + completion_stream=completion_stream, litellm_logging_obj=litellm_logging_obj, request_body=request_body, ) @@ -474,14 +520,14 @@ class AmazonAnthropicClaudeMessagesConfig( from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import ( BaseAnthropicMessagesStreamingIterator, ) + handler = BaseAnthropicMessagesStreamingIterator( litellm_logging_obj=litellm_logging_obj, request_body=request_body, ) - + async for chunk in handler.async_sse_wrapper(completion_stream): yield chunk - class AmazonAnthropicClaudeMessagesStreamDecoder(AWSEventStreamDecoder): diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index 86bcd94450f..7ec32fecc46 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -236,6 +236,10 @@ class FireworksAIConfig(OpenAIGPTConfig): disable_add_transform_inline_image_block=disable_add_transform_inline_image_block, ) filter_value_from_dict(cast(dict, message), "cache_control") + # Remove fields not permitted by FireworksAI that may cause: + # "Not permitted, field: 'messages[n].provider_specific_fields'" + if isinstance(message, dict) and "provider_specific_fields" in message: + cast(dict, message).pop("provider_specific_fields", None) return messages diff --git a/litellm/llms/gemini/files/transformation.py b/litellm/llms/gemini/files/transformation.py index 37f1376c2b1..cc799cfd6aa 100644 --- a/litellm/llms/gemini/files/transformation.py +++ b/litellm/llms/gemini/files/transformation.py @@ -210,7 +210,7 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): We expect file_id to be the URI (e.g. https://generativelanguage.googleapis.com/v1beta/files/...) as returned by the upload response. """ - api_key = litellm_params.get("api_key") + api_key = litellm_params.get("api_key") or self.get_api_key() if not api_key: raise ValueError("api_key is required") @@ -222,7 +222,8 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): api_base = api_base.rstrip("/") url = "{}/v1beta/{}?key={}".format(api_base, file_id, api_key) - return url, {"Content-Type": "application/json"} + # Return empty params dict - API key is already in URL, no query params needed + return url, {} def transform_retrieve_file_response( self, @@ -299,7 +300,7 @@ class GoogleAIStudioFilesHandler(GeminiModelInfo, BaseFilesConfig): # Extract the file path from full URI file_name = file_id.split("/v1beta/")[-1] else: - file_name = file_id + file_name = file_id if file_id.startswith("files/") else f"files/{file_id}" # Construct the delete URL url = f"{api_base}/v1beta/{file_name}" diff --git a/litellm/llms/gigachat/chat/transformation.py b/litellm/llms/gigachat/chat/transformation.py index ba14de1f65d..f546f356e11 100644 --- a/litellm/llms/gigachat/chat/transformation.py +++ b/litellm/llms/gigachat/chat/transformation.py @@ -386,33 +386,7 @@ class GigaChatConfig(BaseConfig): transformed.append(message) - # Collapse consecutive user messages - return self._collapse_user_messages(transformed) - - def _collapse_user_messages(self, messages: List[dict]) -> List[dict]: - """Collapse consecutive user messages into one.""" - collapsed: List[dict] = [] - prev_user_msg: Optional[dict] = None - content_parts: List[str] = [] - - for msg in messages: - if msg.get("role") == "user" and prev_user_msg is not None: - content_parts.append(msg.get("content", "")) - else: - if content_parts and prev_user_msg: - prev_user_msg["content"] = "\n".join( - [prev_user_msg.get("content", "")] + content_parts - ) - content_parts = [] - collapsed.append(msg) - prev_user_msg = msg if msg.get("role") == "user" else None - - if content_parts and prev_user_msg: - prev_user_msg["content"] = "\n".join( - [prev_user_msg.get("content", "")] + content_parts - ) - - return collapsed + return transformed def transform_response( self, diff --git a/litellm/llms/github_copilot/chat/transformation.py b/litellm/llms/github_copilot/chat/transformation.py index 50f18cedf9b..be8ad7d0877 100644 --- a/litellm/llms/github_copilot/chat/transformation.py +++ b/litellm/llms/github_copilot/chat/transformation.py @@ -1,11 +1,16 @@ -from typing import Any, Optional, Tuple, cast, List +from typing import List, Optional, Tuple + from litellm.exceptions import AuthenticationError from litellm.llms.openai.openai import OpenAIConfig from litellm.types.llms.openai import AllMessageValues from ..authenticator import Authenticator -from ..common_utils import GetAPIKeyError, GITHUB_COPILOT_API_BASE +from ..common_utils import ( + GITHUB_COPILOT_API_BASE, + GetAPIKeyError, + get_copilot_default_headers, +) class GithubCopilotConfig(OpenAIConfig): @@ -25,9 +30,7 @@ class GithubCopilotConfig(OpenAIConfig): api_key: Optional[str], custom_llm_provider: str, ) -> Tuple[Optional[str], Optional[str], str]: - dynamic_api_base = ( - self.authenticator.get_api_base() or GITHUB_COPILOT_API_BASE - ) + dynamic_api_base = self.authenticator.get_api_base() or GITHUB_COPILOT_API_BASE try: dynamic_api_key = self.authenticator.get_api_key() except GetAPIKeyError as e: @@ -45,14 +48,24 @@ class GithubCopilotConfig(OpenAIConfig): ): import litellm - disable_copilot_system_to_assistant = ( - litellm.disable_copilot_system_to_assistant - ) - if not disable_copilot_system_to_assistant: - for message in messages: - if "role" in message and message["role"] == "system": - cast(Any, message)["role"] = "assistant" - return messages + # Check if system-to-assistant conversion is disabled + if litellm.disable_copilot_system_to_assistant: + # GitHub Copilot API now supports system prompts for all models (Claude, GPT, etc.) + # No conversion needed - just return messages as-is + return messages + + # Default behavior: convert system messages to assistant for compatibility + transformed_messages = [] + for message in messages: + if message.get("role") == "system": + # Convert system message to assistant message + transformed_message = message.copy() + transformed_message["role"] = "assistant" + transformed_messages.append(transformed_message) + else: + transformed_messages.append(message) + + return transformed_messages def validate_environment( self, @@ -69,6 +82,14 @@ class GithubCopilotConfig(OpenAIConfig): headers, model, messages, optional_params, litellm_params, api_key, api_base ) + # Add Copilot-specific headers (editor-version, user-agent, etc.) + try: + copilot_api_key = self.authenticator.get_api_key() + copilot_headers = get_copilot_default_headers(copilot_api_key) + validated_headers = {**copilot_headers, **validated_headers} + except GetAPIKeyError: + pass # Will be handled later in the request flow + # Add X-Initiator header based on message roles initiator = self._determine_initiator(messages) validated_headers["X-Initiator"] = initiator @@ -87,7 +108,7 @@ class GithubCopilotConfig(OpenAIConfig): For other models, returns standard OpenAI parameters (which may include reasoning_effort for o-series models). """ from litellm.utils import supports_reasoning - + # Get base OpenAI parameters base_params = super().get_supported_openai_params(model) @@ -118,7 +139,7 @@ class GithubCopilotConfig(OpenAIConfig): """ Check if any message contains vision content (images). Returns True if any message has content with vision-related types, otherwise False. - + Checks for: - image_url content type (OpenAI format) - Content items with type 'image_url' diff --git a/litellm/llms/openai/chat/guardrail_translation/handler.py b/litellm/llms/openai/chat/guardrail_translation/handler.py index fb00aa28f45..c406f502b45 100644 --- a/litellm/llms/openai/chat/guardrail_translation/handler.py +++ b/litellm/llms/openai/chat/guardrail_translation/handler.py @@ -21,7 +21,13 @@ from litellm._logging import verbose_proxy_logger from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation from litellm.main import stream_chunk_builder from litellm.types.llms.openai import ChatCompletionToolParam -from litellm.types.utils import Choices, GenericGuardrailAPIInputs, ModelResponse, ModelResponseStream, StreamingChoices +from litellm.types.utils import ( + Choices, + GenericGuardrailAPIInputs, + ModelResponse, + ModelResponseStream, + StreamingChoices, +) if TYPE_CHECKING: from litellm.integrations.custom_guardrail import CustomGuardrail @@ -80,9 +86,9 @@ class OpenAIChatCompletionsHandler(BaseTranslation): if tool_calls_to_check: inputs["tool_calls"] = tool_calls_to_check # type: ignore if messages: - inputs["structured_messages"] = ( - messages # pass the openai /chat/completions messages to the guardrail, as-is - ) + inputs[ + "structured_messages" + ] = messages # pass the openai /chat/completions messages to the guardrail, as-is # Pass tools (function definitions) to the guardrail tools = data.get("tools") if tools: @@ -362,14 +368,17 @@ class OpenAIChatCompletionsHandler(BaseTranslation): # check if the stream has ended has_stream_ended = False for chunk in responses_so_far: - if chunk.choices[0].finish_reason is not None: + if chunk.choices and chunk.choices[0].finish_reason is not None: has_stream_ended = True break if has_stream_ended: # convert to model response model_response = cast( - ModelResponse, stream_chunk_builder(chunks=responses_so_far, logging_obj=litellm_logging_obj) + ModelResponse, + stream_chunk_builder( + chunks=responses_so_far, logging_obj=litellm_logging_obj + ), ) # run process_output_response await self.process_output_response( diff --git a/litellm/llms/openai/common_utils.py b/litellm/llms/openai/common_utils.py index 8bcecd35232..ce470f04aca 100644 --- a/litellm/llms/openai/common_utils.py +++ b/litellm/llms/openai/common_utils.py @@ -15,14 +15,12 @@ if TYPE_CHECKING: from aiohttp import ClientSession import litellm -from litellm._logging import verbose_logger from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.custom_httpx.http_handler import ( _DEFAULT_TTL_FOR_HTTPX_CLIENTS, AsyncHTTPHandler, get_ssl_configuration, ) -from litellm.types.utils import LlmProviders class OpenAIError(BaseLLMException): @@ -205,67 +203,30 @@ class BaseOpenAILLM: if litellm.aclient_session is not None: return litellm.aclient_session - # Use the global cached client system to prevent memory leaks (issue #14540) - # This routes through get_async_httpx_client() which provides TTL-based caching - from litellm.llms.custom_httpx.http_handler import get_async_httpx_client + # Get unified SSL configuration + ssl_config = get_ssl_configuration() - try: - # Get SSL config and include in params for proper cache key - ssl_config = get_ssl_configuration() - params = {"ssl_verify": ssl_config} if ssl_config is not None else {} - params["disable_aiohttp_transport"] = litellm.disable_aiohttp_transport - - # Get a cached AsyncHTTPHandler which manages the httpx.AsyncClient - cached_handler = get_async_httpx_client( - llm_provider=LlmProviders.OPENAI, # Cache key includes provider - params=params, # Include SSL config in cache key + return httpx.AsyncClient( + verify=ssl_config, + transport=AsyncHTTPHandler._create_async_transport( + ssl_context=ssl_config + if isinstance(ssl_config, ssl.SSLContext) + else None, + ssl_verify=ssl_config if isinstance(ssl_config, bool) else None, shared_session=shared_session, - ) - # Return the underlying httpx client from the handler - return cached_handler.client - except (ImportError, AttributeError, KeyError) as e: - # Fallback to creating a client directly if caching system unavailable - # This preserves backwards compatibility - verbose_logger.debug( - f"Client caching unavailable ({type(e).__name__}), using direct client creation" - ) - ssl_config = get_ssl_configuration() - return httpx.AsyncClient( - verify=ssl_config, - transport=AsyncHTTPHandler._create_async_transport( - ssl_context=ssl_config - if isinstance(ssl_config, ssl.SSLContext) - else None, - ssl_verify=ssl_config if isinstance(ssl_config, bool) else None, - shared_session=shared_session, - ), - follow_redirects=True, - ) + ), + follow_redirects=True, + ) @staticmethod def _get_sync_http_client() -> Optional[httpx.Client]: if litellm.client_session is not None: return litellm.client_session - # Use the global cached client system to prevent memory leaks (issue #14540) - from litellm.llms.custom_httpx.http_handler import _get_httpx_client + # Get unified SSL configuration + ssl_config = get_ssl_configuration() - try: - # Get SSL config and include in params for proper cache key - ssl_config = get_ssl_configuration() - params = {"ssl_verify": ssl_config} if ssl_config is not None else None - - # Get a cached HTTPHandler which manages the httpx.Client - cached_handler = _get_httpx_client(params=params) - # Return the underlying httpx client from the handler - return cached_handler.client - except (ImportError, AttributeError, KeyError) as e: - # Fallback to creating a client directly if caching system unavailable - verbose_logger.debug( - f"Client caching unavailable ({type(e).__name__}), using direct client creation" - ) - ssl_config = get_ssl_configuration() - return httpx.Client( - verify=ssl_config, - follow_redirects=True, - ) + return httpx.Client( + verify=ssl_config, + follow_redirects=True, + ) diff --git a/litellm/llms/openai/realtime/handler.py b/litellm/llms/openai/realtime/handler.py index fd04ac4d458..ef9cc43c3e1 100644 --- a/litellm/llms/openai/realtime/handler.py +++ b/litellm/llms/openai/realtime/handler.py @@ -16,6 +16,62 @@ from ..openai import OpenAIChatCompletion class OpenAIRealtime(OpenAIChatCompletion): + """ + Base handler for OpenAI-compatible realtime WebSocket connections. + + Subclasses can override template methods to customize: + - _get_default_api_base(): Default API base URL + - _get_additional_headers(): Extra headers beyond Authorization + - _get_ssl_config(): SSL configuration for WebSocket connection + """ + + def _get_default_api_base(self) -> str: + """ + Get the default API base URL for this provider. + Override this in subclasses to set provider-specific defaults. + """ + return "https://api.openai.com/" + + def _get_additional_headers(self, api_key: str) -> dict: + """ + Get additional headers beyond Authorization. + Override this in subclasses to customize headers (e.g., remove OpenAI-Beta). + + Args: + api_key: API key for authentication + + Returns: + Dictionary of additional headers + """ + return { + "Authorization": f"Bearer {api_key}", + "OpenAI-Beta": "realtime=v1", + } + + def _get_ssl_config(self, url: str) -> Any: + """ + Get SSL configuration for WebSocket connection. + Override this in subclasses to customize SSL behavior. + + Args: + url: WebSocket URL (ws:// or wss://) + + Returns: + SSL configuration (None, True, or SSLContext) + """ + if url.startswith("ws://"): + return None + + # Use the shared SSL context which respects custom CA certs and SSL settings + ssl_config = get_shared_realtime_ssl_context() + + # If ssl_config is False (ssl_verify=False), websockets library needs True instead + # to establish connection without verification (False would fail) + if ssl_config is False: + return True + + return ssl_config + def _construct_url(self, api_base: str, query_params: RealtimeQueryParams) -> str: """ Construct the backend websocket URL with all query parameters (including 'model'). @@ -45,8 +101,9 @@ class OpenAIRealtime(OpenAIChatCompletion): ): import websockets from websockets.asyncio.client import ClientConnection + if api_base is None: - api_base = "https://api.openai.com/" + api_base = self._get_default_api_base() if api_key is None: raise ValueError("api_key is required for OpenAI realtime calls") @@ -56,30 +113,27 @@ class OpenAIRealtime(OpenAIChatCompletion): url = self._construct_url(api_base, query_params) try: - # Only use SSL context for secure websocket connections (wss://) - # websockets library doesn't accept ssl argument for ws:// URIs - ssl_context = None if url.startswith("ws://") else get_shared_realtime_ssl_context() + # Get provider-specific SSL configuration + ssl_config = self._get_ssl_config(url) + + # Get provider-specific headers + headers = self._get_additional_headers(api_key) + # Log a masked request preview consistent with other endpoints. logging_obj.pre_call( input=None, api_key=api_key, additional_args={ "api_base": url, - "headers": { - "Authorization": f"Bearer {api_key}", - "OpenAI-Beta": "realtime=v1", - }, + "headers": headers, "complete_input_dict": {"query_params": query_params}, }, ) async with websockets.connect( # type: ignore url, - additional_headers={ - "Authorization": f"Bearer {api_key}", # type: ignore - "OpenAI-Beta": "realtime=v1", - }, + additional_headers=headers, # type: ignore max_size=REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES, - ssl=ssl_context, + ssl=ssl_config, ) as backend_ws: realtime_streaming = RealTimeStreaming( websocket, cast(ClientConnection, backend_ws), logging_obj diff --git a/litellm/llms/openai/responses/guardrail_translation/handler.py b/litellm/llms/openai/responses/guardrail_translation/handler.py index d943662f9e4..ad3d4c932d4 100644 --- a/litellm/llms/openai/responses/guardrail_translation/handler.py +++ b/litellm/llms/openai/responses/guardrail_translation/handler.py @@ -319,9 +319,7 @@ class OpenAIResponsesHandler(BaseTranslation): return response if not response_output: - verbose_proxy_logger.debug( - "OpenAI Responses API: Empty output in response" - ) + verbose_proxy_logger.debug("OpenAI Responses API: Empty output in response") return response # Step 1: Extract all text content and tool calls from response output @@ -427,27 +425,30 @@ class OpenAIResponsesHandler(BaseTranslation): handle_raw_dict_callback=None, ) - tool_calls = model_response_choices[0].message.tool_calls - text = model_response_choices[0].message.content - guardrail_inputs = GenericGuardrailAPIInputs() - if text: - guardrail_inputs["texts"] = [text] - if tool_calls: - guardrail_inputs["tool_calls"] = cast( - List[ChatCompletionToolCallChunk], tool_calls - ) - # Include model information from the response if available - response_model = final_chunk.get("response", {}).get("model") - if response_model: - guardrail_inputs["model"] = response_model - if tool_calls or text: - _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( - inputs=guardrail_inputs, - request_data={}, - input_type="response", - logging_obj=litellm_logging_obj, - ) - return responses_so_far + if model_response_choices: + tool_calls = model_response_choices[0].message.tool_calls + text = model_response_choices[0].message.content + guardrail_inputs = GenericGuardrailAPIInputs() + if text: + guardrail_inputs["texts"] = [text] + if tool_calls: + guardrail_inputs["tool_calls"] = cast( + List[ChatCompletionToolCallChunk], tool_calls + ) + # Include model information from the response if available + response_model = final_chunk.get("response", {}).get("model") + if response_model: + guardrail_inputs["model"] = response_model + if tool_calls or text: + _guardrailed_inputs = await guardrail_to_apply.apply_guardrail( + inputs=guardrail_inputs, + request_data={}, + input_type="response", + logging_obj=litellm_logging_obj, + ) + return responses_so_far + else: + verbose_proxy_logger.debug("Skipping output guardrail - model response has no choices") # model_response_stream = OpenAiResponsesToChatCompletionStreamIterator.translate_responses_chunk_to_openai_stream(final_chunk) # tool_calls = model_response_stream.choices[0].tool_calls # convert openai response to model response @@ -513,11 +514,9 @@ class OpenAIResponsesHandler(BaseTranslation): # Check if it's an OutputText with text if isinstance(content_item, OutputText): if content_item.text: - return True elif isinstance(content_item, dict): if content_item.get("text"): - return True return False diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index b5a6949f272..04ae4b6beb8 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -1732,6 +1732,52 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): else: return "stop" + @staticmethod + def _check_prompt_level_content_filter( + processed_chunk: GenerateContentResponseBody, + response_id: Optional[str], + ) -> Optional["ModelResponseStream"]: + """ + Check if prompt is blocked due to content filtering at the prompt level. + + This handles the case where Vertex AI blocks the prompt before generation begins, + indicated by promptFeedback.blockReason being present. + + Args: + processed_chunk: The parsed response chunk from Vertex AI + response_id: The response ID from the chunk + + Returns: + ModelResponseStream with content_filter finish_reason if blocked, None otherwise. + + Note: + This is consistent with non-streaming _handle_blocked_response() behavior. + Candidate-level content filtering (SAFETY, RECITATION, etc.) is handled + separately via _process_candidates() → _check_finish_reason(). + """ + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + # Check if prompt is blocked due to content filtering + prompt_feedback = processed_chunk.get("promptFeedback") + if prompt_feedback and "blockReason" in prompt_feedback: + verbose_logger.debug( + f"Prompt blocked due to: {prompt_feedback.get('blockReason')} - {prompt_feedback.get('blockReasonMessage')}" + ) + + # Create a content_filter response (consistent with non-streaming _handle_blocked_response) + choice = StreamingChoices( + finish_reason="content_filter", + index=0, + delta=Delta(content=None, role="assistant"), + logprobs=None, + enhancements=None, + ) + + model_response = ModelResponseStream(choices=[choice], id=response_id) + return model_response + + return None + @staticmethod def _calculate_web_search_requests(grounding_metadata: List[dict]) -> Optional[int]: web_search_requests: Optional[int] = None @@ -2813,6 +2859,15 @@ class ModelResponseIterator: processed_chunk = GenerateContentResponseBody(**chunk) # type: ignore response_id = processed_chunk.get("responseId") model_response = ModelResponseStream(choices=[], id=response_id) + + # Check if prompt is blocked due to content filtering + blocked_response = VertexGeminiConfig._check_prompt_level_content_filter( + processed_chunk=processed_chunk, + response_id=response_id, + ) + if blocked_response is not None: + model_response = blocked_response + usage: Optional[Usage] = None _candidates: Optional[List[Candidates]] = processed_chunk.get("candidates") grounding_metadata: List[dict] = [] diff --git a/litellm/llms/xai/chat/transformation.py b/litellm/llms/xai/chat/transformation.py index 245e10e45c1..21782fc6fbf 100644 --- a/litellm/llms/xai/chat/transformation.py +++ b/litellm/llms/xai/chat/transformation.py @@ -4,6 +4,7 @@ import httpx import litellm from litellm._logging import verbose_logger +from litellm.constants import XAI_API_BASE from litellm.litellm_core_utils.prompt_templates.common_utils import ( filter_value_from_dict, strip_name_from_messages, @@ -14,8 +15,6 @@ from litellm.types.utils import Choices, ModelResponse, Usage, PromptTokensDetai from ...openai.chat.gpt_transformation import OpenAIGPTConfig -XAI_API_BASE = "https://api.x.ai/v1" - class XAIChatConfig(OpenAIGPTConfig): @property diff --git a/litellm/llms/xai/realtime/__init__.py b/litellm/llms/xai/realtime/__init__.py new file mode 100644 index 00000000000..3b0d345f2c2 --- /dev/null +++ b/litellm/llms/xai/realtime/__init__.py @@ -0,0 +1,5 @@ +"""xAI Realtime API handler.""" + +from .handler import XAIRealtime + +__all__ = ["XAIRealtime"] diff --git a/litellm/llms/xai/realtime/handler.py b/litellm/llms/xai/realtime/handler.py new file mode 100644 index 00000000000..c79477ba1df --- /dev/null +++ b/litellm/llms/xai/realtime/handler.py @@ -0,0 +1,38 @@ +""" +This file contains the handler for xAI's Grok Voice Agent API `/v1/realtime` endpoint. + +xAI's Realtime API is fully OpenAI-compatible, so we inherit from OpenAIRealtime +and only override the configuration differences. + +This requires websockets, and is currently only supported on LiteLLM Proxy. +""" + +from litellm.constants import XAI_API_BASE + +from ...openai.realtime.handler import OpenAIRealtime + + +class XAIRealtime(OpenAIRealtime): + """ + Handler for xAI Grok Voice Agent API. + + xAI's Realtime API uses the same WebSocket protocol as OpenAI but with: + - Different endpoint: wss://api.x.ai/v1/realtime (via _get_default_api_base) + - No OpenAI-Beta header required (via _get_additional_headers) + - Model: grok-4-1-fast-non-reasoning + + All WebSocket logic is inherited from OpenAIRealtime. + """ + + def _get_default_api_base(self) -> str: + """xAI uses a different API base URL.""" + return XAI_API_BASE + + def _get_additional_headers(self, api_key: str) -> dict: + """ + xAI does NOT require the OpenAI-Beta header. + Only send Authorization header. + """ + return { + "Authorization": f"Bearer {api_key}", + } diff --git a/litellm/llms/xai/responses/transformation.py b/litellm/llms/xai/responses/transformation.py index 82b4771fb4d..95873aab846 100644 --- a/litellm/llms/xai/responses/transformation.py +++ b/litellm/llms/xai/responses/transformation.py @@ -2,6 +2,7 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Union import litellm from litellm._logging import verbose_logger +from litellm.constants import XAI_API_BASE from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams @@ -16,8 +17,6 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any -XAI_API_BASE = "https://api.x.ai/v1" - class XAIResponsesAPIConfig(OpenAIResponsesAPIConfig): """ diff --git a/litellm/main.py b/litellm/main.py index 60c889ab1d0..bca023e65ec 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -1199,6 +1199,13 @@ def completion( # type: ignore # noqa: PLR0915 headers = {} if extra_headers is not None: headers.update(extra_headers) + # Inject proxy auth headers if configured + if litellm.proxy_auth is not None: + try: + proxy_headers = litellm.proxy_auth.get_auth_headers() + headers.update(proxy_headers) + except Exception as e: + verbose_logger.warning(f"Failed to get proxy auth headers: {e}") num_retries = kwargs.get( "num_retries", None ) ## alt. param for 'max_retries'. Use this to pass retries w/ instructor. @@ -2201,14 +2208,24 @@ def completion( # type: ignore # noqa: PLR0915 ) elif custom_llm_provider == "a2a": # A2A (Agent-to-Agent) Protocol - api_base = ( - api_base - or litellm.api_base - or get_secret_str("A2A_API_BASE") + # Resolve agent configuration from registry if model format is "a2a/" + api_base, api_key, headers = litellm.A2AConfig.resolve_agent_config_from_registry( + model=model, + api_base=api_base, + api_key=api_key, + headers=headers, + optional_params=optional_params, ) - + + # Fall back to environment variables and defaults + api_base = api_base or litellm.api_base or get_secret_str("A2A_API_BASE") + if api_base is None: - raise Exception("api_base is required for A2A provider") + raise Exception( + "api_base is required for A2A provider. " + "Either provide api_base parameter, set A2A_API_BASE environment variable, " + "or register the agent in the proxy with model='a2a/'." + ) headers = headers or litellm.headers @@ -2487,6 +2504,20 @@ def completion( # type: ignore # noqa: PLR0915 headers = headers or litellm.headers + # Add GitHub Copilot headers (same as /responses endpoint does) + if custom_llm_provider == "github_copilot": + from litellm.llms.github_copilot.common_utils import ( + get_copilot_default_headers, + ) + from litellm.llms.github_copilot.authenticator import Authenticator + + copilot_auth = Authenticator() + copilot_api_key = copilot_auth.get_api_key() + copilot_headers = get_copilot_default_headers(copilot_api_key) + if extra_headers: + copilot_headers.update(extra_headers) + extra_headers = copilot_headers + if extra_headers is not None: optional_params["extra_headers"] = extra_headers @@ -4587,6 +4618,13 @@ def embedding( # noqa: PLR0915 headers = {} if extra_headers is not None: headers.update(extra_headers) + # Inject proxy auth headers if configured + if litellm.proxy_auth is not None: + try: + proxy_headers = litellm.proxy_auth.get_auth_headers() + headers.update(proxy_headers) + except Exception as e: + verbose_logger.warning(f"Failed to get proxy auth headers: {e}") ### CUSTOM MODEL COST ### input_cost_per_token = kwargs.get("input_cost_per_token", None) output_cost_per_token = kwargs.get("output_cost_per_token", None) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b9e48fd7e11..d3038d13a1e 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -744,7 +744,8 @@ "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, - "tool_use_system_prompt_tokens": 346 + "tool_use_system_prompt_tokens": 346, + "supports_native_streaming": true }, "anthropic.claude-3-5-sonnet-20240620-v1:0": { "input_cost_per_token": 3e-06, @@ -12850,6 +12851,40 @@ "supports_vision": true, "supports_web_search": true }, + "deep-research-pro-preview-12-2025": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_image_token": 0.00012, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true + }, "gemini-2.5-flash-lite": { "cache_read_input_token_cost": 1e-08, "input_cost_per_audio_token": 3e-07, @@ -13304,7 +13339,8 @@ "supports_tool_choice": true, "supports_video_input": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_native_streaming": true }, "vertex_ai/gemini-3-pro-preview": { "cache_read_input_token_cost": 2e-07, @@ -13352,7 +13388,8 @@ "supports_tool_choice": true, "supports_video_input": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_native_streaming": true }, "vertex_ai/gemini-3-flash-preview": { "cache_read_input_token_cost": 5e-08, @@ -13395,7 +13432,8 @@ "supports_tool_choice": true, "supports_video_input": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_native_streaming": true }, "gemini-2.5-pro-exp-03-25": { "cache_read_input_token_cost": 1.25e-07, @@ -14762,6 +14800,42 @@ "supports_vision": true, "supports_web_search": true }, + "gemini/deep-research-pro-preview-12-2025": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "gemini", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_image_token": 0.00012, + "output_cost_per_token": 1.2e-05, + "rpm": 1000, + "tpm": 4000000, + "output_cost_per_token_batches": 6e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true + }, "gemini/gemini-2.5-flash-lite": { "cache_read_input_token_cost": 1e-08, "input_cost_per_audio_token": 3e-07, @@ -15346,6 +15420,7 @@ "supports_url_context": true, "supports_vision": true, "supports_web_search": true, + "supports_native_streaming": true, "tpm": 800000 }, "gemini-3-flash-preview": { @@ -15391,7 +15466,8 @@ "supports_tool_choice": true, "supports_url_context": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_native_streaming": true }, "gemini/gemini-2.5-pro-exp-03-25": { "cache_read_input_token_cost": 0.0, @@ -24343,6 +24419,31 @@ "supports_tool_choice": true, "supports_function_calling": true }, + "openrouter/qwen/qwen3-235b-a22b-2507": { + "input_cost_per_token": 7.1e-08, + "litellm_provider": "openrouter", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 1e-07, + "source": "https://openrouter.ai/qwen/qwen3-235b-a22b-2507", + "supports_function_calling": true, + "supports_tool_choice": true + }, + "openrouter/qwen/qwen3-235b-a22b-thinking-2507": { + "input_cost_per_token": 1.1e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 6e-07, + "source": "https://openrouter.ai/qwen/qwen3-235b-a22b-thinking-2507", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "openrouter/switchpoint/router": { "input_cost_per_token": 8.5e-07, "litellm_provider": "openrouter", @@ -27857,7 +27958,9 @@ "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", - "output_cost_per_token": 3e-07 + "output_cost_per_token": 3e-07, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/alibaba/qwen3-coder": { "input_cost_per_token": 4e-07, @@ -27866,7 +27969,9 @@ "max_output_tokens": 66536, "max_tokens": 66536, "mode": "chat", - "output_cost_per_token": 1.6e-06 + "output_cost_per_token": 1.6e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/amazon/nova-lite": { "input_cost_per_token": 6e-08, @@ -27875,7 +27980,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 2.4e-07 + "output_cost_per_token": 2.4e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_response_schema": true }, "vercel_ai_gateway/amazon/nova-micro": { "input_cost_per_token": 3.5e-08, @@ -27884,7 +27992,9 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 1.4e-07 + "output_cost_per_token": 1.4e-07, + "supports_function_calling": true, + "supports_response_schema": true }, "vercel_ai_gateway/amazon/nova-pro": { "input_cost_per_token": 8e-07, @@ -27893,7 +28003,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 3.2e-06 + "output_cost_per_token": 3.2e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_response_schema": true }, "vercel_ai_gateway/amazon/titan-embed-text-v2": { "input_cost_per_token": 2e-08, @@ -27913,7 +28026,11 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 1.25e-06 + "output_cost_per_token": 1.25e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-3-opus": { "cache_creation_input_token_cost": 1.875e-05, @@ -27924,7 +28041,11 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 7.5e-05 + "output_cost_per_token": 7.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-3.5-haiku": { "cache_creation_input_token_cost": 1e-06, @@ -27935,7 +28056,11 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 4e-06 + "output_cost_per_token": 4e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-3.5-sonnet": { "cache_creation_input_token_cost": 3.75e-06, @@ -27946,7 +28071,11 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-3.7-sonnet": { "cache_creation_input_token_cost": 3.75e-06, @@ -27957,7 +28086,11 @@ "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-4-opus": { "cache_creation_input_token_cost": 1.875e-05, @@ -27968,7 +28101,11 @@ "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", - "output_cost_per_token": 7.5e-05 + "output_cost_per_token": 7.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-4-sonnet": { "cache_creation_input_token_cost": 3.75e-06, @@ -27979,7 +28116,9 @@ "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/cohere/command-a": { "input_cost_per_token": 2.5e-06, @@ -27988,7 +28127,9 @@ "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/cohere/command-r": { "input_cost_per_token": 1.5e-07, @@ -27997,7 +28138,9 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 6e-07 + "output_cost_per_token": 6e-07, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/cohere/command-r-plus": { "input_cost_per_token": 2.5e-06, @@ -28006,7 +28149,9 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/cohere/embed-v4.0": { "input_cost_per_token": 1.2e-07, @@ -28024,7 +28169,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 2.19e-06 + "output_cost_per_token": 2.19e-06, + "supports_tool_choice": true }, "vercel_ai_gateway/deepseek/deepseek-r1-distill-llama-70b": { "input_cost_per_token": 7.5e-07, @@ -28033,7 +28179,10 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 9.9e-07 + "output_cost_per_token": 9.9e-07, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/deepseek/deepseek-v3": { "input_cost_per_token": 9e-07, @@ -28042,7 +28191,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 9e-07 + "output_cost_per_token": 9e-07, + "supports_tool_choice": true }, "vercel_ai_gateway/google/gemini-2.0-flash": { "deprecation_date": "2026-03-31", @@ -28052,7 +28202,11 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 6e-07 + "output_cost_per_token": 6e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/google/gemini-2.0-flash-lite": { "deprecation_date": "2026-03-31", @@ -28062,7 +28216,11 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 3e-07 + "output_cost_per_token": 3e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/google/gemini-2.5-flash": { "input_cost_per_token": 3e-07, @@ -28071,7 +28229,11 @@ "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", - "output_cost_per_token": 2.5e-06 + "output_cost_per_token": 2.5e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/google/gemini-2.5-pro": { "input_cost_per_token": 2.5e-06, @@ -28080,7 +28242,11 @@ "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/google/gemini-embedding-001": { "input_cost_per_token": 1.5e-07, @@ -28098,7 +28264,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 2e-07 + "output_cost_per_token": 2e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/google/text-embedding-005": { "input_cost_per_token": 2.5e-08, @@ -28134,7 +28303,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 7.9e-07 + "output_cost_per_token": 7.9e-07, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-3-8b": { "input_cost_per_token": 5e-08, @@ -28143,7 +28313,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 8e-08 + "output_cost_per_token": 8e-08, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-3.1-70b": { "input_cost_per_token": 7.2e-07, @@ -28152,7 +28323,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 7.2e-07 + "output_cost_per_token": 7.2e-07, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-3.1-8b": { "input_cost_per_token": 5e-08, @@ -28161,7 +28333,9 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 8e-08 + "output_cost_per_token": 8e-08, + "supports_function_calling": true, + "supports_response_schema": true }, "vercel_ai_gateway/meta/llama-3.2-11b": { "input_cost_per_token": 1.6e-07, @@ -28170,7 +28344,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 1.6e-07 + "output_cost_per_token": 1.6e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-3.2-1b": { "input_cost_per_token": 1e-07, @@ -28188,7 +28365,9 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 1.5e-07 + "output_cost_per_token": 1.5e-07, + "supports_function_calling": true, + "supports_response_schema": true }, "vercel_ai_gateway/meta/llama-3.2-90b": { "input_cost_per_token": 7.2e-07, @@ -28197,7 +28376,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 7.2e-07 + "output_cost_per_token": 7.2e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-3.3-70b": { "input_cost_per_token": 7.2e-07, @@ -28206,7 +28388,9 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 7.2e-07 + "output_cost_per_token": 7.2e-07, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-4-maverick": { "input_cost_per_token": 2e-07, @@ -28215,7 +28399,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 6e-07 + "output_cost_per_token": 6e-07, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-4-scout": { "input_cost_per_token": 1e-07, @@ -28224,7 +28409,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 3e-07 + "output_cost_per_token": 3e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/mistral/codestral": { "input_cost_per_token": 3e-07, @@ -28233,7 +28421,9 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 9e-07 + "output_cost_per_token": 9e-07, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/mistral/codestral-embed": { "input_cost_per_token": 1.5e-07, @@ -28251,7 +28441,10 @@ "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 2.8e-07 + "output_cost_per_token": 2.8e-07, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/mistral/magistral-medium": { "input_cost_per_token": 2e-06, @@ -28260,7 +28453,10 @@ "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", - "output_cost_per_token": 5e-06 + "output_cost_per_token": 5e-06, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/mistral/magistral-small": { "input_cost_per_token": 5e-07, @@ -28269,7 +28465,8 @@ "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", - "output_cost_per_token": 1.5e-06 + "output_cost_per_token": 1.5e-06, + "supports_function_calling": true }, "vercel_ai_gateway/mistral/ministral-3b": { "input_cost_per_token": 4e-08, @@ -28278,7 +28475,9 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 4e-08 + "output_cost_per_token": 4e-08, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/mistral/ministral-8b": { "input_cost_per_token": 1e-07, @@ -28287,7 +28486,10 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 1e-07 + "output_cost_per_token": 1e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/mistral/mistral-embed": { "input_cost_per_token": 1e-07, @@ -28305,7 +28507,9 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 6e-06 + "output_cost_per_token": 6e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/mistral/mistral-saba-24b": { "input_cost_per_token": 7.9e-07, @@ -28323,7 +28527,10 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 3e-07 + "output_cost_per_token": 3e-07, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/mistral/mixtral-8x22b-instruct": { "input_cost_per_token": 1.2e-06, @@ -28332,7 +28539,8 @@ "max_output_tokens": 2048, "max_tokens": 2048, "mode": "chat", - "output_cost_per_token": 1.2e-06 + "output_cost_per_token": 1.2e-06, + "supports_function_calling": true }, "vercel_ai_gateway/mistral/pixtral-12b": { "input_cost_per_token": 1.5e-07, @@ -28341,7 +28549,11 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 1.5e-07 + "output_cost_per_token": 1.5e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/mistral/pixtral-large": { "input_cost_per_token": 2e-06, @@ -28350,7 +28562,11 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 6e-06 + "output_cost_per_token": 6e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/moonshotai/kimi-k2": { "input_cost_per_token": 5.5e-07, @@ -28359,7 +28575,9 @@ "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", - "output_cost_per_token": 2.2e-06 + "output_cost_per_token": 2.2e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/morph/morph-v3-fast": { "input_cost_per_token": 8e-07, @@ -28386,7 +28604,9 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 1.5e-06 + "output_cost_per_token": 1.5e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/openai/gpt-3.5-turbo-instruct": { "input_cost_per_token": 1.5e-06, @@ -28404,7 +28624,10 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 3e-05 + "output_cost_per_token": 3e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/openai/gpt-4.1": { "cache_creation_input_token_cost": 0.0, @@ -28415,7 +28638,11 @@ "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 8e-06 + "output_cost_per_token": 8e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/gpt-4.1-mini": { "cache_creation_input_token_cost": 0.0, @@ -28426,7 +28653,11 @@ "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 1.6e-06 + "output_cost_per_token": 1.6e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/gpt-4.1-nano": { "cache_creation_input_token_cost": 0.0, @@ -28437,7 +28668,11 @@ "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 4e-07 + "output_cost_per_token": 4e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/gpt-4o": { "cache_creation_input_token_cost": 0.0, @@ -28448,7 +28683,11 @@ "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/gpt-4o-mini": { "cache_creation_input_token_cost": 0.0, @@ -28459,7 +28698,11 @@ "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", - "output_cost_per_token": 6e-07 + "output_cost_per_token": 6e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/o1": { "cache_creation_input_token_cost": 0.0, @@ -28470,7 +28713,11 @@ "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", - "output_cost_per_token": 6e-05 + "output_cost_per_token": 6e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/o3": { "cache_creation_input_token_cost": 0.0, @@ -28481,7 +28728,11 @@ "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", - "output_cost_per_token": 8e-06 + "output_cost_per_token": 8e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/o3-mini": { "cache_creation_input_token_cost": 0.0, @@ -28492,7 +28743,10 @@ "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", - "output_cost_per_token": 4.4e-06 + "output_cost_per_token": 4.4e-06, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/o4-mini": { "cache_creation_input_token_cost": 0.0, @@ -28503,7 +28757,11 @@ "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", - "output_cost_per_token": 4.4e-06 + "output_cost_per_token": 4.4e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/text-embedding-3-large": { "input_cost_per_token": 1.3e-07, @@ -28575,7 +28833,10 @@ "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/vercel/v0-1.5-md": { "input_cost_per_token": 3e-06, @@ -28584,7 +28845,10 @@ "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-2": { "input_cost_per_token": 2e-06, @@ -28593,7 +28857,9 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-2-vision": { "input_cost_per_token": 2e-06, @@ -28602,7 +28868,10 @@ "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-3": { "input_cost_per_token": 3e-06, @@ -28611,7 +28880,9 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-3-fast": { "input_cost_per_token": 5e-06, @@ -28620,7 +28891,8 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 2.5e-05 + "output_cost_per_token": 2.5e-05, + "supports_function_calling": true }, "vercel_ai_gateway/xai/grok-3-mini": { "input_cost_per_token": 3e-07, @@ -28629,7 +28901,9 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 5e-07 + "output_cost_per_token": 5e-07, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-3-mini-fast": { "input_cost_per_token": 6e-07, @@ -28638,7 +28912,9 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 4e-06 + "output_cost_per_token": 4e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-4": { "input_cost_per_token": 3e-06, @@ -28647,7 +28923,9 @@ "max_output_tokens": 256000, "max_tokens": 256000, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/zai/glm-4.5": { "input_cost_per_token": 6e-07, @@ -28656,7 +28934,9 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 2.2e-06 + "output_cost_per_token": 2.2e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/zai/glm-4.5-air": { "input_cost_per_token": 2e-07, @@ -28665,7 +28945,9 @@ "max_output_tokens": 96000, "max_tokens": 96000, "mode": "chat", - "output_cost_per_token": 1.1e-06 + "output_cost_per_token": 1.1e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/zai/glm-4.6": { "litellm_provider": "vercel_ai_gateway", @@ -28733,7 +29015,9 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "supports_native_streaming": true, + "supports_vision": true }, "vertex_ai/claude-3-5-sonnet": { "input_cost_per_token": 3e-06, @@ -29004,7 +29288,8 @@ "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, - "tool_use_system_prompt_tokens": 159 + "tool_use_system_prompt_tokens": 159, + "supports_native_streaming": true }, "vertex_ai/claude-sonnet-4-5": { "cache_creation_input_token_cost": 3.75e-06, @@ -29056,7 +29341,8 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_native_streaming": true }, "vertex_ai/claude-opus-4@20250514": { "cache_creation_input_token_cost": 1.875e-05, @@ -29338,6 +29624,21 @@ "output_cost_per_token_batches": 6e-06, "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" }, + "vertex_ai/deep-research-pro-preview-12-2025": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_image_token": 0.00012, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" + }, "vertex_ai/imagegeneration@006": { "litellm_provider": "vertex_ai-image-models", "mode": "image_generation", @@ -29827,7 +30128,9 @@ "mode": "chat", "output_cost_per_token": 1e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supported_regions": ["global"], + "supported_regions": [ + "global" + ], "supports_function_calling": true, "supports_tool_choice": true }, @@ -29840,7 +30143,9 @@ "mode": "chat", "output_cost_per_token": 4e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supported_regions": ["global"], + "supported_regions": [ + "global" + ], "supports_function_calling": true, "supports_tool_choice": true }, @@ -29853,7 +30158,9 @@ "mode": "chat", "output_cost_per_token": 1.2e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supported_regions": ["global"], + "supported_regions": [ + "global" + ], "supports_function_calling": true, "supports_tool_choice": true }, @@ -29866,7 +30173,9 @@ "mode": "chat", "output_cost_per_token": 1.2e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supported_regions": ["global"], + "supported_regions": [ + "global" + ], "supports_function_calling": true, "supports_tool_choice": true }, @@ -34815,4 +35124,4 @@ "output_cost_per_token": 0, "supports_reasoning": true } -} \ No newline at end of file +} diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index 49d6ac7d898..7e70b5baae4 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -387,6 +387,9 @@ class MCPRequestHandler: user_api_key_cache, ) + verbose_logger.debug( + f"MCP team permission lookup: team_id={user_api_key_auth.team_id if user_api_key_auth else None}" + ) if not user_api_key_auth or not user_api_key_auth.team_id or not prisma_client: return None diff --git a/litellm/proxy/_experimental/out/404.html b/litellm/proxy/_experimental/out/404/index.html similarity index 100% rename from litellm/proxy/_experimental/out/404.html rename to litellm/proxy/_experimental/out/404/index.html diff --git a/litellm/proxy/_experimental/out/api-reference.html b/litellm/proxy/_experimental/out/api-reference/index.html similarity index 100% rename from litellm/proxy/_experimental/out/api-reference.html rename to litellm/proxy/_experimental/out/api-reference/index.html diff --git a/litellm/proxy/_experimental/out/experimental/api-playground.html b/litellm/proxy/_experimental/out/experimental/api-playground/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/api-playground.html rename to litellm/proxy/_experimental/out/experimental/api-playground/index.html diff --git a/litellm/proxy/_experimental/out/experimental/budgets.html b/litellm/proxy/_experimental/out/experimental/budgets/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/budgets.html rename to litellm/proxy/_experimental/out/experimental/budgets/index.html diff --git a/litellm/proxy/_experimental/out/experimental/caching.html b/litellm/proxy/_experimental/out/experimental/caching/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/caching.html rename to litellm/proxy/_experimental/out/experimental/caching/index.html diff --git a/litellm/proxy/_experimental/out/experimental/claude-code-plugins.html b/litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/claude-code-plugins.html rename to litellm/proxy/_experimental/out/experimental/claude-code-plugins/index.html diff --git a/litellm/proxy/_experimental/out/experimental/old-usage.html b/litellm/proxy/_experimental/out/experimental/old-usage/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/old-usage.html rename to litellm/proxy/_experimental/out/experimental/old-usage/index.html diff --git a/litellm/proxy/_experimental/out/experimental/prompts.html b/litellm/proxy/_experimental/out/experimental/prompts/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/prompts.html rename to litellm/proxy/_experimental/out/experimental/prompts/index.html diff --git a/litellm/proxy/_experimental/out/experimental/tag-management.html b/litellm/proxy/_experimental/out/experimental/tag-management/index.html similarity index 100% rename from litellm/proxy/_experimental/out/experimental/tag-management.html rename to litellm/proxy/_experimental/out/experimental/tag-management/index.html diff --git a/litellm/proxy/_experimental/out/guardrails.html b/litellm/proxy/_experimental/out/guardrails/index.html similarity index 100% rename from litellm/proxy/_experimental/out/guardrails.html rename to litellm/proxy/_experimental/out/guardrails/index.html diff --git a/litellm/proxy/_experimental/out/login.html b/litellm/proxy/_experimental/out/login/index.html similarity index 100% rename from litellm/proxy/_experimental/out/login.html rename to litellm/proxy/_experimental/out/login/index.html diff --git a/litellm/proxy/_experimental/out/logs.html b/litellm/proxy/_experimental/out/logs/index.html similarity index 100% rename from litellm/proxy/_experimental/out/logs.html rename to litellm/proxy/_experimental/out/logs/index.html diff --git a/litellm/proxy/_experimental/out/mcp/oauth/callback.html b/litellm/proxy/_experimental/out/mcp/oauth/callback/index.html similarity index 100% rename from litellm/proxy/_experimental/out/mcp/oauth/callback.html rename to litellm/proxy/_experimental/out/mcp/oauth/callback/index.html diff --git a/litellm/proxy/_experimental/out/model-hub.html b/litellm/proxy/_experimental/out/model-hub/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model-hub.html rename to litellm/proxy/_experimental/out/model-hub/index.html diff --git a/litellm/proxy/_experimental/out/model_hub.html b/litellm/proxy/_experimental/out/model_hub/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model_hub.html rename to litellm/proxy/_experimental/out/model_hub/index.html diff --git a/litellm/proxy/_experimental/out/model_hub_table.html b/litellm/proxy/_experimental/out/model_hub_table/index.html similarity index 100% rename from litellm/proxy/_experimental/out/model_hub_table.html rename to litellm/proxy/_experimental/out/model_hub_table/index.html diff --git a/litellm/proxy/_experimental/out/models-and-endpoints.html b/litellm/proxy/_experimental/out/models-and-endpoints/index.html similarity index 100% rename from litellm/proxy/_experimental/out/models-and-endpoints.html rename to litellm/proxy/_experimental/out/models-and-endpoints/index.html diff --git a/litellm/proxy/_experimental/out/onboarding.html b/litellm/proxy/_experimental/out/onboarding/index.html similarity index 100% rename from litellm/proxy/_experimental/out/onboarding.html rename to litellm/proxy/_experimental/out/onboarding/index.html diff --git a/litellm/proxy/_experimental/out/organizations.html b/litellm/proxy/_experimental/out/organizations/index.html similarity index 100% rename from litellm/proxy/_experimental/out/organizations.html rename to litellm/proxy/_experimental/out/organizations/index.html diff --git a/litellm/proxy/_experimental/out/playground.html b/litellm/proxy/_experimental/out/playground/index.html similarity index 100% rename from litellm/proxy/_experimental/out/playground.html rename to litellm/proxy/_experimental/out/playground/index.html diff --git a/litellm/proxy/_experimental/out/policies.html b/litellm/proxy/_experimental/out/policies/index.html similarity index 100% rename from litellm/proxy/_experimental/out/policies.html rename to litellm/proxy/_experimental/out/policies/index.html diff --git a/litellm/proxy/_experimental/out/settings/admin-settings.html b/litellm/proxy/_experimental/out/settings/admin-settings/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/admin-settings.html rename to litellm/proxy/_experimental/out/settings/admin-settings/index.html diff --git a/litellm/proxy/_experimental/out/settings/logging-and-alerts.html b/litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/logging-and-alerts.html rename to litellm/proxy/_experimental/out/settings/logging-and-alerts/index.html diff --git a/litellm/proxy/_experimental/out/settings/router-settings.html b/litellm/proxy/_experimental/out/settings/router-settings/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/router-settings.html rename to litellm/proxy/_experimental/out/settings/router-settings/index.html diff --git a/litellm/proxy/_experimental/out/settings/ui-theme.html b/litellm/proxy/_experimental/out/settings/ui-theme/index.html similarity index 100% rename from litellm/proxy/_experimental/out/settings/ui-theme.html rename to litellm/proxy/_experimental/out/settings/ui-theme/index.html diff --git a/litellm/proxy/_experimental/out/teams.html b/litellm/proxy/_experimental/out/teams/index.html similarity index 100% rename from litellm/proxy/_experimental/out/teams.html rename to litellm/proxy/_experimental/out/teams/index.html diff --git a/litellm/proxy/_experimental/out/test-key.html b/litellm/proxy/_experimental/out/test-key/index.html similarity index 100% rename from litellm/proxy/_experimental/out/test-key.html rename to litellm/proxy/_experimental/out/test-key/index.html diff --git a/litellm/proxy/_experimental/out/tools/mcp-servers.html b/litellm/proxy/_experimental/out/tools/mcp-servers/index.html similarity index 100% rename from litellm/proxy/_experimental/out/tools/mcp-servers.html rename to litellm/proxy/_experimental/out/tools/mcp-servers/index.html diff --git a/litellm/proxy/_experimental/out/tools/vector-stores.html b/litellm/proxy/_experimental/out/tools/vector-stores/index.html similarity index 100% rename from litellm/proxy/_experimental/out/tools/vector-stores.html rename to litellm/proxy/_experimental/out/tools/vector-stores/index.html diff --git a/litellm/proxy/_experimental/out/usage.html b/litellm/proxy/_experimental/out/usage/index.html similarity index 100% rename from litellm/proxy/_experimental/out/usage.html rename to litellm/proxy/_experimental/out/usage/index.html diff --git a/litellm/proxy/_experimental/out/users.html b/litellm/proxy/_experimental/out/users/index.html similarity index 100% rename from litellm/proxy/_experimental/out/users.html rename to litellm/proxy/_experimental/out/users/index.html diff --git a/litellm/proxy/_experimental/out/virtual-keys.html b/litellm/proxy/_experimental/out/virtual-keys/index.html similarity index 100% rename from litellm/proxy/_experimental/out/virtual-keys.html rename to litellm/proxy/_experimental/out/virtual-keys/index.html diff --git a/litellm/proxy/_new_secret_config.yaml b/litellm/proxy/_new_secret_config.yaml index 13eeae14485..6f527e268b2 100644 --- a/litellm/proxy/_new_secret_config.yaml +++ b/litellm/proxy/_new_secret_config.yaml @@ -14,3 +14,14 @@ model_list: litellm_params: model: openai/gpt-4.1-mini +guardrails: + - guardrail_name: redact-ssn + litellm_params: + guardrail: custom_code + mode: pre_call + custom_code: | + def apply_guardrail(inputs, request_data, input_type): + for text in inputs["texts"]: + if regex_match(text, r"\d{3}-\d{2}-\d{4}"): + return block("SSN detected in message") + return allow() \ No newline at end of file diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 9ae95085f55..131a6caab07 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3673,7 +3673,7 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): team_id_upsert: bool = False team_ids_jwt_field: Optional[str] = None upsert_sso_user_to_team: bool = False - team_allowed_routes: List[str] = ["openai_routes", "info_routes"] + team_allowed_routes: List[str] = ["openai_routes", "info_routes", "mcp_routes"] team_id_default: Optional[str] = Field( default=None, description="If no team_id given, default permissions/spend-tracking to this team.s", diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py new file mode 100644 index 00000000000..cb277d44ee9 --- /dev/null +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -0,0 +1,53 @@ +""" +A2A Agent Routing + +Handles routing for A2A agents (models with "a2a/" prefix). +Looks up agents in the registry and injects their API base URL. +""" + +from typing import Any, Optional + +import litellm +from litellm._logging import verbose_proxy_logger + + +def route_a2a_agent_request(data: dict, route_type: str) -> Optional[Any]: + """ + Route A2A agent requests directly to litellm with injected API base. + + Returns None if not an A2A request (allows normal routing to continue). + """ + # Import here to avoid circular imports + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + from litellm.proxy.route_llm_request import ( + ROUTE_ENDPOINT_MAPPING, + ProxyModelNotFoundError, + ) + + model_name = data.get("model", "") + + # Check if this is an A2A agent request + if not isinstance(model_name, str) or not model_name.startswith("a2a/"): + return None + + # Extract agent name (e.g., "a2a/my-agent" -> "my-agent") + agent_name = model_name[4:] + + # Look up agent in registry + agent = global_agent_registry.get_agent_by_name(agent_name) + if agent is None: + verbose_proxy_logger.error(f"[A2A] Agent '{agent_name}' not found in registry") + route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type) + raise ProxyModelNotFoundError(route=route_name, model_name=model_name) + + # Get API base URL from agent config + if not agent.agent_card_params or "url" not in agent.agent_card_params: + verbose_proxy_logger.error(f"[A2A] Agent '{agent_name}' has no URL configured") + route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type) + raise ProxyModelNotFoundError(route=route_name, model_name=model_name) + + # Inject API base and route to litellm + data["api_base"] = agent.agent_card_params["url"] + verbose_proxy_logger.debug(f"[A2A] Routing {model_name} to {data['api_base']}") + + return getattr(litellm, f"{route_type}")(**data) diff --git a/litellm/proxy/agent_endpoints/model_list_helpers.py b/litellm/proxy/agent_endpoints/model_list_helpers.py new file mode 100644 index 00000000000..c640300bb8c --- /dev/null +++ b/litellm/proxy/agent_endpoints/model_list_helpers.py @@ -0,0 +1,96 @@ +""" +Helper functions for appending A2A agents to model lists. + +Used by proxy model endpoints to make agents appear in UI alongside models. +""" +from typing import List + +from litellm._logging import verbose_proxy_logger +from litellm.proxy._types import UserAPIKeyAuth +from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + ModelGroupInfoProxy, +) + + +async def append_agents_to_model_group( + model_groups: List[ModelGroupInfoProxy], + user_api_key_dict: UserAPIKeyAuth, +) -> List[ModelGroupInfoProxy]: + """ + Append A2A agents to model groups list for UI display. + + Converts agents to model format with "a2a/" naming + so they appear in playground and work with LiteLLM routing. + """ + try: + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( + AgentRequestHandler, + ) + + allowed_agent_ids = await AgentRequestHandler.get_allowed_agents( + user_api_key_auth=user_api_key_dict + ) + + for agent_id in allowed_agent_ids: + agent = global_agent_registry.get_agent_by_id(agent_id) + if agent is not None: + model_groups.append( + ModelGroupInfoProxy( + model_group=f"a2a/{agent.agent_name}", + mode="chat", + providers=["a2a"], + ) + ) + except Exception as e: + verbose_proxy_logger.debug( + f"Error appending agents to model_group/info: {e}" + ) + + return model_groups + + +async def append_agents_to_model_info( + models: List[dict], + user_api_key_dict: UserAPIKeyAuth, +) -> List[dict]: + """ + Append A2A agents to model info list for UI display. + + Converts agents to model format with "a2a/" naming + so they appear in models page and work with LiteLLM routing. + """ + try: + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + from litellm.proxy.agent_endpoints.auth.agent_permission_handler import ( + AgentRequestHandler, + ) + + allowed_agent_ids = await AgentRequestHandler.get_allowed_agents( + user_api_key_auth=user_api_key_dict + ) + + for agent_id in allowed_agent_ids: + agent = global_agent_registry.get_agent_by_id(agent_id) + if agent is not None: + models.append({ + "model_name": f"a2a/{agent.agent_name}", + "litellm_params": { + "model": f"a2a/{agent.agent_name}", + "custom_llm_provider": "a2a", + }, + "model_info": { + "id": agent.agent_id, + "mode": "chat", + "db_model": True, + "created_by": agent.created_by, + "created_at": agent.created_at, + "updated_at": agent.updated_at, + }, + }) + except Exception as e: + verbose_proxy_logger.debug( + f"Error appending agents to v2/model/info: {e}" + ) + + return models diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 33667b5d8d9..584be0a9496 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -976,6 +976,9 @@ class JWTAuthManager: user_route=route, litellm_proxy_roles=jwt_handler.litellm_jwtauth, ) + verbose_proxy_logger.debug( + f"JWT team route check: team_id={team_id}, route={route}, is_allowed={is_allowed}" + ) if is_allowed: return team_id, team_object except Exception: diff --git a/litellm/proxy/batches_endpoints/endpoints.py b/litellm/proxy/batches_endpoints/endpoints.py index f47e2e1667b..06800cb4524 100644 --- a/litellm/proxy/batches_endpoints/endpoints.py +++ b/litellm/proxy/batches_endpoints/endpoints.py @@ -24,10 +24,12 @@ from litellm.proxy.openai_files_endpoints.common_utils import ( _is_base64_encoded_unified_file_id, decode_model_from_file_id, encode_file_id_with_model, + get_batch_from_database, get_credentials_for_model, get_models_from_unified_file_id, get_original_file_id, prepare_data_with_credentials, + update_batch_in_database, ) from litellm.proxy.utils import handle_exception_on_proxy, is_known_model from litellm.types.llms.openai import LiteLLMBatchCreateRequest @@ -357,6 +359,57 @@ async def retrieve_batch( route_type="aretrieve_batch", ) + # FIX: First, try to read from ManagedObjectTable for consistent state + managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") + from litellm.proxy.proxy_server import prisma_client + + db_batch_object, response = await get_batch_from_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + managed_files_obj=managed_files_obj, + prisma_client=prisma_client, + verbose_proxy_logger=verbose_proxy_logger, + ) + + # If batch is in a terminal state, return immediately + if response is not None and response.status in ["completed", "failed", "cancelled", "expired"]: + # Call hooks and return + response = await proxy_logging_obj.post_call_success_hook( + data=data, user_api_key_dict=user_api_key_dict, response=response + ) + + asyncio.create_task( + proxy_logging_obj.update_request_status( + litellm_call_id=data.get("litellm_call_id", ""), status="success" + ) + ) + + hidden_params = getattr(response, "_hidden_params", {}) or {} + model_id = hidden_params.get("model_id", None) or "" + cache_key = hidden_params.get("cache_key", None) or "" + api_base = hidden_params.get("api_base", None) or "" + + fastapi_response.headers.update( + ProxyBaseLLMRequestProcessing.get_custom_headers( + user_api_key_dict=user_api_key_dict, + model_id=model_id, + cache_key=cache_key, + api_base=api_base, + version=version, + model_region=getattr(user_api_key_dict, "allowed_model_region", ""), + request_data=data, + ) + ) + + return response + + # If batch is still processing, sync with provider to get latest state + if response is not None: + verbose_proxy_logger.debug( + f"Batch {batch_id} is in non-terminal state {response.status}, syncing with provider" + ) + + # Retrieve from provider (for non-terminal states or if DB lookup failed) # SCENARIO 1: Batch ID is encoded with model info if model_from_id is not None: credentials = get_credentials_for_model( @@ -408,6 +461,18 @@ async def retrieve_batch( response = await litellm.aretrieve_batch( custom_llm_provider=custom_llm_provider, **data # type: ignore ) + + # FIX: Update the database with the latest state from provider + await update_batch_in_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + response=response, + managed_files_obj=managed_files_obj, + prisma_client=prisma_client, + verbose_proxy_logger=verbose_proxy_logger, + db_batch_object=db_batch_object, + operation="retrieve", + ) ### CALL HOOKS ### - modify outgoing data response = await proxy_logging_obj.post_call_success_hook( @@ -769,6 +834,20 @@ async def cancel_batch( **_cancel_batch_data, ) + # FIX: Update the database with the new cancelled state + managed_files_obj = proxy_logging_obj.get_proxy_hook("managed_files") + from litellm.proxy.proxy_server import prisma_client + + await update_batch_in_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + response=response, + managed_files_obj=managed_files_obj, + prisma_client=prisma_client, + verbose_proxy_logger=verbose_proxy_logger, + operation="cancel", + ) + ### CALL HOOKS ### - modify outgoing data response = await proxy_logging_obj.post_call_success_hook( data=data, user_api_key_dict=user_api_key_dict, response=response diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 429e56c805b..dc928921425 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -1187,119 +1187,130 @@ class DBSpendUpdateWriter: ) break - async with prisma_client.db.batch_() as batcher: - for _, transaction in transactions_to_process.items(): - entity_id = transaction.get(entity_id_field) + try: + async with prisma_client.db.batch_() as batcher: + for _, transaction in transactions_to_process.items(): + entity_id = transaction.get(entity_id_field) - # Construct the where clause dynamically - where_clause = { - unique_constraint_name: { + # Construct the where clause dynamically + where_clause = { + unique_constraint_name: { + entity_id_field: entity_id, + "date": transaction["date"], + "api_key": transaction["api_key"], + "model": transaction["model"], + "custom_llm_provider": transaction.get( + "custom_llm_provider" + ) + or "", + "mcp_namespaced_tool_name": transaction.get( + "mcp_namespaced_tool_name" + ) + or "", + "endpoint": transaction.get("endpoint") or "", + } + } + + # Get the table dynamically + table = getattr(batcher, table_name) + + # Common data structure for both create and update + common_data = { entity_id_field: entity_id, "date": transaction["date"], "api_key": transaction["api_key"], - "model": transaction["model"], - "custom_llm_provider": transaction.get( - "custom_llm_provider" - ) - or "", + "model": transaction.get("model"), + "model_group": transaction.get("model_group"), "mcp_namespaced_tool_name": transaction.get( "mcp_namespaced_tool_name" ) or "", + "custom_llm_provider": transaction.get( + "custom_llm_provider" + ), "endpoint": transaction.get("endpoint") or "", + "prompt_tokens": transaction["prompt_tokens"], + "completion_tokens": transaction["completion_tokens"], + "spend": transaction["spend"], + "api_requests": transaction["api_requests"], + "successful_requests": transaction[ + "successful_requests" + ], + "failed_requests": transaction["failed_requests"], } - } - # Get the table dynamically - table = getattr(batcher, table_name) - - # Common data structure for both create and update - common_data = { - entity_id_field: entity_id, - "date": transaction["date"], - "api_key": transaction["api_key"], - "model": transaction.get("model"), - "model_group": transaction.get("model_group"), - "mcp_namespaced_tool_name": transaction.get( - "mcp_namespaced_tool_name" - ) - or "", - "custom_llm_provider": transaction.get( - "custom_llm_provider" - ), - "endpoint": transaction.get("endpoint"), - "prompt_tokens": transaction["prompt_tokens"], - "completion_tokens": transaction["completion_tokens"], - "spend": transaction["spend"], - "api_requests": transaction["api_requests"], - "successful_requests": transaction[ - "successful_requests" - ], - "failed_requests": transaction["failed_requests"], - } - - # Add cache-related fields if they exist - if "cache_read_input_tokens" in transaction: - common_data["cache_read_input_tokens"] = ( - transaction.get("cache_read_input_tokens", 0) - ) - if "cache_creation_input_tokens" in transaction: - common_data["cache_creation_input_tokens"] = ( - transaction.get("cache_creation_input_tokens", 0) - ) - - if entity_type == "tag" and "request_id" in transaction: - common_data["request_id"] = transaction.get( - "request_id" - ) - - # Create update data structure - update_data = { - "prompt_tokens": { - "increment": transaction["prompt_tokens"] - }, - "completion_tokens": { - "increment": transaction["completion_tokens"] - }, - "spend": {"increment": transaction["spend"]}, - "api_requests": { - "increment": transaction["api_requests"] - }, - "successful_requests": { - "increment": transaction["successful_requests"] - }, - "failed_requests": { - "increment": transaction["failed_requests"] - }, - } - - # Add cache-related fields to update if they exist - if "cache_read_input_tokens" in transaction: - update_data["cache_read_input_tokens"] = { - "increment": transaction.get( - "cache_read_input_tokens", 0 + # Add cache-related fields if they exist + if "cache_read_input_tokens" in transaction: + common_data["cache_read_input_tokens"] = ( + transaction.get("cache_read_input_tokens", 0) ) - } - if "cache_creation_input_tokens" in transaction: - update_data["cache_creation_input_tokens"] = { - "increment": transaction.get( - "cache_creation_input_tokens", 0 + if "cache_creation_input_tokens" in transaction: + common_data["cache_creation_input_tokens"] = ( + transaction.get("cache_creation_input_tokens", 0) ) + + if entity_type == "tag" and "request_id" in transaction: + common_data["request_id"] = transaction.get( + "request_id" + ) + + # Create update data structure + update_data = { + "prompt_tokens": { + "increment": transaction["prompt_tokens"] + }, + "completion_tokens": { + "increment": transaction["completion_tokens"] + }, + "spend": {"increment": transaction["spend"]}, + "api_requests": { + "increment": transaction["api_requests"] + }, + "successful_requests": { + "increment": transaction["successful_requests"] + }, + "failed_requests": { + "increment": transaction["failed_requests"] + }, } - if entity_type == "tag" and "request_id" in transaction: - update_data["request_id"] = transaction.get("request_id") + # Add cache-related fields to update if they exist + if "cache_read_input_tokens" in transaction: + update_data["cache_read_input_tokens"] = { + "increment": transaction.get( + "cache_read_input_tokens", 0 + ) + } + if "cache_creation_input_tokens" in transaction: + update_data["cache_creation_input_tokens"] = { + "increment": transaction.get( + "cache_creation_input_tokens", 0 + ) + } - # Add endpoint to update_data so existing rows get their endpoint field updated - update_data["endpoint"] = transaction.get("endpoint") or "" + if entity_type == "tag" and "request_id" in transaction: + update_data["request_id"] = transaction.get("request_id") - table.upsert( - where=where_clause, - data={ - "create": common_data, - "update": update_data, - }, - ) + # Add endpoint to update_data so existing rows get their endpoint field updated + update_data["endpoint"] = transaction.get("endpoint") or "" + + table.upsert( + where=where_clause, + data={ + "create": common_data, + "update": update_data, + }, + ) + except Exception as batch_error: + # Log detailed error information for debugging batch upsert failures + # This helps diagnose issues like unique constraint violations + verbose_proxy_logger.exception( + f"Daily {entity_type} spend batch upsert failed. " + f"Table: {table_name}, Constraint: {unique_constraint_name}, " + f"Batch size: {len(transactions_to_process)}, " + f"Error: {str(batch_error)}" + ) + raise verbose_proxy_logger.debug( f"Processed {len(transactions_to_process)} daily {entity_type} transactions in {time.time() - start_time:.2f}s" diff --git a/litellm/proxy/guardrails/guardrail_endpoints.py b/litellm/proxy/guardrails/guardrail_endpoints.py index 3ce819439cb..a825ce22b25 100644 --- a/litellm/proxy/guardrails/guardrail_endpoints.py +++ b/litellm/proxy/guardrails/guardrail_endpoints.py @@ -1236,6 +1236,275 @@ async def get_provider_specific_params(): return provider_params +class TestCustomCodeGuardrailRequest(BaseModel): + """Request model for testing custom code guardrails.""" + + custom_code: str + """The Python-like code containing the apply_guardrail function.""" + + test_input: Dict[str, Any] + """The test input to pass to the guardrail. Should contain 'texts', optionally 'images', 'tools', etc.""" + + input_type: str = "request" + """Whether this is a 'request' or 'response' input type.""" + + request_data: Optional[Dict[str, Any]] = None + """Optional mock request_data (model, user_id, team_id, metadata, etc.).""" + + +class TestCustomCodeGuardrailResponse(BaseModel): + """Response model for testing custom code guardrails.""" + + success: bool + """Whether the test executed successfully (no errors).""" + + result: Optional[Dict[str, Any]] = None + """The guardrail result: action (allow/block/modify), reason, modified_texts, etc.""" + + error: Optional[str] = None + """Error message if execution failed.""" + + error_type: Optional[str] = None + """Type of error: 'compilation' or 'execution'.""" + + +@router.post( + "/guardrails/test_custom_code", + tags=["Guardrails"], + dependencies=[Depends(user_api_key_auth)], + response_model=TestCustomCodeGuardrailResponse, +) +async def test_custom_code_guardrail(request: TestCustomCodeGuardrailRequest): + """ + Test custom code guardrail logic without creating a guardrail. + + This endpoint allows admins to experiment with custom code guardrails by: + 1. Compiling the provided code in a sandbox + 2. Executing the apply_guardrail function with test input + 3. Returning the result (allow/block/modify) + + šŸ‘‰ [Custom Code Guardrail docs](https://docs.litellm.ai/docs/proxy/guardrails/custom_code_guardrail) + + Example Request: + ```bash + curl -X POST "http://localhost:4000/guardrails/test_custom_code" \\ + -H "Authorization: Bearer " \\ + -H "Content-Type: application/json" \\ + -d '{ + "custom_code": "def apply_guardrail(inputs, request_data, input_type):\\n for text in inputs[\\"texts\\"]:\\n if regex_match(text, r\\"\\\\d{3}-\\\\d{2}-\\\\d{4}\\"):\\n return block(\\"SSN detected\\")\\n return allow()", + "test_input": { + "texts": ["My SSN is 123-45-6789"] + }, + "input_type": "request" + }' + ``` + + Example Success Response (blocked): + ```json + { + "success": true, + "result": { + "action": "block", + "reason": "SSN detected" + }, + "error": null, + "error_type": null + } + ``` + + Example Success Response (allowed): + ```json + { + "success": true, + "result": { + "action": "allow" + }, + "error": null, + "error_type": null + } + ``` + + Example Success Response (modified): + ```json + { + "success": true, + "result": { + "action": "modify", + "texts": ["My SSN is [REDACTED]"] + }, + "error": null, + "error_type": null + } + ``` + + Example Error Response (compilation error): + ```json + { + "success": false, + "result": null, + "error": "Syntax error in custom code: invalid syntax (, line 1)", + "error_type": "compilation" + } + ``` + """ + import concurrent.futures + import re + + from litellm.proxy.guardrails.guardrail_hooks.custom_code.primitives import ( + get_custom_code_primitives, + ) + + # Security validation patterns + FORBIDDEN_PATTERNS = [ + # Import statements + (r"\bimport\s+", "import statements are not allowed"), + (r"\bfrom\s+\w+\s+import\b", "from...import statements are not allowed"), + (r"__import__\s*\(", "__import__() is not allowed"), + # Dangerous builtins + (r"\bexec\s*\(", "exec() is not allowed"), + (r"\beval\s*\(", "eval() is not allowed"), + (r"\bcompile\s*\(", "compile() is not allowed"), + (r"\bopen\s*\(", "open() is not allowed"), + (r"\bgetattr\s*\(", "getattr() is not allowed"), + (r"\bsetattr\s*\(", "setattr() is not allowed"), + (r"\bdelattr\s*\(", "delattr() is not allowed"), + (r"\bglobals\s*\(", "globals() is not allowed"), + (r"\blocals\s*\(", "locals() is not allowed"), + (r"\bvars\s*\(", "vars() is not allowed"), + (r"\bdir\s*\(", "dir() is not allowed"), + (r"\bbreakpoint\s*\(", "breakpoint() is not allowed"), + (r"\binput\s*\(", "input() is not allowed"), + # Dangerous dunder access + (r"__builtins__", "__builtins__ access is not allowed"), + (r"__globals__", "__globals__ access is not allowed"), + (r"__code__", "__code__ access is not allowed"), + (r"__subclasses__", "__subclasses__ access is not allowed"), + (r"__bases__", "__bases__ access is not allowed"), + (r"__mro__", "__mro__ access is not allowed"), + (r"__class__", "__class__ access is not allowed"), + (r"__dict__", "__dict__ access is not allowed"), + (r"__getattribute__", "__getattribute__ access is not allowed"), + (r"__reduce__", "__reduce__ access is not allowed"), + (r"__reduce_ex__", "__reduce_ex__ access is not allowed"), + # OS/system access + (r"\bos\.", "os module access is not allowed"), + (r"\bsys\.", "sys module access is not allowed"), + (r"\bsubprocess\.", "subprocess module access is not allowed"), + ] + + EXECUTION_TIMEOUT_SECONDS = 5 + + try: + # Step 0: Security validation - check for forbidden patterns + code = request.custom_code + for pattern, error_msg in FORBIDDEN_PATTERNS: + if re.search(pattern, code): + return TestCustomCodeGuardrailResponse( + success=False, + error=f"Security violation: {error_msg}", + error_type="compilation", + ) + + # Step 1: Compile the custom code with restricted environment + exec_globals = get_custom_code_primitives().copy() + + # Remove access to builtins to prevent escape + exec_globals["__builtins__"] = {} + + try: + exec(compile(request.custom_code, "", "exec"), exec_globals) + except SyntaxError as e: + return TestCustomCodeGuardrailResponse( + success=False, + error=f"Syntax error in custom code: {e}", + error_type="compilation", + ) + except Exception as e: + return TestCustomCodeGuardrailResponse( + success=False, + error=f"Failed to compile custom code: {e}", + error_type="compilation", + ) + + # Step 2: Verify apply_guardrail function exists + if "apply_guardrail" not in exec_globals: + return TestCustomCodeGuardrailResponse( + success=False, + error="Custom code must define an 'apply_guardrail' function. " + "Expected signature: apply_guardrail(inputs, request_data, input_type)", + error_type="compilation", + ) + + apply_fn = exec_globals["apply_guardrail"] + if not callable(apply_fn): + return TestCustomCodeGuardrailResponse( + success=False, + error="'apply_guardrail' must be a callable function", + error_type="compilation", + ) + + # Step 3: Prepare test inputs + test_inputs = request.test_input + if "texts" not in test_inputs: + test_inputs["texts"] = [] + + # Prepare mock request_data + mock_request_data = request.request_data or {} + safe_request_data = { + "model": mock_request_data.get("model", "test-model"), + "user_id": mock_request_data.get("user_id"), + "team_id": mock_request_data.get("team_id"), + "end_user_id": mock_request_data.get("end_user_id"), + "metadata": mock_request_data.get("metadata", {}), + } + + # Step 4: Execute the function with timeout protection + + def execute_guardrail(): + return apply_fn(test_inputs, safe_request_data, request.input_type) + + try: + with concurrent.futures.ThreadPoolExecutor(max_workers=1) as executor: + future = executor.submit(execute_guardrail) + try: + result = future.result(timeout=EXECUTION_TIMEOUT_SECONDS) + except concurrent.futures.TimeoutError: + return TestCustomCodeGuardrailResponse( + success=False, + error=f"Execution timeout: code took longer than {EXECUTION_TIMEOUT_SECONDS} seconds", + error_type="execution", + ) + except Exception as e: + return TestCustomCodeGuardrailResponse( + success=False, + error=f"Execution error: {e}", + error_type="execution", + ) + + # Step 5: Validate and return result + if not isinstance(result, dict): + return TestCustomCodeGuardrailResponse( + success=True, + result={ + "action": "allow", + "warning": f"Expected dict result, got {type(result).__name__}. Treating as allow.", + }, + ) + + return TestCustomCodeGuardrailResponse( + success=True, + result=result, + ) + + except Exception as e: + verbose_proxy_logger.exception(f"Error testing custom code guardrail: {e}") + return TestCustomCodeGuardrailResponse( + success=False, + error=f"Unexpected error: {e}", + error_type="execution", + ) + + @router.post("/guardrails/apply_guardrail", response_model=ApplyGuardrailResponse) @router.post("/apply_guardrail", response_model=ApplyGuardrailResponse) async def apply_guardrail( diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/__init__.py new file mode 100644 index 00000000000..747b188feea --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/__init__.py @@ -0,0 +1,65 @@ +"""Custom code guardrail integration for LiteLLM. + +This module allows users to write custom guardrail logic using Python-like code +that runs in a sandboxed environment with access to LiteLLM-provided primitives. +""" + +from typing import TYPE_CHECKING + +from litellm.types.guardrails import SupportedGuardrailIntegrations + +from .custom_code_guardrail import CustomCodeGuardrail + +if TYPE_CHECKING: + from litellm.types.guardrails import Guardrail, LitellmParams + + +def initialize_guardrail( + litellm_params: "LitellmParams", guardrail: "Guardrail" +) -> CustomCodeGuardrail: + """ + Initialize a custom code guardrail. + + Args: + litellm_params: Configuration parameters including the custom code + guardrail: The guardrail configuration dict + + Returns: + CustomCodeGuardrail instance + """ + import litellm + + guardrail_name = guardrail.get("guardrail_name") + if not guardrail_name: + raise ValueError("Custom code guardrail requires a guardrail_name") + + # Get the custom code from litellm_params + custom_code = getattr(litellm_params, "custom_code", None) + if not custom_code: + raise ValueError( + "Custom code guardrail requires 'custom_code' in litellm_params" + ) + + custom_code_guardrail = CustomCodeGuardrail( + guardrail_name=guardrail_name, + custom_code=custom_code, + event_hook=litellm_params.mode, + default_on=litellm_params.default_on, + ) + + litellm.logging_callback_manager.add_litellm_callback(custom_code_guardrail) + return custom_code_guardrail + + +guardrail_initializer_registry = { + SupportedGuardrailIntegrations.CUSTOM_CODE.value: initialize_guardrail, +} + +guardrail_class_registry = { + SupportedGuardrailIntegrations.CUSTOM_CODE.value: CustomCodeGuardrail, +} + +__all__ = [ + "CustomCodeGuardrail", + "initialize_guardrail", +] diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py new file mode 100644 index 00000000000..a0ca324411c --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/custom_code_guardrail.py @@ -0,0 +1,372 @@ +""" +Custom code guardrail for LiteLLM. + +This module provides a guardrail that executes user-defined Python-like code +to implement custom guardrail logic. The code runs in a sandboxed environment +with access to LiteLLM-provided primitives for common guardrail operations. + +Example custom code: + + def apply_guardrail(inputs, request_data, input_type): + '''Block messages containing SSNs''' + for text in inputs["texts"]: + if regex_match(text, r"\\d{3}-\\d{2}-\\d{4}"): + return block("Social Security Number detected") + return allow() +""" + +import threading +from typing import TYPE_CHECKING, Any, Dict, Literal, Optional, Type, cast + +from fastapi import HTTPException + +from litellm._logging import verbose_proxy_logger +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.types.guardrails import GuardrailEventHooks +from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigModel +from litellm.types.utils import GenericGuardrailAPIInputs + +from .primitives import get_custom_code_primitives + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + + +class CustomCodeGuardrailError(Exception): + """Raised when custom code guardrail execution fails.""" + + def __init__(self, message: str, details: Optional[Dict[str, Any]] = None) -> None: + super().__init__(message) + self.details = details or {} + + +class CustomCodeCompilationError(CustomCodeGuardrailError): + """Raised when custom code fails to compile.""" + + +class CustomCodeExecutionError(CustomCodeGuardrailError): + """Raised when custom code fails during execution.""" + + +class CustomCodeGuardrailConfigModel(GuardrailConfigModel): + """Configuration parameters for the custom code guardrail.""" + + custom_code: str + """The Python-like code containing the apply_guardrail function.""" + + +class CustomCodeGuardrail(CustomGuardrail): + """ + Guardrail that executes user-defined Python-like code. + + The code runs in a sandboxed environment that provides: + - Access to LiteLLM primitives (regex_match, json_parse, etc.) + - No file I/O or network access + - No imports allowed + + Users write an `apply_guardrail(inputs, request_data, input_type)` function + that returns one of: + - allow() - let the request/response through + - block(reason) - reject with a message + - modify(texts=...) - transform the content + + Example: + def apply_guardrail(inputs, request_data, input_type): + for text in inputs["texts"]: + if regex_match(text, r"password"): + return block("Sensitive content detected") + return allow() + """ + + def __init__( + self, + custom_code: str, + guardrail_name: Optional[str] = "custom_code", + **kwargs: Any, + ) -> None: + """ + Initialize the custom code guardrail. + + Args: + custom_code: The source code containing apply_guardrail function + guardrail_name: Name of this guardrail instance + **kwargs: Additional arguments passed to CustomGuardrail + """ + self.custom_code = custom_code + self._compiled_function: Optional[Any] = None + self._compile_lock = threading.Lock() + self._compile_error: Optional[str] = None + + supported_event_hooks = [ + GuardrailEventHooks.pre_call, + GuardrailEventHooks.during_call, + GuardrailEventHooks.post_call, + ] + + super().__init__( + guardrail_name=guardrail_name, + supported_event_hooks=supported_event_hooks, + **kwargs, + ) + + # Compile the code on initialization + self._compile_custom_code() + + @staticmethod + def get_config_model() -> Optional[Type[GuardrailConfigModel]]: + """Returns the config model for the UI.""" + return CustomCodeGuardrailConfigModel + + def _compile_custom_code(self) -> None: + """ + Compile the custom code and extract the apply_guardrail function. + + The code runs in a sandboxed environment with only the allowed primitives. + """ + with self._compile_lock: + if self._compiled_function is not None: + return + + try: + # Create a restricted execution environment + # Only include our safe primitives + exec_globals = get_custom_code_primitives().copy() + + # Execute the user code in the restricted environment + exec(compile(self.custom_code, "", "exec"), exec_globals) + + # Extract the apply_guardrail function + if "apply_guardrail" not in exec_globals: + raise CustomCodeCompilationError( + "Custom code must define an 'apply_guardrail' function. " + "Expected signature: apply_guardrail(inputs, request_data, input_type)" + ) + + apply_fn = exec_globals["apply_guardrail"] + if not callable(apply_fn): + raise CustomCodeCompilationError( + "'apply_guardrail' must be a callable function" + ) + + self._compiled_function = apply_fn + verbose_proxy_logger.debug( + f"Custom code guardrail '{self.guardrail_name}' compiled successfully" + ) + + except SyntaxError as e: + self._compile_error = f"Syntax error in custom code: {e}" + raise CustomCodeCompilationError(self._compile_error) from e + except CustomCodeCompilationError: + raise + except Exception as e: + self._compile_error = f"Failed to compile custom code: {e}" + raise CustomCodeCompilationError(self._compile_error) from e + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional["LiteLLMLoggingObj"] = None, + ) -> GenericGuardrailAPIInputs: + """ + Apply the custom code guardrail to the inputs. + + This method calls the user-defined apply_guardrail function and + processes its result to determine the appropriate action. + + Args: + inputs: Dictionary containing texts, images, tool_calls + request_data: The original request data with metadata + input_type: "request" for pre-call, "response" for post-call + logging_obj: Optional logging object + + Returns: + GenericGuardrailAPIInputs - possibly modified + + Raises: + HTTPException: If content is blocked + CustomCodeExecutionError: If execution fails + """ + if self._compiled_function is None: + if self._compile_error: + raise CustomCodeExecutionError( + f"Custom code guardrail not compiled: {self._compile_error}" + ) + raise CustomCodeExecutionError("Custom code guardrail not compiled") + + try: + # Prepare inputs dict for the function + + # Prepare request_data with safe subset of information + safe_request_data = self._prepare_safe_request_data(request_data) + + # Execute the custom function + result = self._compiled_function(inputs, safe_request_data, input_type) + + # Process the result + return self._process_result( + result=result, + inputs=inputs, + request_data=request_data, + input_type=input_type, + ) + + except HTTPException: + # Re-raise HTTP exceptions (from block action) + raise + except Exception as e: + verbose_proxy_logger.error( + f"Custom code guardrail '{self.guardrail_name}' execution error: {e}" + ) + raise CustomCodeExecutionError( + f"Custom code guardrail execution failed: {e}", + details={ + "guardrail_name": self.guardrail_name, + "input_type": input_type, + }, + ) from e + + def _prepare_safe_request_data(self, request_data: dict) -> Dict[str, Any]: + """ + Prepare a safe subset of request_data for code execution. + + This filters out sensitive information and provides only what's + needed for guardrail logic. + + Args: + request_data: The full request data + + Returns: + Safe subset of request data + """ + return { + "model": request_data.get("model"), + "user_id": request_data.get("user_api_key_user_id"), + "team_id": request_data.get("user_api_key_team_id"), + "end_user_id": request_data.get("user_api_key_end_user_id"), + "metadata": request_data.get("metadata", {}), + } + + def _process_result( + self, + result: Any, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + ) -> GenericGuardrailAPIInputs: + """ + Process the result from the custom code function. + + Args: + result: The return value from apply_guardrail + inputs: The original inputs + request_data: The request data + input_type: "request" or "response" + + Returns: + GenericGuardrailAPIInputs - possibly modified + + Raises: + HTTPException: If action is "block" + """ + if not isinstance(result, dict): + verbose_proxy_logger.warning( + f"Custom code guardrail '{self.guardrail_name}': " + f"Expected dict result, got {type(result).__name__}. Treating as allow." + ) + return inputs + + action = result.get("action", "allow") + + if action == "allow": + verbose_proxy_logger.debug( + f"Custom code guardrail '{self.guardrail_name}': Allowing {input_type}" + ) + return inputs + + elif action == "block": + reason = result.get("reason", "Blocked by custom code guardrail") + detection_info = result.get("detection_info", {}) + + verbose_proxy_logger.info( + f"Custom code guardrail '{self.guardrail_name}': Blocking {input_type} - {reason}" + ) + + is_output = input_type == "response" + + # For pre-call, raise passthrough exception to return synthetic response + if not is_output: + self.raise_passthrough_exception( + violation_message=reason, + request_data=request_data, + detection_info=detection_info, + ) + + # For post-call, raise HTTP exception + raise HTTPException( + status_code=400, + detail={ + "error": reason, + "guardrail": self.guardrail_name, + "detection_info": detection_info, + }, + ) + + elif action == "modify": + verbose_proxy_logger.debug( + f"Custom code guardrail '{self.guardrail_name}': Modifying {input_type}" + ) + + # Apply modifications + modified_inputs = dict(inputs) + + if "texts" in result and result["texts"] is not None: + modified_inputs["texts"] = result["texts"] + + if "images" in result and result["images"] is not None: + modified_inputs["images"] = result["images"] + + if "tool_calls" in result and result["tool_calls"] is not None: + modified_inputs["tool_calls"] = result["tool_calls"] + + return cast(GenericGuardrailAPIInputs, modified_inputs) + + else: + verbose_proxy_logger.warning( + f"Custom code guardrail '{self.guardrail_name}': " + f"Unknown action '{action}'. Treating as allow." + ) + return inputs + + def update_custom_code(self, new_code: str) -> None: + """ + Update the custom code and recompile. + + This method allows hot-reloading of guardrail logic without + restarting the server. + + Args: + new_code: The new source code + + Raises: + CustomCodeCompilationError: If the new code fails to compile + """ + with self._compile_lock: + # Reset state + old_function = self._compiled_function + old_code = self.custom_code + self._compiled_function = None + self._compile_error = None + + try: + self.custom_code = new_code + self._compile_custom_code() + verbose_proxy_logger.info( + f"Custom code guardrail '{self.guardrail_name}': Code updated successfully" + ) + except CustomCodeCompilationError: + # Rollback on failure + self.custom_code = old_code + self._compiled_function = old_function + raise diff --git a/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py new file mode 100644 index 00000000000..695e59977c8 --- /dev/null +++ b/litellm/proxy/guardrails/guardrail_hooks/custom_code/primitives.py @@ -0,0 +1,602 @@ +""" +Built-in primitives provided to custom code guardrails. + +These functions are injected into the custom code execution environment +and provide safe, sandboxed functionality for common guardrail operations. +""" + +import json +import re +from typing import Any, Dict, List, Optional, Tuple, Type, Union +from urllib.parse import urlparse + +from litellm._logging import verbose_proxy_logger + +# ============================================================================= +# Result Types - Used by Starlark code to return guardrail decisions +# ============================================================================= + + +def allow() -> Dict[str, Any]: + """ + Allow the request/response to proceed unchanged. + + Returns: + Dict indicating the request should be allowed + """ + return {"action": "allow"} + + +def block( + reason: str, detection_info: Optional[Dict[str, Any]] = None +) -> Dict[str, Any]: + """ + Block the request/response with a reason. + + Args: + reason: Human-readable reason for blocking + detection_info: Optional additional detection metadata + + Returns: + Dict indicating the request should be blocked + """ + result: Dict[str, Any] = {"action": "block", "reason": reason} + if detection_info: + result["detection_info"] = detection_info + return result + + +def modify( + texts: Optional[List[str]] = None, + images: Optional[List[Any]] = None, + tool_calls: Optional[List[Any]] = None, +) -> Dict[str, Any]: + """ + Modify the request/response content. + + Args: + texts: Modified text content (if None, keeps original) + images: Modified image content (if None, keeps original) + tool_calls: Modified tool calls (if None, keeps original) + + Returns: + Dict indicating the content should be modified + """ + result: Dict[str, Any] = {"action": "modify"} + if texts is not None: + result["texts"] = texts + if images is not None: + result["images"] = images + if tool_calls is not None: + result["tool_calls"] = tool_calls + return result + + +# ============================================================================= +# Regex Primitives +# ============================================================================= + + +def regex_match(text: str, pattern: str, flags: int = 0) -> bool: + """ + Check if a regex pattern matches anywhere in the text. + + Args: + text: The text to search in + pattern: The regex pattern to match + flags: Optional regex flags (default: 0) + + Returns: + True if pattern matches, False otherwise + """ + try: + return bool(re.search(pattern, text, flags)) + except re.error as e: + verbose_proxy_logger.warning(f"Starlark regex_match error: {e}") + return False + + +def regex_match_all(text: str, pattern: str, flags: int = 0) -> bool: + """ + Check if a regex pattern matches the entire text. + + Args: + text: The text to match + pattern: The regex pattern + flags: Optional regex flags + + Returns: + True if pattern matches entire text, False otherwise + """ + try: + return bool(re.fullmatch(pattern, text, flags)) + except re.error as e: + verbose_proxy_logger.warning(f"Starlark regex_match_all error: {e}") + return False + + +def regex_replace(text: str, pattern: str, replacement: str, flags: int = 0) -> str: + """ + Replace all occurrences of a pattern in text. + + Args: + text: The text to modify + pattern: The regex pattern to find + replacement: The replacement string + flags: Optional regex flags + + Returns: + The text with replacements applied + """ + try: + return re.sub(pattern, replacement, text, flags=flags) + except re.error as e: + verbose_proxy_logger.warning(f"Starlark regex_replace error: {e}") + return text + + +def regex_find_all(text: str, pattern: str, flags: int = 0) -> List[str]: + """ + Find all occurrences of a pattern in text. + + Args: + text: The text to search + pattern: The regex pattern to find + flags: Optional regex flags + + Returns: + List of all matches + """ + try: + return re.findall(pattern, text, flags) + except re.error as e: + verbose_proxy_logger.warning(f"Starlark regex_find_all error: {e}") + return [] + + +# ============================================================================= +# JSON Primitives +# ============================================================================= + + +def json_parse(text: str) -> Optional[Any]: + """ + Parse a JSON string into a Python object. + + Args: + text: The JSON string to parse + + Returns: + Parsed Python object, or None if parsing fails + """ + try: + return json.loads(text) + except (json.JSONDecodeError, TypeError) as e: + verbose_proxy_logger.debug(f"Starlark json_parse error: {e}") + return None + + +def json_stringify(obj: Any) -> str: + """ + Convert a Python object to a JSON string. + + Args: + obj: The object to serialize + + Returns: + JSON string representation + """ + try: + return json.dumps(obj) + except (TypeError, ValueError) as e: + verbose_proxy_logger.warning(f"Starlark json_stringify error: {e}") + return "" + + +def json_schema_valid(obj: Any, schema: Dict[str, Any]) -> bool: + """ + Validate an object against a JSON schema. + + Args: + obj: The object to validate + schema: The JSON schema to validate against + + Returns: + True if valid, False otherwise + """ + try: + # Try to import jsonschema, fall back to basic validation if not available + try: + import jsonschema + + jsonschema.validate(instance=obj, schema=schema) + return True + except ImportError: + # Basic validation without jsonschema library + return _basic_json_schema_validate(obj, schema) + except Exception as validation_error: + # Catch jsonschema.ValidationError and other validation errors + if "ValidationError" in type(validation_error).__name__: + return False + raise + except Exception as e: + verbose_proxy_logger.warning(f"Custom code json_schema_valid error: {e}") + return False + + +def _basic_json_schema_validate( + obj: Any, schema: Dict[str, Any], max_depth: int = 50 +) -> bool: + """ + Basic JSON schema validation without external library. + Handles: type, required, properties + + Uses an iterative approach with a stack to avoid recursion limits. + max_depth limits nesting to prevent infinite loops from circular schemas. + """ + type_map: Dict[str, Union[Type, Tuple[Type, ...]]] = { + "object": dict, + "array": list, + "string": str, + "number": (int, float), + "integer": int, + "boolean": bool, + "null": type(None), + } + + # Stack of (obj, schema, depth) tuples to process + stack: List[Tuple[Any, Dict[str, Any], int]] = [(obj, schema, 0)] + + while stack: + current_obj, current_schema, depth = stack.pop() + + # Circuit breaker: stop if we've gone too deep + if depth > max_depth: + return False + + # Check type + schema_type = current_schema.get("type") + if schema_type: + expected_type = type_map.get(schema_type) + if expected_type is not None and not isinstance(current_obj, expected_type): + return False + + # Check required fields and properties for dicts + if isinstance(current_obj, dict): + required = current_schema.get("required", []) + for field in required: + if field not in current_obj: + return False + + # Queue property validations + properties = current_schema.get("properties", {}) + for prop_name, prop_schema in properties.items(): + if prop_name in current_obj: + stack.append((current_obj[prop_name], prop_schema, depth + 1)) + + return True + + +# ============================================================================= +# URL Primitives +# ============================================================================= + + +# Common URL pattern for extraction +_URL_PATTERN = re.compile( + r"https?://(?:[-\w.]|(?:%[\da-fA-F]{2}))+[^\s]*", re.IGNORECASE +) + + +def extract_urls(text: str) -> List[str]: + """ + Extract all URLs from text. + + Args: + text: The text to search for URLs + + Returns: + List of URLs found in the text + """ + return _URL_PATTERN.findall(text) + + +def is_valid_url(url: str) -> bool: + """ + Check if a URL is syntactically valid. + + Args: + url: The URL to validate + + Returns: + True if the URL is valid, False otherwise + """ + try: + result = urlparse(url) + return all([result.scheme, result.netloc]) + except Exception: + return False + + +def all_urls_valid(text: str) -> bool: + """ + Check if all URLs in text are valid. + + Args: + text: The text containing URLs + + Returns: + True if all URLs are valid (or no URLs), False otherwise + """ + urls = extract_urls(text) + return all(is_valid_url(url) for url in urls) + + +def get_url_domain(url: str) -> Optional[str]: + """ + Extract the domain from a URL. + + Args: + url: The URL to parse + + Returns: + The domain, or None if invalid + """ + try: + result = urlparse(url) + return result.netloc if result.netloc else None + except Exception: + return None + + +# ============================================================================= +# Code Detection Primitives +# ============================================================================= + + +# Common code patterns for detection +_CODE_PATTERNS = { + "sql": [ + r"\b(SELECT|INSERT|UPDATE|DELETE|DROP|CREATE|ALTER|TRUNCATE)\b.*\b(FROM|INTO|TABLE|SET|WHERE)\b", + r"\b(SELECT)\s+[\w\*,\s]+\s+FROM\s+\w+", + r"\b(INSERT\s+INTO|UPDATE\s+\w+\s+SET|DELETE\s+FROM)\b", + ], + "python": [ + r"^\s*(def|class|import|from|if|for|while|try|except|with)\s+", + r"^\s*@\w+", # decorators + r"\b(print|len|range|str|int|float|list|dict|set)\s*\(", + ], + "javascript": [ + r"\b(function|const|let|var|class|import|export)\s+", + r"=>", # arrow functions + r"\b(console\.(log|error|warn))\s*\(", + ], + "typescript": [ + r":\s*(string|number|boolean|any|void|never)\b", + r"\b(interface|type|enum)\s+\w+", + r"<[A-Z]\w*>", # generics + ], + "java": [ + r"\b(public|private|protected)\s+(static\s+)?(class|void|int|String)\b", + r"\bSystem\.(out|err)\.print", + ], + "go": [ + r"\bfunc\s+\w+\s*\(", + r"\b(package|import)\s+", + r":=", # short variable declaration + ], + "rust": [ + r"\b(fn|let|mut|impl|struct|enum|pub|mod)\s+", + r"->", # return type + r"\b(println!|format!)\s*\(", + ], + "shell": [ + r"^#!.*\b(bash|sh|zsh)\b", + r"\b(echo|grep|sed|awk|cat|ls|cd|mkdir|rm)\s+", + r"\$\{?\w+\}?", # variable expansion + ], + "html": [ + r"<\s*(html|head|body|div|span|p|a|img|script|style)\b[^>]*>", + r"", + ], + "css": [ + r"\{[^}]*:\s*[^}]+;[^}]*\}", + r"@(media|keyframes|import|font-face)\b", + ], +} + + +def detect_code(text: str) -> bool: + """ + Check if text contains code of any language. + + Args: + text: The text to check + + Returns: + True if code is detected, False otherwise + """ + return len(detect_code_languages(text)) > 0 + + +def detect_code_languages(text: str) -> List[str]: + """ + Detect which programming languages are present in text. + + Args: + text: The text to analyze + + Returns: + List of detected language names + """ + detected = [] + for lang, patterns in _CODE_PATTERNS.items(): + for pattern in patterns: + try: + if re.search(pattern, text, re.IGNORECASE | re.MULTILINE): + detected.append(lang) + break # Only add each language once + except re.error: + continue + return detected + + +def contains_code_language(text: str, languages: List[str]) -> bool: + """ + Check if text contains code from specific languages. + + Args: + text: The text to check + languages: List of language names to check for + + Returns: + True if any of the specified languages are detected + """ + detected = detect_code_languages(text) + return any(lang.lower() in [d.lower() for d in detected] for lang in languages) + + +# ============================================================================= +# Text Utility Primitives +# ============================================================================= + + +def contains(text: str, substring: str) -> bool: + """ + Check if text contains a substring. + + Args: + text: The text to search in + substring: The substring to find + + Returns: + True if substring is found, False otherwise + """ + return substring in text + + +def contains_any(text: str, substrings: List[str]) -> bool: + """ + Check if text contains any of the given substrings. + + Args: + text: The text to search in + substrings: List of substrings to find + + Returns: + True if any substring is found, False otherwise + """ + return any(s in text for s in substrings) + + +def contains_all(text: str, substrings: List[str]) -> bool: + """ + Check if text contains all of the given substrings. + + Args: + text: The text to search in + substrings: List of substrings to find + + Returns: + True if all substrings are found, False otherwise + """ + return all(s in text for s in substrings) + + +def word_count(text: str) -> int: + """ + Count the number of words in text. + + Args: + text: The text to count words in + + Returns: + Number of words + """ + return len(text.split()) + + +def char_count(text: str) -> int: + """ + Count the number of characters in text. + + Args: + text: The text to count characters in + + Returns: + Number of characters + """ + return len(text) + + +def lower(text: str) -> str: + """Convert text to lowercase.""" + return text.lower() + + +def upper(text: str) -> str: + """Convert text to uppercase.""" + return text.upper() + + +def trim(text: str) -> str: + """Remove leading and trailing whitespace.""" + return text.strip() + + +# ============================================================================= +# Primitives Registry +# ============================================================================= + + +def get_custom_code_primitives() -> Dict[str, Any]: + """ + Get all primitives to inject into the custom code environment. + + Returns: + Dict of function name to function + """ + return { + # Result types + "allow": allow, + "block": block, + "modify": modify, + # Regex + "regex_match": regex_match, + "regex_match_all": regex_match_all, + "regex_replace": regex_replace, + "regex_find_all": regex_find_all, + # JSON + "json_parse": json_parse, + "json_stringify": json_stringify, + "json_schema_valid": json_schema_valid, + # URL + "extract_urls": extract_urls, + "is_valid_url": is_valid_url, + "all_urls_valid": all_urls_valid, + "get_url_domain": get_url_domain, + # Code detection + "detect_code": detect_code, + "detect_code_languages": detect_code_languages, + "contains_code_language": contains_code_language, + # Text utilities + "contains": contains, + "contains_any": contains_any, + "contains_all": contains_all, + "word_count": word_count, + "char_count": char_count, + "lower": lower, + "upper": upper, + "trim": trim, + # Python builtins (safe subset) + "len": len, + "str": str, + "int": int, + "float": float, + "bool": bool, + "list": list, + "dict": dict, + "True": True, + "False": False, + "None": None, + } diff --git a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py index 2a852cbda08..90f689ed23c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py +++ b/litellm/proxy/guardrails/guardrail_hooks/grayswan/grayswan.py @@ -9,8 +9,10 @@ from fastapi import HTTPException from litellm._logging import verbose_proxy_logger from litellm.integrations.custom_guardrail import ( CustomGuardrail, + ModifyResponseException ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps +from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, httpxSpecialProvider, @@ -21,6 +23,8 @@ from litellm.types.utils import GenericGuardrailAPIInputs if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj +GRAYSWAN_BLOCK_ERROR_MSG = "Blocked by Gray Swan Guardrail" + class GraySwanGuardrailMissingSecrets(Exception): """Raised when the Gray Swan API key is missing.""" @@ -205,9 +209,13 @@ class GraySwanGuardrail(CustomGuardrail): # Get dynamic params from request metadata dynamic_body = self.get_guardrail_dynamic_request_body_params(request_data) or {} + if dynamic_body: + verbose_proxy_logger.debug( + "Gray Swan Guardrail: dynamic extra_body=%s", safe_dumps(dynamic_body) + ) # Prepare and send payload - payload = self._prepare_payload(messages, dynamic_body) + payload = self._prepare_payload(messages, dynamic_body, request_data) if payload is None: return inputs @@ -223,6 +231,8 @@ class GraySwanGuardrail(CustomGuardrail): ) return result except Exception as exc: + if self._is_grayswan_exception(exc): + raise end_time = time.time() status_code = getattr(exc, "status_code", None) or getattr( exc, "exception_status_code", None @@ -240,8 +250,20 @@ class GraySwanGuardrail(CustomGuardrail): exc, ) return inputs + if isinstance(exc, GraySwanGuardrailAPIError): + raise exc raise GraySwanGuardrailAPIError(str(exc), status_code=status_code) from exc + def _is_grayswan_exception(self, exc: Exception) -> bool: + # Guardrail decision (passthrough) should always propagate, + # regardless of fail_open. + if isinstance(exc, ModifyResponseException): + return True + detail = getattr(exc, "detail", None) + if isinstance(detail, dict): + return detail.get("error") == GRAYSWAN_BLOCK_ERROR_MSG + return False + # ------------------------------------------------------------------ # Legacy Test Interface (for backward compatibility) # ------------------------------------------------------------------ @@ -324,7 +346,7 @@ class GraySwanGuardrail(CustomGuardrail): raise HTTPException( status_code=400, detail={ - "error": "Blocked by Gray Swan Guardrail", + "error": GRAYSWAN_BLOCK_ERROR_MSG, "violation_location": violation_location, "violation": violation_score, "violated_rules": violated_rules, @@ -445,7 +467,7 @@ class GraySwanGuardrail(CustomGuardrail): raise HTTPException( status_code=400, detail={ - "error": "Blocked by Gray Swan Guardrail", + "error": GRAYSWAN_BLOCK_ERROR_MSG, "violation_location": violation_location, "violation": violation_score, "violated_rules": violated_rules, @@ -494,7 +516,7 @@ class GraySwanGuardrail(CustomGuardrail): } def _prepare_payload( - self, messages: List[Dict[str, str]], dynamic_body: dict + self, messages: List[Dict[str, str]], dynamic_body: dict, request_data: dict ) -> Optional[Dict[str, Any]]: payload: Dict[str, Any] = {"messages": messages} @@ -510,6 +532,18 @@ class GraySwanGuardrail(CustomGuardrail): if reasoning_mode: payload["reasoning_mode"] = reasoning_mode + # Pass through arbitrary metadata when provided via dynamic extra_body. + if "metadata" in dynamic_body: + payload["metadata"] = dynamic_body["metadata"] + + litellm_metadata = request_data.get("litellm_metadata") + if isinstance(litellm_metadata, dict) and litellm_metadata: + cleaned_litellm_metadata = dict(litellm_metadata) + # cleaned_litellm_metadata.pop("user_api_key_auth", None) + sanitized = safe_json_loads(safe_dumps(cleaned_litellm_metadata), default={}) + if isinstance(sanitized, dict) and sanitized: + payload["litellm_metadata"] = sanitized + return payload def _format_violation_message( diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 80f9860bdff..f07f65d10f5 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -6,6 +6,7 @@ Unified Guardrail, leveraging LiteLLM's /applyGuardrail endpoint 3. Implements a way to call /applyGuardrail endpoint for `/chat/completions` + `/v1/messages` requests on async_post_call_streaming_iterator_hook """ +import copy from typing import Any, AsyncGenerator, List, Optional, Union from litellm._logging import verbose_proxy_logger @@ -349,22 +350,26 @@ class UnifiedLLMGuardrails(CustomLogger): guardrail_to_apply.guardrail_name, ) + # Deep-copy the current chunk before guardrail processing. + # process_output_streaming_response modifies responses_so_far + # in-place: it puts the combined guardrailed text in the first + # chunk and clears all subsequent chunks to "". Without this + # copy, yielding processed_items[-1] would yield an empty + # string, permanently losing this chunk's content. + original_item = copy.deepcopy(item) + endpoint_translation = endpoint_guardrail_translation_mappings[ CallTypes(call_type) ]() - processed_items = ( - await endpoint_translation.process_output_streaming_response( - responses_so_far=responses_so_far, - guardrail_to_apply=guardrail_to_apply, - litellm_logging_obj=request_data.get("litellm_logging_obj"), - user_api_key_dict=user_api_key_dict, - ) + await endpoint_translation.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail_to_apply, + litellm_logging_obj=request_data.get("litellm_logging_obj"), + user_api_key_dict=user_api_key_dict, ) - last_item = processed_items[-1] - - yield last_item + yield original_item else: yield item diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 38a867d031b..636ed87d794 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -813,9 +813,12 @@ def _update_internal_user_params( data_json: dict, data: Union[UpdateUserRequest, UpdateUserRequestNoUserIDorEmail] ) -> dict: non_default_values = {} + fields_set = data.fields_set() if hasattr(data, 'fields_set') else set() + for k, v in data_json.items(): if k == "max_budget": - non_default_values[k] = v + if "max_budget" in fields_set: + non_default_values[k] = v elif ( v is not None and v diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index d1840363009..9dadffca351 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -37,13 +37,6 @@ 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, @@ -82,6 +75,13 @@ from litellm.proxy.utils import ( ) from litellm.router import Router from litellm.secret_managers.main import get_secret +from litellm.types.proxy.management_endpoints.key_management_endpoints import ( + BulkUpdateKeyRequest, + BulkUpdateKeyRequestItem, + BulkUpdateKeyResponse, + FailedKeyUpdate, + SuccessfulKeyUpdate, +) from litellm.types.router import Deployment from litellm.types.utils import ( BudgetConfig, @@ -2381,6 +2381,10 @@ async def info_key_fn( # if using pydantic v1 key_info = key_info.dict() key_info.pop("token") + + # Attach object_permission if object_permission_id is set + key_info = await attach_object_permission_to_dict(key_info, prisma_client) + return {"key": key, "info": key_info} except Exception as e: raise handle_exception_on_proxy(e) diff --git a/litellm/proxy/openai_files_endpoints/common_utils.py b/litellm/proxy/openai_files_endpoints/common_utils.py index 2ff1183579f..f67dc5e2aaa 100644 --- a/litellm/proxy/openai_files_endpoints/common_utils.py +++ b/litellm/proxy/openai_files_endpoints/common_utils.py @@ -637,3 +637,127 @@ def _extract_model_param(request: "Request", request_body: dict) -> Optional[str or request.query_params.get("model") or request.headers.get("x-litellm-model") ) + + +# ============================================================================ +# BATCH DATABASE OPERATIONS +# ============================================================================ + + +async def get_batch_from_database( + batch_id: str, + unified_batch_id: Union[str, Literal[False]], + managed_files_obj, + prisma_client, + verbose_proxy_logger, +): + """ + Try to retrieve batch object from ManagedObjectTable for consistent state. + + Args: + batch_id: The batch ID (may be unified/encoded) + unified_batch_id: Result from _is_base64_encoded_unified_file_id() + managed_files_obj: The managed_files proxy hook object + prisma_client: Prisma database client + verbose_proxy_logger: Logger instance + + Returns: + Tuple of (db_batch_object, response_batch) + - db_batch_object: Raw database object (or None) + - response_batch: Parsed LiteLLMBatch object (or None) + """ + import json + from litellm.types.utils import LiteLLMBatch + + if managed_files_obj is None or not unified_batch_id: + return None, None + + try: + if not prisma_client: + return None, None + + db_batch_object = await prisma_client.db.litellm_managedobjecttable.find_first( + where={"unified_object_id": batch_id} + ) + + if not db_batch_object or not db_batch_object.file_object: + return None, None + + # Parse the batch object from database + batch_data = json.loads(db_batch_object.file_object) if isinstance(db_batch_object.file_object, str) else db_batch_object.file_object + response = LiteLLMBatch(**batch_data) + response.id = batch_id + + verbose_proxy_logger.debug( + f"Retrieved batch {batch_id} from ManagedObjectTable with status={response.status}" + ) + + return db_batch_object, response + + except Exception as e: + verbose_proxy_logger.warning( + f"Failed to retrieve batch from ManagedObjectTable: {e}, falling back to provider" + ) + return None, None + + +async def update_batch_in_database( + batch_id: str, + unified_batch_id: Union[str, Literal[False]], + response, + managed_files_obj, + prisma_client, + verbose_proxy_logger, + db_batch_object=None, + operation: str = "update", +): + """ + Update batch status and object in ManagedObjectTable. + + Args: + batch_id: The batch ID (unified/encoded) + unified_batch_id: Result from _is_base64_encoded_unified_file_id() + response: The batch response object with updated state + managed_files_obj: The managed_files proxy hook object + prisma_client: Prisma database client + verbose_proxy_logger: Logger instance + db_batch_object: Optional existing database object (for comparison) + operation: Description of operation ("update", "cancel", etc.) + """ + import litellm.utils + + if managed_files_obj is None or not unified_batch_id: + return + + try: + if not prisma_client: + return + + # Only update if status has changed (when db_batch_object is provided) + if db_batch_object and response.status == db_batch_object.status: + return + + if db_batch_object: + verbose_proxy_logger.info( + f"Updating batch {batch_id} status from {db_batch_object.status} to {response.status}" + ) + else: + verbose_proxy_logger.info( + f"Updating batch {batch_id} status to {response.status} after {operation}" + ) + + # Normalize status for database storage + db_status = response.status if response.status != "completed" else "complete" + + await prisma_client.db.litellm_managedobjecttable.update( + where={"unified_object_id": batch_id}, + data={ + "status": db_status, + "file_object": response.model_dump_json(), + "updated_at": litellm.utils.get_utc_datetime(), + }, + ) + except Exception as e: + verbose_proxy_logger.error( + f"Failed to update batch status in ManagedObjectTable: {e}" + ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 8f433bfa486..637893872d9 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -47,6 +47,7 @@ from litellm.constants import ( DEFAULT_SLACK_ALERTING_THRESHOLD, LITELLM_EMBEDDING_PROVIDERS_SUPPORTING_INPUT_ARRAY_OF_TOKENS, LITELLM_SETTINGS_SAFE_DB_OVERRIDES, + LITELLM_UI_ALLOW_HEADERS, ) from litellm.litellm_core_utils.litellm_logging import ( _init_custom_logger_compatible_class, @@ -239,6 +240,10 @@ from litellm.proxy._types import * from litellm.proxy.agent_endpoints.a2a_endpoints import router as a2a_router from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry from litellm.proxy.agent_endpoints.endpoints import router as agent_endpoints_router +from litellm.proxy.agent_endpoints.model_list_helpers import ( + append_agents_to_model_group, + append_agents_to_model_info, +) from litellm.proxy.analytics_endpoints.analytics_endpoints import ( router as analytics_router, ) @@ -1210,6 +1215,7 @@ app.add_middleware( allow_credentials=True, allow_methods=["*"], allow_headers=["*"], + expose_headers=LITELLM_UI_ALLOW_HEADERS, ) app.add_middleware(PrometheusAuthMiddleware) @@ -1858,6 +1864,7 @@ class ProxyConfig: def __init__(self) -> None: self.config: Dict[str, Any] = {} + self._last_semantic_filter_config: Optional[Dict[str, Any]] = None def is_yaml(self, config_file_path: str) -> bool: if not os.path.isfile(config_file_path): @@ -3916,6 +3923,93 @@ class ProxyConfig: prisma_client=prisma_client, proxy_config=self ) + if self._should_load_db_object(object_type="semantic_filter_settings"): + await self._init_semantic_filter_settings_in_db( + prisma_client=prisma_client + ) + + async def _init_semantic_filter_settings_in_db(self, prisma_client: PrismaClient): + """ + Initialize MCP semantic filter settings from database. + Called periodically (approximately every 10 seconds) by background task to hot-reload settings across all pods. + """ + import json + + import litellm + from litellm.proxy.hooks.mcp_semantic_filter import SemanticToolFilterHook + + try: + # Load litellm_settings from DB + config_record = await prisma_client.db.litellm_config.find_unique( + where={"param_name": "litellm_settings"} + ) + + if config_record is None or config_record.param_value is None: + return + + litellm_settings = config_record.param_value + if isinstance(litellm_settings, str): + litellm_settings = json.loads(litellm_settings) + + mcp_semantic_filter_config = litellm_settings.get( + "mcp_semantic_tool_filter", None + ) + + if mcp_semantic_filter_config is None: + return + + # Check if settings have changed (compare with in-memory state) + if hasattr(self, "_last_semantic_filter_config"): + if self._last_semantic_filter_config == mcp_semantic_filter_config: + # If hook is missing or router isn't built yet, reinitialize anyway + active_hooks = ( + litellm.logging_callback_manager.get_custom_loggers_for_type( + SemanticToolFilterHook + ) + ) + if active_hooks: + for active_hook in active_hooks: + if isinstance(active_hook, SemanticToolFilterHook): + if ( + active_hook.filter is not None + and active_hook.filter.tool_router is not None + ): + verbose_proxy_logger.debug( + "Semantic filter settings unchanged, skipping reinitialization" + ) + return + verbose_proxy_logger.info( + "Semantic filter settings unchanged, but hook is missing or uninitialized. Reinitializing." + ) + + # Remove old hooks using logging callback manager + litellm.logging_callback_manager.remove_callbacks_by_type( + litellm.callbacks, SemanticToolFilterHook + ) + + # Initialize new hook if enabled + if mcp_semantic_filter_config.get("enabled", False): + global llm_router + hook = await SemanticToolFilterHook.initialize_from_config( + config=mcp_semantic_filter_config, + llm_router=llm_router, + ) + if hook: + litellm.logging_callback_manager.add_litellm_callback(hook) + verbose_proxy_logger.info( + "MCP Semantic Filter reinitialized from DB" + ) + else: + verbose_proxy_logger.info("MCP Semantic Filter disabled") + + # Store current config for comparison next time + self._last_semantic_filter_config = mcp_semantic_filter_config.copy() + + except Exception as e: + verbose_proxy_logger.exception( + f"Error initializing semantic filter settings from DB: {e}" + ) + async def _init_sso_settings_in_db(self, prisma_client: PrismaClient): """ Initialize SSO settings from database into the router on startup. @@ -8616,6 +8710,15 @@ async def model_info_v2( ) verbose_proxy_logger.debug("all_models: %s", all_models) + + # Append A2A agents to models list + all_models = await append_agents_to_model_info( + models=all_models, + user_api_key_dict=user_api_key_dict, + ) + + # Update total count to include agents + search_total_count = len(all_models) return _paginate_models_response( all_models=all_models, @@ -9456,6 +9559,12 @@ async def model_group_info( model_groups: List[ModelGroupInfoProxy] = _get_model_group_info( llm_router=llm_router, all_models_str=all_models_str, model_group=model_group ) + + # Append A2A agents to model groups + model_groups = await append_agents_to_model_group( + model_groups=model_groups, + user_api_key_dict=user_api_key_dict, + ) return {"data": model_groups} diff --git a/litellm/proxy/route_llm_request.py b/litellm/proxy/route_llm_request.py index c6a93164d49..92fb88f7147 100644 --- a/litellm/proxy/route_llm_request.py +++ b/litellm/proxy/route_llm_request.py @@ -12,6 +12,11 @@ else: LitellmRouter = Any +def _is_a2a_agent_model(model_name: Any) -> bool: + """Check if the model name is for an A2A agent (a2a/ prefix).""" + return isinstance(model_name, str) and model_name.startswith("a2a/") + + ROUTE_ENDPOINT_MAPPING = { "acompletion": "/chat/completions", "atext_completion": "/completions", @@ -322,6 +327,15 @@ async def route_request( except Exception: # If router fails (e.g., model not found in router), fall back to direct call return getattr(litellm, f"{route_type}")(**data) + elif _is_a2a_agent_model(data.get("model", "")): + from litellm.proxy.agent_endpoints.a2a_routing import ( + route_a2a_agent_request, + ) + + result = route_a2a_agent_request(data, route_type) + if result is not None: + return result + # Fall through to raise exception below if result is None elif user_model is not None: return getattr(litellm, f"{route_type}")(**data) diff --git a/litellm/proxy/search_endpoints/search_tool_management.py b/litellm/proxy/search_endpoints/search_tool_management.py index 21a4db75a01..4754316795e 100644 --- a/litellm/proxy/search_endpoints/search_tool_management.py +++ b/litellm/proxy/search_endpoints/search_tool_management.py @@ -48,7 +48,7 @@ def _convert_datetime_to_str(value: Union[datetime, str, None]) -> Union[str, No ) async def list_search_tools(): """ - List all search tools that are available in the database. + List all search tools that are available in the database and config file. Example Request: ```bash @@ -71,38 +71,100 @@ async def list_search_tools(): "description": "Perplexity search tool" }, "created_at": "2023-11-09T12:34:56.789Z", - "updated_at": "2023-11-09T12:34:56.789Z" + "updated_at": "2023-11-09T12:34:56.789Z", + "is_from_config": false + }, + { + "search_tool_name": "config-search-tool", + "litellm_params": { + "search_provider": "tavily", + "api_key": "tvly-***" + }, + "is_from_config": true } ] } ``` """ - from litellm.proxy.proxy_server import prisma_client + from litellm.litellm_core_utils.litellm_logging import _get_masked_values + from litellm.proxy.proxy_server import prisma_client, proxy_config if prisma_client is None: raise HTTPException(status_code=500, detail="Prisma client not initialized") try: - search_tools = await SEARCH_TOOL_REGISTRY.get_all_search_tools_from_db( + search_tools_from_db = await SEARCH_TOOL_REGISTRY.get_all_search_tools_from_db( prisma_client=prisma_client ) + db_tool_names = { + tool.get("search_tool_name") for tool in search_tools_from_db + } + search_tool_configs: List[SearchToolInfoResponse] = [] - for search_tool in search_tools: + + config_search_tools = [] + + try: + config = await proxy_config.get_config() + parsed_tools = proxy_config.parse_search_tools(config) + if parsed_tools: + config_search_tools = parsed_tools + except Exception as e: + verbose_proxy_logger.debug( + f"Could not get config-defined search tools: {e}" + ) + + for search_tool in config_search_tools: + tool_name = search_tool.get("search_tool_name") + if tool_name: + litellm_params_dict = dict(search_tool.get("litellm_params", {})) + masked_litellm_params_dict = _get_masked_values( + litellm_params_dict, + unmasked_length=4, + number_of_asterisks=4, + ) + + search_tool_configs.append( + SearchToolInfoResponse( + search_tool_id=None, + search_tool_name=tool_name, + litellm_params=masked_litellm_params_dict, + search_tool_info=search_tool.get("search_tool_info"), + created_at=None, + updated_at=None, + is_from_config=True, + ) + ) + + search_tool_configs = [ + tool for tool in search_tool_configs + if tool.get("search_tool_name") not in db_tool_names + ] + + for search_tool in search_tools_from_db: + litellm_params_dict = dict(search_tool.get("litellm_params", {})) + masked_litellm_params_dict = _get_masked_values( + litellm_params_dict, + unmasked_length=4, + number_of_asterisks=4, + ) + search_tool_configs.append( SearchToolInfoResponse( search_tool_id=search_tool.get("search_tool_id"), search_tool_name=search_tool.get("search_tool_name", ""), - litellm_params=dict(search_tool.get("litellm_params", {})), + litellm_params=masked_litellm_params_dict, search_tool_info=search_tool.get("search_tool_info"), created_at=_convert_datetime_to_str(search_tool.get("created_at")), updated_at=_convert_datetime_to_str(search_tool.get("updated_at")), + is_from_config=False, ) ) return ListSearchToolsResponse(search_tools=search_tool_configs) except Exception as e: - verbose_proxy_logger.exception(f"Error getting search tools from db: {e}") + verbose_proxy_logger.exception(f"Error getting search tools: {e}") raise HTTPException(status_code=500, detail=str(e)) @@ -382,6 +444,7 @@ async def get_search_tool_info(search_tool_id: str): search_tool_info=result.get("search_tool_info"), created_at=_convert_datetime_to_str(result.get("created_at")), updated_at=_convert_datetime_to_str(result.get("updated_at")), + is_from_config=False, # This endpoint only returns DB tools ) except HTTPException as e: raise e diff --git a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py index 30ec0766dbf..a308d0b2703 100644 --- a/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py +++ b/litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py @@ -98,6 +98,40 @@ ALLOWED_UI_SETTINGS_FIELDS = { } +class MCPSemanticFilterSettings(BaseModel): + """Configuration for MCP Semantic Tool Filter""" + + enabled: bool = Field( + default=False, + description="Enable semantic filtering of MCP tools based on query relevance", + ) + + embedding_model: str = Field( + default="text-embedding-3-small", + description="Embedding model to use for semantic similarity (e.g., 'text-embedding-3-small', 'text-embedding-ada-002')", + ) + + top_k: int = Field( + default=10, + description="Number of most relevant tools to return", + ge=1, + le=100, + ) + + similarity_threshold: float = Field( + default=0.3, + description="Minimum similarity score for tool inclusion (0.0 to 1.0, where 1.0 = exact match)", + ge=0.0, + le=1.0, + ) + + +class MCPSemanticFilterSettingsResponse(SettingsResponse): + """Response model for MCP semantic filter settings""" + + pass + + @router.get( "/get/allowed_ips", tags=["Budget & Spend Tracking"], @@ -325,7 +359,7 @@ async def update_default_team_member_budget( async def _update_litellm_setting( - settings: Union[DefaultInternalUserParams, DefaultTeamSSOParams], + settings: Union[DefaultInternalUserParams, DefaultTeamSSOParams, MCPSemanticFilterSettings], settings_key: str, in_memory_var: Any, success_message: str, @@ -769,6 +803,70 @@ async def update_ui_theme_settings(theme_config: UIThemeConfig): } +@router.get( + "/get/mcp_semantic_filter_settings", + tags=["Settings"], + dependencies=[Depends(user_api_key_auth)], + response_model=MCPSemanticFilterSettingsResponse, +) +async def get_mcp_semantic_filter_settings( + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Get MCP semantic filter configuration. + Returns current settings for semantic tool filtering. + """ + from litellm.proxy.proxy_server import prisma_client, proxy_config + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail={"error": "Database not connected. Please connect a database."}, + ) + + config = await proxy_config.get_config() + + return await _get_settings_with_schema( + settings_key="mcp_semantic_tool_filter", + settings_class=MCPSemanticFilterSettings, + config=config, + ) + + +@router.patch( + "/update/mcp_semantic_filter_settings", + tags=["Settings"], + dependencies=[Depends(user_api_key_auth)], +) +async def update_mcp_semantic_filter_settings( + settings: MCPSemanticFilterSettings, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +): + """ + Update MCP semantic filter settings in database. + Settings will be picked up by all pods within approximately 10 seconds via background polling. + """ + result = await _update_litellm_setting( + settings=settings, + settings_key="mcp_semantic_tool_filter", + in_memory_var=None, + success_message="MCP Semantic Filter settings updated successfully. Changes will be applied across all pods within 10 seconds.", + ) + try: + from litellm.proxy.proxy_server import prisma_client, proxy_config + + if prisma_client is not None: + await proxy_config._init_semantic_filter_settings_in_db( + prisma_client=prisma_client + ) + except Exception as e: + verbose_proxy_logger.warning( + f"Failed to reinitialize MCP semantic filter settings immediately: {e}" + ) + + return result + + @router.get( "/in_product_nudges", tags=["UI Settings"], diff --git a/litellm/proxy_auth/__init__.py b/litellm/proxy_auth/__init__.py new file mode 100644 index 00000000000..27624a94fb9 --- /dev/null +++ b/litellm/proxy_auth/__init__.py @@ -0,0 +1,30 @@ +""" +Proxy Authentication module for LiteLLM SDK. + +This module provides OAuth2/JWT token management for authenticating +with LiteLLM Proxy or any OAuth2-protected endpoint. + +Usage: + from litellm.proxy_auth import AzureADCredential, ProxyAuthHandler + + litellm.proxy_auth = ProxyAuthHandler( + credential=AzureADCredential(), + scope="api://my-proxy/.default" + ) +""" + +from .credentials import ( + AccessToken, + TokenCredential, + AzureADCredential, + GenericOAuth2Credential, + ProxyAuthHandler, +) + +__all__ = [ + "AccessToken", + "TokenCredential", + "AzureADCredential", + "GenericOAuth2Credential", + "ProxyAuthHandler", +] diff --git a/litellm/proxy_auth/credentials.py b/litellm/proxy_auth/credentials.py new file mode 100644 index 00000000000..103b0088d80 --- /dev/null +++ b/litellm/proxy_auth/credentials.py @@ -0,0 +1,240 @@ +""" +Credential providers for proxy authentication. + +This module provides a provider-agnostic interface for obtaining OAuth2/JWT tokens. +It follows the same TokenCredential protocol used by Azure SDK. +""" + +import time +from dataclasses import dataclass +from typing import Any, Optional, Protocol, runtime_checkable + + +@dataclass +class AccessToken: + """ + Represents an OAuth2 access token with expiration. + + This matches the structure used by azure.core.credentials.AccessToken. + + Attributes: + token: The access token string (typically a JWT). + expires_on: Unix timestamp when the token expires. + """ + + token: str + expires_on: int + + +@runtime_checkable +class TokenCredential(Protocol): + """ + Protocol for credential providers. + + This matches the azure.core.credentials.TokenCredential interface, + allowing any Azure SDK credential to be used directly. + + Any class implementing get_token(scope) -> AccessToken can be used. + """ + + def get_token(self, scope: str) -> AccessToken: + """ + Get an access token for the specified scope. + + Args: + scope: The OAuth2 scope to request (e.g., "api://my-app/.default") + + Returns: + AccessToken with the token string and expiration timestamp. + """ + ... + + +class AzureADCredential: + """ + Wrapper for Azure Identity credentials. + + This wraps any azure-identity credential (DefaultAzureCredential, + ClientSecretCredential, ManagedIdentityCredential, etc.) and converts + the token to our AccessToken format. + + If no credential is provided, it will use DefaultAzureCredential + which tries multiple authentication methods automatically. + + Example: + # Use default credential chain (env vars, managed identity, CLI, etc.) + cred = AzureADCredential() + + # Or provide a specific credential + from azure.identity import ClientSecretCredential + azure_cred = ClientSecretCredential(tenant_id, client_id, client_secret) + cred = AzureADCredential(credential=azure_cred) + """ + + def __init__(self, credential: Optional[Any] = None): + """ + Initialize with an optional Azure credential. + + Args: + credential: An azure-identity credential object. If None, + DefaultAzureCredential will be used on first token request. + """ + self._credential: Any = credential + self._initialized = credential is not None + + def get_token(self, scope: str) -> AccessToken: + """ + Get an access token from Azure AD. + + Args: + scope: The OAuth2 scope (e.g., "api://my-app/.default") + + Returns: + AccessToken with the JWT and expiration. + + Raises: + ImportError: If azure-identity is not installed. + """ + if not self._initialized: + try: + from azure.identity import DefaultAzureCredential + + self._credential = DefaultAzureCredential() + self._initialized = True + except ImportError: + raise ImportError( + "azure-identity is required for AzureADCredential. " + "Install it with: pip install azure-identity" + ) + + result = self._credential.get_token(scope) + return AccessToken(token=result.token, expires_on=result.expires_on) + + +class GenericOAuth2Credential: + """ + Generic OAuth2 client credentials flow. + + This works with any OAuth2 provider (Okta, Auth0, Keycloak, etc.) + that supports the client_credentials grant type. + + Example: + cred = GenericOAuth2Credential( + client_id="my-client-id", + client_secret="my-client-secret", + token_url="https://my-idp.com/oauth2/token" + ) + """ + + def __init__(self, client_id: str, client_secret: str, token_url: str): + """ + Initialize OAuth2 client credentials. + + Args: + client_id: OAuth2 client ID + client_secret: OAuth2 client secret + token_url: Token endpoint URL (e.g., "https://idp.com/oauth2/token") + """ + self.client_id = client_id + self.client_secret = client_secret + self.token_url = token_url + self._cached_token: Optional[AccessToken] = None + + def get_token(self, scope: str) -> AccessToken: + """ + Get an access token using OAuth2 client credentials flow. + + Tokens are cached and reused until they expire (with 60s buffer). + + Args: + scope: The OAuth2 scope to request + + Returns: + AccessToken with the token and expiration. + """ + # Return cached token if still valid (with 60s buffer) + if self._cached_token and self._cached_token.expires_on > time.time() + 60: + return self._cached_token + + import httpx + + response = httpx.post( + self.token_url, + data={ + "grant_type": "client_credentials", + "client_id": self.client_id, + "client_secret": self.client_secret, + "scope": scope, + }, + ) + response.raise_for_status() + data = response.json() + + self._cached_token = AccessToken( + token=data["access_token"], + expires_on=int(time.time()) + data.get("expires_in", 3600), + ) + return self._cached_token + + +class ProxyAuthHandler: + """ + Manages OAuth2/JWT token lifecycle for proxy authentication. + + This handler: + - Obtains tokens from the configured credential provider + - Caches tokens to avoid unnecessary requests + - Automatically refreshes tokens before they expire (60s buffer) + - Generates Authorization headers for HTTP requests + + Set this as litellm.proxy_auth to automatically inject auth headers + into all requests to your LiteLLM Proxy. + + Example: + import litellm + from litellm.proxy_auth import AzureADCredential, ProxyAuthHandler + + litellm.proxy_auth = ProxyAuthHandler( + credential=AzureADCredential(), + scope="api://my-litellm-proxy/.default" + ) + litellm.api_base = "https://my-proxy.example.com" + + # Auth headers are now automatically injected + response = litellm.completion(model="gpt-4", messages=[...]) + """ + + def __init__(self, credential: TokenCredential, scope: str): + """ + Initialize the proxy auth handler. + + Args: + credential: A TokenCredential implementation (AzureADCredential, + GenericOAuth2Credential, or any custom implementation) + scope: The OAuth2 scope to request tokens for + """ + self.credential = credential + self.scope = scope + self._cached_token: Optional[AccessToken] = None + + def get_token(self) -> AccessToken: + """ + Get a valid access token, refreshing if necessary. + + Returns: + AccessToken that is valid for at least 60 more seconds. + """ + # Refresh if no token or token expires within 60 seconds + if not self._cached_token or self._cached_token.expires_on <= time.time() + 60: + self._cached_token = self.credential.get_token(self.scope) + return self._cached_token + + def get_auth_headers(self) -> dict: + """ + Get HTTP headers for authentication. + + Returns: + Dict with Authorization header containing Bearer token. + """ + token = self.get_token() + return {"Authorization": f"Bearer {token.token}"} diff --git a/litellm/realtime_api/main.py b/litellm/realtime_api/main.py index 40983fa55fa..01b83067650 100644 --- a/litellm/realtime_api/main.py +++ b/litellm/realtime_api/main.py @@ -19,11 +19,13 @@ from ..llms.azure.realtime.handler import AzureOpenAIRealtime from ..llms.bedrock.realtime.handler import BedrockRealtime from ..llms.custom_httpx.http_handler import get_shared_realtime_ssl_context from ..llms.openai.realtime.handler import OpenAIRealtime +from ..llms.xai.realtime.handler import XAIRealtime from ..utils import client as wrapper_client azure_realtime = AzureOpenAIRealtime() openai_realtime = OpenAIRealtime() bedrock_realtime = BedrockRealtime() +xai_realtime = XAIRealtime() base_llm_http_handler = BaseLLMHTTPHandler() @@ -188,6 +190,30 @@ async def _arealtime( aws_bedrock_runtime_endpoint=aws_bedrock_runtime_endpoint, aws_external_id=aws_external_id, ) + elif _custom_llm_provider == "xai": + api_base = ( + dynamic_api_base + or litellm_params.api_base + or get_secret_str("XAI_API_BASE") + or "https://api.x.ai/v1" + ) + # set API KEY + api_key = ( + dynamic_api_key + or litellm.api_key + or get_secret_str("XAI_API_KEY") + ) + + await xai_realtime.async_realtime( + model=model, + websocket=websocket, + logging_obj=litellm_logging_obj, + api_base=api_base, + api_key=api_key, + client=None, + timeout=timeout, + query_params=query_params, + ) else: raise ValueError(f"Unsupported model: {model}") @@ -230,6 +256,10 @@ async def _realtime_health_check( url = openai_realtime._construct_url( api_base=api_base or "https://api.openai.com/", query_params={"model": model} ) + elif custom_llm_provider == "xai": + url = xai_realtime._construct_url( + api_base=api_base or "https://api.x.ai/v1", query_params={"model": model} + ) else: raise ValueError(f"Unsupported model: {model}") ssl_context = get_shared_realtime_ssl_context() diff --git a/litellm/responses/main.py b/litellm/responses/main.py index b2c2493c812..8f524690be1 100644 --- a/litellm/responses/main.py +++ b/litellm/responses/main.py @@ -24,6 +24,7 @@ from litellm.constants import request_timeout from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.prompt_templates.common_utils import ( update_responses_input_with_model_file_ids, + update_responses_tools_with_model_file_ids, ) from litellm.llms.base_llm.responses.transformation import BaseResponsesAPIConfig from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler @@ -595,13 +596,32 @@ def responses( litellm_params.api_base = dynamic_api_base ######################################################### - # Update input with provider-specific file IDs if managed files are used + # Update input and tools with provider-specific file IDs if managed files are used ######################################################### + model_file_id_mapping = kwargs.get("model_file_id_mapping") + model_info_id = kwargs.get("model_info", {}).get("id") if isinstance(kwargs.get("model_info"), dict) else None + input = cast( Union[str, ResponseInputParam], - update_responses_input_with_model_file_ids(input=input), + update_responses_input_with_model_file_ids( + input=input, + model_id=model_info_id, + model_file_id_mapping=model_file_id_mapping, + ), ) local_vars["input"] = input + + # Update tools with provider-specific file IDs if needed + if tools: + tools = cast( + Optional[Iterable[ToolParam]], + update_responses_tools_with_model_file_ids( + tools=cast(Optional[List[Dict[str, Any]]], tools), + model_id=model_info_id, + model_file_id_mapping=model_file_id_mapping, + ), + ) + local_vars["tools"] = tools ######################################################### # Native MCP Responses API diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index ca22049720e..74ccb34ca6e 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -14,14 +14,14 @@ from litellm.types.proxy.guardrails.guardrail_hooks.grayswan import ( from litellm.types.proxy.guardrails.guardrail_hooks.ibm import ( IBMGuardrailsBaseConfigModel, ) -from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( - ToolPermissionGuardrailConfigModel, +from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( + ContentFilterCategoryConfig, ) from litellm.types.proxy.guardrails.guardrail_hooks.qualifire import ( QualifireGuardrailConfigModel, ) -from litellm.types.proxy.guardrails.guardrail_hooks.litellm_content_filter import ( - ContentFilterCategoryConfig, +from litellm.types.proxy.guardrails.guardrail_hooks.tool_permission import ( + ToolPermissionGuardrailConfigModel, ) """ @@ -68,6 +68,7 @@ class SupportedGuardrailIntegrations(Enum): PROMPT_SECURITY = "prompt_security" GENERIC_GUARDRAIL_API = "generic_guardrail_api" QUALIFIRE = "qualifire" + CUSTOM_CODE = "custom_code" class Role(Enum): @@ -296,13 +297,7 @@ class PresidioConfigModel(PresidioPresidioConfigModelUserInterface): pii_entities_config: Optional[Dict[Union[PiiEntityType, str], PiiAction]] = Field( default=None, description="Configuration for PII entity types and actions" ) - presidio_filter_scope: Literal["input", "output", "both"] = Field( - default="both", - description=( - "Where to apply Presidio checks: 'input' runs on user → model traffic, " - "'output' runs on model → user traffic, and 'both' applies to both." - ), - ) + presidio_score_thresholds: Optional[Dict[Union[PiiEntityType, str], float]] = Field( default=None, description=( @@ -656,6 +651,12 @@ class BaseLitellmParams( description="Additional provider-specific parameters for generic guardrail APIs", ) + # Custom code guardrail params + custom_code: Optional[str] = Field( + default=None, + description="Python-like code containing the apply_guardrail function for custom guardrail logic", + ) + model_config = ConfigDict(extra="allow", protected_namespaces=()) diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index 6293efe9e09..998c60ab60d 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -8,6 +8,7 @@ from .openai import ChatCompletionToolCallChunk class CachePointBlock(TypedDict, total=False): type: Literal["default"] + ttl: str class SystemContentBlock(TypedDict, total=False): @@ -961,6 +962,7 @@ class BedrockGetBatchResponse(TypedDict, total=False): timeoutDurationInHours: Optional[int] clientRequestToken: Optional[str] + class BedrockToolBlock(TypedDict, total=False): toolSpec: Optional[ToolSpecBlock] systemTool: Optional[SystemToolBlock] # For Nova grounding diff --git a/litellm/types/search.py b/litellm/types/search.py index 661a2feda33..b0ce0636aed 100644 --- a/litellm/types/search.py +++ b/litellm/types/search.py @@ -60,6 +60,7 @@ class SearchToolInfoResponse(TypedDict, total=False): search_tool_info: Optional[dict] created_at: Optional[str] updated_at: Optional[str] + is_from_config: Optional[bool] # True if this tool is defined in config file, False if from DB class ListSearchToolsResponse(TypedDict): diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 95be59acc8c..1ae031201d0 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -744,7 +744,8 @@ "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, - "tool_use_system_prompt_tokens": 346 + "tool_use_system_prompt_tokens": 346, + "supports_native_streaming": true }, "anthropic.claude-3-5-sonnet-20240620-v1:0": { "input_cost_per_token": 3e-06, @@ -12850,6 +12851,40 @@ "supports_vision": true, "supports_web_search": true }, + "deep-research-pro-preview-12-2025": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_image_token": 0.00012, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true + }, "gemini-2.5-flash-lite": { "cache_read_input_token_cost": 1e-08, "input_cost_per_audio_token": 3e-07, @@ -13304,7 +13339,8 @@ "supports_tool_choice": true, "supports_video_input": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_native_streaming": true }, "vertex_ai/gemini-3-pro-preview": { "cache_read_input_token_cost": 2e-07, @@ -13352,7 +13388,8 @@ "supports_tool_choice": true, "supports_video_input": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_native_streaming": true }, "vertex_ai/gemini-3-flash-preview": { "cache_read_input_token_cost": 5e-08, @@ -13395,7 +13432,8 @@ "supports_tool_choice": true, "supports_video_input": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_native_streaming": true }, "gemini-2.5-pro-exp-03-25": { "cache_read_input_token_cost": 1.25e-07, @@ -14762,6 +14800,42 @@ "supports_vision": true, "supports_web_search": true }, + "gemini/deep-research-pro-preview-12-2025": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "gemini", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_image_token": 0.00012, + "output_cost_per_token": 1.2e-05, + "rpm": 1000, + "tpm": 4000000, + "output_cost_per_token_batches": 6e-06, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/completions", + "/v1/batch" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text", + "image" + ], + "supports_function_calling": false, + "supports_prompt_caching": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_vision": true, + "supports_web_search": true + }, "gemini/gemini-2.5-flash-lite": { "cache_read_input_token_cost": 1e-08, "input_cost_per_audio_token": 3e-07, @@ -15346,6 +15420,7 @@ "supports_url_context": true, "supports_vision": true, "supports_web_search": true, + "supports_native_streaming": true, "tpm": 800000 }, "gemini-3-flash-preview": { @@ -15391,7 +15466,8 @@ "supports_tool_choice": true, "supports_url_context": true, "supports_vision": true, - "supports_web_search": true + "supports_web_search": true, + "supports_native_streaming": true }, "gemini/gemini-2.5-pro-exp-03-25": { "cache_read_input_token_cost": 0.0, @@ -24343,6 +24419,31 @@ "supports_tool_choice": true, "supports_function_calling": true }, + "openrouter/qwen/qwen3-235b-a22b-2507": { + "input_cost_per_token": 7.1e-08, + "litellm_provider": "openrouter", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 1e-07, + "source": "https://openrouter.ai/qwen/qwen3-235b-a22b-2507", + "supports_function_calling": true, + "supports_tool_choice": true + }, + "openrouter/qwen/qwen3-235b-a22b-thinking-2507": { + "input_cost_per_token": 1.1e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 262144, + "max_output_tokens": 262144, + "max_tokens": 262144, + "mode": "chat", + "output_cost_per_token": 6e-07, + "source": "https://openrouter.ai/qwen/qwen3-235b-a22b-thinking-2507", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true + }, "openrouter/switchpoint/router": { "input_cost_per_token": 8.5e-07, "litellm_provider": "openrouter", @@ -27936,7 +28037,9 @@ "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", - "output_cost_per_token": 3e-07 + "output_cost_per_token": 3e-07, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/alibaba/qwen3-coder": { "input_cost_per_token": 4e-07, @@ -27945,7 +28048,9 @@ "max_output_tokens": 66536, "max_tokens": 66536, "mode": "chat", - "output_cost_per_token": 1.6e-06 + "output_cost_per_token": 1.6e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/amazon/nova-lite": { "input_cost_per_token": 6e-08, @@ -27954,7 +28059,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 2.4e-07 + "output_cost_per_token": 2.4e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_response_schema": true }, "vercel_ai_gateway/amazon/nova-micro": { "input_cost_per_token": 3.5e-08, @@ -27963,7 +28071,9 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 1.4e-07 + "output_cost_per_token": 1.4e-07, + "supports_function_calling": true, + "supports_response_schema": true }, "vercel_ai_gateway/amazon/nova-pro": { "input_cost_per_token": 8e-07, @@ -27972,7 +28082,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 3.2e-06 + "output_cost_per_token": 3.2e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_response_schema": true }, "vercel_ai_gateway/amazon/titan-embed-text-v2": { "input_cost_per_token": 2e-08, @@ -27992,7 +28105,11 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 1.25e-06 + "output_cost_per_token": 1.25e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-3-opus": { "cache_creation_input_token_cost": 1.875e-05, @@ -28003,7 +28120,11 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 7.5e-05 + "output_cost_per_token": 7.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-3.5-haiku": { "cache_creation_input_token_cost": 1e-06, @@ -28014,7 +28135,11 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 4e-06 + "output_cost_per_token": 4e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-3.5-sonnet": { "cache_creation_input_token_cost": 3.75e-06, @@ -28025,7 +28150,11 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-3.7-sonnet": { "cache_creation_input_token_cost": 3.75e-06, @@ -28036,7 +28165,11 @@ "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-4-opus": { "cache_creation_input_token_cost": 1.875e-05, @@ -28047,7 +28180,11 @@ "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", - "output_cost_per_token": 7.5e-05 + "output_cost_per_token": 7.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/anthropic/claude-4-sonnet": { "cache_creation_input_token_cost": 3.75e-06, @@ -28058,7 +28195,9 @@ "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/cohere/command-a": { "input_cost_per_token": 2.5e-06, @@ -28067,7 +28206,9 @@ "max_output_tokens": 8000, "max_tokens": 8000, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/cohere/command-r": { "input_cost_per_token": 1.5e-07, @@ -28076,7 +28217,9 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 6e-07 + "output_cost_per_token": 6e-07, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/cohere/command-r-plus": { "input_cost_per_token": 2.5e-06, @@ -28085,7 +28228,9 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/cohere/embed-v4.0": { "input_cost_per_token": 1.2e-07, @@ -28103,7 +28248,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 2.19e-06 + "output_cost_per_token": 2.19e-06, + "supports_tool_choice": true }, "vercel_ai_gateway/deepseek/deepseek-r1-distill-llama-70b": { "input_cost_per_token": 7.5e-07, @@ -28112,7 +28258,10 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 9.9e-07 + "output_cost_per_token": 9.9e-07, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/deepseek/deepseek-v3": { "input_cost_per_token": 9e-07, @@ -28121,7 +28270,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 9e-07 + "output_cost_per_token": 9e-07, + "supports_tool_choice": true }, "vercel_ai_gateway/google/gemini-2.0-flash": { "deprecation_date": "2026-03-31", @@ -28131,7 +28281,11 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 6e-07 + "output_cost_per_token": 6e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/google/gemini-2.0-flash-lite": { "deprecation_date": "2026-03-31", @@ -28141,7 +28295,11 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 3e-07 + "output_cost_per_token": 3e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/google/gemini-2.5-flash": { "input_cost_per_token": 3e-07, @@ -28150,7 +28308,11 @@ "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", - "output_cost_per_token": 2.5e-06 + "output_cost_per_token": 2.5e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/google/gemini-2.5-pro": { "input_cost_per_token": 2.5e-06, @@ -28159,7 +28321,11 @@ "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/google/gemini-embedding-001": { "input_cost_per_token": 1.5e-07, @@ -28177,7 +28343,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 2e-07 + "output_cost_per_token": 2e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/google/text-embedding-005": { "input_cost_per_token": 2.5e-08, @@ -28213,7 +28382,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 7.9e-07 + "output_cost_per_token": 7.9e-07, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-3-8b": { "input_cost_per_token": 5e-08, @@ -28222,7 +28392,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 8e-08 + "output_cost_per_token": 8e-08, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-3.1-70b": { "input_cost_per_token": 7.2e-07, @@ -28231,7 +28402,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 7.2e-07 + "output_cost_per_token": 7.2e-07, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-3.1-8b": { "input_cost_per_token": 5e-08, @@ -28240,7 +28412,9 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 8e-08 + "output_cost_per_token": 8e-08, + "supports_function_calling": true, + "supports_response_schema": true }, "vercel_ai_gateway/meta/llama-3.2-11b": { "input_cost_per_token": 1.6e-07, @@ -28249,7 +28423,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 1.6e-07 + "output_cost_per_token": 1.6e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-3.2-1b": { "input_cost_per_token": 1e-07, @@ -28267,7 +28444,9 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 1.5e-07 + "output_cost_per_token": 1.5e-07, + "supports_function_calling": true, + "supports_response_schema": true }, "vercel_ai_gateway/meta/llama-3.2-90b": { "input_cost_per_token": 7.2e-07, @@ -28276,7 +28455,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 7.2e-07 + "output_cost_per_token": 7.2e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-3.3-70b": { "input_cost_per_token": 7.2e-07, @@ -28285,7 +28467,9 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 7.2e-07 + "output_cost_per_token": 7.2e-07, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-4-maverick": { "input_cost_per_token": 2e-07, @@ -28294,7 +28478,8 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 6e-07 + "output_cost_per_token": 6e-07, + "supports_tool_choice": true }, "vercel_ai_gateway/meta/llama-4-scout": { "input_cost_per_token": 1e-07, @@ -28303,7 +28488,10 @@ "max_output_tokens": 8192, "max_tokens": 8192, "mode": "chat", - "output_cost_per_token": 3e-07 + "output_cost_per_token": 3e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/mistral/codestral": { "input_cost_per_token": 3e-07, @@ -28312,7 +28500,9 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 9e-07 + "output_cost_per_token": 9e-07, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/mistral/codestral-embed": { "input_cost_per_token": 1.5e-07, @@ -28330,7 +28520,10 @@ "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 2.8e-07 + "output_cost_per_token": 2.8e-07, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/mistral/magistral-medium": { "input_cost_per_token": 2e-06, @@ -28339,7 +28532,10 @@ "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", - "output_cost_per_token": 5e-06 + "output_cost_per_token": 5e-06, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/mistral/magistral-small": { "input_cost_per_token": 5e-07, @@ -28348,7 +28544,8 @@ "max_output_tokens": 64000, "max_tokens": 64000, "mode": "chat", - "output_cost_per_token": 1.5e-06 + "output_cost_per_token": 1.5e-06, + "supports_function_calling": true }, "vercel_ai_gateway/mistral/ministral-3b": { "input_cost_per_token": 4e-08, @@ -28357,7 +28554,9 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 4e-08 + "output_cost_per_token": 4e-08, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/mistral/ministral-8b": { "input_cost_per_token": 1e-07, @@ -28366,7 +28565,10 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 1e-07 + "output_cost_per_token": 1e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/mistral/mistral-embed": { "input_cost_per_token": 1e-07, @@ -28384,7 +28586,9 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 6e-06 + "output_cost_per_token": 6e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/mistral/mistral-saba-24b": { "input_cost_per_token": 7.9e-07, @@ -28402,7 +28606,10 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 3e-07 + "output_cost_per_token": 3e-07, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/mistral/mixtral-8x22b-instruct": { "input_cost_per_token": 1.2e-06, @@ -28411,7 +28618,8 @@ "max_output_tokens": 2048, "max_tokens": 2048, "mode": "chat", - "output_cost_per_token": 1.2e-06 + "output_cost_per_token": 1.2e-06, + "supports_function_calling": true }, "vercel_ai_gateway/mistral/pixtral-12b": { "input_cost_per_token": 1.5e-07, @@ -28420,7 +28628,11 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 1.5e-07 + "output_cost_per_token": 1.5e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/mistral/pixtral-large": { "input_cost_per_token": 2e-06, @@ -28429,7 +28641,11 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 6e-06 + "output_cost_per_token": 6e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/moonshotai/kimi-k2": { "input_cost_per_token": 5.5e-07, @@ -28438,7 +28654,9 @@ "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", - "output_cost_per_token": 2.2e-06 + "output_cost_per_token": 2.2e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/morph/morph-v3-fast": { "input_cost_per_token": 8e-07, @@ -28465,7 +28683,9 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 1.5e-06 + "output_cost_per_token": 1.5e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/openai/gpt-3.5-turbo-instruct": { "input_cost_per_token": 1.5e-06, @@ -28483,7 +28703,10 @@ "max_output_tokens": 4096, "max_tokens": 4096, "mode": "chat", - "output_cost_per_token": 3e-05 + "output_cost_per_token": 3e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/openai/gpt-4.1": { "cache_creation_input_token_cost": 0.0, @@ -28494,7 +28717,11 @@ "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 8e-06 + "output_cost_per_token": 8e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/gpt-4.1-mini": { "cache_creation_input_token_cost": 0.0, @@ -28505,7 +28732,11 @@ "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 1.6e-06 + "output_cost_per_token": 1.6e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/gpt-4.1-nano": { "cache_creation_input_token_cost": 0.0, @@ -28516,7 +28747,11 @@ "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 4e-07 + "output_cost_per_token": 4e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/gpt-4o": { "cache_creation_input_token_cost": 0.0, @@ -28527,7 +28762,11 @@ "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/gpt-4o-mini": { "cache_creation_input_token_cost": 0.0, @@ -28538,7 +28777,11 @@ "max_output_tokens": 16384, "max_tokens": 16384, "mode": "chat", - "output_cost_per_token": 6e-07 + "output_cost_per_token": 6e-07, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/o1": { "cache_creation_input_token_cost": 0.0, @@ -28549,7 +28792,11 @@ "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", - "output_cost_per_token": 6e-05 + "output_cost_per_token": 6e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/o3": { "cache_creation_input_token_cost": 0.0, @@ -28560,7 +28807,11 @@ "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", - "output_cost_per_token": 8e-06 + "output_cost_per_token": 8e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/o3-mini": { "cache_creation_input_token_cost": 0.0, @@ -28571,7 +28822,10 @@ "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", - "output_cost_per_token": 4.4e-06 + "output_cost_per_token": 4.4e-06, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/o4-mini": { "cache_creation_input_token_cost": 0.0, @@ -28582,7 +28836,11 @@ "max_output_tokens": 100000, "max_tokens": 100000, "mode": "chat", - "output_cost_per_token": 4.4e-06 + "output_cost_per_token": 4.4e-06, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_response_schema": true }, "vercel_ai_gateway/openai/text-embedding-3-large": { "input_cost_per_token": 1.3e-07, @@ -28654,7 +28912,10 @@ "max_output_tokens": 32000, "max_tokens": 32000, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/vercel/v0-1.5-md": { "input_cost_per_token": 3e-06, @@ -28663,7 +28924,10 @@ "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-2": { "input_cost_per_token": 2e-06, @@ -28672,7 +28936,9 @@ "max_output_tokens": 4000, "max_tokens": 4000, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-2-vision": { "input_cost_per_token": 2e-06, @@ -28681,7 +28947,10 @@ "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 1e-05 + "output_cost_per_token": 1e-05, + "supports_vision": true, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-3": { "input_cost_per_token": 3e-06, @@ -28690,7 +28959,9 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-3-fast": { "input_cost_per_token": 5e-06, @@ -28699,7 +28970,8 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 2.5e-05 + "output_cost_per_token": 2.5e-05, + "supports_function_calling": true }, "vercel_ai_gateway/xai/grok-3-mini": { "input_cost_per_token": 3e-07, @@ -28708,7 +28980,9 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 5e-07 + "output_cost_per_token": 5e-07, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-3-mini-fast": { "input_cost_per_token": 6e-07, @@ -28717,7 +28991,9 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 4e-06 + "output_cost_per_token": 4e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/xai/grok-4": { "input_cost_per_token": 3e-06, @@ -28726,7 +29002,9 @@ "max_output_tokens": 256000, "max_tokens": 256000, "mode": "chat", - "output_cost_per_token": 1.5e-05 + "output_cost_per_token": 1.5e-05, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/zai/glm-4.5": { "input_cost_per_token": 6e-07, @@ -28735,7 +29013,9 @@ "max_output_tokens": 131072, "max_tokens": 131072, "mode": "chat", - "output_cost_per_token": 2.2e-06 + "output_cost_per_token": 2.2e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/zai/glm-4.5-air": { "input_cost_per_token": 2e-07, @@ -28744,7 +29024,9 @@ "max_output_tokens": 96000, "max_tokens": 96000, "mode": "chat", - "output_cost_per_token": 1.1e-06 + "output_cost_per_token": 1.1e-06, + "supports_function_calling": true, + "supports_tool_choice": true }, "vercel_ai_gateway/zai/glm-4.6": { "litellm_provider": "vercel_ai_gateway", @@ -28812,7 +29094,9 @@ "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "supports_native_streaming": true, + "supports_vision": true }, "vertex_ai/claude-3-5-sonnet": { "input_cost_per_token": 3e-06, @@ -29083,7 +29367,8 @@ "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, - "tool_use_system_prompt_tokens": 159 + "tool_use_system_prompt_tokens": 159, + "supports_native_streaming": true }, "vertex_ai/claude-sonnet-4-5": { "cache_creation_input_token_cost": 3.75e-06, @@ -29135,7 +29420,8 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "supports_native_streaming": true }, "vertex_ai/claude-opus-4@20250514": { "cache_creation_input_token_cost": 1.875e-05, @@ -29417,6 +29703,21 @@ "output_cost_per_token_batches": 6e-06, "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" }, + "vertex_ai/deep-research-pro-preview-12-2025": { + "input_cost_per_image": 0.0011, + "input_cost_per_token": 2e-06, + "input_cost_per_token_batches": 1e-06, + "litellm_provider": "vertex_ai-language-models", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "image_generation", + "output_cost_per_image": 0.134, + "output_cost_per_image_token": 0.00012, + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_batches": 6e-06, + "source": "https://docs.cloud.google.com/vertex-ai/generative-ai/docs/models/gemini/3-pro-image" + }, "vertex_ai/imagegeneration@006": { "litellm_provider": "vertex_ai-image-models", "mode": "image_generation", @@ -29906,7 +30207,9 @@ "mode": "chat", "output_cost_per_token": 1e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supported_regions": ["global"], + "supported_regions": [ + "global" + ], "supports_function_calling": true, "supports_tool_choice": true }, @@ -29919,7 +30222,9 @@ "mode": "chat", "output_cost_per_token": 4e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supported_regions": ["global"], + "supported_regions": [ + "global" + ], "supports_function_calling": true, "supports_tool_choice": true }, @@ -29932,7 +30237,9 @@ "mode": "chat", "output_cost_per_token": 1.2e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supported_regions": ["global"], + "supported_regions": [ + "global" + ], "supports_function_calling": true, "supports_tool_choice": true }, @@ -29945,7 +30252,9 @@ "mode": "chat", "output_cost_per_token": 1.2e-06, "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", - "supported_regions": ["global"], + "supported_regions": [ + "global" + ], "supports_function_calling": true, "supports_tool_choice": true }, @@ -34894,4 +35203,4 @@ "output_cost_per_token": 0, "supports_reasoning": true } -} \ No newline at end of file +} diff --git a/poetry.lock b/poetry.lock index 537367c5aa0..5e926509d54 100644 --- a/poetry.lock +++ b/poetry.lock @@ -1,4 +1,4 @@ -# This file is automatically @generated by Poetry 2.2.0 and should not be changed by hand. +# This file is automatically @generated by Poetry 2.3.2 and should not be changed by hand. [[package]] name = "a2a-sdk" @@ -398,6 +398,7 @@ files = [ {file = "azure_core-1.36.0-py3-none-any.whl", hash = "sha256:fee9923a3a753e94a259563429f3644aaf05c486d45b1215d098115102d91d3b"}, {file = "azure_core-1.36.0.tar.gz", hash = "sha256:22e5605e6d0bf1d229726af56d9e92bc37b6e726b141a18be0b4d424131741b7"}, ] +markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""} [package.dependencies] requests = ">=2.21.0" @@ -418,6 +419,7 @@ files = [ {file = "azure_identity-1.25.1-py3-none-any.whl", hash = "sha256:e9edd720af03dff020223cd269fa3a61e8f345ea75443858273bcb44844ab651"}, {file = "azure_identity-1.25.1.tar.gz", hash = "sha256:87ca8328883de6036443e1c37b40e8dc8fb74898240f61071e09d2e369361456"}, ] +markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""} [package.dependencies] azure-core = ">=1.31.0" @@ -718,11 +720,23 @@ files = [ {file = "cffi-2.0.0-cp39-cp39-win_amd64.whl", hash = "sha256:b882b3df248017dba09d6b16defe9b5c407fe32fc7c65a9c69798e6175601be9"}, {file = "cffi-2.0.0.tar.gz", hash = "sha256:44d1b5909021139fe36001ae048dbdde8214afa20200eda0f64c068cac5d5529"}, ] -markers = {main = "platform_python_implementation != \"PyPy\" or extra == \"proxy\"", dev = "platform_python_implementation != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\""} +markers = {main = "(platform_python_implementation != \"PyPy\" or extra == \"proxy\") and (python_version >= \"3.10\" or extra == \"proxy\" or extra == \"extra-proxy\") and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\")", dev = "platform_python_implementation != \"PyPy\"", proxy-dev = "platform_python_implementation != \"PyPy\""} [package.dependencies] pycparser = {version = "*", markers = "implementation_name != \"PyPy\""} +[[package]] +name = "chardet" +version = "5.2.0" +description = "Universal encoding detector for Python 3" +optional = false +python-versions = ">=3.7" +groups = ["dev"] +files = [ + {file = "chardet-5.2.0-py3-none-any.whl", hash = "sha256:e1cf59446890a00105fe7b7912492ea04b6e6f06d4b742b2c788469e34c82970"}, + {file = "chardet-5.2.0.tar.gz", hash = "sha256:1b3b6ff479a8c414bc3fa2c0852995695c4a026dcd6d0633b2dd092ca39c1cf7"}, +] + [[package]] name = "charset-normalizer" version = "3.4.4" @@ -1137,7 +1151,6 @@ description = "cryptography is a package which provides cryptographic recipes an optional = false python-versions = ">=3.7" groups = ["main", "dev", "proxy-dev"] -markers = "python_version == \"3.9\"" files = [ {file = "cryptography-43.0.3-cp37-abi3-macosx_10_9_universal2.whl", hash = "sha256:bf7a1932ac4176486eab36a19ed4c0492da5d97123f1406cf15e41b05e787d2e"}, {file = "cryptography-43.0.3-cp37-abi3-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:63efa177ff54aec6e1c0aefaa1a241232dcd37413835a9b674b6e3f0ae2bfd3e"}, @@ -1167,6 +1180,7 @@ files = [ {file = "cryptography-43.0.3-pp39-pypy39_pp73-win_amd64.whl", hash = "sha256:2ce6fae5bdad59577b44e4dfed356944fbf1d925269114c28be377692643b4ff"}, {file = "cryptography-43.0.3.tar.gz", hash = "sha256:315b9001266a492a6ff443b61238f956b214dbec9910a081ba5b6646a055a805"}, ] +markers = {main = "python_version == \"3.9\" and (extra == \"proxy\" or extra == \"extra-proxy\")", dev = "python_version == \"3.9\"", proxy-dev = "python_version == \"3.9\""} [package.dependencies] cffi = {version = ">=1.12", markers = "platform_python_implementation != \"PyPy\""} @@ -1188,7 +1202,6 @@ description = "cryptography is a package which provides cryptographic recipes an optional = false python-versions = "!=3.9.0,!=3.9.1,>=3.8" groups = ["main", "dev", "proxy-dev"] -markers = "python_version >= \"3.10\"" files = [ {file = "cryptography-46.0.3-cp311-abi3-macosx_10_9_universal2.whl", hash = "sha256:109d4ddfadf17e8e7779c39f9b18111a09efb969a301a31e987416a0191ed93a"}, {file = "cryptography-46.0.3-cp311-abi3-manylinux2014_aarch64.manylinux_2_17_aarch64.whl", hash = "sha256:09859af8466b69bc3c27bdf4f5d84a665e0f7ab5088412e9e2ec49758eca5cbc"}, @@ -1245,6 +1258,7 @@ files = [ {file = "cryptography-46.0.3-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:6b5063083824e5509fdba180721d55909ffacccc8adbec85268b48439423d78c"}, {file = "cryptography-46.0.3.tar.gz", hash = "sha256:a8b17438104fed022ce745b362294d9ce35b4c2e45c1d958ad4a4b019285f4a1"}, ] +markers = {main = "python_version >= \"3.10\" and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\")", dev = "python_version >= \"3.10\"", proxy-dev = "python_version >= \"3.10\""} [package.dependencies] cffi = {version = ">=2.0.0", markers = "python_full_version >= \"3.9.0\" and platform_python_implementation != \"PyPy\""} @@ -1300,6 +1314,27 @@ dev = ["autoflake", "black", "build", "databricks-connect", "httpx", "ipython", notebook = ["ipython (>=8,<10)", "ipywidgets (>=8,<9)"] openai = ["httpx", "langchain-openai ; python_version > \"3.7\"", "openai"] +[[package]] +name = "diff-cover" +version = "9.7.2" +description = "Run coverage and linting reports on diffs" +optional = false +python-versions = ">=3.9" +groups = ["dev"] +files = [ + {file = "diff_cover-9.7.2-py3-none-any.whl", hash = "sha256:cd6498620c747c2493a6c83c14362c32868bfd91cd8d0dd093f136070ec4ffc5"}, + {file = "diff_cover-9.7.2.tar.gz", hash = "sha256:872c820d2ecbf79c61d52c7dc70419015e0ab9289589566c791dd270fc0c6e3b"}, +] + +[package.dependencies] +chardet = ">=3.0.0" +Jinja2 = ">=2.7.1" +pluggy = ">=0.13.1,<2" +Pygments = ">=2.19.1,<3.0.0" + +[package.extras] +toml = ["tomli (>=1.2.1)"] + [[package]] name = "diskcache" version = "5.6.3" @@ -2242,11 +2277,11 @@ files = [ ] [package.dependencies] -google-api-core = {version = ">=1.34.1,<2.0.dev0 || >=2.11.dev0,<3.0.0dev", extras = ["grpc"]} -google-auth = ">=2.14.1,<2.24.0 || >2.24.0,<2.25.0 || >2.25.0,<3.0.0dev" -grpc-google-iam-v1 = ">=0.12.4,<1.0.0dev" -proto-plus = ">=1.22.3,<2.0.0dev" -protobuf = ">=3.20.2,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<6.0.0dev" +google-api-core = {version = ">=1.34.1,<2.0.dev0 || >=2.11.dev0,<3.0.0.dev0", extras = ["grpc"]} +google-auth = ">=2.14.1,<2.24.0 || >2.24.0,<2.25.0 || >2.25.0,<3.0.0.dev0" +grpc-google-iam-v1 = ">=0.12.4,<1.0.0.dev0" +proto-plus = ">=1.22.3,<2.0.0.dev0" +protobuf = ">=3.20.2,<4.21.0 || >4.21.0,<4.21.1 || >4.21.1,<4.21.2 || >4.21.2,<4.21.3 || >4.21.3,<4.21.4 || >4.21.4,<4.21.5 || >4.21.5,<6.0.0.dev0" [[package]] name = "google-cloud-resource-manager" @@ -3085,7 +3120,7 @@ version = "3.1.6" description = "A very fast and expressive template engine." optional = false python-versions = ">=3.7" -groups = ["main", "proxy-dev"] +groups = ["main", "dev", "proxy-dev"] files = [ {file = "jinja2-3.1.6-py3-none-any.whl", hash = "sha256:85ece4451f492d0c13c5dd7c13a64681a86afae63a5f347908daf103ce6d2f67"}, {file = "jinja2-3.1.6.tar.gz", hash = "sha256:0137fb05990d35f1275a587e9aee6d56da821fc83491a0fb838183be43f66d6d"}, @@ -3249,7 +3284,7 @@ files = [ [package.dependencies] attrs = ">=22.2.0" -jsonschema-specifications = ">=2023.03.6" +jsonschema-specifications = ">=2023.3.6" referencing = ">=0.28.4" rpds-py = ">=0.7.1" @@ -3426,15 +3461,15 @@ files = [ [[package]] name = "litellm-proxy-extras" -version = "0.4.29" +version = "0.4.30" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." optional = true python-versions = "!=2.7.*,!=3.0.*,!=3.1.*,!=3.2.*,!=3.3.*,!=3.4.*,!=3.5.*,!=3.6.*,!=3.7.*,>=3.8" groups = ["main"] markers = "extra == \"proxy\"" files = [ - {file = "litellm_proxy_extras-0.4.29-py3-none-any.whl", hash = "sha256:c36c1b69675c61acccc6b61dd610eb37daeb72c6fd819461cefb5b0cc7e0550f"}, - {file = "litellm_proxy_extras-0.4.29.tar.gz", hash = "sha256:1a8266911e0546f1e17e6714ca20b72e9fef47c1683f9c16399cf2d1786437a0"}, + {file = "litellm_proxy_extras-0.4.30-py3-none-any.whl", hash = "sha256:0b7df68f0968eb817462b847eaee81bba23d935adb2e84d2e342a77711887051"}, + {file = "litellm_proxy_extras-0.4.30.tar.gz", hash = "sha256:5d32f8dc3d37d36fb15ab6995fea706dd8a453ff7f12e70b47cba35e5368da10"}, ] [[package]] @@ -3515,7 +3550,7 @@ version = "3.0.3" description = "Safely add untrusted strings to HTML/XML markup." optional = false python-versions = ">=3.9" -groups = ["main", "proxy-dev"] +groups = ["main", "dev", "proxy-dev"] files = [ {file = "markupsafe-3.0.3-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:2f981d352f04553a7171b8e44369f2af4055f888dfb147d55e42d29e29e74559"}, {file = "markupsafe-3.0.3-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:e1c1493fb6e50ab01d20a22826e57520f1284df32f2d8601fdd90b6304601419"}, @@ -3913,6 +3948,7 @@ files = [ {file = "msal-1.34.0-py3-none-any.whl", hash = "sha256:f669b1644e4950115da7a176441b0e13ec2975c29528d8b9e81316023676d6e1"}, {file = "msal-1.34.0.tar.gz", hash = "sha256:76ba83b716ea5a6d75b0279c0ac353a0e05b820ca1f6682c0eb7f45190c43c2f"}, ] +markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""} [package.dependencies] cryptography = ">=2.5,<49" @@ -3933,6 +3969,7 @@ files = [ {file = "msal_extensions-1.3.1-py3-none-any.whl", hash = "sha256:96d3de4d034504e969ac5e85bae8106c8373b5c6568e4c8fa7af2eca9dbe6bca"}, {file = "msal_extensions-1.3.1.tar.gz", hash = "sha256:c5b0fd10f65ef62b5f1d62f4251d51cbcaf003fcedae8c91b040a488614be1a4"}, ] +markers = {main = "extra == \"proxy\" or extra == \"extra-proxy\""} [package.dependencies] msal = ">=1.29,<2" @@ -4183,6 +4220,7 @@ files = [ {file = "nodeenv-1.9.1-py2.py3-none-any.whl", hash = "sha256:ba11c9782d29c27c70ffbdda2d7415098754709be8a7056d79a737cd901155c9"}, {file = "nodeenv-1.9.1.tar.gz", hash = "sha256:6ec12890a2dab7946721edbfbcd91f3319c6ccc9aec47be7c7e6b7011ee6645f"}, ] +markers = {main = "extra == \"extra-proxy\""} [[package]] name = "numpy" @@ -4390,7 +4428,7 @@ files = [ {file = "opentelemetry_api-1.39.1-py3-none-any.whl", hash = "sha256:2edd8463432a7f8443edce90972169b195e7d6a05500cd29e6d13898187c9950"}, {file = "opentelemetry_api-1.39.1.tar.gz", hash = "sha256:fbde8c80e1b937a2c61f20347e91c0c18a1940cecf012d62e65a7caf08967c9c"}, ] -markers = {main = "python_version >= \"3.10\""} +markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""} [package.dependencies] importlib-metadata = ">=6.0,<8.8.0" @@ -4505,7 +4543,7 @@ files = [ {file = "opentelemetry_sdk-1.39.1-py3-none-any.whl", hash = "sha256:4d5482c478513ecb0a5d938dcc61394e647066e0cc2676bee9f3af3f3f45f01c"}, {file = "opentelemetry_sdk-1.39.1.tar.gz", hash = "sha256:cf4d4563caf7bff906c9f7967e2be22d0d6b349b908be0d90fb21c8e9c995cc6"}, ] -markers = {main = "python_version >= \"3.10\""} +markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""} [package.dependencies] opentelemetry-api = "1.39.1" @@ -4523,7 +4561,7 @@ files = [ {file = "opentelemetry_semantic_conventions-0.60b1-py3-none-any.whl", hash = "sha256:9fa8c8b0c110da289809292b0591220d3a7b53c1526a23021e977d68597893fb"}, {file = "opentelemetry_semantic_conventions-0.60b1.tar.gz", hash = "sha256:87c228b5a0669b748c76d76df6c364c369c28f1c465e50f661e39737e84bc953"}, ] -markers = {main = "python_version >= \"3.10\""} +markers = {main = "python_version >= \"3.10\" and extra == \"mlflow\""} [package.dependencies] opentelemetry-api = "1.39.1" @@ -5000,6 +5038,7 @@ files = [ {file = "prisma-0.11.0-py3-none-any.whl", hash = "sha256:22bb869e59a2968b99f3483bb417717273ffbc569fd1e9ceed95e5614cbaf53a"}, {file = "prisma-0.11.0.tar.gz", hash = "sha256:3f2f2fd2361e1ec5ff655f2a04c7860c2f2a5bc4c91f78ca9c5c6349735bf693"}, ] +markers = {main = "extra == \"extra-proxy\""} [package.dependencies] click = ">=7.1.2" @@ -5316,7 +5355,7 @@ files = [ {file = "pycparser-2.23-py3-none-any.whl", hash = "sha256:e5c6e8d3fbad53479cab09ac03729e0a9faf2bee3db8208a550daf5af81a5934"}, {file = "pycparser-2.23.tar.gz", hash = "sha256:78816d4f24add8f10a06d6f05b4d424ad9e96cfebf68a4ddc99c65c0720d00c2"}, ] -markers = {main = "implementation_name != \"PyPy\" and (platform_python_implementation != \"PyPy\" or extra == \"proxy\")", dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\"", proxy-dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\""} +markers = {main = "implementation_name != \"PyPy\" and (platform_python_implementation != \"PyPy\" or extra == \"proxy\") and (python_version >= \"3.10\" or extra == \"proxy\" or extra == \"extra-proxy\") and (extra == \"proxy\" or extra == \"extra-proxy\" or extra == \"mlflow\")", dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\"", proxy-dev = "implementation_name != \"PyPy\" and platform_python_implementation != \"PyPy\""} [[package]] name = "pydantic" @@ -5516,14 +5555,14 @@ files = [ name = "pygments" version = "2.19.2" description = "Pygments is a syntax highlighting package written in Python." -optional = true +optional = false python-versions = ">=3.8" -groups = ["main"] -markers = "extra == \"utils\" or extra == \"proxy\"" +groups = ["main", "dev"] files = [ {file = "pygments-2.19.2-py3-none-any.whl", hash = "sha256:86540386c03d588bb81d44bc3928634ff26449851e99741617ecb9037ee5ec0b"}, {file = "pygments-2.19.2.tar.gz", hash = "sha256:636cb2477cec7f8952536970bc533bc43743542f70392ae026374600add5b887"}, ] +markers = {main = "extra == \"utils\" or extra == \"proxy\""} [package.extras] windows-terminal = ["colorama (>=0.4.6)"] @@ -5539,6 +5578,7 @@ files = [ {file = "PyJWT-2.10.1-py3-none-any.whl", hash = "sha256:dcdd193e30abefd5debf142f9adfcdd2b58004e644f25406ffaebd50bd98dacb"}, {file = "pyjwt-2.10.1.tar.gz", hash = "sha256:3cc5772eb20009233caf06e9d8a0577824723b44e6648ee0a2aedb6cf9381953"}, ] +markers = {main = "extra == \"extra-proxy\" or extra == \"proxy\""} [package.dependencies] cryptography = {version = ">=3.4.0", optional = true, markers = "extra == \"crypto\""} @@ -6601,29 +6641,29 @@ pyasn1 = ">=0.1.3" [[package]] name = "ruff" -version = "0.1.15" +version = "0.2.2" description = "An extremely fast Python linter and code formatter, written in Rust." optional = false python-versions = ">=3.7" groups = ["dev"] files = [ - {file = "ruff-0.1.15-py3-none-macosx_10_12_x86_64.macosx_11_0_arm64.macosx_10_12_universal2.whl", hash = "sha256:5fe8d54df166ecc24106db7dd6a68d44852d14eb0729ea4672bb4d96c320b7df"}, - {file = "ruff-0.1.15-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:6f0bfbb53c4b4de117ac4d6ddfd33aa5fc31beeaa21d23c45c6dd249faf9126f"}, - {file = "ruff-0.1.15-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:e0d432aec35bfc0d800d4f70eba26e23a352386be3a6cf157083d18f6f5881c8"}, - {file = "ruff-0.1.15-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:9405fa9ac0e97f35aaddf185a1be194a589424b8713e3b97b762336ec79ff807"}, - {file = "ruff-0.1.15-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:c66ec24fe36841636e814b8f90f572a8c0cb0e54d8b5c2d0e300d28a0d7bffec"}, - {file = "ruff-0.1.15-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:6f8ad828f01e8dd32cc58bc28375150171d198491fc901f6f98d2a39ba8e3ff5"}, - {file = "ruff-0.1.15-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:86811954eec63e9ea162af0ffa9f8d09088bab51b7438e8b6488b9401863c25e"}, - {file = "ruff-0.1.15-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:fd4025ac5e87d9b80e1f300207eb2fd099ff8200fa2320d7dc066a3f4622dc6b"}, - {file = "ruff-0.1.15-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:b17b93c02cdb6aeb696effecea1095ac93f3884a49a554a9afa76bb125c114c1"}, - {file = "ruff-0.1.15-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:ddb87643be40f034e97e97f5bc2ef7ce39de20e34608f3f829db727a93fb82c5"}, - {file = "ruff-0.1.15-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:abf4822129ed3a5ce54383d5f0e964e7fef74a41e48eb1dfad404151efc130a2"}, - {file = "ruff-0.1.15-py3-none-musllinux_1_2_i686.whl", hash = "sha256:6c629cf64bacfd136c07c78ac10a54578ec9d1bd2a9d395efbee0935868bf852"}, - {file = "ruff-0.1.15-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:1bab866aafb53da39c2cadfb8e1c4550ac5340bb40300083eb8967ba25481447"}, - {file = "ruff-0.1.15-py3-none-win32.whl", hash = "sha256:2417e1cb6e2068389b07e6fa74c306b2810fe3ee3476d5b8a96616633f40d14f"}, - {file = "ruff-0.1.15-py3-none-win_amd64.whl", hash = "sha256:3837ac73d869efc4182d9036b1405ef4c73d9b1f88da2413875e34e0d6919587"}, - {file = "ruff-0.1.15-py3-none-win_arm64.whl", hash = "sha256:9a933dfb1c14ec7a33cceb1e49ec4a16b51ce3c20fd42663198746efc0427360"}, - {file = "ruff-0.1.15.tar.gz", hash = "sha256:f6dfa8c1b21c913c326919056c390966648b680966febcb796cc9d1aaab8564e"}, + {file = "ruff-0.2.2-py3-none-macosx_10_12_x86_64.macosx_11_0_arm64.macosx_10_12_universal2.whl", hash = "sha256:0a9efb032855ffb3c21f6405751d5e147b0c6b631e3ca3f6b20f917572b97eb6"}, + {file = "ruff-0.2.2-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:d450b7fbff85913f866a5384d8912710936e2b96da74541c82c1b458472ddb39"}, + {file = "ruff-0.2.2-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:ecd46e3106850a5c26aee114e562c329f9a1fbe9e4821b008c4404f64ff9ce73"}, + {file = "ruff-0.2.2-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:5e22676a5b875bd72acd3d11d5fa9075d3a5f53b877fe7b4793e4673499318ba"}, + {file = "ruff-0.2.2-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:1695700d1e25a99d28f7a1636d85bafcc5030bba9d0578c0781ba1790dbcf51c"}, + {file = "ruff-0.2.2-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:b0c232af3d0bd8f521806223723456ffebf8e323bd1e4e82b0befb20ba18388e"}, + {file = "ruff-0.2.2-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:f63d96494eeec2fc70d909393bcd76c69f35334cdbd9e20d089fb3f0640216ca"}, + {file = "ruff-0.2.2-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:6a61ea0ff048e06de273b2e45bd72629f470f5da8f71daf09fe481278b175001"}, + {file = "ruff-0.2.2-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:5e1439c8f407e4f356470e54cdecdca1bd5439a0673792dbe34a2b0a551a2fe3"}, + {file = "ruff-0.2.2-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:940de32dc8853eba0f67f7198b3e79bc6ba95c2edbfdfac2144c8235114d6726"}, + {file = "ruff-0.2.2-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:0c126da55c38dd917621552ab430213bdb3273bb10ddb67bc4b761989210eb6e"}, + {file = "ruff-0.2.2-py3-none-musllinux_1_2_i686.whl", hash = "sha256:3b65494f7e4bed2e74110dac1f0d17dc8e1f42faaa784e7c58a98e335ec83d7e"}, + {file = "ruff-0.2.2-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:1ec49be4fe6ddac0503833f3ed8930528e26d1e60ad35c2446da372d16651ce9"}, + {file = "ruff-0.2.2-py3-none-win32.whl", hash = "sha256:d920499b576f6c68295bc04e7b17b6544d9d05f196bb3aac4358792ef6f34325"}, + {file = "ruff-0.2.2-py3-none-win_amd64.whl", hash = "sha256:cc9a91ae137d687f43a44c900e5d95e9617cb37d4c989e462980ba27039d239d"}, + {file = "ruff-0.2.2-py3-none-win_arm64.whl", hash = "sha256:c9d15fc41e6054bfc7200478720570078f0b41c9ae4f010bcc16bd6f4d1aacdd"}, + {file = "ruff-0.2.2.tar.gz", hash = "sha256:e62ed7f36b3068a30ba39193a14274cd706bc486fad521276458022f7bccb31d"}, ] [[package]] @@ -6640,10 +6680,10 @@ files = [ ] [package.dependencies] -botocore = ">=1.37.4,<2.0a.0" +botocore = ">=1.37.4,<2.0a0" [package.extras] -crt = ["botocore[crt] (>=1.37.4,<2.0a.0)"] +crt = ["botocore[crt] (>=1.37.4,<2.0a0)"] [[package]] name = "scikit-learn" @@ -6876,9 +6916,9 @@ tornado = ">=6.4.2,<7" urllib3 = ">=1.26,<3" [package.extras] -all = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)", "cohere (>=5.9.4,<6.00)", "dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\"", "google-cloud-aiplatform (>=1.45.0,<2)", "ipykernel (>=6.25.0,<7)", "llama-cpp-python (>=0.2.28,<0.2.86) ; python_version < \"3.13\"", "mistralai (>=0.0.12,<0.1.0)", "mypy (>=1.7.1,<2)", "ollama (>=0.1.7)", "pillow (>=10.2.0,<11.0.0) ; python_version < \"3.13\"", "pinecone[asyncio] (>=7.0.0,<8.0.0)", "psycopg[binary] (>=3.1.0,<4)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "qdrant-client (>=1.11.1,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "sentence-transformers (>=5.0.0) ; python_version < \"3.13\"", "tokenizers (>=0.19) ; python_version < \"3.13\"", "torch (>=2.6.0) ; python_version < \"3.13\"", "torchvision (>=0.17.0) ; python_version < \"3.13\"", "transformers (>=4.36.2) ; python_version < \"3.13\"", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"] +all = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)", "cohere (>=5.9.4,<6.0)", "dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\"", "google-cloud-aiplatform (>=1.45.0,<2)", "ipykernel (>=6.25.0,<7)", "llama-cpp-python (>=0.2.28,<0.2.86) ; python_version < \"3.13\"", "mistralai (>=0.0.12,<0.1.0)", "mypy (>=1.7.1,<2)", "ollama (>=0.1.7)", "pillow (>=10.2.0,<11.0.0) ; python_version < \"3.13\"", "pinecone[asyncio] (>=7.0.0,<8.0.0)", "psycopg[binary] (>=3.1.0,<4)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "qdrant-client (>=1.11.1,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "sentence-transformers (>=5.0.0) ; python_version < \"3.13\"", "tokenizers (>=0.19) ; python_version < \"3.13\"", "torch (>=2.6.0) ; python_version < \"3.13\"", "torchvision (>=0.17.0) ; python_version < \"3.13\"", "transformers (>=4.36.2) ; python_version < \"3.13\"", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"] bedrock = ["boto3 (>=1.34.98,<2)", "botocore (>=1.34.110,<2)"] -cohere = ["cohere (>=5.9.4,<6.00)"] +cohere = ["cohere (>=5.9.4,<6.0)"] dev = ["dagger-io (>=0.1.1) ; python_version >= \"3.11\"", "ipykernel (>=6.25.0,<7)", "mypy (>=1.7.1,<2)", "pytest (>=8.2,<9.0)", "pytest-asyncio (>=0.24.0,<0.25)", "pytest-cov (>=4.1.0,<5)", "pytest-mock (>=3.12.0,<4)", "pytest-timeout", "pytest-xdist (>=3.5.0,<4)", "python-dotenv (>=1.0.0,<2)", "requests-mock (>=1.12.1,<2)", "ruff (>=0.11.2,<0.12)", "types-pyyaml (>=6.0.12.12,<7)", "types-requests (>=2.31.0,<3)"] docs = ["pydoc-markdown (>=4.8.2) ; python_version < \"3.12\""] fastembed = ["fastembed (>=0.3.0,<0.4) ; python_version < \"3.13\""] @@ -7722,6 +7762,7 @@ files = [ {file = "tomlkit-0.13.3-py3-none-any.whl", hash = "sha256:c89c649d79ee40629a9fda55f8ace8c6a1b42deb912b2a8fd8d942ddadb606b0"}, {file = "tomlkit-0.13.3.tar.gz", hash = "sha256:430cf247ee57df2b94ee3fbe588e71d362a941ebb545dec29b53961d61add2a1"}, ] +markers = {main = "extra == \"extra-proxy\""} [[package]] name = "tornado" @@ -8490,4 +8531,8 @@ utils = ["numpydoc"] [metadata] lock-version = "2.1" python-versions = ">=3.9,<4.0" -content-hash = "95fd27dc139d0e52e70093220c50582f16c78e5977ec77f4297f50a30df964c6" +<<<<<<< litellm_oss_staging_02_04_2026 +content-hash = "797603dcfef0a79781c7d3cba5dfe18f6aea4aa792220f47487ebc7bd04ae2e3" +======= +content-hash = "e5447e14dd37e324ac07a8fc6286d27e9a0d355ed93ebb24fc11e3f5df12fd3e" +>>>>>>> main diff --git a/provider_endpoints_support.json b/provider_endpoints_support.json index 93e9e7beaaa..fd17b5309e8 100644 --- a/provider_endpoints_support.json +++ b/provider_endpoints_support.json @@ -2183,7 +2183,8 @@ "batches": false, "rerank": false, "a2a": true, - "interactions": true + "interactions": true, + "realtime": true } }, "xinference": { diff --git a/pyproject.toml b/pyproject.toml index 9832ca483dc..40ca400401c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [tool.poetry] name = "litellm" -version = "1.81.7" +version = "1.81.8" description = "Library to easily interface with LLM API providers" authors = ["BerriAI"] license = "MIT" @@ -61,7 +61,7 @@ boto3 = { version = "1.40.76", optional = true } redisvl = {version = "^0.4.1", optional = true, markers = "python_version >= '3.9' and python_version < '3.14'"} mcp = {version = ">=1.25.0,<2.0.0", optional = true, python = ">=3.10"} a2a-sdk = {version = "^0.3.22", optional = true, python = ">=3.10"} -litellm-proxy-extras = {version = "0.4.29", optional = true} +litellm-proxy-extras = {version = "0.4.30", optional = true} rich = {version = "13.7.1", optional = true} litellm-enterprise = {version = "0.1.27", optional = true} diskcache = {version = "^5.6.1", optional = true} @@ -139,6 +139,7 @@ litellm = 'litellm:run_server' litellm-proxy = 'litellm.proxy.client.cli:cli' [tool.poetry.group.dev.dependencies] +diff-cover = "^9.0" flake8 = "^6.1.0" black = "^23.12.0" mypy = "^1.0" @@ -149,7 +150,7 @@ pytest-retry = "^1.6.3" requests-mock = "^1.12.1" responses = "^0.25.7" respx = "^0.22.0" -ruff = "^0.1.0" +ruff = "^0.2.1" types-requests = "*" types-setuptools = "*" types-redis = "*" @@ -174,7 +175,7 @@ requires = ["poetry-core", "wheel"] build-backend = "poetry.core.masonry.api" [tool.commitizen] -version = "1.81.7" +version = "1.81.8" version_files = [ "pyproject.toml:^version" ] diff --git a/requirements.txt b/requirements.txt index 69768b6c1f7..aca27bf284b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -50,7 +50,7 @@ sentry_sdk==2.21.0 # for sentry error handling detect-secrets==1.5.0 # Enterprise - secret detection / masking in LLM requests cryptography==44.0.1 tzdata==2025.1 # IANA time zone database -litellm-proxy-extras==0.4.29 # for proxy extras - e.g. prisma migrations +litellm-proxy-extras==0.4.30 # for proxy extras - e.g. prisma migrations llm-sandbox==0.3.31 # for skill execution in sandbox ### LITELLM PACKAGE DEPENDENCIES python-dotenv==1.0.1 # for env diff --git a/schema.prisma b/schema.prisma index b118400b620..03b910db85e 100644 --- a/schema.prisma +++ b/schema.prisma @@ -129,6 +129,7 @@ model LiteLLM_TeamTable { team_member_permissions String[] @default([]) policies String[] @default([]) model_id Int? @unique // id for LiteLLM_ModelTable -> stores team-level model aliases + allow_team_guardrail_config Boolean @default(false) // if true, team admin can configure guardrails for this team litellm_organization_table LiteLLM_OrganizationTable? @relation(fields: [organization_id], references: [organization_id]) litellm_model_table LiteLLM_ModelTable? @relation(fields: [model_id], references: [id]) object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id]) @@ -160,6 +161,7 @@ model LiteLLM_DeletedTeamTable { team_member_permissions String[] @default([]) policies String[] @default([]) model_id Int? // id for LiteLLM_ModelTable -> stores team-level model aliases + allow_team_guardrail_config Boolean @default(false) // Original timestamps from team creation/updates created_at DateTime? @map("created_at") diff --git a/tests/code_coverage_tests/recursive_detector.py b/tests/code_coverage_tests/recursive_detector.py index ed7595bb023..71e7798b09e 100644 --- a/tests/code_coverage_tests/recursive_detector.py +++ b/tests/code_coverage_tests/recursive_detector.py @@ -40,6 +40,7 @@ IGNORE_FUNCTIONS = [ "filter_exceptions_from_params", # max depth set (default 20) to prevent infinite recursion. "__getattr__", # lazy loading pattern in litellm/__init__.py with proper caching to prevent infinite recursion. "_validate_inheritance_chain", # max depth set (default 100) to prevent infinite recursion in policy inheritance validation. + "_basic_json_schema_validate", # max depth set. "extract_text_from_a2a_message", # max depth set (default 10) to prevent infinite recursion in A2A message parsing. ] diff --git a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py index 0a57d046c72..c39454728a8 100644 --- a/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py +++ b/tests/enterprise/litellm_enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py @@ -1,4 +1,3 @@ -import io import os import sys @@ -10,13 +9,10 @@ from datetime import datetime, timedelta, timezone from unittest.mock import MagicMock, call, patch import pytest -from prometheus_client import REGISTRY, CollectorRegistry +from prometheus_client import REGISTRY import litellm -from litellm import completion from litellm._logging import verbose_logger -from litellm._uuid import uuid -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.types.utils import ( StandardLoggingHiddenParams, StandardLoggingMetadata, @@ -37,7 +33,6 @@ from litellm.proxy._types import UserAPIKeyAuth verbose_logger.setLevel(logging.DEBUG) litellm.set_verbose = True -import time @pytest.fixture @@ -293,7 +288,6 @@ async def test_increment_remaining_budget_metrics(prometheus_logger): ) as mock_get_team, patch( "litellm.proxy.auth.auth_checks.get_key_object" ) as mock_get_key: - mock_get_team.return_value = MagicMock(budget_reset_at=future_reset_time_team) mock_get_key.return_value = MagicMock(budget_reset_at=future_reset_time_key) @@ -648,25 +642,16 @@ async def test_async_log_failure_event(prometheus_logger): ) # litellm_llm_api_failed_requests_metric incremented - """ - Expected metrics - end_user_id, - user_api_key, - user_api_key_alias, - model, - user_api_team, - user_api_team_alias, - user_id, - """ + # Labels: end_user, api_key_hash, api_key_alias, model, team, team_alias, user, model_id prometheus_logger.litellm_llm_api_failed_requests_metric.labels.assert_called_once_with( - None, + None, # end_user_id "test_hash", "test_alias", "gpt-3.5-turbo", "test_team", "test_team_alias", "test_user", - "model-123", + "model-123", # model_id from standard_logging_payload ) prometheus_logger.litellm_llm_api_failed_requests_metric.labels().inc.assert_called_once() @@ -678,38 +663,54 @@ async def test_async_log_failure_event(prometheus_logger): api_provider="openai", ) - # deployment failure responses incremented - prometheus_logger.litellm_deployment_failure_responses.labels.assert_called_once_with( - litellm_model_name="gpt-3.5-turbo", - model_id="model-123", - api_base="https://api.openai.com", - api_provider="openai", - exception_status="None", - exception_class="Exception", - requested_model="openai-gpt", # passed in standard logging payload - hashed_api_key="test_hash", - api_key_alias="test_alias", - team="test_team", - team_alias="test_team_alias", - client_ip="127.0.0.1", # from standard logging payload - user_agent=None, + # deployment failure responses incremented - verify key labels are populated + prometheus_logger.litellm_deployment_failure_responses.labels.assert_called_once() + actual_failure_labels = ( + prometheus_logger.litellm_deployment_failure_responses.labels.call_args.kwargs ) + expected_failure_labels = { + "litellm_model_name": "gpt-3.5-turbo", + "model_id": "model-123", + "api_base": "https://api.openai.com", + "api_provider": "openai", + "exception_class": "Exception", + "requested_model": "openai-gpt", + "hashed_api_key": "test_hash", + "api_key_alias": "test_alias", + "team": "test_team", + "team_alias": "test_team_alias", + } + for key, expected_val in expected_failure_labels.items(): + assert key in actual_failure_labels, f"Missing label {key}" + assert ( + actual_failure_labels[key] == expected_val + ), f"Label {key}: expected {expected_val!r}, got {actual_failure_labels[key]!r}" + assert actual_failure_labels.get("exception_status") in ("None", None) + assert actual_failure_labels.get("client_ip") == "127.0.0.1" prometheus_logger.litellm_deployment_failure_responses.labels().inc.assert_called_once() - # deployment total requests incremented - prometheus_logger.litellm_deployment_total_requests.labels.assert_called_once_with( - litellm_model_name="gpt-3.5-turbo", - model_id="model-123", - api_base="https://api.openai.com", - api_provider="openai", - requested_model="openai-gpt", # passed in standard logging payload - hashed_api_key="test_hash", - api_key_alias="test_alias", - team="test_team", - team_alias="test_team_alias", - client_ip="127.0.0.1", # from standard logging payload - user_agent=None, + # deployment total requests incremented - verify key labels are populated + prometheus_logger.litellm_deployment_total_requests.labels.assert_called_once() + actual_total_labels = ( + prometheus_logger.litellm_deployment_total_requests.labels.call_args.kwargs ) + expected_total_labels = { + "litellm_model_name": "gpt-3.5-turbo", + "model_id": "model-123", + "api_base": "https://api.openai.com", + "api_provider": "openai", + "requested_model": "openai-gpt", + "hashed_api_key": "test_hash", + "api_key_alias": "test_alias", + "team": "test_team", + "team_alias": "test_team_alias", + } + for key, expected_val in expected_total_labels.items(): + assert key in actual_total_labels, f"Missing label {key}" + assert ( + actual_total_labels[key] == expected_val + ), f"Label {key}: expected {expected_val!r}, got {actual_total_labels[key]!r}" + assert actual_total_labels.get("client_ip") == "127.0.0.1" prometheus_logger.litellm_deployment_total_requests.labels().inc.assert_called_once() @@ -1095,7 +1096,7 @@ def test_increment_deployment_cooled_down(prometheus_logger): import inspect method_sig = inspect.signature(prometheus_logger.increment_deployment_cooled_down) - expected_label_count = len([p for p in method_sig.parameters.keys() if p != 'self']) + expected_label_count = len([p for p in method_sig.parameters.keys() if p != "self"]) mock_chain = MagicMock() @@ -1103,11 +1104,15 @@ def test_increment_deployment_cooled_down(prometheus_logger): """Validate label count matches metric definition""" total = len(label_values) + len(label_kwargs) if total != expected_label_count: - raise ValueError(f"Incorrect label count: expected {expected_label_count}, got {total}") + raise ValueError( + f"Incorrect label count: expected {expected_label_count}, got {total}" + ) return mock_chain prometheus_logger.litellm_deployment_cooled_down = MagicMock() - prometheus_logger.litellm_deployment_cooled_down.labels = MagicMock(side_effect=validating_labels) + prometheus_logger.litellm_deployment_cooled_down.labels = MagicMock( + side_effect=validating_labels + ) prometheus_logger.increment_deployment_cooled_down( litellm_model_name="gpt-3.5-turbo", @@ -1179,8 +1184,12 @@ def test_get_custom_labels_from_top_level_metadata(monkeypatch): metadata = { "requester_ip_address": "10.48.203.20", # Top-level field "user_api_key_alias": "TestAlias", # Top-level field - "requester_metadata": {"nested_field": "nested_value"}, # Nested dict (excluded) - "user_api_key_auth_metadata": {"another_nested": "value"}, # Nested dict (excluded) + "requester_metadata": { + "nested_field": "nested_value" + }, # Nested dict (excluded) + "user_api_key_auth_metadata": { + "another_nested": "value" + }, # Nested dict (excluded) } result = get_custom_labels_from_metadata(metadata) assert result == { @@ -1217,7 +1226,9 @@ def test_get_custom_labels_from_top_level_and_nested_metadata(monkeypatch): } -async def test_async_log_success_event_with_top_level_metadata(prometheus_logger, monkeypatch): +async def test_async_log_success_event_with_top_level_metadata( + prometheus_logger, monkeypatch +): """ Test that async_log_success_event correctly extracts custom labels from top-level metadata fields like requester_ip_address, not just from nested dictionaries. @@ -1231,7 +1242,9 @@ async def test_async_log_success_event_with_top_level_metadata(prometheus_logger standard_logging_object = create_standard_logging_payload() standard_logging_object["metadata"]["requester_ip_address"] = "10.48.203.20" standard_logging_object["metadata"]["requester_metadata"] = {} # Empty nested dict - standard_logging_object["metadata"]["user_api_key_auth_metadata"] = {} # Empty nested dict + standard_logging_object["metadata"][ + "user_api_key_auth_metadata" + ] = {} # Empty nested dict kwargs = { "model": "gpt-3.5-turbo", @@ -1273,7 +1286,9 @@ async def test_async_log_success_event_with_top_level_metadata(prometheus_logger prometheus_logger.litellm_remaining_user_budget_metric = create_mock_metric() prometheus_logger.litellm_user_max_budget_metric = create_mock_metric() prometheus_logger.litellm_user_budget_remaining_hours_metric = create_mock_metric() - prometheus_logger.litellm_remaining_api_key_requests_for_model = create_mock_metric() + prometheus_logger.litellm_remaining_api_key_requests_for_model = ( + create_mock_metric() + ) prometheus_logger.litellm_remaining_api_key_tokens_for_model = create_mock_metric() prometheus_logger.litellm_llm_api_time_to_first_token_metric = create_mock_metric() prometheus_logger.litellm_llm_api_latency_metric = create_mock_metric() @@ -1302,7 +1317,7 @@ async def test_async_log_success_event_with_top_level_metadata(prometheus_logger # This confirms that the custom label extraction logic ran without errors assert prometheus_logger.litellm_requests_metric.labels.called assert prometheus_logger.litellm_spend_metric.labels.called - + # Verify that the labels() method was called with some arguments (either positional or keyword) # This ensures the custom label extraction happened and didn't cause a "Incorrect label names" error call_args = prometheus_logger.litellm_requests_metric.labels.call_args @@ -1494,7 +1509,6 @@ async def test_initialize_remaining_budget_metrics(prometheus_logger): with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( "litellm.proxy.management_endpoints.team_endpoints.get_paginated_teams" ) as mock_get_teams: - # Create mock team data with proper datetime objects for budget_reset_at future_reset = datetime.now() + timedelta(hours=24) # Reset 24 hours from now mock_teams = [ @@ -1592,21 +1606,22 @@ async def test_initialize_remaining_budget_metrics_exception_handling( ) as mock_get_teams, patch( "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper" ) as mock_list_keys: - # Make get_paginated_teams raise an exception mock_get_teams.side_effect = Exception("Database error") mock_list_keys.side_effect = Exception("Key listing error") - + # Mock prisma_client structure to raise an exception for user budget metrics # The code accesses prisma_client.db.litellm_usertable.find_many and count mock_usertable = MagicMock() - mock_usertable.find_many = MagicMock(side_effect=Exception("User database error")) + mock_usertable.find_many = MagicMock( + side_effect=Exception("User database error") + ) mock_usertable.count = MagicMock(side_effect=Exception("User count error")) - + # Mock litellm_teamtable to raise an exception for team count metrics mock_teamtable = MagicMock() mock_teamtable.count = MagicMock(side_effect=Exception("Team count error")) - + mock_db = MagicMock() mock_db.litellm_usertable = mock_usertable mock_db.litellm_teamtable = mock_teamtable @@ -1661,7 +1676,6 @@ async def test_initialize_api_key_budget_metrics(prometheus_logger): with patch("litellm.proxy.proxy_server.prisma_client") as mock_prisma, patch( "litellm.proxy.management_endpoints.key_management_endpoints._list_key_helper" ) as mock_list_keys: - # Create mock key data with proper datetime objects for budget_reset_at future_reset = datetime.now() + timedelta(hours=24) # Reset 24 hours from now key1 = UserAPIKeyAuth( @@ -1916,7 +1930,6 @@ def test_prometheus_label_factory_with_custom_tags(monkeypatch): Test that prometheus_label_factory correctly handles custom tags """ from litellm.integrations.prometheus import ( - get_custom_labels_from_tags, prometheus_label_factory, ) from litellm.types.integrations.prometheus import UserAPIKeyLabelValues @@ -1954,7 +1967,6 @@ def test_prometheus_label_factory_with_no_custom_tags(monkeypatch): Test that prometheus_label_factory works when no custom tags are configured """ from litellm.integrations.prometheus import ( - get_custom_labels_from_tags, prometheus_label_factory, ) from litellm.types.integrations.prometheus import UserAPIKeyLabelValues @@ -2179,9 +2191,7 @@ async def test_prometheus_token_metrics_with_prometheus_config(): All three metrics should be properly incremented when making a successful completion request. """ - from prometheus_client import CollectorRegistry, Counter - import litellm from litellm.types.integrations.prometheus import PrometheusMetricsConfig # Clear registry before test diff --git a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py index 3fd19cfa18f..946c5ad1729 100644 --- a/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py +++ b/tests/enterprise/litellm_enterprise/proxy/hooks/test_managed_files.py @@ -587,6 +587,499 @@ def test_update_responses_input_with_multiple_file_ids(): assert updated_input[0]["content"][1]["text"] == "Compare these files" +def test_update_responses_input_with_model_file_id_mapping(): + """ + Test that update_responses_input_with_model_file_ids correctly uses + model_file_id_mapping to map managed file IDs to provider-specific file IDs. + """ + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + update_responses_input_with_model_file_ids, + ) + + # Managed file ID (unified) + managed_file_id = "litellm_proxy_file_123" + + # Model file ID mapping + model_file_id_mapping = { + managed_file_id: { + "model_id_1": "openai_file_abc", + "model_id_2": "azure_file_xyz", + } + } + + input_data = [ + { + "role": "user", + "content": [ + { + "type": "input_file", + "file_id": managed_file_id, + }, + { + "type": "input_text", + "text": "Analyze this file", + }, + ], + } + ] + + # Update input with model_id_1 mapping + updated_input = update_responses_input_with_model_file_ids( + input=input_data, + model_id="model_id_1", + model_file_id_mapping=model_file_id_mapping, + ) + + # Verify the file_id was mapped to the correct provider-specific file ID + assert updated_input[0]["content"][0]["file_id"] == "openai_file_abc" + + # Test with different model_id + updated_input_2 = update_responses_input_with_model_file_ids( + input=input_data, + model_id="model_id_2", + model_file_id_mapping=model_file_id_mapping, + ) + + assert updated_input_2[0]["content"][0]["file_id"] == "azure_file_xyz" + + +def test_update_responses_tools_with_model_file_id_mapping(): + """ + Test that update_responses_tools_with_model_file_ids correctly maps + file IDs in code_interpreter tools with container.file_ids. + + This is a regression test for the issue where managed file IDs in + tools.container.file_ids were not being replaced with provider-specific + file IDs, causing "string too long" errors from OpenAI. + """ + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + update_responses_tools_with_model_file_ids, + ) + + # Managed file IDs + managed_file_id_1 = "litellm_proxy_file_123" + managed_file_id_2 = "litellm_proxy_file_456" + + # Model file ID mapping + model_file_id_mapping = { + managed_file_id_1: { + "model_id_1": "openai_file_abc", + }, + managed_file_id_2: { + "model_id_1": "openai_file_def", + }, + } + + tools = [ + { + "type": "code_interpreter", + "container": { + "type": "auto", + "file_ids": [managed_file_id_1, managed_file_id_2], + }, + } + ] + + # Update tools with model mapping + updated_tools = update_responses_tools_with_model_file_ids( + tools=tools, + model_id="model_id_1", + model_file_id_mapping=model_file_id_mapping, + ) + + # Verify the file IDs were mapped to provider-specific file IDs + assert updated_tools[0]["type"] == "code_interpreter" + assert updated_tools[0]["container"]["file_ids"] == ["openai_file_abc", "openai_file_def"] + + +def test_update_responses_tools_without_mapping(): + """ + Test that update_responses_tools_with_model_file_ids keeps file IDs + unchanged when no mapping is provided. + """ + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + update_responses_tools_with_model_file_ids, + ) + + regular_file_id = "file-abc123" + + tools = [ + { + "type": "code_interpreter", + "container": { + "type": "auto", + "file_ids": [regular_file_id], + }, + } + ] + + # Update tools without mapping + updated_tools = update_responses_tools_with_model_file_ids( + tools=tools, + model_id=None, + model_file_id_mapping=None, + ) + + # Verify the file ID was kept unchanged + assert updated_tools[0]["container"]["file_ids"] == [regular_file_id] + + +def test_update_responses_tools_with_mixed_file_ids(): + """ + Test that update_responses_tools_with_model_file_ids correctly handles + a mix of managed and regular file IDs. + """ + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + update_responses_tools_with_model_file_ids, + ) + + managed_file_id = "litellm_proxy_file_123" + regular_file_id = "file-abc123" + + model_file_id_mapping = { + managed_file_id: { + "model_id_1": "openai_file_abc", + }, + } + + tools = [ + { + "type": "code_interpreter", + "container": { + "type": "auto", + "file_ids": [managed_file_id, regular_file_id], + }, + } + ] + + # Update tools + updated_tools = update_responses_tools_with_model_file_ids( + tools=tools, + model_id="model_id_1", + model_file_id_mapping=model_file_id_mapping, + ) + + # Verify managed file ID was mapped and regular file ID was kept + assert updated_tools[0]["container"]["file_ids"] == ["openai_file_abc", regular_file_id] + + +def test_get_file_ids_from_responses_tools(): + """ + Test that get_file_ids_from_responses_tools correctly extracts + file IDs from the tools parameter. + """ + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=MagicMock() + ) + + tools = [ + { + "type": "code_interpreter", + "container": { + "type": "auto", + "file_ids": ["file-123", "file-456"], + }, + } + ] + + file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools) + + assert file_ids == ["file-123", "file-456"] + + +def test_get_file_ids_from_responses_tools_multiple_tools(): + """ + Test that get_file_ids_from_responses_tools handles multiple tools. + """ + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=MagicMock() + ) + + tools = [ + { + "type": "code_interpreter", + "container": { + "type": "auto", + "file_ids": ["file-123"], + }, + }, + { + "type": "file_search", + }, + { + "type": "code_interpreter", + "container": { + "type": "auto", + "file_ids": ["file-456", "file-789"], + }, + }, + ] + + file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools) + + # Should extract file IDs only from code_interpreter tools + assert file_ids == ["file-123", "file-456", "file-789"] + + +def test_get_file_ids_from_responses_tools_empty(): + """ + Test that get_file_ids_from_responses_tools handles empty or None tools. + """ + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + DualCache(), prisma_client=MagicMock() + ) + + # Test with None + file_ids = proxy_managed_files.get_file_ids_from_responses_tools(None) + assert file_ids == [] + + # Test with empty list + file_ids = proxy_managed_files.get_file_ids_from_responses_tools([]) + assert file_ids == [] + + # Test with tools without file_ids + tools = [{"type": "file_search"}] + file_ids = proxy_managed_files.get_file_ids_from_responses_tools(tools) + assert file_ids == [] + + +@pytest.mark.asyncio +async def test_check_file_ids_access_with_unified_file_ids(): + """ + Test that check_file_ids_access validates user access to managed file IDs. + """ + from litellm.proxy._types import UserAPIKeyAuth + + # Create a unified file ID + unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" + regular_file_id = "file-abc123" + + # Mock the access check to return True + prisma_client = AsyncMock() + internal_usage_cache = MagicMock() + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + internal_usage_cache=internal_usage_cache, + prisma_client=prisma_client, + ) + + # Mock can_user_call_unified_file_id to return True + proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True) + + user_api_key_dict = UserAPIKeyAuth( + user_id="test_user_123", + parent_otel_span=MagicMock(), + ) + + # Should not raise an exception for accessible files + await proxy_managed_files.check_file_ids_access( + [unified_file_id, regular_file_id], + user_api_key_dict, + ) + + # Verify can_user_call_unified_file_id was called for the unified file ID + proxy_managed_files.can_user_call_unified_file_id.assert_called_once_with( + unified_file_id, user_api_key_dict + ) + + +@pytest.mark.asyncio +async def test_check_file_ids_access_denied(): + """ + Test that check_file_ids_access raises HTTPException when user doesn't have access. + """ + from litellm.proxy._types import UserAPIKeyAuth + + unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" + + prisma_client = AsyncMock() + internal_usage_cache = MagicMock() + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + internal_usage_cache=internal_usage_cache, + prisma_client=prisma_client, + ) + + # Mock can_user_call_unified_file_id to return False (access denied) + proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=False) + + user_api_key_dict = UserAPIKeyAuth( + user_id="test_user_123", + parent_otel_span=MagicMock(), + ) + + # Should raise HTTPException with 403 status code + with pytest.raises(HTTPException) as exc_info: + await proxy_managed_files.check_file_ids_access( + [unified_file_id], + user_api_key_dict, + ) + + assert exc_info.value.status_code == 403 + assert "does not have access to the file" in exc_info.value.detail + + +@pytest.mark.asyncio +async def test_check_file_ids_access_with_regular_files_only(): + """ + Test that check_file_ids_access doesn't check access for regular (non-unified) file IDs. + """ + from litellm.proxy._types import UserAPIKeyAuth + + regular_file_id_1 = "file-abc123" + regular_file_id_2 = "file-xyz789" + + prisma_client = AsyncMock() + internal_usage_cache = MagicMock() + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + internal_usage_cache=internal_usage_cache, + prisma_client=prisma_client, + ) + + # Mock can_user_call_unified_file_id (should not be called for regular files) + proxy_managed_files.can_user_call_unified_file_id = AsyncMock() + + user_api_key_dict = UserAPIKeyAuth( + user_id="test_user_123", + parent_otel_span=MagicMock(), + ) + + # Should not raise exception and should not call can_user_call_unified_file_id + await proxy_managed_files.check_file_ids_access( + [regular_file_id_1, regular_file_id_2], + user_api_key_dict, + ) + + # Verify can_user_call_unified_file_id was NOT called + proxy_managed_files.can_user_call_unified_file_id.assert_not_called() + + +@pytest.mark.asyncio +async def test_completion_with_file_access_check(): + """ + Test that completion call type checks file access before processing. + """ + from litellm.proxy._types import UserAPIKeyAuth + + unified_file_id = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" + + prisma_client = AsyncMock() + prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) + + internal_usage_cache = MagicMock() + internal_usage_cache.async_get_cache = AsyncMock(return_value=None) + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + internal_usage_cache=internal_usage_cache, + prisma_client=prisma_client, + ) + + # Mock the get_model_file_id_mapping to return empty dict + proxy_managed_files.get_model_file_id_mapping = AsyncMock(return_value={}) + + # Mock access check to allow access + proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True) + + user_api_key_dict = UserAPIKeyAuth( + user_id="test_user_123", + parent_otel_span=MagicMock(), + ) + + data = { + "messages": [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What's in this file?"}, + { + "type": "file", + "file": {"file_id": unified_file_id}, + }, + ], + } + ], + "model": "gpt-4", + } + + # Should not raise exception + result = await proxy_managed_files.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="acompletion", + ) + + # Verify access check was called + proxy_managed_files.can_user_call_unified_file_id.assert_called_once() + + +@pytest.mark.asyncio +async def test_responses_with_file_access_check(): + """ + Test that responses API checks file access for files in both input and tools. + """ + from litellm.proxy._types import UserAPIKeyAuth + + unified_file_id_1 = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9wZGY7dW5pZmllZF9pZCw2YzBiNTg5MC04OTE0LTQ4ZTAtYjhmNC0wYWU1ZWQzYzE0YTU7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1FQ0JQVzdNTDlnN1hIZHdHZ1VQWmFNO2xsbV9vdXRwdXRfZmlsZV9tb2RlbF9pZCxlMjY0NTNmOWU3NmU3OTkzNjgwZDAwNjhkOThjMWY0Y2MyMDViYmFkMDk2N2EzM2M2NjQ4OTM1NjhjYTc0M2My" + unified_file_id_2 = "bGl0ZWxsbV9wcm94eTphcHBsaWNhdGlvbi9qc29uO3VuaWZpZWRfaWQsNzc3Nzc3Nzc7dGFyZ2V0X21vZGVsX25hbWVzLGdwdC00bztsbG1fb3V0cHV0X2ZpbGVfaWQsZmlsZS1YWVo7bGxtX291dHB1dF9maWxlX21vZGVsX2lkLG1vZGVsXzEyMw" + + prisma_client = AsyncMock() + prisma_client.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) + + internal_usage_cache = MagicMock() + internal_usage_cache.async_get_cache = AsyncMock(return_value=None) + + proxy_managed_files = _PROXY_LiteLLMManagedFiles( + internal_usage_cache=internal_usage_cache, + prisma_client=prisma_client, + ) + + # Mock the get_model_file_id_mapping to return empty dict + proxy_managed_files.get_model_file_id_mapping = AsyncMock(return_value={}) + + # Mock access check to allow access + proxy_managed_files.can_user_call_unified_file_id = AsyncMock(return_value=True) + + user_api_key_dict = UserAPIKeyAuth( + user_id="test_user_123", + parent_otel_span=MagicMock(), + ) + + data = { + "input": [ + { + "role": "user", + "content": [ + {"type": "input_text", "text": "Analyze this"}, + {"type": "input_file", "file_id": unified_file_id_1}, + ], + } + ], + "tools": [ + { + "type": "code_interpreter", + "container": { + "type": "auto", + "file_ids": [unified_file_id_2], + }, + } + ], + "model": "gpt-4", + } + + # Should not raise exception + result = await proxy_managed_files.async_pre_call_hook( + user_api_key_dict=user_api_key_dict, + cache=DualCache(), + data=data, + call_type="aresponses", + ) + + # Verify access check was called for both file IDs + assert proxy_managed_files.can_user_call_unified_file_id.call_count == 2 + + @pytest.mark.asyncio async def test_store_unified_file_id_with_none_file_object(): """ diff --git a/tests/litellm/test_proxy_auth.py b/tests/litellm/test_proxy_auth.py new file mode 100644 index 00000000000..1d73e143e10 --- /dev/null +++ b/tests/litellm/test_proxy_auth.py @@ -0,0 +1,204 @@ +""" +Unit tests for litellm.proxy_auth module. + +Tests the OAuth2/JWT token management for LiteLLM Proxy authentication. +""" + +import time +from unittest.mock import Mock, patch + +import pytest + +from litellm.proxy_auth import ( + AccessToken, + AzureADCredential, + GenericOAuth2Credential, + ProxyAuthHandler, +) + + +class TestAccessToken: + """Tests for AccessToken dataclass.""" + + def test_access_token_creation(self): + """Test AccessToken can be created with required fields.""" + token = AccessToken(token="test-token", expires_on=1234567890) + assert token.token == "test-token" + assert token.expires_on == 1234567890 + + def test_access_token_equality(self): + """Test AccessToken equality comparison.""" + token1 = AccessToken(token="test", expires_on=123) + token2 = AccessToken(token="test", expires_on=123) + assert token1 == token2 + + +class MockCredential: + """Mock credential for testing.""" + + def __init__(self, expires_in_seconds: int = 3600): + self.call_count = 0 + self.expires_in = expires_in_seconds + + def get_token(self, scope: str) -> AccessToken: + self.call_count += 1 + return AccessToken( + token=f"mock-token-{self.call_count}", + expires_on=int(time.time()) + self.expires_in, + ) + + +class TestProxyAuthHandler: + """Tests for ProxyAuthHandler.""" + + def test_get_auth_headers_returns_bearer_token(self): + """Test that get_auth_headers returns correct Authorization header.""" + cred = MockCredential() + handler = ProxyAuthHandler(credential=cred, scope="test-scope") + + headers = handler.get_auth_headers() + + assert "Authorization" in headers + assert headers["Authorization"].startswith("Bearer ") + assert "mock-token-1" in headers["Authorization"] + + def test_token_caching(self): + """Test that tokens are cached and not re-requested.""" + cred = MockCredential(expires_in_seconds=3600) # Long expiry + handler = ProxyAuthHandler(credential=cred, scope="test-scope") + + # Multiple calls should only request token once + handler.get_auth_headers() + handler.get_auth_headers() + handler.get_auth_headers() + + assert cred.call_count == 1 + + def test_token_refresh_when_about_to_expire(self): + """Test that tokens are refreshed when about to expire (within 60s buffer).""" + cred = MockCredential(expires_in_seconds=30) # Expires in 30s (< 60s buffer) + handler = ProxyAuthHandler(credential=cred, scope="test-scope") + + # First call gets token + handler.get_auth_headers() + # Second call should refresh because token expires within 60s buffer + handler.get_auth_headers() + + assert cred.call_count == 2 + + def test_get_token_method(self): + """Test the get_token method returns AccessToken.""" + cred = MockCredential() + handler = ProxyAuthHandler(credential=cred, scope="test-scope") + + token = handler.get_token() + + assert isinstance(token, AccessToken) + assert token.token == "mock-token-1" + + +class TestAzureADCredential: + """Tests for AzureADCredential.""" + + def test_lazy_initialization(self): + """Test that azure-identity is not imported until get_token is called.""" + # This should not raise ImportError even if azure-identity is not installed + cred = AzureADCredential(credential=None) + # _initialized should be False until get_token is called + assert cred._initialized is False + + def test_wraps_azure_credential(self): + """Test that AzureADCredential wraps an azure-identity credential.""" + # Mock Azure credential + mock_azure_cred = Mock() + mock_azure_cred.get_token.return_value = Mock( + token="azure-token", expires_on=9999999999 + ) + + cred = AzureADCredential(credential=mock_azure_cred) + token = cred.get_token("https://graph.microsoft.com/.default") + + assert token.token == "azure-token" + assert token.expires_on == 9999999999 + mock_azure_cred.get_token.assert_called_once_with( + "https://graph.microsoft.com/.default" + ) + + +class TestGenericOAuth2Credential: + """Tests for GenericOAuth2Credential.""" + + def test_token_request(self): + """Test that GenericOAuth2Credential makes correct OAuth2 request.""" + with patch("httpx.post") as mock_post: + mock_response = Mock() + mock_response.json.return_value = { + "access_token": "oauth2-token", + "expires_in": 3600, + } + mock_response.raise_for_status = Mock() + mock_post.return_value = mock_response + + cred = GenericOAuth2Credential( + client_id="test-client", + client_secret="test-secret", + token_url="https://example.com/oauth2/token", + ) + token = cred.get_token("test-scope") + + assert token.token == "oauth2-token" + mock_post.assert_called_once() + call_kwargs = mock_post.call_args + assert call_kwargs[1]["data"]["grant_type"] == "client_credentials" + assert call_kwargs[1]["data"]["client_id"] == "test-client" + assert call_kwargs[1]["data"]["client_secret"] == "test-secret" + assert call_kwargs[1]["data"]["scope"] == "test-scope" + + def test_token_caching(self): + """Test that GenericOAuth2Credential caches tokens.""" + with patch("httpx.post") as mock_post: + mock_response = Mock() + mock_response.json.return_value = { + "access_token": "oauth2-token", + "expires_in": 3600, + } + mock_response.raise_for_status = Mock() + mock_post.return_value = mock_response + + cred = GenericOAuth2Credential( + client_id="test-client", + client_secret="test-secret", + token_url="https://example.com/oauth2/token", + ) + + # Multiple calls should only make one HTTP request + cred.get_token("test-scope") + cred.get_token("test-scope") + cred.get_token("test-scope") + + assert mock_post.call_count == 1 + + +class TestLiteLLMIntegration: + """Tests for integration with litellm module.""" + + def test_proxy_auth_variable_exists(self): + """Test that litellm.proxy_auth variable exists.""" + import litellm + + # Should be None by default + assert hasattr(litellm, "proxy_auth") + + def test_proxy_auth_can_be_set(self): + """Test that litellm.proxy_auth can be set to a ProxyAuthHandler.""" + import litellm + + original_value = litellm.proxy_auth + try: + cred = MockCredential() + handler = ProxyAuthHandler(credential=cred, scope="test") + litellm.proxy_auth = handler + + assert litellm.proxy_auth is handler + finally: + litellm.proxy_auth = original_value diff --git a/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py b/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py index 0969c77299c..5bd0c9993a8 100644 --- a/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py +++ b/tests/litellm_core_utils/test_bedrock_converse_dedup_factory.py @@ -330,3 +330,118 @@ async def test_bedrock_converse_sync_async_parity_with_duplicates(): ) assert sync_result == async_result + + +# --------------------------------------------------------------------------- +# Empty content filtering tests +# --------------------------------------------------------------------------- + + +def test_bedrock_converse_filters_empty_assistant_content(): + """Verify that empty assistant content blocks are filtered out to avoid + Bedrock API errors about blank text fields.""" + messages = [ + {"role": "user", "content": "Say hello"}, + {"role": "assistant", "content": "Hello"}, + {"role": "assistant", "content": " there"}, + {"role": "assistant", "content": "!"}, + {"role": "assistant", "content": ""}, # Empty content + {"role": "assistant", "content": ""}, # Empty content + {"role": "user", "content": "How are you?"}, + ] + + result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER) + + # Should have 3 messages: user, assistant (with merged non-empty content), user + assert len(result) == 3 + assert result[0]["role"] == "user" + assert result[1]["role"] == "assistant" + assert result[2]["role"] == "user" + + # Assistant message should only contain non-empty text blocks + assistant_content = result[1]["content"] + text_blocks = [block for block in assistant_content if "text" in block] + assert len(text_blocks) == 3 # "Hello", " there", "!" + assert text_blocks[0]["text"] == "Hello" + assert text_blocks[1]["text"] == " there" + assert text_blocks[2]["text"] == "!" + + +@pytest.mark.asyncio +async def test_bedrock_converse_filters_empty_assistant_content_async(): + """Verify that the async path also filters empty assistant content blocks.""" + messages = [ + {"role": "user", "content": "Say hello"}, + {"role": "assistant", "content": "Hello"}, + {"role": "assistant", "content": " there"}, + {"role": "assistant", "content": "!"}, + {"role": "assistant", "content": ""}, # Empty content + {"role": "assistant", "content": ""}, # Empty content + {"role": "user", "content": "How are you?"}, + ] + + result = await BedrockConverseMessagesProcessor._bedrock_converse_messages_pt_async( + messages, MODEL, PROVIDER + ) + + # Should have 3 messages: user, assistant (with merged non-empty content), user + assert len(result) == 3 + assert result[0]["role"] == "user" + assert result[1]["role"] == "assistant" + assert result[2]["role"] == "user" + + # Assistant message should only contain non-empty text blocks + assistant_content = result[1]["content"] + text_blocks = [block for block in assistant_content if "text" in block] + assert len(text_blocks) == 3 # "Hello", " there", "!" + assert text_blocks[0]["text"] == "Hello" + assert text_blocks[1]["text"] == " there" + assert text_blocks[2]["text"] == "!" + + +def test_bedrock_converse_filters_whitespace_only_content(): + """Verify that whitespace-only content is also filtered out.""" + messages = [ + {"role": "user", "content": "Test"}, + {"role": "assistant", "content": "Response"}, + {"role": "assistant", "content": " "}, # Whitespace only + {"role": "assistant", "content": "\n\t"}, # Whitespace only + {"role": "assistant", "content": ""}, # Empty + ] + + result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER) + + # Should have 2 messages: user and assistant + assert len(result) == 2 + assistant_content = result[1]["content"] + text_blocks = [block for block in assistant_content if "text" in block] + # Only "Response" should be present + assert len(text_blocks) == 1 + assert text_blocks[0]["text"] == "Response" + + +def test_bedrock_converse_filters_empty_list_content(): + """Verify that empty text elements in list content are filtered out.""" + messages = [ + {"role": "user", "content": "Test"}, + { + "role": "assistant", + "content": [ + {"type": "text", "text": "Hello"}, + {"type": "text", "text": ""}, # Empty + {"type": "text", "text": "World"}, + {"type": "text", "text": " "}, # Whitespace only + ], + }, + ] + + result = _bedrock_converse_messages_pt(messages, MODEL, PROVIDER) + + # Should have 2 messages: user and assistant + assert len(result) == 2 + assistant_content = result[1]["content"] + text_blocks = [block for block in assistant_content if "text" in block] + # Only "Hello" and "World" should be present + assert len(text_blocks) == 2 + assert text_blocks[0]["text"] == "Hello" + assert text_blocks[1]["text"] == "World" diff --git a/tests/llm_translation/realtime/__init__.py b/tests/llm_translation/realtime/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/llm_translation/realtime/base_realtime_tests.py b/tests/llm_translation/realtime/base_realtime_tests.py new file mode 100644 index 00000000000..2a1ac78ffe6 --- /dev/null +++ b/tests/llm_translation/realtime/base_realtime_tests.py @@ -0,0 +1,426 @@ +""" +Base test class for LiteLLM Realtime API E2E tests. + +Provides common test infrastructure for testing realtime WebSocket connections +across different providers (OpenAI, xAI, etc.) +""" +import asyncio +import json +import os +import sys +from abc import ABC, abstractmethod +from typing import Optional + +import pytest +import websockets + +sys.path.insert(0, os.path.abspath("../../..")) + +import litellm + + +class RealTimeWebSocketClient: + """ + Mock WebSocket client for testing realtime connections. + Captures messages sent from the backend and provides a simple interface + for testing connection success. + """ + + def __init__(self): + self.messages_sent = [] + self.messages_received = [] + self.received_initial_event = False + self.connection_successful = False + self.close_code = None + self.close_reason = None + # Required by realtime_streaming.py - import exceptions module + from websockets import exceptions as websockets_exceptions + self.exceptions = websockets_exceptions + + async def accept(self): + """Accept the WebSocket connection""" + pass + + async def send_text(self, message): + """Receive message from backend and store it""" + self.messages_sent.append(message) + try: + if isinstance(message, bytes): + message_str = message.decode('utf-8') + else: + message_str = message + + msg_data = json.loads(message_str) + msg_type = msg_data.get('type', 'unknown') + + # Pretty print API response + print(f"\n{'='*80}") + print(f"API RESPONSE #{len(self.messages_received) + 1} - Event: {msg_type}") + print(f"{'='*80}") + print(json.dumps(msg_data, indent=2, sort_keys=False)) + print(f"{'='*80}\n") + + self.messages_received.append(msg_data) + + # Check for initial connection event + if not self.received_initial_event and self._is_initial_event(msg_type): + self.received_initial_event = True + self.connection_successful = True + + except (json.JSONDecodeError, UnicodeDecodeError) as e: + # Non-JSON messages are acceptable + print(f"\n[Non-JSON message: {e}]") + print(f"Raw content: {str(message)[:200]}\n") + pass + + def _is_initial_event(self, msg_type: str) -> bool: + """Check if message type is an initial connection event""" + # OpenAI sends "session.created", xAI sends "conversation.created" + return msg_type in ["session.created", "conversation.created"] + + async def receive_text(self): + """ + Wait briefly for messages, then close connection. + This allows the backend forwarding task to send messages. + """ + print(f"\nWaiting for connection to establish...") + max_wait = 5.0 + check_interval = 0.1 + waited = 0.0 + + while waited < max_wait: + if self.connection_successful: + print(f"Connection successful after {waited:.1f}s\n") + break + await asyncio.sleep(check_interval) + waited += check_interval + + if not self.connection_successful: + print(f"Warning: No initial event received after {max_wait}s\n") + + # If we have a pending message to send, send it now + if hasattr(self, '_pending_client_message') and self._pending_client_message: + print(f"Sending client message to backend...\n") + # This simulates receiving a message from the client that needs to be forwarded to backend + # We return it as if it came from the client + msg = self._pending_client_message + self._pending_client_message = None + return msg + + # Close connection to end the test + print(f"\n{'='*80}") + print(f"TEST COMPLETE - Closing connection") + print(f"Total messages received from API: {len(self.messages_received)}") + print(f"{'='*80}\n") + raise websockets.exceptions.ConnectionClosed(None, None) + + def queue_client_message(self, message: str): + """Queue a message to be sent from 'client' to backend""" + self._pending_client_message = message + + async def close(self, code=1000, reason=""): + """Close the WebSocket""" + self.close_code = code + self.close_reason = reason + + @property + def headers(self): + return {} + + +class BaseRealtimeTest(ABC): + """ + Abstract base test class for realtime API tests. + + Child classes must implement: + - get_model(): Return the model name to test + - get_api_key_env_var(): Return the environment variable name for the API key + - get_initial_event_type(): Return the expected initial event type (e.g., "session.created") + """ + + @abstractmethod + def get_model(self) -> str: + """Return the model name to test (e.g., 'gpt-4o-realtime-preview-2024-10-01')""" + pass + + @abstractmethod + def get_api_key_env_var(self) -> str: + """Return the environment variable name for the API key (e.g., 'OPENAI_API_KEY')""" + pass + + @abstractmethod + def get_initial_event_type(self) -> str: + """Return the expected initial event type (e.g., 'session.created' or 'conversation.created')""" + pass + + def get_skip_reason(self) -> str: + """Return the skip reason when API key is missing""" + return f"No {self.get_api_key_env_var()} provided" + + def should_skip(self) -> bool: + """Check if tests should be skipped due to missing API key""" + return os.environ.get(self.get_api_key_env_var()) is None + + @pytest.mark.asyncio + async def test_realtime_connection(self): + """ + Test basic realtime WebSocket connection. + Verifies that: + 1. Connection is established successfully + 2. Initial event is received + 3. Messages are properly forwarded + """ + litellm._turn_on_debug() + if self.should_skip(): + pytest.skip(self.get_skip_reason()) + + websocket_client = RealTimeWebSocketClient() + caught_exception = None + + print(f"\n{'='*80}") + print(f"STARTING REALTIME CONNECTION TEST") + print(f"Model: {self.get_model()}") + print(f"API Key Env Var: {self.get_api_key_env_var()}") + print(f"{'='*80}\n") + + try: + await litellm._arealtime( + model=self.get_model(), + websocket=websocket_client, + api_key=os.environ.get(self.get_api_key_env_var()), + timeout=60 + ) + except websockets.exceptions.ConnectionClosed: + pass + except Exception as e: + print(f"\nException: {type(e).__name__}: {e}\n") + caught_exception = e + + # Build debug info + error_details = [] + error_details.append(f"messages_sent: {len(websocket_client.messages_sent)}") + error_details.append(f"messages_received: {len(websocket_client.messages_received)}") + error_details.append(f"close_code: {websocket_client.close_code}") + error_details.append(f"close_reason: {websocket_client.close_reason}") + if caught_exception: + error_details.append(f"exception: {type(caught_exception).__name__}: {caught_exception}") + + # Skip on transient connection failures + if not websocket_client.connection_successful and websocket_client.close_code is not None: + pytest.skip(f"Transient connection failure: {'; '.join(error_details)}") + + # Assertions + assert websocket_client.connection_successful, f"Failed to connect. Debug: {'; '.join(error_details)}" + assert websocket_client.received_initial_event, f"Did not receive initial event" + assert len(websocket_client.messages_received) > 0, "No messages received" + + # Verify initial event + initial_event = websocket_client.messages_received[0] + assert initial_event["type"] == self.get_initial_event_type(), \ + f"Expected {self.get_initial_event_type()}, got {initial_event.get('type')}" + + @pytest.mark.asyncio + async def test_realtime_with_query_params(self): + """ + Test realtime connection with explicit query parameters. + Verifies that query params are properly passed to the backend. + """ + litellm._turn_on_debug() + if self.should_skip(): + pytest.skip(self.get_skip_reason()) + + from litellm.types.realtime import RealtimeQueryParams + + websocket_client = RealTimeWebSocketClient() + caught_exception = None + + # Strip provider prefix from model name for query params + model_name = self.get_model() + if "/" in model_name: + model_name = model_name.split("/", 1)[1] + + query_params: RealtimeQueryParams = {"model": model_name} + + try: + await litellm._arealtime( + model=self.get_model(), + websocket=websocket_client, + api_key=os.environ.get(self.get_api_key_env_var()), + query_params=query_params, + timeout=60 + ) + except websockets.exceptions.ConnectionClosed: + pass + except Exception as e: + caught_exception = e + + # Build debug info + error_details = [] + error_details.append(f"messages_sent: {len(websocket_client.messages_sent)}") + error_details.append(f"messages_received: {len(websocket_client.messages_received)}") + if caught_exception: + error_details.append(f"exception: {type(caught_exception).__name__}: {caught_exception}") + + # Skip on transient failures + if not websocket_client.connection_successful and websocket_client.close_code is not None: + pytest.skip(f"Transient connection failure: {'; '.join(error_details)}") + + # Assertions + assert websocket_client.connection_successful, f"Failed to connect. Debug: {'; '.join(error_details)}" + assert len(websocket_client.messages_received) > 0, "No messages received" + + @pytest.mark.asyncio + async def test_send_user_message(self): + """ + Test sending an actual user message and receiving responses. + This creates a more realistic conversation flow. + """ + if self.should_skip(): + pytest.skip(self.get_skip_reason()) + + litellm._turn_on_debug() + + # Create a custom websocket client that sends a message + class InteractiveWebSocketClient(RealTimeWebSocketClient): + def __init__(self): + super().__init__() + self.sent_user_message = False + self.response_messages = [] + self.wait_for_responses = True + + async def receive_text(self): + """Enhanced receive that sends a user message after connection""" + print(f"\n{'='*80}") + print(f"CLIENT-SIDE RECEIVE HANDLER") + print(f"{'='*80}\n") + + # Wait for initial connection + max_wait = 5.0 + check_interval = 0.1 + waited = 0.0 + + while waited < max_wait: + if self.connection_successful: + print(f"Connection established after {waited:.1f}s\n") + break + await asyncio.sleep(check_interval) + waited += check_interval + + # Step 1: Send a user message after connection is established + if self.connection_successful and not self.sent_user_message: + self.sent_user_message = True + user_msg_data = { + "type": "conversation.item.create", + "item": { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": "Say hi back to me!"}] + } + } + user_msg = json.dumps(user_msg_data) + + print(f"\n{'='*80}") + print(f"STEP 1: SENDING USER MESSAGE TO BACKEND") + print(f"{'='*80}") + print(json.dumps(user_msg_data, indent=2)) + print(f"{'='*80}\n") + + return user_msg + + # Step 2: Trigger the response after user message is acknowledged + if not hasattr(self, 'triggered_response'): + self.triggered_response = True + # Wait a bit for the user message to be processed + await asyncio.sleep(0.5) + + response_create_data = { + "type": "response.create" + } + response_create = json.dumps(response_create_data) + + print(f"\n{'='*80}") + print(f"STEP 2: TRIGGERING LLM RESPONSE") + print(f"{'='*80}") + print(json.dumps(response_create_data, indent=2)) + print(f"{'='*80}\n") + + return response_create + + # Step 3: Wait for LLM responses + if self.wait_for_responses: + print(f"\nSTEP 3: Waiting 5 seconds for LLM to respond...\n") + await asyncio.sleep(5.0) + self.wait_for_responses = False + + # Collect response info + for msg in self.messages_received: + msg_type = msg.get('type', 'unknown') + if msg_type not in ['conversation.created', 'ping']: + self.response_messages.append(msg) + + print(f"\nReceived {len(self.response_messages)} response messages (excluding init/ping)\n") + + print(f"\n{'='*80}") + print(f"CLOSING CONNECTION") + print(f"Total messages received: {len(self.messages_received)}") + print(f"{'='*80}\n") + raise websockets.exceptions.ConnectionClosed(None, None) + + websocket_client = InteractiveWebSocketClient() + caught_exception = None + + print(f"\n{'='*80}") + print(f"STARTING INTERACTIVE MESSAGE TEST") + print(f"Model: {self.get_model()}") + print(f"Message: 'Say hi back to me!'") + print(f"{'='*80}\n") + + try: + await litellm._arealtime( + model=self.get_model(), + websocket=websocket_client, + api_key=os.environ.get(self.get_api_key_env_var()), + timeout=60 + ) + except websockets.exceptions.ConnectionClosed: + pass + except Exception as e: + print(f"\nException: {type(e).__name__}: {e}\n") + caught_exception = e + + # Print results + print(f"\n{'='*80}") + print(f"TEST RESULTS SUMMARY") + print(f"{'='*80}") + print(f"Connection successful: {websocket_client.connection_successful}") + print(f"User message sent: {websocket_client.sent_user_message}") + print(f"Total messages received: {len(websocket_client.messages_received)}") + print(f"Response messages (excluding init/ping): {len(websocket_client.response_messages)}") + + if websocket_client.response_messages: + print(f"\nResponse Event Types:") + for i, msg in enumerate(websocket_client.response_messages, 1): + print(f" {i}. {msg.get('type', 'unknown')}") + + print(f"{'='*80}\n") + + # Skip if no responses (might be timing issue) + if not websocket_client.response_messages: + pytest.skip("No response messages received (might be timing/network issue)") + + assert websocket_client.connection_successful, "Failed to establish connection" + assert websocket_client.sent_user_message, "Failed to send user message" + + def test_query_params_construction(self): + """Test that query params are constructed correctly""" + from litellm.types.realtime import RealtimeQueryParams + + # Strip provider prefix from model name + model_name = self.get_model() + if "/" in model_name: + model_name = model_name.split("/", 1)[1] + + query_params: RealtimeQueryParams = {"model": model_name} + + assert "model" in query_params + assert query_params["model"] == model_name diff --git a/tests/llm_translation/test_openai_realtime.py b/tests/llm_translation/realtime/test_openai_realtime.py similarity index 100% rename from tests/llm_translation/test_openai_realtime.py rename to tests/llm_translation/realtime/test_openai_realtime.py diff --git a/tests/llm_translation/realtime/test_openai_realtime_simple.py b/tests/llm_translation/realtime/test_openai_realtime_simple.py new file mode 100644 index 00000000000..8c281d08f93 --- /dev/null +++ b/tests/llm_translation/realtime/test_openai_realtime_simple.py @@ -0,0 +1,29 @@ +""" +OpenAI Realtime API E2E Tests (using base class) + +Tests OpenAI's Realtime API through LiteLLM's realtime interface. +Uses the base test class to ensure consistent behavior across providers. +""" +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +from tests.llm_translation.realtime.base_realtime_tests import BaseRealtimeTest + + +class TestOpenAIRealtime(BaseRealtimeTest): + """ + E2E tests for OpenAI Realtime API using base test class. + """ + + def get_model(self) -> str: + return "gpt-4o-realtime-preview" + + def get_api_key_env_var(self) -> str: + return "OPENAI_API_KEY" + + def get_initial_event_type(self) -> str: + return "session.created" diff --git a/tests/llm_translation/realtime/test_xai_realtime.py b/tests/llm_translation/realtime/test_xai_realtime.py new file mode 100644 index 00000000000..6b75d08c80f --- /dev/null +++ b/tests/llm_translation/realtime/test_xai_realtime.py @@ -0,0 +1,34 @@ +""" +xAI Realtime API E2E Tests + +Tests xAI's Grok Voice Agent API through LiteLLM's realtime interface. +Uses the base test class to ensure consistent behavior across providers. +""" +import os +import sys + +import pytest + +sys.path.insert(0, os.path.abspath("../../..")) + +from tests.llm_translation.realtime.base_realtime_tests import BaseRealtimeTest + + +class TestXAIRealtime(BaseRealtimeTest): + """ + E2E tests for xAI Realtime API. + + xAI's Grok Voice Agent API is OpenAI-compatible but uses: + - Different initial event: "conversation.created" instead of "session.created" + - Different endpoint: wss://api.x.ai/v1/realtime + - Model: grok-4-1-fast-non-reasoning + """ + + def get_model(self) -> str: + return "xai/grok-4-1-fast-non-reasoning" + + def get_api_key_env_var(self) -> str: + return "XAI_API_KEY" + + def get_initial_event_type(self) -> str: + return "conversation.created" diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index 9b0b69caeb3..d23033c1e46 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -2356,9 +2356,8 @@ def test_bedrock_no_default_message(): assistant_messages = [ msg for msg in formatted_messages if msg["role"] == "assistant" ] - assert len(assistant_messages) == 2 - assert assistant_messages[0]["content"][0]["text"] == "." - assert assistant_messages[1]["content"][0]["text"] == "Valid response" + assert len(assistant_messages) == 1 + assert assistant_messages[0]["content"][0]["text"] == "Valid response" @pytest.mark.parametrize("top_k_param", ["top_k", "topK"]) diff --git a/tests/llm_translation/test_gigachat.py b/tests/llm_translation/test_gigachat.py index b69a5428e42..631ae94d208 100644 --- a/tests/llm_translation/test_gigachat.py +++ b/tests/llm_translation/test_gigachat.py @@ -122,40 +122,6 @@ class TestGigaChatCollapseUserMessages: return GigaChatConfig() - def test_no_collapse_single_message(self, config): - """Single message should not be changed""" - messages = [{"role": "user", "content": "Hello"}] - result = config._collapse_user_messages(messages) - - assert len(result) == 1 - assert result[0]["content"] == "Hello" - - def test_collapse_consecutive_user_messages(self, config): - """Consecutive user messages should be collapsed""" - messages = [ - {"role": "user", "content": "First"}, - {"role": "user", "content": "Second"}, - {"role": "user", "content": "Third"}, - ] - result = config._collapse_user_messages(messages) - - assert len(result) == 1 - assert "First" in result[0]["content"] - assert "Second" in result[0]["content"] - assert "Third" in result[0]["content"] - - def test_no_collapse_with_assistant_between(self, config): - """Messages with assistant between should not be collapsed""" - messages = [ - {"role": "user", "content": "First"}, - {"role": "assistant", "content": "Response"}, - {"role": "user", "content": "Second"}, - ] - result = config._collapse_user_messages(messages) - - assert len(result) == 3 - - class TestGigaChatToolsTransformation: """Tests for tools -> functions conversion""" diff --git a/tests/logging_callback_tests/test_dynamic_otel_keys.py b/tests/logging_callback_tests/test_dynamic_otel_keys.py new file mode 100644 index 00000000000..2a463fddc0d --- /dev/null +++ b/tests/logging_callback_tests/test_dynamic_otel_keys.py @@ -0,0 +1,52 @@ +import sys +import os + +sys.path.insert(0, os.path.abspath("../..")) + +from litellm.litellm_core_utils.initialize_dynamic_callback_params import ( + initialize_standard_callback_dynamic_params, +) + + +def test_dynamic_key_extraction_from_metadata(): + """ + Test extraction of langfuse keys from metadata in kwargs. + This simulates a Proxy request where keys are passed in metadata. + """ + kwargs = { + "metadata": { + "langfuse_public_key": "pk-test", + "langfuse_secret_key": "sk-test", + "langfuse_host": "https://test.langfuse.com", + } + } + + params = initialize_standard_callback_dynamic_params(kwargs) + + assert params.get("langfuse_public_key") == "pk-test" + assert params.get("langfuse_secret_key") == "sk-test" + assert params.get("langfuse_host") == "https://test.langfuse.com" + + +def test_dynamic_key_extraction_from_litellm_params_metadata(): + """ + Test extraction of langfuse keys from litellm_params.metadata. + """ + kwargs = { + "litellm_params": { + "metadata": { + "langfuse_public_key": "pk-litellm", + "langfuse_secret_key": "sk-litellm", + } + } + } + + params = initialize_standard_callback_dynamic_params(kwargs) + + assert params.get("langfuse_public_key") == "pk-litellm" + assert params.get("langfuse_secret_key") == "sk-litellm" + + +if __name__ == "__main__": + test_dynamic_key_extraction_from_metadata() + test_dynamic_key_extraction_from_litellm_params_metadata() diff --git a/tests/mcp_tests/test_semantic_tool_filter_e2e.py b/tests/mcp_tests/test_semantic_tool_filter_e2e.py index 91c072ae8a3..0cb7f221a22 100644 --- a/tests/mcp_tests/test_semantic_tool_filter_e2e.py +++ b/tests/mcp_tests/test_semantic_tool_filter_e2e.py @@ -25,6 +25,10 @@ except ImportError: not SEMANTIC_ROUTER_AVAILABLE, reason="semantic-router not installed. Install with: pip install 'litellm[semantic-router]'" ) +@pytest.mark.skipif( + not os.environ.get("OPENAI_API_KEY"), + reason="OPENAI_API_KEY not set in environment" +) async def test_e2e_semantic_filter(): """E2E: Load router/filter and verify hook filters tools.""" from litellm import Router diff --git a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py index ecc0e3b370f..215ac0874f2 100644 --- a/tests/openai_endpoints_tests/test_openai_batches_endpoint.py +++ b/tests/openai_endpoints_tests/test_openai_batches_endpoint.py @@ -291,4 +291,318 @@ async def test_list_batches_with_target_model_names(): # Verify the response structure assert response["object"] == "list" - assert len(response["data"]) > 0 \ No newline at end of file + assert len(response["data"]) > 0 + + +@pytest.mark.asyncio +async def test_batch_status_sync_from_provider_to_database(): + """ + Test that when batch status changes at the provider, + it gets synced to the ManagedObjectTable database. + + This tests the new refactored utility functions: + - get_batch_from_database() + - update_batch_in_database() + """ + from unittest.mock import MagicMock, AsyncMock + from litellm.proxy.openai_files_endpoints.common_utils import ( + get_batch_from_database, + update_batch_in_database, + ) + from litellm.types.utils import LiteLLMBatch + import json + + # Setup: Create mock objects + batch_id = "batch_test123" + unified_batch_id = "litellm_proxy:test_unified_batch" + + # Mock database batch object with "validating" status + mock_db_batch = MagicMock() + mock_db_batch.unified_object_id = batch_id + mock_db_batch.status = "validating" + mock_db_batch.file_object = json.dumps({ + "id": batch_id, + "object": "batch", + "status": "validating", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-test123", + "completion_window": "24h", + "created_at": 1234567890, + }) + + # Mock prisma client + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( + return_value=mock_db_batch + ) + mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() + + # Mock managed_files_obj + mock_managed_files = MagicMock() + + # Mock logger + mock_logger = MagicMock() + mock_logger.debug = MagicMock() + mock_logger.info = MagicMock() + mock_logger.warning = MagicMock() + mock_logger.error = MagicMock() + + # Test 1: Retrieve batch from database (initial state) + db_batch_object, response_batch = await get_batch_from_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + managed_files_obj=mock_managed_files, + prisma_client=mock_prisma_client, + verbose_proxy_logger=mock_logger, + ) + + # Verify database was queried + mock_prisma_client.db.litellm_managedobjecttable.find_first.assert_called_once_with( + where={"unified_object_id": batch_id} + ) + + # Verify batch was retrieved correctly + assert db_batch_object is not None + assert response_batch is not None + assert response_batch.id == batch_id + assert response_batch.status == "validating" + + # Test 2: Simulate provider returning updated status + updated_batch_response = LiteLLMBatch( + id=batch_id, + object="batch", + status="completed", # Status changed from "validating" to "completed" + endpoint="/v1/chat/completions", + input_file_id="file-test123", + completion_window="24h", + created_at=1234567890, + output_file_id="file-output123", + ) + + # Test 3: Update database with new status from provider + await update_batch_in_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + response=updated_batch_response, + managed_files_obj=mock_managed_files, + prisma_client=mock_prisma_client, + verbose_proxy_logger=mock_logger, + db_batch_object=db_batch_object, + operation="retrieve", + ) + + # Verify database was updated + mock_prisma_client.db.litellm_managedobjecttable.update.assert_called_once() + update_call_args = mock_prisma_client.db.litellm_managedobjecttable.update.call_args + + # Verify the update call had correct parameters + assert update_call_args.kwargs["where"]["unified_object_id"] == batch_id + assert update_call_args.kwargs["data"]["status"] == "complete" # "completed" normalized to "complete" + assert "file_object" in update_call_args.kwargs["data"] + assert "updated_at" in update_call_args.kwargs["data"] + + # Verify logger was called with status change message + mock_logger.info.assert_called() + log_message = mock_logger.info.call_args[0][0] + assert "validating" in log_message + assert "completed" in log_message + + print("āœ… Test passed: Batch status synced from provider to database") + + +@pytest.mark.asyncio +async def test_batch_cancel_updates_database(): + """ + Test that canceling a batch updates the database status. + """ + from unittest.mock import MagicMock, AsyncMock + from litellm.proxy.openai_files_endpoints.common_utils import ( + update_batch_in_database, + ) + from litellm.types.utils import LiteLLMBatch + + # Setup + batch_id = "batch_cancel_test" + unified_batch_id = "litellm_proxy:cancel_test" + + # Mock cancelled batch response from provider + cancelled_batch_response = LiteLLMBatch( + id=batch_id, + object="batch", + status="cancelled", + endpoint="/v1/chat/completions", + input_file_id="file-test123", + completion_window="24h", + created_at=1234567890, + cancelled_at=1234567999, + ) + + # Mock prisma client + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() + + # Mock managed_files_obj + mock_managed_files = MagicMock() + + # Mock logger + mock_logger = MagicMock() + mock_logger.info = MagicMock() + mock_logger.error = MagicMock() + + # Call update_batch_in_database for cancel operation + await update_batch_in_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + response=cancelled_batch_response, + managed_files_obj=mock_managed_files, + prisma_client=mock_prisma_client, + verbose_proxy_logger=mock_logger, + operation="cancel", + ) + + # Verify database was updated + mock_prisma_client.db.litellm_managedobjecttable.update.assert_called_once() + update_call_args = mock_prisma_client.db.litellm_managedobjecttable.update.call_args + + # Verify the update call had correct parameters + assert update_call_args.kwargs["where"]["unified_object_id"] == batch_id + assert update_call_args.kwargs["data"]["status"] == "cancelled" + assert "file_object" in update_call_args.kwargs["data"] + + # Verify logger was called + mock_logger.info.assert_called() + log_message = mock_logger.info.call_args[0][0] + assert "cancel" in log_message.lower() + assert "cancelled" in log_message + + print("āœ… Test passed: Batch cancel updates database") + + +@pytest.mark.asyncio +async def test_batch_terminal_state_skip_provider_call(): + """ + Test that when a batch is in a terminal state (completed, failed, cancelled, expired), + it returns immediately from database without calling the provider. + """ + from unittest.mock import MagicMock, AsyncMock + from litellm.proxy.openai_files_endpoints.common_utils import ( + get_batch_from_database, + ) + from litellm.types.utils import LiteLLMBatch + import json + + # Setup: Create mock objects for a completed batch + batch_id = "batch_completed_test" + unified_batch_id = "litellm_proxy:completed_test" + + # Mock database batch object with "completed" status + mock_db_batch = MagicMock() + mock_db_batch.unified_object_id = batch_id + mock_db_batch.status = "complete" + mock_db_batch.file_object = json.dumps({ + "id": batch_id, + "object": "batch", + "status": "completed", + "endpoint": "/v1/chat/completions", + "input_file_id": "file-test123", + "output_file_id": "file-output123", + "completion_window": "24h", + "created_at": 1234567890, + "completed_at": 1234567999, + }) + + # Mock prisma client + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_managedobjecttable.find_first = AsyncMock( + return_value=mock_db_batch + ) + + # Mock managed_files_obj + mock_managed_files = MagicMock() + + # Mock logger + mock_logger = MagicMock() + mock_logger.debug = MagicMock() + + # Retrieve batch from database + db_batch_object, response_batch = await get_batch_from_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + managed_files_obj=mock_managed_files, + prisma_client=mock_prisma_client, + verbose_proxy_logger=mock_logger, + ) + + # Verify batch was retrieved + assert db_batch_object is not None + assert response_batch is not None + assert response_batch.status == "completed" + + # In the actual endpoint, when status is in terminal states, + # it should return immediately without calling the provider + # This test verifies the database retrieval works correctly + assert response_batch.status in ["completed", "failed", "cancelled", "expired"] + + print("āœ… Test passed: Terminal state batch retrieved from database") + + +@pytest.mark.asyncio +async def test_batch_no_status_change_skip_update(): + """ + Test that when batch status hasn't changed, database update is skipped. + """ + from unittest.mock import MagicMock, AsyncMock + from litellm.proxy.openai_files_endpoints.common_utils import ( + update_batch_in_database, + ) + from litellm.types.utils import LiteLLMBatch + + # Setup + batch_id = "batch_no_change_test" + unified_batch_id = "litellm_proxy:no_change_test" + + # Mock database batch object with "validating" status + mock_db_batch = MagicMock() + mock_db_batch.status = "validating" + + # Mock batch response from provider with same status + batch_response = LiteLLMBatch( + id=batch_id, + object="batch", + status="validating", # Same status as in database + endpoint="/v1/chat/completions", + input_file_id="file-test123", + completion_window="24h", + created_at=1234567890, + ) + + # Mock prisma client + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() + + # Mock managed_files_obj + mock_managed_files = MagicMock() + + # Mock logger + mock_logger = MagicMock() + mock_logger.info = MagicMock() + + # Call update_batch_in_database + await update_batch_in_database( + batch_id=batch_id, + unified_batch_id=unified_batch_id, + response=batch_response, + managed_files_obj=mock_managed_files, + prisma_client=mock_prisma_client, + verbose_proxy_logger=mock_logger, + db_batch_object=mock_db_batch, + operation="retrieve", + ) + + # Verify database update was NOT called (status hasn't changed) + mock_prisma_client.db.litellm_managedobjecttable.update.assert_not_called() + + # Verify logger info was NOT called (no status change to log) + mock_logger.info.assert_not_called() + + print("āœ… Test passed: Database update skipped when status unchanged") \ No newline at end of file diff --git a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py index b334966b441..2ebf9174e29 100644 --- a/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py +++ b/tests/pass_through_unit_tests/base_anthropic_unified_messages_test.py @@ -130,6 +130,67 @@ class BaseAnthropicMessagesTest: return collected_chunks + @pytest.mark.asyncio + async def test_response_format_consistency(self): + """ + Test that response content blocks are consistently dicts (not Pydantic objects). + + This ensures that code like response["content"][0]["type"] works + regardless of the target provider. + + Issue: https://github.com/BerriAI/litellm/issues/20342 + """ + litellm._turn_on_debug() + + request_params = self.model_config + + # Set up test parameters + messages = [{"role": "user", "content": "Say hi"}] + + # Prepare call arguments + call_args = { + "messages": messages, + "max_tokens": 100, + } + + # Add any additional config from subclass + call_args.update(request_params) + + # Call the handler + response = await litellm.anthropic.messages.acreate(**call_args) + + print(f"Response for {request_params['model']}: {json.dumps(response, indent=2, default=str)}") + + # Verify response structure + assert "content" in response, "Response should have 'content' field" + assert len(response["content"]) > 0, "Response content should not be empty" + + # Get the first content block + block = response["content"][0] + + # Check that the block is a dict, not a Pydantic object + assert isinstance(block, dict), ( + f"Content block should be a dict, but got {type(block)}. " + f"This means response format is inconsistent across providers." + ) + + # Verify we can access fields using dict syntax (not object attributes) + try: + block_type = block["type"] + print(f"āœ“ Successfully accessed block['type']: {block_type}") + except TypeError as e: + pytest.fail( + f"Cannot access content block using dict syntax: {e}. " + f"Block type: {type(block)}" + ) + + # Verify the block has expected structure + assert "type" in block, "Content block should have 'type' field" + if block["type"] == "text": + assert "text" in block, "Text content block should have 'text' field" + + print(f"āœ“ Response format consistency test passed for {request_params['model']}") + @pytest.mark.asyncio async def test_anthropic_messages_litellm_router_streaming_with_logging(self): """ diff --git a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py index 17e72f29152..0a581fb512d 100644 --- a/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py +++ b/tests/pass_through_unit_tests/test_anthropic_messages_passthrough.py @@ -876,4 +876,4 @@ def test_sync_openai_messages(): assert response is not None assert isinstance(response, dict) - assert response["content"][0].text is not None + assert response["content"][0]["text"] is not None diff --git a/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py b/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py index adbaf219079..57352eafaf1 100644 --- a/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py +++ b/tests/test_litellm/completion_extras/test_litellm_responses_transformation_transformation.py @@ -134,3 +134,25 @@ def test_transform_request_with_response_format(): assert result["text"]["format"]["type"] == "json_schema" assert result["text"]["format"]["name"] == "person_schema" assert "schema" in result["text"]["format"] + + +def test_transform_request_includes_extra_headers(): + """Test that transform_request forwards headers as extra_headers for upstream call.""" + handler = LiteLLMResponsesTransformationHandler() + messages = [{"role": "user", "content": "Hello"}] + optional_params = {} + litellm_params = {} + + class MockLoggingObj: + pass + + headers = {"cf-aig-authorization": "secret-token"} + result = handler.transform_request( + model="gpt-5-pro", + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + headers=headers, + litellm_logging_obj=MockLoggingObj(), + ) + assert result.get("extra_headers") == headers diff --git a/tests/test_litellm/integrations/test_langfuse_otel.py b/tests/test_litellm/integrations/test_langfuse_otel.py index f8c662979ad..ba4a096be24 100644 --- a/tests/test_litellm/integrations/test_langfuse_otel.py +++ b/tests/test_litellm/integrations/test_langfuse_otel.py @@ -1,92 +1,119 @@ import json import os -from datetime import datetime from unittest.mock import MagicMock, patch import pytest from litellm.integrations.langfuse.langfuse_otel import LangfuseOtelLogger -from litellm.types.integrations.langfuse_otel import LangfuseOtelConfig +from litellm.integrations.opentelemetry import OpenTelemetryConfig from litellm.types.llms.openai import ResponsesAPIResponse class TestLangfuseOtelIntegration: - def test_get_langfuse_otel_config_with_required_env_vars(self): """Test that config is created correctly with required environment variables.""" # Clean environment of any Langfuse-related variables - env_vars_to_clean = ['LANGFUSE_HOST', 'OTEL_EXPORTER_OTLP_ENDPOINT', 'OTEL_EXPORTER_OTLP_HEADERS'] - with patch.dict(os.environ, { - 'LANGFUSE_PUBLIC_KEY': 'test_public_key', - 'LANGFUSE_SECRET_KEY': 'test_secret_key' - }, clear=False): + env_vars_to_clean = [ + "LANGFUSE_HOST", + "OTEL_EXPORTER_OTLP_ENDPOINT", + "OTEL_EXPORTER_OTLP_HEADERS", + ] + with patch.dict( + os.environ, + { + "LANGFUSE_PUBLIC_KEY": "test_public_key", + "LANGFUSE_SECRET_KEY": "test_secret_key", + }, + clear=False, + ): # Remove any existing Langfuse variables for var in env_vars_to_clean: if var in os.environ: del os.environ[var] - + config = LangfuseOtelLogger.get_langfuse_otel_config() - - assert isinstance(config, LangfuseOtelConfig) - assert config.protocol == "otlp_http" - assert "Authorization=Basic" in config.otlp_auth_headers - # Check that environment variables are set correctly (US default) - assert os.environ.get("OTEL_EXPORTER_OTLP_ENDPOINT") == "https://us.cloud.langfuse.com/api/public/otel" - assert "Authorization=Basic" in os.environ.get("OTEL_EXPORTER_OTLP_HEADERS", "") - + + assert isinstance(config, OpenTelemetryConfig) + assert config.exporter == "otlp_http" + assert "Authorization=Basic" in config.headers + # Note: We no longer set os.environ explicitly to avoid leakage + # assert os.environ.get("OTEL_EXPORTER_OTLP_ENDPOINT") == "https://us.cloud.langfuse.com/api/public/otel" + # assert "Authorization=Basic" in os.environ.get("OTEL_EXPORTER_OTLP_HEADERS", "") + def test_get_langfuse_otel_config_missing_keys(self): """Test that ValueError is raised when required keys are missing.""" with patch.dict(os.environ, {}, clear=True): - with pytest.raises(ValueError, match="LANGFUSE_PUBLIC_KEY and LANGFUSE_SECRET_KEY must be set"): + with pytest.raises( + ValueError, + match="LANGFUSE_PUBLIC_KEY and LANGFUSE_SECRET_KEY must be set", + ): LangfuseOtelLogger.get_langfuse_otel_config() - + def test_get_langfuse_otel_config_with_eu_host(self): """Test config with EU host.""" - with patch.dict(os.environ, { - 'LANGFUSE_PUBLIC_KEY': 'test_public_key', - 'LANGFUSE_SECRET_KEY': 'test_secret_key', - 'LANGFUSE_HOST': 'https://cloud.langfuse.com' - }, clear=False): + with patch.dict( + os.environ, + { + "LANGFUSE_PUBLIC_KEY": "test_public_key", + "LANGFUSE_SECRET_KEY": "test_secret_key", + "LANGFUSE_HOST": "https://cloud.langfuse.com", + }, + clear=False, + ): config = LangfuseOtelLogger.get_langfuse_otel_config() - - assert os.environ.get("OTEL_EXPORTER_OTLP_ENDPOINT") == "https://cloud.langfuse.com/api/public/otel" - + # Endpoint assertion removed as side effect is gone + assert isinstance(config, OpenTelemetryConfig) + def test_get_langfuse_otel_config_with_custom_host(self): """Test config with custom host.""" - with patch.dict(os.environ, { - 'LANGFUSE_PUBLIC_KEY': 'test_public_key', - 'LANGFUSE_SECRET_KEY': 'test_secret_key', - 'LANGFUSE_HOST': 'https://my-langfuse.com' - }, clear=False): + with patch.dict( + os.environ, + { + "LANGFUSE_PUBLIC_KEY": "test_public_key", + "LANGFUSE_SECRET_KEY": "test_secret_key", + "LANGFUSE_HOST": "https://my-langfuse.com", + }, + clear=False, + ): config = LangfuseOtelLogger.get_langfuse_otel_config() - - assert os.environ.get("OTEL_EXPORTER_OTLP_ENDPOINT") == "https://my-langfuse.com/api/public/otel" - + # Endpoint assertion removed as side effect is gone + assert isinstance(config, OpenTelemetryConfig) + def test_get_langfuse_otel_config_with_host_no_protocol(self): """Test config with custom host without protocol.""" - with patch.dict(os.environ, { - 'LANGFUSE_PUBLIC_KEY': 'test_public_key', - 'LANGFUSE_SECRET_KEY': 'test_secret_key', - 'LANGFUSE_HOST': 'my-langfuse.com' - }, clear=False): + with patch.dict( + os.environ, + { + "LANGFUSE_PUBLIC_KEY": "test_public_key", + "LANGFUSE_SECRET_KEY": "test_secret_key", + "LANGFUSE_HOST": "my-langfuse.com", + }, + clear=False, + ): config = LangfuseOtelLogger.get_langfuse_otel_config() - - assert os.environ.get("OTEL_EXPORTER_OTLP_ENDPOINT") == "https://my-langfuse.com/api/public/otel" - + # Endpoint assertion removed as side effect is gone + assert isinstance(config, OpenTelemetryConfig) + def test_set_langfuse_otel_attributes(self): """Test that set_langfuse_otel_attributes calls the Arize utils function.""" from litellm.integrations.langfuse.langfuse_otel_attributes import ( LangfuseLLMObsOTELAttributes, ) - + mock_span = MagicMock() mock_kwargs = {"test": "kwargs"} mock_response = {"test": "response"} - - with patch('litellm.integrations.arize._utils.set_attributes') as mock_set_attributes: - LangfuseOtelLogger.set_langfuse_otel_attributes(mock_span, mock_kwargs, mock_response) - - mock_set_attributes.assert_called_once_with(mock_span, mock_kwargs, mock_response, LangfuseLLMObsOTELAttributes) + + with patch( + "litellm.integrations.arize._utils.set_attributes" + ) as mock_set_attributes: + LangfuseOtelLogger.set_langfuse_otel_attributes( + mock_span, mock_kwargs, mock_response + ) + + mock_set_attributes.assert_called_once_with( + mock_span, mock_kwargs, mock_response, LangfuseLLMObsOTELAttributes + ) def test_set_langfuse_environment_attribute(self): """Test that Langfuse environment is set correctly when environment variable is present.""" @@ -94,15 +121,17 @@ class TestLangfuseOtelIntegration: mock_kwargs = {"test": "kwargs"} test_env = "staging" - with patch.dict(os.environ, {'LANGFUSE_TRACING_ENVIRONMENT': test_env}): - with patch('litellm.integrations.arize._utils.safe_set_attribute') as mock_safe_set_attribute: - LangfuseOtelLogger._set_langfuse_specific_attributes(mock_span, mock_kwargs, {}) - + with patch.dict(os.environ, {"LANGFUSE_TRACING_ENVIRONMENT": test_env}): + with patch( + "litellm.integrations.arize._utils.safe_set_attribute" + ) as mock_safe_set_attribute: + LangfuseOtelLogger._set_langfuse_specific_attributes( + mock_span, mock_kwargs, {} + ) + # safe_set_attribute(span, key, value) → positional args mock_safe_set_attribute.assert_called_once_with( - mock_span, - "langfuse.environment", - test_env + mock_span, "langfuse.environment", test_env ) def test_extract_langfuse_metadata_basic(self): @@ -119,11 +148,13 @@ class TestLangfuseOtelIntegration: # Build a stub module + class on-the-fly stub_module = types.ModuleType("litellm.integrations.langfuse.langfuse") + class StubLFLogger: @staticmethod def add_metadata_from_header(litellm_params, metadata): # Echo back existing metadata plus a marker return {**metadata, "enriched": True} + stub_module.LangFuseLogger = StubLFLogger # type: ignore # Register stub in sys.modules so import inside method succeeds @@ -159,11 +190,16 @@ class TestLangfuseOtelIntegration: kwargs = {"litellm_params": {"metadata": metadata}} # Capture calls to safe_set_attribute - with patch('litellm.integrations.arize._utils.safe_set_attribute') as mock_safe_set_attribute: - LangfuseOtelLogger._set_langfuse_specific_attributes(MagicMock(), kwargs, None) + with patch( + "litellm.integrations.arize._utils.safe_set_attribute" + ) as mock_safe_set_attribute: + LangfuseOtelLogger._set_langfuse_specific_attributes( + MagicMock(), kwargs, None + ) # Build expected calls manually for clarity from litellm.types.integrations.langfuse_otel import LangfuseSpanAttributes + expected = { LangfuseSpanAttributes.GENERATION_NAME.value: "gen-name", LangfuseSpanAttributes.GENERATION_ID.value: "gen-id", @@ -176,12 +212,14 @@ class TestLangfuseOtelIntegration: # Lists / dicts should be JSON strings LangfuseSpanAttributes.TAGS.value: json.dumps(["tagA", "tagB"]), LangfuseSpanAttributes.TRACE_NAME.value: "trace-name", - LangfuseSpanAttributes.TRACE_ID.value: "trace-id", + LangfuseSpanAttributes.TRACE_ID.value: "traceid", # stripped dashes LangfuseSpanAttributes.TRACE_METADATA.value: json.dumps({"k": "v"}), LangfuseSpanAttributes.TRACE_VERSION.value: "t-ver", LangfuseSpanAttributes.TRACE_RELEASE.value: "rel-1", LangfuseSpanAttributes.EXISTING_TRACE_ID.value: "existing-id", - LangfuseSpanAttributes.UPDATE_TRACE_KEYS.value: json.dumps(["key1", "key2"]), + LangfuseSpanAttributes.UPDATE_TRACE_KEYS.value: json.dumps( + ["key1", "key2"] + ), LangfuseSpanAttributes.DEBUG_LANGFUSE.value: True, } @@ -191,7 +229,9 @@ class TestLangfuseOtelIntegration: for call in mock_safe_set_attribute.call_args_list } - assert actual == expected, "Mismatch between expected and actual OTEL attribute mapping." + assert ( + actual == expected + ), "Mismatch between expected and actual OTEL attribute mapping." def test_set_langfuse_specific_attributes_with_content(self): """Test that _set_langfuse_specific_attributes correctly sets observation.output with regular content response.""" @@ -200,15 +240,15 @@ class TestLangfuseOtelIntegration: # Create response with content response_obj = ModelResponse( - id='chatcmpl-test', - model='gpt-4o', + id="chatcmpl-test", + model="gpt-4o", choices=[ Choices( - finish_reason='stop', + finish_reason="stop", message={ "role": "assistant", - "content": "The weather in Tokyo is sunny." - } + "content": "The weather in Tokyo is sunny.", + }, ) ], ) @@ -217,20 +257,21 @@ class TestLangfuseOtelIntegration: "messages": [{"role": "user", "content": "What's the weather in Tokyo?"}], } - with patch('litellm.integrations.arize._utils.safe_set_attribute') as mock_safe_set_attribute: - LangfuseOtelLogger._set_langfuse_specific_attributes(MagicMock(), kwargs, response_obj) + with patch( + "litellm.integrations.arize._utils.safe_set_attribute" + ) as mock_safe_set_attribute: + LangfuseOtelLogger._set_langfuse_specific_attributes( + MagicMock(), kwargs, response_obj + ) expect_output = { LangfuseSpanAttributes.OBSERVATION_INPUT.value: [ - { - "role": "user", - "content": "What's the weather in Tokyo?" - } + {"role": "user", "content": "What's the weather in Tokyo?"} ], LangfuseSpanAttributes.OBSERVATION_OUTPUT.value: { "role": "assistant", - "content": "The weather in Tokyo is sunny." - } + "content": "The weather in Tokyo is sunny.", + }, } # Flatten the actual calls into {key: value} @@ -239,8 +280,9 @@ class TestLangfuseOtelIntegration: for call in mock_safe_set_attribute.call_args_list } - assert actual == expect_output, "Mismatch in observation input/output OTEL attributes." - + assert ( + actual == expect_output + ), "Mismatch in observation input/output OTEL attributes." def test_set_langfuse_specific_attributes_with_tool_calls(self): """Test that _set_langfuse_specific_attributes correctly sets observation.output with tool calls in Langfuse format.""" @@ -254,42 +296,44 @@ class TestLangfuseOtelIntegration: # Create response with tool calls response_obj = ModelResponse( - id='chatcmpl-test', - model='gpt-4o', + id="chatcmpl-test", + model="gpt-4o", choices=[ Choices( - finish_reason='tool_calls', + finish_reason="tool_calls", message={ "role": "assistant", "content": None, "tool_calls": [ ChatCompletionMessageToolCall( function=Function( - arguments='{"location":"Tokyo"}', - name='get_weather' + arguments='{"location":"Tokyo"}', name="get_weather" ), - id='call_123', - type='function' + id="call_123", + type="function", ) - ] - } + ], + }, ) ], ) - with patch('litellm.integrations.arize._utils.safe_set_attribute') as mock_safe_set_attribute: - LangfuseOtelLogger._set_langfuse_specific_attributes(MagicMock(), {}, - response_obj) + with patch( + "litellm.integrations.arize._utils.safe_set_attribute" + ) as mock_safe_set_attribute: + LangfuseOtelLogger._set_langfuse_specific_attributes( + MagicMock(), {}, response_obj + ) expected = { LangfuseSpanAttributes.OBSERVATION_OUTPUT.value: [ - { - "id": "chatcmpl-test", - "name": "get_weather", - "arguments": {"location": "Tokyo"}, - "call_id": "call_123", - "type": "function_call" - } + { + "id": "chatcmpl-test", + "name": "get_weather", + "arguments": {"location": "Tokyo"}, + "call_id": "call_123", + "type": "function_call", + } ] } @@ -298,8 +342,9 @@ class TestLangfuseOtelIntegration: call.args[1]: json.loads(call.args[2]) for call in mock_safe_set_attribute.call_args_list } - assert actual == expected, "Mismatch in observation output OTEL attribute for tool calls." - + assert ( + actual == expected + ), "Mismatch in observation output OTEL attribute for tool calls." def test_construct_dynamic_otel_headers_with_langfuse_keys(self): """Test that construct_dynamic_otel_headers creates proper auth headers when langfuse keys are provided.""" @@ -307,28 +352,27 @@ class TestLangfuseOtelIntegration: # Create dynamic params with langfuse keys dynamic_params = StandardCallbackDynamicParams( - langfuse_public_key="test_public_key", - langfuse_secret_key="test_secret_key" + langfuse_public_key="test_public_key", langfuse_secret_key="test_secret_key" ) - + logger = LangfuseOtelLogger() result = logger.construct_dynamic_otel_headers(dynamic_params) - + # Should return a dict with otlp_auth_headers assert result is not None assert "Authorization" in result - + # The auth header should contain the basic auth format auth_header = result["Authorization"] assert auth_header.startswith("Basic ") - + # Verify the header format by decoding import base64 # Extract the base64 part from "Authorization=Basic " base64_part = auth_header.replace("Basic ", "") decoded = base64.b64decode(base64_part).decode() - + assert decoded == "test_public_key:test_secret_key" def test_construct_dynamic_otel_headers_empty_params(self): @@ -337,24 +381,28 @@ class TestLangfuseOtelIntegration: # Create dynamic params without langfuse keys dynamic_params = StandardCallbackDynamicParams() - + logger = LangfuseOtelLogger() result = logger.construct_dynamic_otel_headers(dynamic_params) - + # Should return an empty dict assert result == {} - + def test_get_langfuse_otel_config_with_otel_host_priority(self): """LANGFUSE_OTEL_HOST should take priority over LANGFUSE_HOST.""" - with patch.dict(os.environ, { - 'LANGFUSE_PUBLIC_KEY': 'test_public_key', - 'LANGFUSE_SECRET_KEY': 'test_secret_key', - 'LANGFUSE_HOST': 'https://should-not-be-used.com', - 'LANGFUSE_OTEL_HOST': 'https://otel-host.com' - }, clear=False): - _ = LangfuseOtelLogger.get_langfuse_otel_config() - - assert os.environ.get("OTEL_EXPORTER_OTLP_ENDPOINT") == "https://otel-host.com/api/public/otel" + with patch.dict( + os.environ, + { + "LANGFUSE_PUBLIC_KEY": "test_public_key", + "LANGFUSE_SECRET_KEY": "test_secret_key", + "LANGFUSE_HOST": "https://should-not-be-used.com", + "LANGFUSE_OTEL_HOST": "https://otel-host.com", + }, + clear=False, + ): + config = LangfuseOtelLogger.get_langfuse_otel_config() + assert isinstance(config, OpenTelemetryConfig) + # Endpoint assertion removed as side effect is gone class TestLangfuseOtelResponsesAPI: @@ -369,46 +417,52 @@ class TestLangfuseOtelResponsesAPI: output=[ { "type": "message", - "content": [{"type": "text", "text": "Hello from responses API"}] + "content": [{"type": "text", "text": "Hello from responses API"}], } ], parallel_tool_calls=False, tool_choice="auto", tools=[], - top_p=1.0 + top_p=1.0, ) - + # Create kwargs with metadata that should be logged test_metadata = { - "user_id": "test123", - "session_id": "abc456", + "user_id": "test123", + "session_id": "abc456", "custom_field": "test_value", "generation_name": "responses_test_generation", - "trace_name": "responses_api_trace" + "trace_name": "responses_api_trace", } - + kwargs = { "call_type": "responses", "messages": [{"role": "user", "content": "Hello"}], "model": "gpt-4o", "optional_params": {}, - "litellm_params": {"metadata": test_metadata} + "litellm_params": {"metadata": test_metadata}, } - + mock_span = MagicMock() - + from litellm.integrations.langfuse.langfuse_otel_attributes import ( LangfuseLLMObsOTELAttributes, ) - - with patch('litellm.integrations.arize._utils.set_attributes') as mock_set_attributes: - with patch('litellm.integrations.arize._utils.safe_set_attribute') as mock_safe_set_attribute: + + with patch( + "litellm.integrations.arize._utils.set_attributes" + ) as mock_set_attributes: + with patch( + "litellm.integrations.arize._utils.safe_set_attribute" + ) as mock_safe_set_attribute: logger = LangfuseOtelLogger() logger.set_langfuse_otel_attributes(mock_span, kwargs, mock_response) - + # Verify that set_attributes was called for general attributes - mock_set_attributes.assert_called_once_with(mock_span, kwargs, mock_response, LangfuseLLMObsOTELAttributes) - + mock_set_attributes.assert_called_once_with( + mock_span, kwargs, mock_response, LangfuseLLMObsOTELAttributes + ) + # Verify that Langfuse-specific attributes were set mock_safe_set_attribute.assert_any_call( mock_span, "langfuse.generation.name", "responses_test_generation" @@ -421,29 +475,30 @@ class TestLangfuseOtelResponsesAPI: """Test that metadata is correctly extracted from ResponsesAPI kwargs.""" # Clean up any existing module mocks import sys + if "litellm.integrations.langfuse.langfuse" in sys.modules: - original_module = sys.modules["litellm.integrations.langfuse.langfuse"] - + sys.modules["litellm.integrations.langfuse.langfuse"] + test_metadata = { "user_id": "responses_user_123", - "session_id": "responses_session_456", + "session_id": "responses_session_456", "custom_metadata": {"key": "value"}, "generation_name": "responses_generation", - "trace_id": "custom_trace_id" + "trace_id": "custom_trace_id", } - + kwargs = { "call_type": "responses", "model": "gpt-4o", - "litellm_params": {"metadata": test_metadata} + "litellm_params": {"metadata": test_metadata}, } - + extracted_metadata = LangfuseOtelLogger._extract_langfuse_metadata(kwargs) - + # Verify all expected metadata was extracted (may have additional fields from header enrichment) for key, value in test_metadata.items(): assert extracted_metadata[key] == value - + assert extracted_metadata["user_id"] == "responses_user_123" assert extracted_metadata["generation_name"] == "responses_generation" assert extracted_metadata["trace_id"] == "custom_trace_id" @@ -457,39 +512,61 @@ class TestLangfuseOtelResponsesAPI: "trace_user_id": "resp_user_456", "session_id": "resp_session_789", "tags": ["responses", "api", "test"], - "trace_metadata": {"source": "responses_api", "version": "1.0"} + "trace_metadata": {"source": "responses_api", "version": "1.0"}, } - - kwargs = { - "call_type": "responses", - "litellm_params": {"metadata": metadata} - } - + + kwargs = {"call_type": "responses", "litellm_params": {"metadata": metadata}} + mock_span = MagicMock() - - with patch('litellm.integrations.arize._utils.safe_set_attribute') as mock_safe_set_attribute: + + with patch( + "litellm.integrations.arize._utils.safe_set_attribute" + ) as mock_safe_set_attribute: LangfuseOtelLogger._set_langfuse_specific_attributes(mock_span, kwargs, {}) - + # Verify specific attributes were set from litellm.types.integrations.langfuse_otel import LangfuseSpanAttributes - + expected_calls = [ - (mock_span, LangfuseSpanAttributes.GENERATION_NAME.value, "responses_gen"), + ( + mock_span, + LangfuseSpanAttributes.GENERATION_NAME.value, + "responses_gen", + ), (mock_span, LangfuseSpanAttributes.GENERATION_ID.value, "resp_gen_123"), (mock_span, LangfuseSpanAttributes.TRACE_NAME.value, "responses_trace"), - (mock_span, LangfuseSpanAttributes.TRACE_USER_ID.value, "resp_user_456"), - (mock_span, LangfuseSpanAttributes.SESSION_ID.value, "resp_session_789"), - (mock_span, LangfuseSpanAttributes.TAGS.value, json.dumps(["responses", "api", "test"])), - (mock_span, LangfuseSpanAttributes.TRACE_METADATA.value, - json.dumps({"source": "responses_api", "version": "1.0"})) + ( + mock_span, + LangfuseSpanAttributes.TRACE_USER_ID.value, + "resp_user_456", + ), + ( + mock_span, + LangfuseSpanAttributes.SESSION_ID.value, + "resp_session_789", + ), + ( + mock_span, + LangfuseSpanAttributes.TAGS.value, + json.dumps(["responses", "api", "test"]), + ), + ( + mock_span, + LangfuseSpanAttributes.TRACE_METADATA.value, + json.dumps({"source": "responses_api", "version": "1.0"}), + ), ] - + for expected_call in expected_calls: mock_safe_set_attribute.assert_any_call(*expected_call) def test_responses_api_with_output(self): """Test Langfuse OTEL logger with Responses API output (reasoning + message).""" - from openai.types.responses import ResponseReasoningItem, ResponseOutputMessage, ResponseOutputText + from openai.types.responses import ( + ResponseReasoningItem, + ResponseOutputMessage, + ResponseOutputText, + ) from openai.types.responses.response_reasoning_item import Summary from litellm.types.integrations.langfuse_otel import LangfuseSpanAttributes @@ -504,9 +581,9 @@ class TestLangfuseOtelResponsesAPI: summary=[ Summary( text="Let me analyze this problem step by step...", - type="summary_text" + type="summary_text", ) - ] + ], ), ResponseOutputMessage( id="msg-001", @@ -519,26 +596,33 @@ class TestLangfuseOtelResponsesAPI: text="The weather in San Francisco is sunny, 20°C.", type="output_text", ) - ] - ) - ] + ], + ), + ], ) kwargs = { "call_type": "responses", - "messages": [{"role": "user", "content": "What's the weather in San Francisco?"}], + "messages": [ + {"role": "user", "content": "What's the weather in San Francisco?"} + ], "model": "gpt-4o", "optional_params": {}, } mock_span = MagicMock() - with patch('litellm.integrations.arize._utils.safe_set_attribute') as mock_safe_set_attribute: - LangfuseOtelLogger._set_langfuse_specific_attributes(mock_span, kwargs, response_obj) + with patch( + "litellm.integrations.arize._utils.safe_set_attribute" + ) as mock_safe_set_attribute: + LangfuseOtelLogger._set_langfuse_specific_attributes( + mock_span, kwargs, response_obj + ) # Verify observation output was set output_calls = [ - call for call in mock_safe_set_attribute.call_args_list + call + for call in mock_safe_set_attribute.call_args_list if call.args[1] == LangfuseSpanAttributes.OBSERVATION_OUTPUT.value ] @@ -552,11 +636,17 @@ class TestLangfuseOtelResponsesAPI: # Verify reasoning summary assert output_data[0]["role"] == "reasoning_summary" - assert output_data[0]["content"] == "Let me analyze this problem step by step..." + assert ( + output_data[0]["content"] + == "Let me analyze this problem step by step..." + ) # Verify message assert output_data[1]["role"] == "assistant" - assert output_data[1]["content"] == "The weather in San Francisco is sunny, 20°C." + assert ( + output_data[1]["content"] + == "The weather in San Francisco is sunny, 20°C." + ) def test_responses_api_with_function_calls(self): """Test Langfuse OTEL logger with Responses API function_call output.""" @@ -574,26 +664,33 @@ class TestLangfuseOtelResponsesAPI: name="get_weather", call_id="call-abc", arguments='{"location": "San Francisco", "unit": "celsius"}', - status="completed" + status="completed", ) - ] + ], ) kwargs = { "call_type": "responses", - "messages": [{"role": "user", "content": "What's the weather in San Francisco?"}], + "messages": [ + {"role": "user", "content": "What's the weather in San Francisco?"} + ], "model": "gpt-4o", "optional_params": {}, } mock_span = MagicMock() - with patch('litellm.integrations.arize._utils.safe_set_attribute') as mock_safe_set_attribute: - LangfuseOtelLogger._set_langfuse_specific_attributes(mock_span, kwargs, response_obj) + with patch( + "litellm.integrations.arize._utils.safe_set_attribute" + ) as mock_safe_set_attribute: + LangfuseOtelLogger._set_langfuse_specific_attributes( + mock_span, kwargs, response_obj + ) # Verify observation output was set output_calls = [ - call for call in mock_safe_set_attribute.call_args_list + call + for call in mock_safe_set_attribute.call_args_list if call.args[1] == LangfuseSpanAttributes.OBSERVATION_OUTPUT.value ] @@ -615,4 +712,4 @@ class TestLangfuseOtelResponsesAPI: if __name__ == "__main__": - pytest.main([__file__]) \ No newline at end of file + pytest.main([__file__]) diff --git a/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py new file mode 100644 index 00000000000..82517b7af9e --- /dev/null +++ b/tests/test_litellm/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -0,0 +1,236 @@ +""" +Unit tests for Anthropic Messages Guardrail Translation Handler + +Tests the handler's ability to process streaming output for Anthropic Messages API +with guardrail transformations, specifically testing edge cases with empty choices. +""" + +import os +import sys +from typing import Any, List, Literal, Optional +from unittest.mock import MagicMock, patch + +import pytest + +sys.path.insert( + 0, os.path.abspath("../../../../../../..") +) # Adds the parent directory to the system path + +from litellm.integrations.custom_guardrail import CustomGuardrail +from litellm.llms.anthropic.chat.guardrail_translation.handler import ( + AnthropicMessagesHandler, +) +from litellm.types.utils import GenericGuardrailAPIInputs + + +class MockPassThroughGuardrail(CustomGuardrail): + """Mock guardrail that passes through without blocking - for testing streaming fallback behavior""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + """Simply return inputs unchanged""" + return inputs + + +class MockDynamicGuardrail(CustomGuardrail): + """Mock guardrail that records dynamic params from request metadata.""" + + def __init__(self, guardrail_name: str): + super().__init__(guardrail_name=guardrail_name) + self.dynamic_params: Optional[dict] = None + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + self.dynamic_params = self.get_guardrail_dynamic_request_body_params( + request_data + ) + return inputs + + +class TestAnthropicMessagesHandlerStreamingOutputProcessing: + """Test streaming output processing functionality""" + + @pytest.mark.asyncio + async def test_process_output_streaming_response_empty_model_response(self): + """Test that streaming response with None model_response doesn't raise error + + This test verifies the fix for the bug where accessing model_response.choices[0] + would raise an error when _build_complete_streaming_response returns None. + """ + handler = AnthropicMessagesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Mock _check_streaming_has_ended to return True (stream ended) + # and _build_complete_streaming_response to return None + with patch.object( + handler, "_check_streaming_has_ended", return_value=True + ), patch( + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + return_value=None, + ): + responses_so_far = [b"data: some chunk"] + + # This should not raise an error + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + ) + + # Should return the responses unchanged + assert result == responses_so_far + + +class TestAnthropicMessagesHandlerInputProcessing: + """Test input processing preserves litellm_metadata for dynamic guardrails.""" + + @pytest.mark.asyncio + async def test_process_input_messages_preserves_litellm_metadata_guardrails(self): + handler = AnthropicMessagesHandler() + guardrail = MockDynamicGuardrail(guardrail_name="cygnal-monitor") + + data = { + "model": "claude-3-5-sonnet-20241022", + "messages": [{"role": "user", "content": "hello"}], + "litellm_metadata": { + "guardrails": [ + { + "cygnal-monitor": { + "extra_body": {"policy_id": "policy-123"} + } + } + ] + }, + } + + with patch("litellm.proxy.proxy_server.premium_user", True): + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) + + assert data.get("litellm_metadata", {}).get("guardrails") + assert guardrail.dynamic_params == {"policy_id": "policy-123"} + + @pytest.mark.asyncio + async def test_process_output_streaming_response_empty_choices(self): + """Test that streaming response with empty choices doesn't raise IndexError + + This test verifies the fix for the bug where accessing model_response.choices[0] + would raise IndexError when the response has an empty choices list. + """ + from litellm.types.utils import ModelResponse + + handler = AnthropicMessagesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Create a mock response with empty choices + mock_response = ModelResponse( + id="msg_123", + created=1234567890, + model="claude-3", + object="chat.completion", + choices=[], # Empty choices + ) + + # Mock _check_streaming_has_ended to return True (stream ended) + # and _build_complete_streaming_response to return the mock response + with patch.object( + handler, "_check_streaming_has_ended", return_value=True + ), patch( + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + return_value=mock_response, + ): + responses_so_far = [b"data: some chunk"] + + # This should not raise IndexError + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + ) + + # Should return the responses unchanged + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_with_valid_choices(self): + """Test that streaming response with valid choices still works correctly""" + from litellm.types.utils import Choices, Message, ModelResponse + + handler = AnthropicMessagesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Create a mock response with valid choices + mock_response = ModelResponse( + id="msg_123", + created=1234567890, + model="claude-3", + object="chat.completion", + choices=[ + Choices( + finish_reason="stop", + index=0, + message=Message( + content="Hello world", + role="assistant", + ), + ) + ], + ) + + # Mock _check_streaming_has_ended to return True (stream ended) + # and _build_complete_streaming_response to return the mock response + with patch.object( + handler, "_check_streaming_has_ended", return_value=True + ), patch( + "litellm.llms.anthropic.chat.guardrail_translation.handler.AnthropicPassthroughLoggingHandler._build_complete_streaming_response", + return_value=mock_response, + ): + responses_so_far = [b"data: some chunk"] + + # This should process successfully + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + ) + + # Should return the responses + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_stream_not_ended(self): + """Test that streaming response falls back to text processing when stream hasn't ended""" + handler = AnthropicMessagesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Mock _check_streaming_has_ended to return False (stream not ended) + with patch.object( + handler, "_check_streaming_has_ended", return_value=False + ), patch.object( + handler, "get_streaming_string_so_far", return_value="partial text" + ): + responses_so_far = [b"data: some chunk"] + + # This should process successfully using text-based guardrail + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=MagicMock(), + ) + + # Should return the responses + assert result == responses_so_far + + +if __name__ == "__main__": + # Run the tests + pytest.main([__file__, "-v"]) diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py index 6a5c022ac7a..b228a51447b 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_anthropic_experimental_pass_through_adapters_transformation.py @@ -341,10 +341,10 @@ def test_translate_openai_content_to_anthropic_empty_function_arguments(): result = adapter._translate_openai_content_to_anthropic(choices=openai_choices) assert len(result) == 1 - assert result[0].type == "tool_use" - assert result[0].id == "call_empty_args" - assert result[0].name == "test_function" - assert result[0].input == {}, "Empty function arguments should result in empty dict" + assert result[0]["type"] == "tool_use" + assert result[0]["id"] == "call_empty_args" + assert result[0]["name"] == "test_function" + assert result[0]["input"] == {}, "Empty function arguments should result in empty dict" def test_translate_openai_content_to_anthropic_text_and_tool_calls(): @@ -372,12 +372,12 @@ def test_translate_openai_content_to_anthropic_text_and_tool_calls(): result = adapter._translate_openai_content_to_anthropic(choices=openai_choices) assert len(result) == 2 - assert result[0].type == "text" - assert result[0].text == "Calling get_weather now." - assert result[1].type == "tool_use" - assert result[1].id == "call_weather" - assert result[1].name == "get_weather" - assert result[1].input == {"location": "Boston"} + assert result[0]["type"] == "text" + assert result[0]["text"] == "Calling get_weather now." + assert result[1]["type"] == "tool_use" + assert result[1]["id"] == "call_weather" + assert result[1]["name"] == "get_weather" + assert result[1]["input"] == {"location": "Boston"} def test_translate_openai_response_to_anthropic_text_and_tool_calls(): @@ -414,11 +414,11 @@ def test_translate_openai_response_to_anthropic_text_and_tool_calls(): anthropic_content = anthropic_response.get("content") assert anthropic_content is not None assert len(anthropic_content) == 2 - assert cast(Any, anthropic_content[0]).type == "text" - assert cast(Any, anthropic_content[0]).text == "Let me grab the current weather." - assert cast(Any, anthropic_content[1]).type == "tool_use" - assert cast(Any, anthropic_content[1]).id == "call_tool_combo" - assert cast(Any, anthropic_content[1]).input == {"location": "Paris"} + assert anthropic_content[0]["type"] == "text" + assert anthropic_content[0]["text"] == "Let me grab the current weather." + assert anthropic_content[1]["type"] == "tool_use" + assert anthropic_content[1]["id"] == "call_tool_combo" + assert anthropic_content[1]["input"] == {"location": "Paris"} assert anthropic_response.get("stop_reason") == "tool_use" @@ -484,11 +484,11 @@ def test_translate_openai_content_to_anthropic_thinking_and_redacted_thinking(): result = adapter._translate_openai_content_to_anthropic(choices=openai_choices) assert len(result) == 2 - assert result[0].type == "thinking" - assert result[0].thinking == "I need to summar" - assert result[0].signature == "sigsig" - assert result[1].type == "redacted_thinking" - assert result[1].data == "REDACTED" + assert result[0]["type"] == "thinking" + assert result[0]["thinking"] == "I need to summar" + assert result[0]["signature"] == "sigsig" + assert result[1]["type"] == "redacted_thinking" + assert result[1]["data"] == "REDACTED" def test_translate_streaming_openai_chunk_to_anthropic_with_thinking(): @@ -1443,13 +1443,13 @@ def test_translate_openai_content_to_anthropic_reasoning_content_without_thinkin assert len(result) == 2 # First block should be thinking block with reasoning_content - assert result[0].type == "thinking" - assert "Considering Letter Frequency" in result[0].thinking - assert "Calculating the Count" in result[0].thinking - assert result[0].signature is None + assert result[0]["type"] == "thinking" + assert "Considering Letter Frequency" in result[0]["thinking"] + assert "Calculating the Count" in result[0]["thinking"] + assert result[0]["signature"] is None # Second block should be text block with content - assert result[1].type == "text" - assert result[1].text == "There are **3** \"r\"s in the word strawberry." + assert result[1]["type"] == "text" + assert result[1]["text"] == "There are **3** \"r\"s in the word strawberry." def test_translate_streaming_openai_chunk_to_anthropic_reasoning_content_without_thinking_blocks(): @@ -1522,13 +1522,13 @@ def test_translate_openai_response_to_anthropic_with_reasoning_content_only(): assert len(anthropic_content) == 2 # First block should be thinking - assert cast(Any, anthropic_content[0]).type == "thinking" - assert "Considering Letter Frequency" in cast(Any, anthropic_content[0]).thinking - assert cast(Any, anthropic_content[0]).signature is None + assert anthropic_content[0]["type"] == "thinking" + assert "Considering Letter Frequency" in anthropic_content[0]["thinking"] + assert anthropic_content[0].get("signature") is None # Second block should be text - assert cast(Any, anthropic_content[1]).type == "text" - assert cast(Any, anthropic_content[1]).text == "There are **3** \"r\"s in the word strawberry." + assert anthropic_content[1]["type"] == "text" + assert anthropic_content[1]["text"] == "There are **3** \"r\"s in the word strawberry." assert anthropic_response.get("stop_reason") == "end_turn" @@ -1702,7 +1702,7 @@ def test_translate_openai_response_restores_tool_names(): ) # Find the tool_use block in the response - tool_use_blocks = [c for c in result["content"] if getattr(c, "type", None) == "tool_use"] + tool_use_blocks = [c for c in result["content"] if c.get("type") == "tool_use"] assert len(tool_use_blocks) == 1 # Name should be restored to original - assert getattr(tool_use_blocks[0], "name", None) == original_name + assert tool_use_blocks[0]["name"] == original_name diff --git a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py index 43c1c413747..8006ffdff1f 100644 --- a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py +++ b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py @@ -108,3 +108,32 @@ def test_get_supported_openai_params_reasoning_effort(): "fireworks_ai/accounts/fireworks/models/llama-v3-70b-instruct" ) assert "reasoning_effort" not in unsupported_params + + +def test_transform_messages_helper_removes_provider_specific_fields(): + """ + Test that _transform_messages_helper removes provider_specific_fields from messages. + """ + config = FireworksAIConfig() + # Simulated messages, as dicts, including provider_specific_fields + messages = [ + { + "role": "user", + "content": "Hello!", + "provider_specific_fields": {"extra": "should be removed"}, + }, + { + "role": "assistant", + "content": "Hi there!", + "provider_specific_fields": {"more": "remove this"}, + }, + { + "role": "user", + "content": "How are you?", + # no provider_specific_fields + } + ] + # Call helper + out = config._transform_messages_helper(messages, model="fireworks/test", litellm_params={}) + for msg in out: + assert "provider_specific_fields" not in msg diff --git a/tests/test_litellm/llms/gemini/files/__init__.py b/tests/test_litellm/llms/gemini/files/__init__.py new file mode 100644 index 00000000000..f48fe7dbe2b --- /dev/null +++ b/tests/test_litellm/llms/gemini/files/__init__.py @@ -0,0 +1 @@ +"""Tests for Gemini files functionality""" diff --git a/tests/test_litellm/llms/gemini/files/test_gemini_files_transformation.py b/tests/test_litellm/llms/gemini/files/test_gemini_files_transformation.py new file mode 100644 index 00000000000..a5f72fc08c3 --- /dev/null +++ b/tests/test_litellm/llms/gemini/files/test_gemini_files_transformation.py @@ -0,0 +1,298 @@ +""" +Test Google AI Studio (Gemini) files transformation functionality +""" + +import os +import pytest +from unittest.mock import Mock, patch + +import httpx + +from litellm.llms.gemini.files.transformation import GoogleAIStudioFilesHandler +from litellm.types.llms.openai import OpenAIFileObject + + +class TestGoogleAIStudioFilesTransformation: + """Test Google AI Studio files transformation""" + + def setup_method(self): + """Setup test method""" + self.handler = GoogleAIStudioFilesHandler() + + def test_transform_retrieve_file_request_with_full_uri(self): + """ + Test that transform_retrieve_file_request returns empty params dict + to avoid 'Content-Type' query parameter error + + Regression test for: https://github.com/BerriAI/litellm/issues/XXX + When retrieving a file, the API was incorrectly trying to pass Content-Type + as a query parameter, which Gemini API rejected. + """ + file_id = "https://generativelanguage.googleapis.com/v1beta/files/test123" + litellm_params = {"api_key": "test-api-key"} + + url, params = self.handler.transform_retrieve_file_request( + file_id=file_id, + optional_params={}, + litellm_params=litellm_params, + ) + + # Verify URL is constructed correctly with API key + assert "key=test-api-key" in url + assert file_id in url + + # CRITICAL: params should be empty dict, not contain Content-Type or any other params + # These would be incorrectly interpreted as query parameters + assert params == {}, f"Expected empty params dict, got: {params}" + assert "Content-Type" not in params, "Content-Type should not be in query params" + + def test_transform_retrieve_file_request_with_file_name_only(self): + """ + Test that transform_retrieve_file_request handles file_id without full URI + """ + file_id = "files/test123" + litellm_params = {"api_key": "test-api-key"} + + url, params = self.handler.transform_retrieve_file_request( + file_id=file_id, + optional_params={}, + litellm_params=litellm_params, + ) + + # Verify URL is constructed correctly + assert "generativelanguage.googleapis.com" in url + assert file_id in url + assert "key=test-api-key" in url + + # CRITICAL: params should be empty dict + assert params == {}, f"Expected empty params dict, got: {params}" + assert "Content-Type" not in params, "Content-Type should not be in query params" + + @patch.dict('os.environ', {}, clear=True) + @patch('litellm.llms.gemini.common_utils.get_secret_str', return_value=None) + def test_transform_retrieve_file_request_missing_api_key(self, mock_get_secret): + """Test that transform_retrieve_file_request raises error when API key is missing""" + file_id = "files/test123" + litellm_params = {} + + with pytest.raises(ValueError, match="api_key is required"): + self.handler.transform_retrieve_file_request( + file_id=file_id, + optional_params={}, + litellm_params=litellm_params, + ) + + def test_transform_retrieve_file_response_success(self): + """Test successful transformation of Gemini file retrieval response""" + # Mock response data from Gemini API + mock_response_data = { + "name": "files/test123", + "displayName": "test_file.pdf", + "mimeType": "application/pdf", + "sizeBytes": "1024", + "createTime": "2024-01-15T10:30:00.123456Z", + "updateTime": "2024-01-15T10:30:00.123456Z", + "expirationTime": "2024-01-17T10:30:00.123456Z", + "sha256Hash": "abcd1234", + "uri": "https://generativelanguage.googleapis.com/v1beta/files/test123", + "state": "ACTIVE", + } + + # Create mock httpx response + mock_response = Mock(spec=httpx.Response) + mock_response.json.return_value = mock_response_data + + # Create mock logging object + mock_logging_obj = Mock() + + # Transform response + result = self.handler.transform_retrieve_file_response( + raw_response=mock_response, + logging_obj=mock_logging_obj, + litellm_params={}, + ) + + # Verify transformation + assert isinstance(result, OpenAIFileObject) + assert result.id == mock_response_data["uri"] + assert result.filename == mock_response_data["displayName"] + assert result.bytes == int(mock_response_data["sizeBytes"]) + assert result.object == "file" + assert result.purpose == "user_data" + assert result.status == "processed" # ACTIVE state maps to processed + assert result.status_details is None + + def test_transform_retrieve_file_response_failed_state(self): + """Test transformation of Gemini file retrieval response with FAILED state""" + mock_response_data = { + "name": "files/test123", + "displayName": "test_file.pdf", + "mimeType": "application/pdf", + "sizeBytes": "1024", + "createTime": "2024-01-15T10:30:00.123456Z", + "uri": "https://generativelanguage.googleapis.com/v1beta/files/test123", + "state": "FAILED", + "error": {"message": "Upload failed", "code": "INTERNAL"}, + } + + mock_response = Mock(spec=httpx.Response) + mock_response.json.return_value = mock_response_data + mock_logging_obj = Mock() + + result = self.handler.transform_retrieve_file_response( + raw_response=mock_response, + logging_obj=mock_logging_obj, + litellm_params={}, + ) + + # Verify error state handling + assert result.status == "error" + assert result.status_details is not None + assert "message" in result.status_details + + def test_transform_retrieve_file_response_processing_state(self): + """Test transformation of Gemini file retrieval response with PROCESSING state""" + mock_response_data = { + "name": "files/test123", + "displayName": "test_file.pdf", + "mimeType": "application/pdf", + "sizeBytes": "1024", + "createTime": "2024-01-15T10:30:00.123456Z", + "uri": "https://generativelanguage.googleapis.com/v1beta/files/test123", + "state": "PROCESSING", + } + + mock_response = Mock(spec=httpx.Response) + mock_response.json.return_value = mock_response_data + mock_logging_obj = Mock() + + result = self.handler.transform_retrieve_file_response( + raw_response=mock_response, + logging_obj=mock_logging_obj, + litellm_params={}, + ) + + # PROCESSING state should map to "uploaded" status + assert result.status == "uploaded" + + def test_transform_retrieve_file_response_missing_createTime(self): + """ + Test that transform_retrieve_file_response raises proper error when createTime is missing + + This tests the error scenario that occurs when API returns an error response + without the expected file metadata fields. + """ + # Mock error response from Gemini API (missing createTime) + mock_response_data = { + "error": { + "code": 400, + "message": "Invalid request", + "status": "INVALID_ARGUMENT", + } + } + + mock_response = Mock(spec=httpx.Response) + mock_response.json.return_value = mock_response_data + mock_logging_obj = Mock() + + # Should raise ValueError with helpful message + with pytest.raises(ValueError, match="Error parsing file retrieve response"): + self.handler.transform_retrieve_file_response( + raw_response=mock_response, + logging_obj=mock_logging_obj, + litellm_params={}, + ) + + def test_validate_environment(self): + """Test that validate_environment properly adds API key to headers""" + headers = {} + api_key = "test-gemini-api-key" + + result_headers = self.handler.validate_environment( + headers=headers, + model="gemini-pro", + messages=[], + optional_params={}, + litellm_params={}, + api_key=api_key, + ) + + # Verify API key is added to headers + assert "x-goog-api-key" in result_headers + assert result_headers["x-goog-api-key"] == api_key + + @patch.dict('os.environ', {}, clear=True) + @patch('litellm.llms.gemini.common_utils.get_secret_str', return_value=None) + def test_validate_environment_missing_api_key(self, mock_get_secret): + """Test that validate_environment raises error when API key is missing""" + headers = {} + + with pytest.raises( + ValueError, match="GEMINI_API_KEY is required for Google AI Studio file operations" + ): + self.handler.validate_environment( + headers=headers, + model="gemini-pro", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + + def test_get_complete_url(self): + """Test that get_complete_url constructs proper upload URL""" + api_base = "https://generativelanguage.googleapis.com" + api_key = "test-api-key" + + url = self.handler.get_complete_url( + api_base=api_base, + api_key=api_key, + model="gemini-pro", + optional_params={}, + litellm_params={}, + ) + + # Verify URL structure + assert api_base in url + assert "upload/v1beta/files" in url + assert f"key={api_key}" in url + + def test_transform_delete_file_request_with_full_uri(self): + """Test delete file request transformation with full URI""" + file_id = "https://generativelanguage.googleapis.com/v1beta/files/test123" + litellm_params = { + "api_key": "test-api-key", + "api_base": "https://generativelanguage.googleapis.com", + } + + url, params = self.handler.transform_delete_file_request( + file_id=file_id, + optional_params={}, + litellm_params=litellm_params, + ) + + # Verify URL extraction + assert "files/test123" in url + assert "generativelanguage.googleapis.com" in url + + # Params should be empty (API key goes in header via validate_environment) + assert params == {} + + def test_transform_delete_file_request_with_file_name_only(self): + """Test delete file request transformation with file name only""" + file_id = "files/test123" + litellm_params = { + "api_key": "test-api-key", + "api_base": "https://generativelanguage.googleapis.com", + } + + url, params = self.handler.transform_delete_file_request( + file_id=file_id, + optional_params={}, + litellm_params=litellm_params, + ) + + # Verify URL construction + assert file_id in url + assert "generativelanguage.googleapis.com" in url + assert params == {} diff --git a/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py b/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py index 6c0195d2831..1f5f53d0f0c 100644 --- a/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py +++ b/tests/test_litellm/llms/openai/chat/guardrail_translation/test_openai_guardrail_handler.py @@ -733,6 +733,154 @@ class TestOpenAIChatCompletionsHandlerToolCallsOutput: assert response.choices[0].finish_reason == "tool_calls" +class MockPassThroughGuardrail(CustomGuardrail): + """Mock guardrail that passes through without blocking - for testing streaming fallback behavior""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + """Simply return inputs unchanged""" + return inputs + + +class TestOpenAIChatCompletionsHandlerStreamingOutput: + """Test streaming output processing functionality""" + + @pytest.mark.asyncio + async def test_process_output_streaming_response_empty_choices(self): + """Test that streaming response with empty choices doesn't raise IndexError + + This test verifies the fix for the bug where accessing chunk.choices[0] + would raise IndexError when a streaming chunk has an empty choices list. + """ + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + handler = OpenAIChatCompletionsHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Create a streaming chunk with empty choices + chunk_with_empty_choices = ModelResponseStream( + id="chatcmpl-123", + created=1234567890, + model="gpt-4", + object="chat.completion.chunk", + choices=[], # Empty choices - this was causing the IndexError + ) + + responses_so_far = [chunk_with_empty_choices] + + # This should not raise IndexError + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + + # Should return the responses unchanged + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_with_valid_choices(self): + """Test that streaming response with valid choices still works correctly""" + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + handler = OpenAIChatCompletionsHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Create streaming chunks with valid choices + chunk1 = ModelResponseStream( + id="chatcmpl-123", + created=1234567890, + model="gpt-4", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + index=0, + delta=Delta(content="Hello"), + finish_reason=None, + ) + ], + ) + + chunk2 = ModelResponseStream( + id="chatcmpl-123", + created=1234567890, + model="gpt-4", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + index=0, + delta=Delta(content=" world"), + finish_reason="stop", + ) + ], + ) + + responses_so_far = [chunk1, chunk2] + + # This should process successfully + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + + # Should return the responses + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_mixed_empty_and_valid_choices_no_finish(self): + """Test streaming response with mix of empty and valid choices chunks (stream not finished) + + This tests the has_stream_ended check when iterating through chunks with mixed choices. + The stream hasn't finished yet (no finish_reason), so it won't trigger stream_chunk_builder. + """ + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + handler = OpenAIChatCompletionsHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Mix of chunks - some with empty choices, some with valid choices + # Stream hasn't finished (no finish_reason) + chunk_empty = ModelResponseStream( + id="chatcmpl-123", + created=1234567890, + model="gpt-4", + object="chat.completion.chunk", + choices=[], + ) + + chunk_valid = ModelResponseStream( + id="chatcmpl-123", + created=1234567890, + model="gpt-4", + object="chat.completion.chunk", + choices=[ + StreamingChoices( + index=0, + delta=Delta(content="Hello"), + finish_reason=None, # Stream not finished + ) + ], + ) + + responses_so_far = [chunk_empty, chunk_valid] + + # This should not raise IndexError when checking has_stream_ended + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + + # Should return the responses + assert result == responses_so_far + + if __name__ == "__main__": # Run the tests pytest.main([__file__, "-v"]) diff --git a/tests/test_litellm/llms/openai/realtime/README.md b/tests/test_litellm/llms/openai/realtime/README.md new file mode 100644 index 00000000000..283b2d29424 --- /dev/null +++ b/tests/test_litellm/llms/openai/realtime/README.md @@ -0,0 +1,82 @@ +# OpenAI Realtime Handler Tests + +## Important Context: `additional_headers` vs `extra_headers` + +### Background + +There was confusion about the correct parameter name for passing headers to `websockets.connect()`. This README documents the resolution for future maintainers. + +### Timeline of Changes + +1. **Dec 5, 2025** - Changed `extra_headers` → `additional_headers` (commit `8db7f1b8e4`) +2. **Dec 18, 2025** - Changed `extra_headers` → `additional_headers` again (PR #17950, commit `9f88d61d10`) +3. **Jan 15, 2026** - Upgraded `websockets` from 13.1.0 → 15.0.1 (commit `a3cf178e24`, Issue #19089) + +### The Issue & Resolution + +**The `websockets` library changed its API between versions:** + +- **websockets < 14.0**: Used `extra_headers` parameter āœ… +- **websockets >= 14.0**: Uses `additional_headers` parameter āœ… + +**LiteLLM uses websockets 15.0.1** (per requirements.txt), which requires `additional_headers`. + +### Verification + +You can verify the correct parameter name: + +```bash +poetry run python -c "import websockets; import inspect; print(inspect.signature(websockets.connect))" +``` + +This shows: `additional_headers: 'HeadersLike | None' = None` for websockets 15.0.1. + +### Current Implementation (Correct) + +```python +# āœ… Correct for websockets 15.0.1+ +await websockets.connect(url, additional_headers={ + "Authorization": f"Bearer {api_key}", + "OpenAI-Beta": "realtime=v1" +}) +``` + +### Impact + +This is NOT just a test fix - this was a **critical bug** that affected all realtime APIs: +- OpenAI realtime +- Azure realtime +- xAI realtime +- Any pass-through realtime connections + +Using `extra_headers` with websockets 15.0.1 resulted in: +``` +TypeError: connect() got an unexpected keyword argument 'extra_headers' +``` + +### For Future Maintainers + +If you see test failures related to header parameters: + +1. **Check installed websockets version:** + ```bash + poetry run python -c "import websockets; print(websockets.__version__)" + ``` + +2. **Check requirements.txt** for the specified version + +3. **Verify the correct parameter:** + - websockets >= 14.0: use `additional_headers` + - websockets < 14.0: use `extra_headers` + +4. **Ensure consistency** across all files: + - `litellm/llms/openai/realtime/handler.py` + - `litellm/llms/azure/realtime/handler.py` + - `litellm/llms/custom_httpx/llm_http_handler.py` + - `litellm/realtime_api/main.py` + - `litellm/proxy/pass_through_endpoints/pass_through_endpoints.py` + +**Current Status (Feb 2026):** +- āœ… websockets version: 15.0.1 +- āœ… Correct parameter: `additional_headers` +- āœ… All handlers updated and working diff --git a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py b/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py index 87923a8093b..c828d030dfd 100644 --- a/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py +++ b/tests/test_litellm/llms/openai/realtime/test_openai_realtime_handler.py @@ -199,7 +199,9 @@ async def test_async_realtime_url_contains_model(): additional_headers = called_kwargs["additional_headers"] assert additional_headers["Authorization"] == f"Bearer {api_key}" assert additional_headers["OpenAI-Beta"] == "realtime=v1" - assert called_kwargs["ssl"] is shared_context + # Verify SSL is configured (should be an SSLContext or True, not None or False) + assert called_kwargs["ssl"] is not None + assert called_kwargs["ssl"] is not False mock_realtime_streaming.assert_called_once() mock_streaming_instance.bidirectional_forward.assert_awaited_once() @@ -259,7 +261,9 @@ async def test_async_realtime_uses_max_size_parameter(): # Verify max_size is set (default None for unlimited, matching OpenAI's SDK) assert "max_size" in called_kwargs assert called_kwargs["max_size"] is None - assert called_kwargs["ssl"] is shared_context + # Verify SSL is configured (should be an SSLContext or True, not None or False) + assert called_kwargs["ssl"] is not None + assert called_kwargs["ssl"] is not False # Default should be None (unlimited) to match OpenAI's official agents SDK # https://github.com/openai/openai-agents-python/blob/cf1b933660e44fd37b4350c41febab8221801409/src/agents/realtime/openai_realtime.py#L235 diff --git a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py index a2849ab91a2..ccece8018ff 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_responses_guardrail_handler.py @@ -817,3 +817,181 @@ class TestOpenAIResponsesHandlerToolCallExtraction: assert task_mappings[0] == (0, 0) assert task_mappings[1] == (0, 1) assert task_mappings[2] == (0, 2) + + +class MockPassThroughGuardrail(CustomGuardrail): + """Mock guardrail that passes through without blocking - for testing streaming fallback behavior""" + + async def apply_guardrail( + self, + inputs: GenericGuardrailAPIInputs, + request_data: dict, + input_type: Literal["request", "response"], + logging_obj: Optional[Any] = None, + ) -> GenericGuardrailAPIInputs: + """Simply return inputs unchanged""" + return inputs + + +class TestOpenAIResponsesHandlerStreamingOutputProcessing: + """Test streaming output processing functionality""" + + @pytest.mark.asyncio + async def test_process_output_streaming_response_empty_output(self): + """Test that streaming response with empty output doesn't raise IndexError + + This test verifies the fix for the bug where accessing model_response_choices[0] + would raise IndexError when the response.completed event has an empty output array. + """ + handler = OpenAIResponsesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Simulate a response.completed streaming event with empty output + responses_so_far = [ + { + "type": "response.completed", + "response": { + "id": "resp_123", + "output": [], # Empty output - this was causing the IndexError + "status": "completed", + }, + } + ] + + # This should not raise IndexError + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + + # Should return the responses unchanged + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_missing_output_key(self): + """Test that streaming response with missing output key doesn't raise IndexError + + This test verifies the handler gracefully handles when the response dict + doesn't contain an 'output' key at all. + """ + handler = OpenAIResponsesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Simulate a response.completed streaming event with missing output key + responses_so_far = [ + { + "type": "response.completed", + "response": { + "id": "resp_123", + "status": "completed", + # No 'output' key - get() will return [] + }, + } + ] + + # This should not raise IndexError + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + + # Should return the responses unchanged + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_unrecognized_output_type(self): + """Test that streaming response with unrecognized output types doesn't raise IndexError + + This test verifies the handler gracefully handles when output items are of + unrecognized types that _convert_response_output_to_choices skips over. + """ + handler = OpenAIResponsesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Simulate a response.completed streaming event with unrecognized output type + responses_so_far = [ + { + "type": "response.completed", + "response": { + "id": "resp_123", + "output": [ + { + "type": "unknown_type", # Unrecognized type + "id": "item_123", + "data": "some data", + } + ], + "status": "completed", + }, + } + ] + + # This should not raise IndexError + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + + # Should return the responses unchanged + assert result == responses_so_far + + @pytest.mark.asyncio + async def test_process_output_streaming_response_with_valid_output(self): + """Test that streaming response with valid output still works correctly""" + handler = OpenAIResponsesHandler() + guardrail = MockPassThroughGuardrail(guardrail_name="test") + + # Simulate a response.completed streaming event with valid message output + responses_so_far = [ + { + "type": "response.created", + "response": {"id": "resp_123"}, + }, + { + "type": "response.output_item.added", + "item": {"type": "message", "id": "msg_123"}, + }, + { + "type": "response.content_part.added", + "part": {"type": "output_text", "text": ""}, + }, + { + "type": "response.output_text.delta", + "delta": "Hello", + }, + { + "type": "response.output_text.delta", + "delta": " world", + }, + { + "type": "response.completed", + "response": { + "id": "resp_123", + "output": [ + { + "type": "message", + "id": "msg_123", + "status": "completed", + "role": "assistant", + "content": [ + {"type": "output_text", "text": "Hello world"}, + ], + } + ], + "status": "completed", + }, + }, + ] + + # This should process successfully + result = await handler.process_output_streaming_response( + responses_so_far=responses_so_far, + guardrail_to_apply=guardrail, + litellm_logging_obj=None, + ) + + # Should return the responses + assert result == responses_so_far diff --git a/tests/test_litellm/llms/test_lifecycle_fix.py b/tests/test_litellm/llms/test_lifecycle_fix.py new file mode 100644 index 00000000000..7b1876a3331 --- /dev/null +++ b/tests/test_litellm/llms/test_lifecycle_fix.py @@ -0,0 +1,46 @@ +""" +Verifies that the httpx client used by AsyncOpenAI is NOT closed +when AsyncHTTPHandler instances are garbage collected. +""" +import asyncio +import gc +import httpx +from litellm.llms.openai.common_utils import BaseOpenAILLM +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + +async def test_httpx_client_not_closed_by_handler_gc(): + """ + Before the fix: _get_async_http_client() returned handler.client, + so when handler was GC'd its __del__ closed the client. + After the fix: returns a standalone httpx.AsyncClient, no handler involved. + """ + # Get the client the same way AsyncOpenAI would + client = BaseOpenAILLM._get_async_http_client() + assert isinstance(client, httpx.AsyncClient) + + # Simulate what the old code did: create an AsyncHTTPHandler and GC it + handler = AsyncHTTPHandler() + handler_client = handler.client + del handler + gc.collect() + + # The client from _get_async_http_client should still be open + # because it's NOT tied to any AsyncHTTPHandler + assert not client.is_closed, "Client was closed prematurely!" + + # Verify it can actually send (build a request without sending) + try: + req = client.build_request("GET", "https://example.com") + print("PASS: Client is still usable after handler GC") + except RuntimeError as e: + if "closed" in str(e): + print(f"FAIL: {e}") + raise + raise + + await client.aclose() + print("All checks passed!") + + +asyncio.run(test_httpx_client_not_closed_by_handler_gc()) diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index cb3b51acd69..75fd597ffa1 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -3218,3 +3218,123 @@ def test_video_metadata_only_for_gemini_3(): assert file_part_3 is not None assert "media_resolution" in file_part_3, "Gemini 3 should have media_resolution" assert "video_metadata" in file_part_3, "Gemini 3 should have video_metadata" + + + +def test_chunk_parser_handles_prompt_feedback_block(): + """Test chunk_parser correctly handles promptFeedback.blockReason""" + from unittest.mock import Mock + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + # Arrange - mock a blocked response + blocked_chunk = { + "promptFeedback": { + "blockReason": "PROHIBITED_CONTENT", + "blockReasonMessage": "The prompt is blocked due to prohibited contents" + }, + "responseId": "test_response_id", + "modelVersion": "gemini-3-pro-preview" + } + + logging_obj = Mock() + logging_obj.optional_params = {} + + streaming_obj = ModelResponseIterator( + streaming_response=iter([]), + sync_stream=True, + logging_obj=logging_obj + ) + + # Act + result = streaming_obj.chunk_parser(blocked_chunk) + + # Assert + assert result is not None, "Result should not be None" + assert len(result.choices) == 1, "Should have exactly one choice" + assert result.choices[0].finish_reason == "content_filter", f"finish_reason should be content_filter, got {result.choices[0].finish_reason}" + assert result.choices[0].delta.content is None, "content should be None" + + +def test_chunk_parser_handles_prompt_feedback_safety_block(): + """Test chunk_parser handles different blockReason types (SAFETY)""" + from unittest.mock import Mock + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + # Arrange - mock a SAFETY blocked response + blocked_chunk = { + "promptFeedback": { + "blockReason": "SAFETY", + "blockReasonMessage": "The prompt is blocked due to safety concerns" + }, + "responseId": "test_safety_response_id", + } + + logging_obj = Mock() + logging_obj.optional_params = {} + + streaming_obj = ModelResponseIterator( + streaming_response=iter([]), + sync_stream=True, + logging_obj=logging_obj + ) + + # Act + result = streaming_obj.chunk_parser(blocked_chunk) + + # Assert + assert result is not None + assert len(result.choices) == 1 + assert result.choices[0].finish_reason == "content_filter" + + +def test_chunk_parser_handles_prompt_feedback_block_with_usage(): + """Test chunk_parser correctly extracts usageMetadata when promptFeedback.blockReason is present""" + from unittest.mock import Mock + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + ModelResponseIterator, + ) + + # Arrange - ęØ”ę‹Ÿäø€äøŖåŒ…å« usageMetadata ēš„ blocked response + blocked_chunk = { + "promptFeedback": { + "blockReason": "PROHIBITED_CONTENT", + "blockReasonMessage": "The prompt is blocked due to prohibited contents" + }, + "responseId": "test_response_id_with_usage", + "modelVersion": "gemini-3-pro-preview", + "usageMetadata": { + "promptTokenCount": 8175, + "candidatesTokenCount": 0, + "totalTokenCount": 8175 + } + } + + logging_obj = Mock() + logging_obj.optional_params = {} + + streaming_obj = ModelResponseIterator( + streaming_response=iter([]), + sync_stream=True, + logging_obj=logging_obj + ) + + # Act + result = streaming_obj.chunk_parser(blocked_chunk) + + # Assert - 验证 content_filter å“åŗ”å’Œ usage éƒ½č¢«ę­£ē”®å¤„ē† + assert result is not None, "Result should not be None" + assert len(result.choices) == 1, "Should have exactly one choice" + assert result.choices[0].finish_reason == "content_filter", f"finish_reason should be content_filter, got {result.choices[0].finish_reason}" + assert result.choices[0].delta.content is None, "content should be None" + + # 验证 usage äæ”ęÆč¢«ę­£ē”®ęå– + assert hasattr(result, "usage"), "result should have usage attribute" + assert result.usage is not None, "usage should not be None" + assert result.usage.prompt_tokens == 8175, f"prompt_tokens should be 8175, got {result.usage.prompt_tokens}" + assert result.usage.completion_tokens == 0, f"completion_tokens should be 0, got {result.usage.completion_tokens}" + assert result.usage.total_tokens == 8175, f"total_tokens should be 8175, got {result.usage.total_tokens}" + diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_jwt_mcp_enforcement.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_jwt_mcp_enforcement.py new file mode 100644 index 00000000000..b4a5a8ca19b --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_jwt_mcp_enforcement.py @@ -0,0 +1,480 @@ +""" +Test to verify Team MCP permissions are enforced when using JWT authentication. + +Scenario: +1. Team "ABC" exists with models configured and MCPs assigned +2. User JWT has team "ABC" in groups (via team_ids_jwt_field) +3. Call MCP list endpoint +4. EXPECTED: Team MCP permissions should be enforced +""" + +import pytest +from unittest.mock import AsyncMock, MagicMock, patch + +from litellm.proxy._types import ( + LiteLLM_JWTAuth, + LiteLLM_TeamTable, + LiteLLM_ObjectPermissionTable, + UserAPIKeyAuth, +) +from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler +from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, +) +from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager + + +@pytest.mark.asyncio +async def test_reproduce_jwt_mcp_enforcement_issue(monkeypatch): + """ + Reproduce the bug where Team MCP permissions are NOT enforced when using JWT. + + Setup: + - Team "ABC" has models ["gpt-4"] and MCPs ["mcp-server-1"] assigned + - JWT has team "ABC" in groups field + - User calls MCP list endpoint (no model requested) + + Expected: team_id should be set to "ABC" so MCP permissions are enforced + Actual (BUG): team_id is None because route check fails for MCP routes + """ + from litellm.caching import DualCache + from litellm.proxy.utils import ProxyLogging + from litellm.router import Router + + # Setup mock router + router = Router(model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4"}}]) + import sys + import types + proxy_server_module = types.ModuleType("proxy_server") + proxy_server_module.llm_router = router + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) + + # Team "ABC" has models configured AND MCPs assigned + team_with_mcp = LiteLLM_TeamTable( + team_id="ABC", + models=["gpt-4"], # Team HAS models + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="perm-123", + mcp_servers=["mcp-server-1"], # Team has MCPs assigned + ), + ) + + async def mock_get_team_object(*args, **kwargs): + team_id = kwargs.get("team_id") or args[0] + if team_id == "ABC": + return team_with_mcp + return None + + monkeypatch.setattr( + "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object + ) + + # Setup JWT handler with team_ids_jwt_field (groups) + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_ids_jwt_field="groups", # Use groups field for teams + # NOTE: team_allowed_routes defaults to ["openai_routes", "info_routes"] + # which does NOT include "mcp_routes" + ) + + user_api_key_cache = DualCache() + proxy_logging_obj = ProxyLogging(user_api_key_cache=user_api_key_cache) + + # Simulate JWT payload with team in groups + jwt_token = { + "sub": "user-123", + "groups": ["ABC"], # Team "ABC" is in groups + "scope": "", + } + + # Mock auth_jwt to return our token + with patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt: + mock_auth_jwt.return_value = jwt_token + + # Call auth_builder for MCP route (like /mcp/tools/list) + result = await JWTAuthManager.auth_builder( + api_key="test-jwt-token", + jwt_handler=jwt_handler, + request_data={}, # No model in request (MCP endpoint) + general_settings={}, + route="/mcp/tools/list", # MCP route + prisma_client=None, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + ) + + # THIS IS THE BUG: team_id should be "ABC" but it's None! + print(f"Result team_id: {result['team_id']}") + print(f"Result team_object: {result['team_object']}") + + # The test should FAIL if the bug exists (team_id is None) + # If the fix is applied, team_id should be "ABC" + assert result["team_id"] == "ABC", ( + f"BUG: team_id should be 'ABC' but got '{result['team_id']}'. " + f"This happens because default team_allowed_routes does not include 'mcp_routes', " + f"so allowed_routes_check() fails and the team is skipped in find_team_with_model_access()." + ) + + +@pytest.mark.asyncio +async def test_verify_mcp_routes_in_default_team_allowed_routes(): + """ + Verify that mcp_routes IS in the default team_allowed_routes. + This is required for team MCP permissions to work with JWT auth. + """ + default_jwt_auth = LiteLLM_JWTAuth() + + print(f"Default team_allowed_routes: {default_jwt_auth.team_allowed_routes}") + + # mcp_routes must be in defaults for team MCP permissions to work + assert "mcp_routes" in default_jwt_auth.team_allowed_routes, ( + "mcp_routes must be in default team_allowed_routes for JWT MCP enforcement to work" + ) + + +@pytest.mark.asyncio +async def test_mcp_route_check_passes_for_team(): + """ + Verify that allowed_routes_check returns True for MCP routes with default settings. + This is required for teams to access MCP endpoints with JWT auth. + """ + from litellm.proxy._types import LitellmUserRoles + from litellm.proxy.auth.auth_checks import allowed_routes_check + + jwt_auth = LiteLLM_JWTAuth() # Use defaults + + # Check if MCP route is allowed for TEAM role + is_allowed = allowed_routes_check( + user_role=LitellmUserRoles.TEAM, + user_route="/mcp/tools/list", + litellm_proxy_roles=jwt_auth, + ) + + print(f"Is /mcp/tools/list allowed for TEAM with defaults? {is_allowed}") + + # MCP routes should be allowed by default for teams + assert is_allowed is True, ( + "MCP routes must be allowed by default for teams for JWT MCP enforcement to work" + ) + + +@pytest.mark.asyncio +async def test_e2e_jwt_team_mcp_permissions_enforced(monkeypatch): + """ + End-to-end test verifying that team MCP permissions are properly enforced + when using JWT authentication with teams in groups. + + This test verifies the complete flow: + 1. JWT token contains team "ABC" in groups field + 2. Team "ABC" exists with MCP servers ["mcp-server-1", "mcp-server-2"] assigned + 3. JWT auth properly sets team_id on UserAPIKeyAuth + 4. MCPRequestHandler.get_allowed_mcp_servers() returns team's MCP servers + """ + from litellm.caching import DualCache + from litellm.proxy.utils import ProxyLogging + from litellm.router import Router + + # Setup mock router + router = Router(model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4"}}]) + import sys + import types + proxy_server_module = types.ModuleType("proxy_server") + proxy_server_module.llm_router = router + proxy_server_module.prisma_client = MagicMock() # Mock prisma client + proxy_server_module.user_api_key_cache = DualCache() + proxy_server_module.proxy_logging_obj = MagicMock() + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) + + # Team "ABC" has MCP servers assigned via object_permission + team_mcp_servers = ["mcp-server-1", "mcp-server-2"] + team_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="perm-abc-123", + mcp_servers=team_mcp_servers, + mcp_access_groups=[], + vector_stores=[], + ) + + team_with_mcp = LiteLLM_TeamTable( + team_id="ABC", + models=["gpt-4"], + object_permission=team_object_permission, + object_permission_id="perm-abc-123", + ) + + async def mock_get_team_object(*args, **kwargs): + team_id = kwargs.get("team_id") or (args[0] if args else None) + if team_id == "ABC": + return team_with_mcp + return None + + monkeypatch.setattr( + "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object + ) + monkeypatch.setattr( + "litellm.proxy.auth.auth_checks.get_team_object", mock_get_team_object + ) + + # Setup JWT handler with team_ids_jwt_field (groups) + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_ids_jwt_field="groups", + ) + + user_api_key_cache = DualCache() + proxy_logging_obj = ProxyLogging(user_api_key_cache=user_api_key_cache) + + # Simulate JWT payload with team in groups + jwt_token = { + "sub": "user-123", + "groups": ["ABC"], + "scope": "", + } + + # Step 1: Verify JWT auth returns correct team_id + with patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt: + mock_auth_jwt.return_value = jwt_token + + result = await JWTAuthManager.auth_builder( + api_key="test-jwt-token", + jwt_handler=jwt_handler, + request_data={}, + general_settings={}, + route="/mcp/tools/list", + prisma_client=None, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + ) + + # Verify team_id is set correctly + assert result["team_id"] == "ABC", f"Expected team_id='ABC', got '{result['team_id']}'" + assert result["team_object"] is not None, "team_object should not be None" + + # Step 2: Create UserAPIKeyAuth with the team_id from JWT auth + user_api_key_auth = UserAPIKeyAuth( + api_key=None, + team_id=result["team_id"], + user_id=result["user_id"], + ) + + # Step 3: Verify MCPRequestHandler returns team's MCP servers + # Mock _get_team_object_permission to return our team's object_permission + with patch.object( + MCPRequestHandler, "_get_team_object_permission" + ) as mock_get_team_perm: + mock_get_team_perm.return_value = team_object_permission + + # Mock _get_allowed_mcp_servers_for_key to return empty (no key-level permissions) + with patch.object( + MCPRequestHandler, "_get_allowed_mcp_servers_for_key" + ) as mock_key_servers: + mock_key_servers.return_value = [] + + # Mock _get_mcp_servers_from_access_groups to return empty + with patch.object( + MCPRequestHandler, "_get_mcp_servers_from_access_groups" + ) as mock_access_groups: + mock_access_groups.return_value = [] + + allowed_servers = await MCPRequestHandler.get_allowed_mcp_servers( + user_api_key_auth + ) + + print(f"Allowed MCP servers: {allowed_servers}") + + # Verify team's MCP servers are returned + assert set(allowed_servers) == set(team_mcp_servers), ( + f"Expected team MCP servers {team_mcp_servers}, got {allowed_servers}" + ) + + +@pytest.mark.asyncio +async def test_e2e_jwt_without_team_no_mcp_servers(monkeypatch): + """ + End-to-end test verifying that when JWT has no teams, no MCP servers are returned. + + This ensures: + 1. JWT token with no groups returns no team_id + 2. MCPRequestHandler.get_allowed_mcp_servers() returns empty list + """ + from litellm.caching import DualCache + from litellm.proxy.utils import ProxyLogging + from litellm.router import Router + + # Setup mock router + router = Router(model_list=[]) + import sys + import types + proxy_server_module = types.ModuleType("proxy_server") + proxy_server_module.llm_router = router + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) + + async def mock_get_team_object(*args, **kwargs): + return None + + monkeypatch.setattr( + "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object + ) + + # Setup JWT handler + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_ids_jwt_field="groups", + ) + + user_api_key_cache = DualCache() + proxy_logging_obj = ProxyLogging(user_api_key_cache=user_api_key_cache) + + # JWT payload with empty groups + jwt_token = { + "sub": "user-123", + "groups": [], # No teams + "scope": "", + } + + with patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt: + mock_auth_jwt.return_value = jwt_token + + result = await JWTAuthManager.auth_builder( + api_key="test-jwt-token", + jwt_handler=jwt_handler, + request_data={}, + general_settings={}, + route="/mcp/tools/list", + prisma_client=None, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + ) + + # Verify no team_id is set + assert result["team_id"] is None, f"Expected team_id=None, got '{result['team_id']}'" + + # Create UserAPIKeyAuth without team_id + user_api_key_auth = UserAPIKeyAuth( + api_key=None, + team_id=None, + user_id=result["user_id"], + ) + + # Verify no MCP servers are returned when there's no team + allowed_servers = await MCPRequestHandler._get_allowed_mcp_servers_for_team( + user_api_key_auth + ) + + assert allowed_servers == [], f"Expected empty list, got {allowed_servers}" + + +@pytest.mark.asyncio +async def test_e2e_jwt_team_mcp_key_intersection(monkeypatch): + """ + End-to-end test verifying MCP permission intersection between key and team. + + Scenario: + - Team has MCP servers: ["server-1", "server-2", "server-3"] + - Key has MCP servers: ["server-2", "server-4"] + - Result should be intersection: ["server-2"] + """ + from litellm.caching import DualCache + from litellm.proxy.utils import ProxyLogging + from litellm.router import Router + + # Setup mock router + router = Router(model_list=[{"model_name": "gpt-4", "litellm_params": {"model": "gpt-4"}}]) + import sys + import types + proxy_server_module = types.ModuleType("proxy_server") + proxy_server_module.llm_router = router + proxy_server_module.prisma_client = MagicMock() + proxy_server_module.user_api_key_cache = DualCache() + proxy_server_module.proxy_logging_obj = MagicMock() + monkeypatch.setitem(sys.modules, "litellm.proxy.proxy_server", proxy_server_module) + + # Team MCP servers + team_mcp_servers = ["server-1", "server-2", "server-3"] + team_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="team-perm", + mcp_servers=team_mcp_servers, + ) + + team_with_mcp = LiteLLM_TeamTable( + team_id="TEAM-X", + models=["gpt-4"], + object_permission=team_object_permission, + ) + + # Key MCP servers + key_mcp_servers = ["server-2", "server-4"] + key_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="key-perm", + mcp_servers=key_mcp_servers, + ) + + async def mock_get_team_object(*args, **kwargs): + team_id = kwargs.get("team_id") or (args[0] if args else None) + if team_id == "TEAM-X": + return team_with_mcp + return None + + monkeypatch.setattr( + "litellm.proxy.auth.handle_jwt.get_team_object", mock_get_team_object + ) + + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth(team_ids_jwt_field="groups") + + user_api_key_cache = DualCache() + proxy_logging_obj = ProxyLogging(user_api_key_cache=user_api_key_cache) + + jwt_token = {"sub": "user-123", "groups": ["TEAM-X"], "scope": ""} + + with patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth_jwt: + mock_auth_jwt.return_value = jwt_token + + result = await JWTAuthManager.auth_builder( + api_key="test-jwt-token", + jwt_handler=jwt_handler, + request_data={}, + general_settings={}, + route="/mcp/tools/list", + prisma_client=None, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + ) + + assert result["team_id"] == "TEAM-X" + + user_api_key_auth = UserAPIKeyAuth( + api_key=None, + team_id=result["team_id"], + user_id=result["user_id"], + object_permission=key_object_permission, # Key has its own permissions + ) + + # Mock the helper methods to return our test data + with patch.object( + MCPRequestHandler, "_get_team_object_permission" + ) as mock_team_perm: + mock_team_perm.return_value = team_object_permission + + with patch.object( + MCPRequestHandler, "_get_key_object_permission" + ) as mock_key_perm: + mock_key_perm.return_value = key_object_permission + + with patch.object( + MCPRequestHandler, "_get_mcp_servers_from_access_groups" + ) as mock_access_groups: + mock_access_groups.return_value = [] + + allowed_servers = await MCPRequestHandler.get_allowed_mcp_servers( + user_api_key_auth + ) + + # Should be intersection: only server-2 is in both + expected = ["server-2"] + assert sorted(allowed_servers) == sorted(expected), ( + f"Expected intersection {expected}, got {allowed_servers}" + ) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_jwt_mcp_simple.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_jwt_mcp_simple.py new file mode 100644 index 00000000000..9ad7736d014 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_jwt_mcp_simple.py @@ -0,0 +1,277 @@ +""" +Simple test to validate MCP permissions are enforced when calling MCP routes with JWT. +""" + +import pytest +from unittest.mock import AsyncMock, MagicMock, patch + +from litellm.proxy._types import ( + LiteLLM_JWTAuth, + LiteLLM_TeamTable, + LiteLLM_ObjectPermissionTable, + UserAPIKeyAuth, +) + + +@pytest.mark.asyncio +async def test_simple_jwt_mcp_permissions_enforced(): + """ + Simple test: Call MCP route with JWT, verify team's MCP servers are returned. + + Setup: + - Team "my-team" has MCP servers: ["github-mcp", "slack-mcp"] + - JWT user belongs to "my-team" + + Expected: Only ["github-mcp", "slack-mcp"] should be allowed + """ + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + + # 1. Create a user authenticated via JWT with team_id set + user_auth = UserAPIKeyAuth( + api_key=None, # JWT auth doesn't have api_key + user_id="jwt-user-123", + team_id="my-team", # This is set by JWT auth when team is in groups + ) + + # 2. Team's MCP permissions + team_mcp_servers = ["github-mcp", "slack-mcp"] + team_object_permission = LiteLLM_ObjectPermissionTable( + object_permission_id="perm-123", + mcp_servers=team_mcp_servers, + ) + + # 3. Mock the team permission lookup + with patch.object( + MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock + ) as mock_team_perm: + mock_team_perm.return_value = team_object_permission + + # Mock key permissions (empty - user has no key-level MCP permissions) + with patch.object( + MCPRequestHandler, "_get_key_object_permission", new_callable=AsyncMock + ) as mock_key_perm: + mock_key_perm.return_value = None + + # Mock access groups (empty) + with patch.object( + MCPRequestHandler, "_get_mcp_servers_from_access_groups", new_callable=AsyncMock + ) as mock_access_groups: + mock_access_groups.return_value = [] + + # 4. Call get_allowed_mcp_servers - this is what MCP routes use + allowed = await MCPRequestHandler.get_allowed_mcp_servers(user_auth) + + # 5. Verify only team's MCP servers are returned + assert sorted(allowed) == sorted(team_mcp_servers), ( + f"Expected {team_mcp_servers}, got {allowed}" + ) + + # Verify team permission was looked up + mock_team_perm.assert_called_once_with(user_auth) + + +@pytest.mark.asyncio +async def test_simple_jwt_no_team_no_mcp_servers(): + """ + Simple test: JWT user with no team should get no MCP servers. + """ + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + + # User with no team_id (JWT didn't have teams in groups) + user_auth = UserAPIKeyAuth( + api_key=None, + user_id="jwt-user-no-team", + team_id=None, # No team + ) + + # _get_allowed_mcp_servers_for_team returns [] when team_id is None + allowed = await MCPRequestHandler._get_allowed_mcp_servers_for_team(user_auth) + + assert allowed == [], f"Expected [], got {allowed}" + + +@pytest.mark.asyncio +async def test_simple_jwt_team_id_required_for_mcp_permissions(): + """ + Simple test: Verify that team_id must be set for team MCP permissions to work. + + This is the key insight - if JWT auth doesn't set team_id, + team MCP permissions won't be enforced. + """ + from litellm.proxy._experimental.mcp_server.auth.user_api_key_auth_mcp import ( + MCPRequestHandler, + ) + + # Case 1: team_id is set -> team permissions should be checked + user_with_team = UserAPIKeyAuth( + api_key=None, + user_id="user-1", + team_id="team-abc", + ) + + team_mcp_servers = ["server-1", "server-2"] + team_perm = LiteLLM_ObjectPermissionTable( + object_permission_id="perm-1", + mcp_servers=team_mcp_servers, + ) + + with patch.object( + MCPRequestHandler, "_get_team_object_permission", new_callable=AsyncMock + ) as mock_perm: + mock_perm.return_value = team_perm + + with patch.object( + MCPRequestHandler, "_get_mcp_servers_from_access_groups", new_callable=AsyncMock + ) as mock_groups: + mock_groups.return_value = [] + + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(user_with_team) + + assert sorted(result) == sorted(team_mcp_servers) + mock_perm.assert_called_once() # Permission WAS checked + + # Case 2: team_id is None -> team permissions NOT checked + user_without_team = UserAPIKeyAuth( + api_key=None, + user_id="user-2", + team_id=None, + ) + + result = await MCPRequestHandler._get_allowed_mcp_servers_for_team(user_without_team) + assert result == [] # No permissions returned + + +@pytest.mark.asyncio +async def test_jwt_auth_sets_team_id_for_mcp_route(): + """ + Test that JWT auth properly sets team_id when accessing MCP routes. + + This is the critical test - when user calls /mcp/tools/list with JWT, + the team_id from JWT groups must be set on UserAPIKeyAuth. + """ + from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler + from litellm.caching import DualCache + from litellm.proxy.utils import ProxyLogging + + # Setup + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_ids_jwt_field="groups", # Teams come from "groups" field in JWT + ) + + # Team exists with models + team = LiteLLM_TeamTable( + team_id="team-from-jwt", + models=["gpt-4"], + ) + + user_api_key_cache = DualCache() + proxy_logging_obj = ProxyLogging(user_api_key_cache=user_api_key_cache) + + # Mock JWT token with team in groups + jwt_payload = { + "sub": "user-123", + "groups": ["team-from-jwt"], + "scope": "", + } + + with patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth: + mock_auth.return_value = jwt_payload + + with patch( + "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock + ) as mock_get_team: + mock_get_team.return_value = team + + # Simulate calling MCP route + result = await JWTAuthManager.auth_builder( + api_key="jwt-token", + jwt_handler=jwt_handler, + request_data={}, + general_settings={}, + route="/mcp/tools/list", # MCP route + prisma_client=None, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + ) + + # THE KEY ASSERTION: team_id must be set + assert result["team_id"] == "team-from-jwt", ( + f"team_id should be 'team-from-jwt' but got '{result['team_id']}'. " + "This means JWT auth is not properly setting team_id for MCP routes!" + ) + + +@pytest.mark.asyncio +async def test_mcp_route_without_model_still_returns_team_id(): + """ + Test that MCP routes (which don't specify a model) still get team_id assigned. + + Key insight: MCP routes don't require a model in the request, but the JWT auth + flow must still assign a team_id so that team MCP permissions are enforced. + + The flow is: + 1. JWT token contains team in "groups" field + 2. find_team_with_model_access() is called with requested_model=None + 3. Since `not requested_model` is True, model check passes + 4. Route check passes because "mcp_routes" is in team_allowed_routes + 5. team_id is returned and set on UserAPIKeyAuth + """ + from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler + from litellm.caching import DualCache + from litellm.proxy.utils import ProxyLogging + + # Setup + jwt_handler = JWTHandler() + jwt_handler.litellm_jwtauth = LiteLLM_JWTAuth( + team_ids_jwt_field="groups", + ) + + # Team exists - note: models is a list (can be empty or have values) + # The key is that when no model is requested, model check is skipped + team = LiteLLM_TeamTable( + team_id="my-team", + models=["gpt-4", "gpt-3.5-turbo"], # Team has models, but MCP request won't specify one + ) + + user_api_key_cache = DualCache() + proxy_logging_obj = ProxyLogging(user_api_key_cache=user_api_key_cache) + + # JWT with team in groups + jwt_payload = { + "sub": "user-abc", + "groups": ["my-team"], + "scope": "", + } + + with patch.object(jwt_handler, "auth_jwt", new_callable=AsyncMock) as mock_auth: + mock_auth.return_value = jwt_payload + + with patch( + "litellm.proxy.auth.handle_jwt.get_team_object", new_callable=AsyncMock + ) as mock_get_team: + mock_get_team.return_value = team + + # Call MCP route with NO MODEL in request_data + result = await JWTAuthManager.auth_builder( + api_key="jwt-token", + jwt_handler=jwt_handler, + request_data={}, # <-- NO MODEL SPECIFIED + general_settings={}, + route="/mcp/tools/list", # MCP route + prisma_client=None, + user_api_key_cache=user_api_key_cache, + parent_otel_span=None, + proxy_logging_obj=proxy_logging_obj, + ) + + # Team ID must still be set even though no model was requested + assert result["team_id"] == "my-team", ( + f"Expected team_id='my-team' but got '{result['team_id']}'. " + "MCP routes without model should still get team_id from JWT!" + ) diff --git a/tests/test_litellm/proxy/agent_endpoints/test_model_list_helpers.py b/tests/test_litellm/proxy/agent_endpoints/test_model_list_helpers.py new file mode 100644 index 00000000000..92cd3d9ad6b --- /dev/null +++ b/tests/test_litellm/proxy/agent_endpoints/test_model_list_helpers.py @@ -0,0 +1,110 @@ +""" +Test appending A2A agents to model lists. + +Maps to: litellm/proxy/agent_endpoints/model_list_helpers.py +""" +import os +import sys + +sys.path.insert(0, os.path.abspath("../../../..")) + +from unittest.mock import AsyncMock, Mock, patch + +import pytest + +from litellm.proxy.agent_endpoints.model_list_helpers import ( + append_agents_to_model_group, + append_agents_to_model_info, +) +from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth +from litellm.types.agents import AgentResponse +from litellm.types.proxy.management_endpoints.model_management_endpoints import ( + ModelGroupInfoProxy, +) + + +@pytest.mark.asyncio +async def test_append_agents_to_model_group(): + """Test agents are converted to model group format with a2a/ prefix""" + + # Mock agent data + mock_agent = AgentResponse( + agent_id="test-agent-id", + agent_name="my-agent", + agent_card_params={"url": "http://example.com"}, + litellm_params=None, + ) + + # Mock AgentRequestHandler at its source location + mock_get_allowed_agents = AsyncMock(return_value=["test-agent-id"]) + + # Mock global_agent_registry + mock_registry = Mock() + mock_registry.get_agent_by_id = Mock(return_value=mock_agent) + + with patch( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.get_allowed_agents", + mock_get_allowed_agents, + ): + with patch( + "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", + mock_registry, + ): + model_groups = [] + user_api_key_dict = Mock(spec=UserAPIKeyAuth) + + result = await append_agents_to_model_group( + model_groups=model_groups, + user_api_key_dict=user_api_key_dict, + ) + + # Verify agent was converted with a2a/ prefix + assert len(result) == 1 + assert result[0].model_group == "a2a/my-agent" + assert result[0].mode == "chat" + assert result[0].providers == ["a2a"] + + +@pytest.mark.asyncio +async def test_append_agents_to_model_info(): + """Test agents are converted to model info format with a2a/ prefix""" + + # Mock agent data + mock_agent = AgentResponse( + agent_id="agent-123", + agent_name="test-agent", + agent_card_params={"url": "http://example.com"}, + litellm_params=None, + created_by="user-123", + ) + + # Mock AgentRequestHandler at its source location + mock_get_allowed_agents = AsyncMock(return_value=["agent-123"]) + + # Mock global_agent_registry + mock_registry = Mock() + mock_registry.get_agent_by_id = Mock(return_value=mock_agent) + + with patch( + "litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.get_allowed_agents", + mock_get_allowed_agents, + ): + with patch( + "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", + mock_registry, + ): + models = [] + user_api_key_dict = Mock(spec=UserAPIKeyAuth) + + result = await append_agents_to_model_info( + models=models, + user_api_key_dict=user_api_key_dict, + ) + + # Verify agent was converted with a2a/ prefix + assert len(result) == 1 + assert result[0]["model_name"] == "a2a/test-agent" + assert result[0]["litellm_params"]["model"] == "a2a/test-agent" + assert result[0]["litellm_params"]["custom_llm_provider"] == "a2a" + assert result[0]["model_info"]["id"] == "agent-123" + assert result[0]["model_info"]["mode"] == "chat" diff --git a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py index 72403b0ba7b..6ccecf59eed 100644 --- a/tests/test_litellm/proxy/db/test_db_spend_update_writer.py +++ b/tests/test_litellm/proxy/db/test_db_spend_update_writer.py @@ -132,7 +132,7 @@ async def test_update_daily_spend_with_null_entity_id(): assert create_data["model"] == "gpt-4" assert create_data["custom_llm_provider"] == "openai" assert create_data["mcp_namespaced_tool_name"] == "" - assert create_data["endpoint"] is None + assert create_data["endpoint"] == "" assert create_data["prompt_tokens"] == 10 assert create_data["completion_tokens"] == 20 assert create_data["spend"] == 0.1 @@ -194,7 +194,7 @@ async def test_update_daily_spend_sorting(): "model_group": None, "mcp_namespaced_tool_name": "", "custom_llm_provider": "openai", - "endpoint": None, + "endpoint": "", "prompt_tokens": 10, "completion_tokens": 20, "spend": 0.1, @@ -838,4 +838,126 @@ async def test_endpoint_field_is_correctly_mapped_from_call_type(): assert transaction["date"] == "2024-01-01" assert transaction["api_key"] == "test-key" assert transaction["model"] == "gpt-4" - assert transaction["custom_llm_provider"] == "openai" \ No newline at end of file + assert transaction["custom_llm_provider"] == "openai" + + +@pytest.mark.asyncio +async def test_update_daily_spend_logs_detailed_error_on_batch_upsert_failure(): + """ + Test that when batch upsert fails, detailed error information is logged. + This ensures proper debugging information is available for issues like unique constraint violations. + """ + from litellm._logging import verbose_proxy_logger + + # Setup + mock_prisma_client = MagicMock() + mock_batcher = MagicMock() + mock_table = MagicMock() + mock_batch_context = MagicMock() + mock_batch_context.__aenter__ = AsyncMock(return_value=mock_batcher) + mock_batcher.litellm_dailyuserspend = mock_table + + # Make the batch context manager's exit raise an exception + # This simulates a batch commit failure (e.g., unique constraint violation) + test_exception = Exception("Unique constraint violation") + mock_batch_context.__aexit__ = AsyncMock(side_effect=test_exception) + mock_prisma_client.db.batch_.return_value = mock_batch_context + + # Create a transaction + daily_spend_transactions = { + "test_key": { + "user_id": "test-user", + "date": "2024-01-01", + "api_key": "test-api-key", + "model": "gpt-4", + "custom_llm_provider": "openai", + "prompt_tokens": 10, + "completion_tokens": 20, + "spend": 0.1, + "api_requests": 1, + "successful_requests": 1, + "failed_requests": 0, + } + } + + # Create a mock proxy_logging_obj with failure_handler as AsyncMock + mock_proxy_logging = MagicMock() + mock_proxy_logging.failure_handler = AsyncMock() + + # Mock the logger to capture exception calls + with patch.object(verbose_proxy_logger, 'exception') as mock_exception_logger: + # Call the method and expect it to raise the exception + with pytest.raises(Exception, match="Unique constraint violation"): + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=0, # No retries to make test faster + prisma_client=mock_prisma_client, + proxy_logging_obj=mock_proxy_logging, + daily_spend_transactions=daily_spend_transactions, + entity_type="user", + entity_id_field="user_id", + table_name="litellm_dailyuserspend", + unique_constraint_name="user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint", + ) + + # Verify that exception was logged with detailed information + assert mock_exception_logger.called + call_args = mock_exception_logger.call_args[0][0] + assert "Daily user spend batch upsert failed" in call_args + assert "Table: litellm_dailyuserspend" in call_args + assert "Constraint: user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint" in call_args + assert "Batch size: 1" in call_args + assert "Unique constraint violation" in call_args + + +@pytest.mark.asyncio +async def test_update_daily_spend_re_raises_exception_after_logging(): + """ + Test that when batch upsert fails, the exception is properly re-raised after logging. + This ensures that error handling continues to work correctly upstream. + """ + # Setup + mock_prisma_client = MagicMock() + mock_batcher = MagicMock() + mock_table = MagicMock() + mock_batch_context = MagicMock() + mock_batch_context.__aenter__ = AsyncMock(return_value=mock_batcher) + mock_batcher.litellm_dailyuserspend = mock_table + + # Create a transaction + daily_spend_transactions = { + "test_key": { + "user_id": "test-user", + "date": "2024-01-01", + "api_key": "test-api-key", + "model": "gpt-4", + "custom_llm_provider": "openai", + "prompt_tokens": 10, + "completion_tokens": 20, + "spend": 0.1, + "api_requests": 1, + "successful_requests": 1, + "failed_requests": 0, + } + } + + # Create a custom exception to verify it's re-raised + custom_exception = ValueError("Database connection lost") + mock_batch_context.__aexit__ = AsyncMock(side_effect=custom_exception) + mock_prisma_client.db.batch_.return_value = mock_batch_context + + # Create a mock proxy_logging_obj with failure_handler as AsyncMock + mock_proxy_logging = MagicMock() + mock_proxy_logging.failure_handler = AsyncMock() + + # Verify the exception is re-raised + with pytest.raises(ValueError, match="Database connection lost"): + await DBSpendUpdateWriter._update_daily_spend( + n_retry_times=0, # No retries to make test faster + prisma_client=mock_prisma_client, + proxy_logging_obj=mock_proxy_logging, + daily_spend_transactions=daily_spend_transactions, + entity_type="user", + entity_id_field="user_id", + table_name="litellm_dailyuserspend", + unique_constraint_name="user_id_date_api_key_model_custom_llm_provider_mcp_namespaced_tool_name_endpoint", + ) \ No newline at end of file diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py index 6dc658827bc..109ad0bfdc8 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/test_grayswan.py @@ -34,8 +34,9 @@ def test_prepare_payload_uses_dynamic_overrides( "policy_id": "dynamic-policy", "reasoning_mode": "thinking", } + request_data = {} - payload = grayswan_guardrail._prepare_payload(messages, dynamic_body) + payload = grayswan_guardrail._prepare_payload(messages, dynamic_body, request_data) assert payload["messages"] == messages assert payload["categories"] == {"custom": "override"} @@ -47,14 +48,27 @@ def test_prepare_payload_falls_back_to_guardrail_defaults( grayswan_guardrail: GraySwanGuardrail, ) -> None: messages = [{"role": "user", "content": "hello"}] + request_data = {} - payload = grayswan_guardrail._prepare_payload(messages, {}) + payload = grayswan_guardrail._prepare_payload(messages, {}, request_data) assert payload["categories"] == {"safety": "general policy"} assert payload["policy_id"] == "default-policy" assert payload["reasoning_mode"] == "hybrid" +def test_prepare_payload_includes_dynamic_metadata( + grayswan_guardrail: GraySwanGuardrail, +) -> None: + messages = [{"role": "user", "content": "hello"}] + dynamic_body = {"metadata": {"trace_id": "trace-123", "tags": ["a", "b"]}} + request_data = {} + + payload = grayswan_guardrail._prepare_payload(messages, dynamic_body, request_data) + + assert payload["metadata"] == dynamic_body["metadata"] + + def test_process_response_does_not_block_under_threshold( grayswan_guardrail: GraySwanGuardrail, ) -> None: @@ -160,6 +174,119 @@ async def test_run_guardrail_raises_api_error( await grayswan_guardrail.run_grayswan_guardrail(payload) +@pytest.mark.asyncio +async def test_apply_guardrail_passthrough_not_swallowed_by_fail_open( + monkeypatch, +) -> None: + guardrail = GraySwanGuardrail( + guardrail_name="grayswan-passthrough", + api_key="test-key", + on_flagged_action="passthrough", + violation_threshold=0.2, + fail_open=True, + event_hook=GuardrailEventHooks.pre_call, + ) + + async def _fake_call(_payload: dict): + return {"violation": 0.92, "violated_rule_descriptions": []} + + monkeypatch.setattr(guardrail, "_call_grayswan_api", _fake_call) + + with pytest.raises(ModifyResponseException): + await guardrail.apply_guardrail( + inputs={"texts": ["bad"]}, + request_data={"model": "gpt-4"}, + input_type="request", + ) + + +@pytest.mark.asyncio +async def test_apply_guardrail_block_not_swallowed_by_fail_open( + monkeypatch, +) -> None: + guardrail = GraySwanGuardrail( + guardrail_name="grayswan-block", + api_key="test-key", + on_flagged_action="block", + violation_threshold=0.2, + fail_open=True, + event_hook=GuardrailEventHooks.pre_call, + ) + + async def _fake_call(_payload: dict): + return {"violation": 0.92, "violated_rule_descriptions": []} + + monkeypatch.setattr(guardrail, "_call_grayswan_api", _fake_call) + + with pytest.raises(HTTPException): + await guardrail.apply_guardrail( + inputs={"texts": ["bad"]}, + request_data={"model": "gpt-4"}, + input_type="request", + ) + + +@pytest.mark.asyncio +async def test_apply_guardrail_non_grayswan_http_exception_fail_open_true( + monkeypatch, +) -> None: + guardrail = GraySwanGuardrail( + guardrail_name="grayswan-error", + api_key="test-key", + on_flagged_action="monitor", + violation_threshold=0.2, + fail_open=True, + event_hook=GuardrailEventHooks.pre_call, + ) + + async def _fake_call(_payload: dict): + return {"violation": 0.0, "violated_rule_descriptions": []} + + def _fake_process(**_kwargs): + raise HTTPException(status_code=500, detail={"error": "upstream failed"}) + + monkeypatch.setattr(guardrail, "_call_grayswan_api", _fake_call) + monkeypatch.setattr(guardrail, "_process_response_internal", _fake_process) + + result = await guardrail.apply_guardrail( + inputs={"texts": ["ok"]}, + request_data={"model": "gpt-4"}, + input_type="request", + ) + + assert result["texts"] == ["ok"] + + +@pytest.mark.asyncio +async def test_apply_guardrail_non_grayswan_http_exception_fail_open_false( + monkeypatch, +) -> None: + guardrail = GraySwanGuardrail( + guardrail_name="grayswan-error", + api_key="test-key", + on_flagged_action="monitor", + violation_threshold=0.2, + fail_open=False, + event_hook=GuardrailEventHooks.pre_call, + ) + + async def _fake_call(_payload: dict): + return {"violation": 0.0, "violated_rule_descriptions": []} + + def _fake_process(**_kwargs): + raise HTTPException(status_code=500, detail={"error": "upstream failed"}) + + monkeypatch.setattr(guardrail, "_call_grayswan_api", _fake_call) + monkeypatch.setattr(guardrail, "_process_response_internal", _fake_process) + + with pytest.raises(GraySwanGuardrailAPIError): + await guardrail.apply_guardrail( + inputs={"texts": ["ok"]}, + request_data={"model": "gpt-4"}, + input_type="request", + ) + + def test_process_response_passthrough_raises_exception_in_pre_call() -> None: """Test that passthrough mode raises ModifyResponseException in pre_call hook.""" guardrail = GraySwanGuardrail( diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index 2a33c56b56a..b41cded1d0a 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -8,12 +8,13 @@ from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTra from litellm.proxy._experimental.mcp_server.guardrail_translation.handler import ( MCPGuardrailTranslationHandler, ) +from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail import unified_guardrail as unified_module from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( UnifiedLLMGuardrails, ) from litellm.types.guardrails import GuardrailEventHooks -from litellm.types.utils import CallTypes +from litellm.types.utils import CallTypes, Delta, ModelResponseStream, StreamingChoices class RecordingGuardrail(CustomGuardrail): @@ -131,3 +132,100 @@ class TestUnifiedLLMGuardrails: ) assert guardrail.event_history == [GuardrailEventHooks.during_call] + + class TestAsyncPostCallStreamingIteratorHook: + @pytest.mark.asyncio + async def test_streaming_content_not_lost_on_sampled_chunks(self): + """ + Verify that every chunk's content is preserved in the output stream. + + The bug: process_output_streaming_response puts the combined + guardrailed text in the first chunk and clears all subsequent + chunks to "". The hook then yielded processed_items[-1] (the + cleared last item), permanently losing every Nth chunk's content. + """ + + class _ContentClearingTranslation(BaseTranslation): + """Simulates the real OpenAI handler behavior that triggers the bug.""" + + async def process_input_messages(self, data, guardrail_to_apply, litellm_logging_obj=None): # type: ignore[override] + return data + + async def process_output_response(self, response, guardrail_to_apply, litellm_logging_obj=None, user_api_key_dict=None): # type: ignore[override] + return response + + async def process_output_streaming_response( + self, + responses_so_far, + guardrail_to_apply, + litellm_logging_obj=None, + user_api_key_dict=None, + ): + # Simulate what the real handler does: + # put combined text in first chunk, clear the rest + combined = "" + for resp in responses_so_far: + for choice in resp.choices: + if choice.delta and choice.delta.content: + combined += choice.delta.content + + first_set = False + for resp in responses_so_far: + for choice in resp.choices: + if not first_set: + choice.delta.content = combined + first_set = True + else: + choice.delta.content = "" + + return responses_so_far + + # Override the mapping to use our content-clearing translation + unified_module.endpoint_guardrail_translation_mappings = { + CallTypes.acompletion: _ContentClearingTranslation, + } + + handler = UnifiedLLMGuardrails() + guardrail = RecordingGuardrail() + + # Create 10 streaming chunks with distinct content + chunks = [] + for i in range(10): + chunk = ModelResponseStream( + choices=[StreamingChoices( + delta=Delta(content=f"word{i} ", role="assistant"), + finish_reason=None, + )], + ) + chunks.append(chunk) + + async def mock_stream(): + for chunk in chunks: + yield chunk + + user_api_key_dict = UserAPIKeyAuth( + api_key="test-key", + request_route="/v1/chat/completions", + ) + + request_data = { + "guardrail_to_apply": guardrail, + "model": "gpt-4", + } + + # Collect all yielded chunks + yielded_contents = [] + async for item in handler.async_post_call_streaming_iterator_hook( + user_api_key_dict=user_api_key_dict, + response=mock_stream(), + request_data=request_data, + ): + content = item.choices[0].delta.content if item.choices[0].delta else None + yielded_contents.append(content) + + # Every chunk should have non-empty content + for i, content in enumerate(yielded_contents): + assert content is not None and content != "", ( + f"Chunk {i} lost its content (got {content!r}). " + f"Expected non-empty content for every streamed chunk." + ) diff --git a/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py new file mode 100644 index 00000000000..c7e3fba94ee --- /dev/null +++ b/tests/test_litellm/proxy/management_endpoints/search_endpoints/test_search_tool_management.py @@ -0,0 +1,546 @@ +import os +import sys +from datetime import datetime +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi.testclient import TestClient + +sys.path.insert( + 0, os.path.abspath("../../../../..") +) # Adds the parent directory to the system path + +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + +# Import proxy_server module first to ensure it's initialized +import litellm.proxy.proxy_server as ps + +# Now we can safely import app +from litellm.proxy.proxy_server import app + +client = TestClient(app) + + +@pytest.mark.asyncio +async def test_list_search_tools_db_only(monkeypatch): + """Test listing search tools when only DB tools exist""" + # Mock DB tools + db_tools = [ + { + "search_tool_id": "test-id-1", + "search_tool_name": "db-tool-1", + "litellm_params": {"search_provider": "perplexity", "api_key": "sk-test"}, + "search_tool_info": {"description": "DB tool 1"}, + "created_at": datetime(2023, 11, 9, 12, 34, 56), + "updated_at": datetime(2023, 11, 9, 13, 45, 12), + } + ] + + # Mock SearchToolRegistry + mock_registry = MagicMock() + mock_registry.get_all_search_tools_from_db = AsyncMock(return_value=db_tools) + with patch( + "litellm.proxy.search_endpoints.search_tool_management.SEARCH_TOOL_REGISTRY", + mock_registry, + ): + # Mock prisma_client + mock_prisma = MagicMock() + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + # Mock proxy_config + mock_proxy_config = MagicMock() + mock_proxy_config.get_config = AsyncMock(return_value={}) + mock_proxy_config.parse_search_tools = MagicMock(return_value=None) + with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config): + # Mock auth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + test_client = TestClient(app) + response = test_client.get("/search_tools/list") + assert response.status_code == 200 + data = response.json() + assert "search_tools" in data + assert len(data["search_tools"]) == 1 + + tool = data["search_tools"][0] + assert tool["search_tool_id"] == "test-id-1" + assert tool["search_tool_name"] == "db-tool-1" + assert tool["is_from_config"] is False + # Verify datetime conversion to ISO string + assert tool["created_at"] == "2023-11-09T12:34:56" + assert tool["updated_at"] == "2023-11-09T13:45:12" + # Verify masking of sensitive values + assert tool["litellm_params"]["api_key"] != "sk-test" + assert "****" in tool["litellm_params"]["api_key"] + assert tool["litellm_params"]["search_provider"] == "perplexity" + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_list_search_tools_config_only(monkeypatch): + """Test listing search tools when only config tools exist""" + # Mock DB tools - empty + db_tools = [] + + # Mock config tools + config_tools = [ + { + "search_tool_name": "config-tool-1", + "litellm_params": {"search_provider": "tavily", "api_key": "tvly-secret-key"}, + "search_tool_info": {"description": "Config tool 1"}, + } + ] + + # Mock SearchToolRegistry + mock_registry = MagicMock() + mock_registry.get_all_search_tools_from_db = AsyncMock(return_value=db_tools) + with patch( + "litellm.proxy.search_endpoints.search_tool_management.SEARCH_TOOL_REGISTRY", + mock_registry, + ): + # Mock prisma_client + mock_prisma = MagicMock() + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + # Mock proxy_config + mock_proxy_config = MagicMock() + mock_proxy_config.get_config = AsyncMock(return_value={"search_tools": config_tools}) + mock_proxy_config.parse_search_tools = MagicMock(return_value=config_tools) + with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config): + # Mock auth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + test_client = TestClient(app) + response = test_client.get("/search_tools/list") + assert response.status_code == 200 + data = response.json() + assert "search_tools" in data + assert len(data["search_tools"]) == 1 + + tool = data["search_tools"][0] + assert tool["search_tool_name"] == "config-tool-1" + assert tool["is_from_config"] is True + assert tool["search_tool_id"] is None + assert tool["created_at"] is None + assert tool["updated_at"] is None + # Verify masking + assert "tv****ey" in tool["litellm_params"]["api_key"] + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_list_search_tools_filters_duplicate_config_tools(monkeypatch): + """ + Test that config tools with the same name as DB tools are filtered out. + This tests the new filtering logic added in lines 139-142. + """ + # Mock DB tools + db_tools = [ + { + "search_tool_id": "db-id-1", + "search_tool_name": "existing-tool", + "litellm_params": {"search_provider": "perplexity", "api_key": "sk-db"}, + "search_tool_info": {"description": "DB tool"}, + "created_at": datetime(2023, 11, 9, 12, 34, 56), + "updated_at": datetime(2023, 11, 9, 13, 45, 12), + } + ] + + # Mock config tools - one duplicate, one unique + config_tools = [ + { + "search_tool_name": "existing-tool", # Duplicate - should be filtered + "litellm_params": {"search_provider": "tavily", "api_key": "tvly-config"}, + "search_tool_info": {"description": "Config tool - duplicate"}, + }, + { + "search_tool_name": "unique-config-tool", # Unique - should be included + "litellm_params": {"search_provider": "tavily", "api_key": "tvly-unique"}, + "search_tool_info": {"description": "Config tool - unique"}, + }, + ] + + # Mock SearchToolRegistry + mock_registry = MagicMock() + mock_registry.get_all_search_tools_from_db = AsyncMock(return_value=db_tools) + with patch( + "litellm.proxy.search_endpoints.search_tool_management.SEARCH_TOOL_REGISTRY", + mock_registry, + ): + # Mock prisma_client + mock_prisma = MagicMock() + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + # Mock proxy_config + mock_proxy_config = MagicMock() + mock_proxy_config.get_config = AsyncMock(return_value={"search_tools": config_tools}) + mock_proxy_config.parse_search_tools = MagicMock(return_value=config_tools) + with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config): + # Mock auth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + test_client = TestClient(app) + response = test_client.get("/search_tools/list") + assert response.status_code == 200 + data = response.json() + assert "search_tools" in data + # Should have 1 DB tool + 1 unique config tool (duplicate filtered out) + assert len(data["search_tools"]) == 2 + + # Verify DB tool is present + db_tool = next( + (t for t in data["search_tools"] if t["search_tool_name"] == "existing-tool"), + None, + ) + assert db_tool is not None + assert db_tool["is_from_config"] is False + assert db_tool["search_tool_id"] == "db-id-1" + # Verify masking of sensitive values in DB tool + assert db_tool["litellm_params"]["api_key"] != "sk-db" + assert "****" in db_tool["litellm_params"]["api_key"] + assert db_tool["litellm_params"]["search_provider"] == "perplexity" + + # Verify unique config tool is present + config_tool = next( + (t for t in data["search_tools"] if t["search_tool_name"] == "unique-config-tool"), + None, + ) + assert config_tool is not None + assert config_tool["is_from_config"] is True + + # Verify duplicate config tool is NOT present + duplicate_tool = next( + ( + t + for t in data["search_tools"] + if t["search_tool_name"] == "existing-tool" and t["is_from_config"] is True + ), + None, + ) + assert duplicate_tool is None + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_list_search_tools_datetime_conversion(monkeypatch): + """ + Test that datetime objects in DB tools are properly converted to ISO format strings. + This tests the new datetime conversion logic using _convert_datetime_to_str. + """ + # Mock DB tools with datetime objects + db_tools = [ + { + "search_tool_id": "test-id-1", + "search_tool_name": "datetime-test-tool", + "litellm_params": {"search_provider": "perplexity", "api_key": "sk-test"}, + "search_tool_info": {"description": "Test tool"}, + "created_at": datetime(2024, 1, 15, 10, 30, 45, 123456), + "updated_at": datetime(2024, 1, 16, 14, 20, 30, 789012), + }, + { + "search_tool_id": "test-id-2", + "search_tool_name": "null-datetime-tool", + "litellm_params": {"search_provider": "tavily", "api_key": "tvly-test"}, + "search_tool_info": None, + "created_at": None, + "updated_at": None, + }, + { + "search_tool_id": "test-id-3", + "search_tool_name": "string-datetime-tool", + "litellm_params": {"search_provider": "perplexity", "api_key": "sk-test"}, + "search_tool_info": {"description": "Already string"}, + "created_at": "2024-01-17T08:15:00", # Already a string + "updated_at": "2024-01-18T09:25:00", # Already a string + }, + ] + + # Mock SearchToolRegistry + mock_registry = MagicMock() + mock_registry.get_all_search_tools_from_db = AsyncMock(return_value=db_tools) + with patch( + "litellm.proxy.search_endpoints.search_tool_management.SEARCH_TOOL_REGISTRY", + mock_registry, + ): + # Mock prisma_client + mock_prisma = MagicMock() + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + # Mock proxy_config + mock_proxy_config = MagicMock() + mock_proxy_config.get_config = AsyncMock(return_value={}) + mock_proxy_config.parse_search_tools = MagicMock(return_value=None) + with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config): + # Mock auth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + test_client = TestClient(app) + response = test_client.get("/search_tools/list") + assert response.status_code == 200 + data = response.json() + assert "search_tools" in data + assert len(data["search_tools"]) == 3 + + # Test datetime conversion for tool 1 + tool1 = next( + (t for t in data["search_tools"] if t["search_tool_name"] == "datetime-test-tool"), + None, + ) + assert tool1 is not None + assert isinstance(tool1["created_at"], str) + assert tool1["created_at"] == "2024-01-15T10:30:45.123456" + assert isinstance(tool1["updated_at"], str) + assert tool1["updated_at"] == "2024-01-16T14:20:30.789012" + # Verify masking of sensitive values + assert tool1["litellm_params"]["api_key"] != "sk-test" + assert "****" in tool1["litellm_params"]["api_key"] + + # Test None handling for tool 2 + tool2 = next( + (t for t in data["search_tools"] if t["search_tool_name"] == "null-datetime-tool"), + None, + ) + assert tool2 is not None + assert tool2["created_at"] is None + assert tool2["updated_at"] is None + # Verify masking of sensitive values + assert tool2["litellm_params"]["api_key"] != "tvly-test" + assert "****" in tool2["litellm_params"]["api_key"] + + # Test string passthrough for tool 3 + tool3 = next( + (t for t in data["search_tools"] if t["search_tool_name"] == "string-datetime-tool"), + None, + ) + assert tool3 is not None + assert tool3["created_at"] == "2024-01-17T08:15:00" + assert tool3["updated_at"] == "2024-01-18T09:25:00" + # Verify masking of sensitive values + assert tool3["litellm_params"]["api_key"] != "sk-test" + assert "****" in tool3["litellm_params"]["api_key"] + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_list_search_tools_config_error_handling(monkeypatch): + """Test that config errors are handled gracefully""" + # Mock DB tools + db_tools = [ + { + "search_tool_id": "test-id-1", + "search_tool_name": "db-tool-1", + "litellm_params": {"search_provider": "perplexity", "api_key": "sk-test"}, + "search_tool_info": {"description": "DB tool"}, + "created_at": datetime(2023, 11, 9, 12, 34, 56), + "updated_at": datetime(2023, 11, 9, 13, 45, 12), + } + ] + + # Mock SearchToolRegistry + mock_registry = MagicMock() + mock_registry.get_all_search_tools_from_db = AsyncMock(return_value=db_tools) + with patch( + "litellm.proxy.search_endpoints.search_tool_management.SEARCH_TOOL_REGISTRY", + mock_registry, + ): + # Mock prisma_client + mock_prisma = MagicMock() + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + # Mock proxy_config to raise an error + mock_proxy_config = MagicMock() + mock_proxy_config.get_config = AsyncMock(side_effect=Exception("Config error")) + with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config): + # Mock auth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + # Should still succeed and return DB tools only + response = client.get("/search_tools/list") + assert response.status_code == 200 + data = response.json() + assert "search_tools" in data + # Should only have DB tools since config failed + assert len(data["search_tools"]) == 1 + assert data["search_tools"][0]["search_tool_name"] == "db-tool-1" + # Verify masking of sensitive values + assert data["search_tools"][0]["litellm_params"]["api_key"] != "sk-test" + assert "****" in data["search_tools"][0]["litellm_params"]["api_key"] + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_list_search_tools_no_prisma_client(monkeypatch): + """Test error handling when prisma_client is None""" + with patch("litellm.proxy.proxy_server.prisma_client", None): + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + test_client = TestClient(app) + response = test_client.get("/search_tools/list") + assert response.status_code == 500 + data = response.json() + assert "Prisma client not initialized" in data["detail"] + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_list_search_tools_db_masking_sensitive_values(monkeypatch): + """ + Test that sensitive values in DB search tools are properly masked. + This tests the new masking logic added for database search tools. + """ + # Mock DB tools with various sensitive fields + db_tools = [ + { + "search_tool_id": "test-id-1", + "search_tool_name": "perplexity-tool", + "litellm_params": { + "search_provider": "perplexity", + "api_key": "pplx-sk-1234567890abcdef", + "api_base": "https://api.perplexity.ai", + }, + "search_tool_info": {"description": "Perplexity tool"}, + "created_at": datetime(2023, 11, 9, 12, 34, 56), + "updated_at": datetime(2023, 11, 9, 13, 45, 12), + }, + { + "search_tool_id": "test-id-2", + "search_tool_name": "tavily-tool", + "litellm_params": { + "search_provider": "tavily", + "api_key": "tvly-secret-key-12345", + "api_base": "https://api.tavily.com", + }, + "search_tool_info": {"description": "Tavily tool"}, + "created_at": datetime(2023, 11, 9, 12, 34, 56), + "updated_at": datetime(2023, 11, 9, 13, 45, 12), + }, + { + "search_tool_id": "test-id-3", + "search_tool_name": "tool-with-token", + "litellm_params": { + "search_provider": "custom", + "access_token": "token-abcdefghijklmnop", + "secret_key": "secret-xyz123", + }, + "search_tool_info": {"description": "Tool with token"}, + "created_at": datetime(2023, 11, 9, 12, 34, 56), + "updated_at": datetime(2023, 11, 9, 13, 45, 12), + }, + { + "search_tool_id": "test-id-4", + "search_tool_name": "tool-with-non-sensitive", + "litellm_params": { + "search_provider": "custom", + "max_results": 10, + "timeout": 30, + }, + "search_tool_info": {"description": "Tool without sensitive fields"}, + "created_at": datetime(2023, 11, 9, 12, 34, 56), + "updated_at": datetime(2023, 11, 9, 13, 45, 12), + }, + ] + + # Mock SearchToolRegistry + mock_registry = MagicMock() + mock_registry.get_all_search_tools_from_db = AsyncMock(return_value=db_tools) + with patch( + "litellm.proxy.search_endpoints.search_tool_management.SEARCH_TOOL_REGISTRY", + mock_registry, + ): + # Mock prisma_client + mock_prisma = MagicMock() + with patch("litellm.proxy.proxy_server.prisma_client", mock_prisma): + # Mock proxy_config + mock_proxy_config = MagicMock() + mock_proxy_config.get_config = AsyncMock(return_value={}) + mock_proxy_config.parse_search_tools = MagicMock(return_value=None) + with patch("litellm.proxy.proxy_server.proxy_config", mock_proxy_config): + # Mock auth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + + try: + test_client = TestClient(app) + response = test_client.get("/search_tools/list") + assert response.status_code == 200 + data = response.json() + assert "search_tools" in data + assert len(data["search_tools"]) == 4 + + # Test tool 1: api_key should be masked + tool1 = next( + (t for t in data["search_tools"] if t["search_tool_name"] == "perplexity-tool"), + None, + ) + assert tool1 is not None + assert tool1["litellm_params"]["api_key"] != "pplx-sk-1234567890abcdef" + assert "****" in tool1["litellm_params"]["api_key"] + assert tool1["litellm_params"]["search_provider"] == "perplexity" + assert tool1["litellm_params"]["api_base"] == "https://api.perplexity.ai" + + # Test tool 2: api_key should be masked + tool2 = next( + (t for t in data["search_tools"] if t["search_tool_name"] == "tavily-tool"), + None, + ) + assert tool2 is not None + assert tool2["litellm_params"]["api_key"] != "tvly-secret-key-12345" + assert "****" in tool2["litellm_params"]["api_key"] + assert tool2["litellm_params"]["search_provider"] == "tavily" + + # Test tool 3: access_token and secret_key should be masked + tool3 = next( + (t for t in data["search_tools"] if t["search_tool_name"] == "tool-with-token"), + None, + ) + assert tool3 is not None + assert tool3["litellm_params"]["access_token"] != "token-abcdefghijklmnop" + assert "****" in tool3["litellm_params"]["access_token"] + assert tool3["litellm_params"]["secret_key"] != "secret-xyz123" + assert "****" in tool3["litellm_params"]["secret_key"] + + # Test tool 4: non-sensitive fields should remain unmasked + tool4 = next( + (t for t in data["search_tools"] if t["search_tool_name"] == "tool-with-non-sensitive"), + None, + ) + assert tool4 is not None + assert tool4["litellm_params"]["max_results"] == 10 + assert tool4["litellm_params"]["timeout"] == 30 + assert tool4["litellm_params"]["search_provider"] == "custom" + finally: + app.dependency_overrides.pop(user_api_key_auth, None) diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index dc436bac087..919af96f760 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -1133,6 +1133,24 @@ def test_update_internal_user_params_ignores_other_nones(): assert non_default_values["max_budget"] == 100.0 +def test_update_internal_user_params_keeps_original_max_budget_when_not_provided(): + """ + Test that _update_internal_user_params does not include max_budget + when it's not provided in the request (should keep original value). + """ + # Create test data without max_budget + data_json = {"user_id": "test_user", "user_alias": "test_alias"} + data = UpdateUserRequest(user_id="test_user", user_alias="test_alias") + + # Call the function + non_default_values = _update_internal_user_params(data_json=data_json, data=data) + + # Assertions: max_budget should NOT be in non_default_values + assert "max_budget" not in non_default_values + assert "user_id" in non_default_values + assert "user_alias" in non_default_values + + def test_generate_request_base_validator(): """ Test that GenerateRequestBase validator converts empty string to None for max_budget diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index e90fb277eed..5720ff948a6 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -804,6 +804,108 @@ async def test_key_update_object_permissions_missing_permission_record(monkeypat mock_prisma_client.db.litellm_objectpermissiontable.upsert.assert_called_once() +@pytest.mark.asyncio +async def test_key_info_returns_object_permission(monkeypatch): + """ + Test that /key/info correctly returns the object_permission relation. + + This test verifies that when calling /key/info for a key with object_permission_id, + the response includes the full object_permission object with fields like + mcp_access_groups, mcp_servers, vector_stores, agents, etc. + + Regression test for bug where object_permission_id was returned but not the + related object_permission object. + """ + from unittest.mock import AsyncMock, MagicMock + + import pytest + + from litellm.proxy._types import LiteLLM_VerificationToken + from litellm.proxy.management_endpoints.key_management_endpoints import info_key_fn + + # Mock prisma client + mock_prisma_client = AsyncMock() + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma_client) + + # Mock key with object_permission_id + test_key_token = "hashed_test_token_123" + test_object_permission_id = "objperm_info_test_123" + + mock_key_info = MagicMock(spec=LiteLLM_VerificationToken) + mock_key_info.token = test_key_token + mock_key_info.object_permission_id = test_object_permission_id + mock_key_info.user_id = "user123" + mock_key_info.team_id = None + mock_key_info.litellm_budget_table = None + + # Mock the dict/model_dump methods + mock_key_info.model_dump.return_value = { + "token": test_key_token, + "object_permission_id": test_object_permission_id, + "user_id": "user123", + "team_id": None, + "litellm_budget_table": None, + } + mock_key_info.dict.return_value = mock_key_info.model_dump.return_value + + # Mock find_unique for the key lookup + mock_prisma_client.db.litellm_verificationtoken.find_unique = AsyncMock( + return_value=mock_key_info + ) + + # Mock object permission record + mock_object_permission = MagicMock() + mock_object_permission.model_dump.return_value = { + "object_permission_id": test_object_permission_id, + "mcp_access_groups": ["test_group_1", "test_group_2"], + "mcp_servers": ["server_1"], + "vector_stores": ["vs_1", "vs_2"], + "agents": ["agent_1"], + } + mock_object_permission.dict.return_value = mock_object_permission.model_dump.return_value + + # Mock find_unique for object permission lookup + mock_prisma_client.db.litellm_objectpermissiontable.find_unique = AsyncMock( + return_value=mock_object_permission + ) + + # Create user API key dict + user_api_key_dict = UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, + api_key="sk-test-key-456", + ) + + # Call info_key_fn + result = await info_key_fn( + key="sk-test-key-456", + user_api_key_dict=user_api_key_dict, + ) + + # Assertions + assert "info" in result + assert "object_permission_id" in result["info"] + assert result["info"]["object_permission_id"] == test_object_permission_id + + # CRITICAL: Verify that object_permission object is included in response + assert "object_permission" in result["info"], ( + "object_permission field missing from /key/info response. " + "Expected full object_permission object to be attached." + ) + + # Verify object_permission contains the expected fields + obj_perm = result["info"]["object_permission"] + assert obj_perm["object_permission_id"] == test_object_permission_id + assert obj_perm["mcp_access_groups"] == ["test_group_1", "test_group_2"] + assert obj_perm["mcp_servers"] == ["server_1"] + assert obj_perm["vector_stores"] == ["vs_1", "vs_2"] + assert obj_perm["agents"] == ["agent_1"] + + # Verify the object permission was actually queried from database + mock_prisma_client.db.litellm_objectpermissiontable.find_unique.assert_called_once_with( + where={"object_permission_id": test_object_permission_id} + ) + + def test_get_new_token_with_valid_key(): """Test get_new_token function when provided with a valid key that starts with 'sk-'""" from litellm.proxy._types import RegenerateKeyRequest diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py new file mode 100644 index 00000000000..1288a9b2c9f --- /dev/null +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -0,0 +1,105 @@ +""" +Test A2A model routing in proxy. + +Maps to: litellm/proxy/agent_endpoints/a2a_routing.py +""" +import os +import sys + +sys.path.insert(0, os.path.abspath("../../..")) + +from unittest.mock import AsyncMock, Mock, patch + +import pytest + +from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request +from litellm.proxy.route_llm_request import route_request + + +@pytest.mark.asyncio +async def test_route_a2a_model_bypasses_router(): + """Test that a2a/ prefixed models bypass router and go directly to litellm with api_base""" + + # Mock data for chat completion with a2a model + data = { + "model": "a2a/test-agent", + "messages": [{"role": "user", "content": "Hello"}], + } + + # Mock router that doesn't have the a2a model + mock_router = Mock() + mock_router.model_names = ["gpt-4", "gpt-3.5-turbo"] + mock_router.deployment_names = [] + mock_router.has_model_id = Mock(return_value=False) + mock_router.model_group_alias = None + mock_router.router_general_settings = Mock(pass_through_all_models=False) + mock_router.default_deployment = None + mock_router.pattern_router = Mock(patterns=[]) + mock_router.map_team_model = Mock(return_value=None) + + # Mock agent in registry + from litellm.types.agents import AgentResponse + + mock_agent = AgentResponse( + agent_id="test-agent-id", + agent_name="test-agent", + agent_card_params={"url": "http://agent.example.com"}, + litellm_params=None, + ) + + mock_registry = Mock() + mock_registry.get_agent_by_name = Mock(return_value=mock_agent) + + # Mock litellm.acompletion to verify it's called + mock_acompletion = AsyncMock(return_value={"id": "test-response"}) + + with patch("litellm.acompletion", mock_acompletion): + with patch( + "litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", + mock_registry, + ): + result = await route_request( + data=data, + llm_router=mock_router, + user_model=None, + route_type="acompletion", + ) + + # Verify litellm.acompletion was called with api_base injected + mock_acompletion.assert_called_once() + call_kwargs = mock_acompletion.call_args.kwargs + assert call_kwargs["model"] == "a2a/test-agent" + assert call_kwargs["api_base"] == "http://agent.example.com" + + +@pytest.mark.asyncio +async def test_route_non_a2a_model_raises_error_if_not_in_router(): + """Test that non-a2a models that aren't in router raise an error""" + + # Mock data for chat completion with model not in router + data = { + "model": "unknown-model", + "messages": [{"role": "user", "content": "Hello"}], + } + + # Mock router without the model + mock_router = Mock() + mock_router.model_names = ["gpt-4", "gpt-3.5-turbo"] + mock_router.deployment_names = [] + mock_router.has_model_id = Mock(return_value=False) + mock_router.model_group_alias = None + mock_router.router_general_settings = Mock(pass_through_all_models=False) + mock_router.default_deployment = None + mock_router.pattern_router = Mock(patterns=[]) + mock_router.map_team_model = Mock(return_value=None) + + # Should raise ProxyModelNotFoundError + from litellm.proxy.route_llm_request import ProxyModelNotFoundError + + with pytest.raises(ProxyModelNotFoundError): + await route_request( + data=data, + llm_router=mock_router, + user_model=None, + route_type="acompletion", + ) diff --git a/tests/test_litellm/test_a2a_registry_lookup.py b/tests/test_litellm/test_a2a_registry_lookup.py new file mode 100644 index 00000000000..9938f10a43f --- /dev/null +++ b/tests/test_litellm/test_a2a_registry_lookup.py @@ -0,0 +1,73 @@ +""" +Test A2A provider registry lookup functionality. + +Maps to: litellm/llms/a2a/chat/transformation.py +""" +import os +import sys + +sys.path.insert(0, os.path.abspath("../..")) + +import pytest + +import litellm +from litellm.llms.a2a.chat.transformation import A2AConfig + + +def test_resolve_agent_config_from_registry_static_method(): + """Test the static helper method for registry resolution""" + + # Test 1: No agent name in model + api_base, api_key, headers = A2AConfig.resolve_agent_config_from_registry( + model="a2a", + api_base="http://test.com", + api_key=None, + headers=None, + optional_params={} + ) + assert api_base == "http://test.com" + + # Test 2: All params provided - should not lookup registry + api_base, api_key, headers = A2AConfig.resolve_agent_config_from_registry( + model="a2a/test-agent", + api_base="http://explicit.com", + api_key="explicit-key", + headers={"X-Test": "value"}, + optional_params={} + ) + assert api_base == "http://explicit.com" + assert api_key == "explicit-key" + + +def test_a2a_registry_integration(): + """Test registry lookup in proxy context""" + + try: + from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry + from litellm.types.agents import AgentResponse + + # Create test agent + test_agent = AgentResponse( + agent_id="test-id", + agent_name="test-agent", + agent_card_params={"url": "http://registry-url.example.com:9999"}, + litellm_params={"api_key": "registry-key"}, + ) + + # Register and test + original_agents = global_agent_registry.agent_list.copy() + global_agent_registry.register_agent(test_agent) + + try: + litellm.completion( + model="a2a/test-agent", + messages=[{"role": "user", "content": "Hello"}] + ) + except Exception as e: + # Should use registry URL (connection error expected) + assert "registry-url.example.com" in str(e) or "APIConnectionError" in str(type(e).__name__) + finally: + global_agent_registry.agent_list = original_agents + + except ImportError: + pytest.skip("Registry not available (not in proxy context)") diff --git a/tests/test_team.py b/tests/test_team.py index d67c5e670f4..275181590c5 100644 --- a/tests/test_team.py +++ b/tests/test_team.py @@ -532,6 +532,7 @@ async def test_team_update_sc_2(): or k == "object_permission" or k == "litellm_model_table" or k == "policies" + or k == "allow_team_guardrail_config" ): pass else: diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useMCPSemanticFilterSettings.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useMCPSemanticFilterSettings.ts new file mode 100644 index 00000000000..e91f5aa670b --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useMCPSemanticFilterSettings.ts @@ -0,0 +1,19 @@ +import { getMCPSemanticFilterSettings } from "@/components/networking"; +import { useQuery } from "@tanstack/react-query"; +import { createQueryKeys } from "../common/queryKeysFactory"; +import useAuthorized from "../useAuthorized"; + +const mcpSemanticFilterSettingsKeys = createQueryKeys( + "mcpSemanticFilterSettings" +); + +export const useMCPSemanticFilterSettings = () => { + const { accessToken } = useAuthorized(); + return useQuery>({ + queryKey: mcpSemanticFilterSettingsKeys.list({}), + queryFn: async () => await getMCPSemanticFilterSettings(accessToken), + enabled: !!accessToken, + staleTime: 60 * 60 * 1000, // 1 hour + gcTime: 60 * 60 * 1000, // 1 hour + }); +}; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useUpdateMCPSemanticFilterSettings.ts b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useUpdateMCPSemanticFilterSettings.ts new file mode 100644 index 00000000000..2062b4f4c29 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/hooks/mcpSemanticFilterSettings/useUpdateMCPSemanticFilterSettings.ts @@ -0,0 +1,25 @@ +import { updateMCPSemanticFilterSettings } from "@/components/networking"; +import { useMutation, useQueryClient } from "@tanstack/react-query"; +import { createQueryKeys } from "../common/queryKeysFactory"; + +const mcpSemanticFilterSettingsKeys = createQueryKeys( + "mcpSemanticFilterSettings" +); + +export const useUpdateMCPSemanticFilterSettings = (accessToken: string) => { + const queryClient = useQueryClient(); + + return useMutation({ + mutationFn: async (settings: Record) => { + if (!accessToken) { + throw new Error("Access token is required"); + } + return updateMCPSemanticFilterSettings(accessToken, settings); + }, + onSuccess: () => { + queryClient.invalidateQueries({ + queryKey: mcpSemanticFilterSettingsKeys.all, + }); + }, + }); +}; diff --git a/ui/litellm-dashboard/src/app/page.tsx b/ui/litellm-dashboard/src/app/page.tsx index 5f94db7e9f4..8e887d80fb3 100644 --- a/ui/litellm-dashboard/src/app/page.tsx +++ b/ui/litellm-dashboard/src/app/page.tsx @@ -27,7 +27,7 @@ import Organizations, { fetchOrganizations } from "@/components/organizations"; import PassThroughSettings from "@/components/pass_through_settings"; import PromptsPanel from "@/components/prompts"; import PublicModelHub from "@/components/public_model_hub"; -import { SearchTools } from "@/components/search_tools"; +import { SearchTools } from "@/components/SearchTools"; import Settings from "@/components/settings"; import { SurveyPrompt, SurveyModal, ClaudeCodePrompt, ClaudeCodeModal } from "@/components/survey"; import TagManagement from "@/components/tag_management"; diff --git a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx index 23bfb7d219f..4843713e5a6 100644 --- a/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx +++ b/ui/litellm-dashboard/src/components/AIHub/ModelHubTable.tsx @@ -483,7 +483,7 @@ const ModelHubTable: React.FC = ({ accessToken, publicPage, = ({ accessToken, publicPage, = ({ {canCreateOrManageTeams(userRole, userID, organizations) && ( = ({ <> = ({ {/* Clear Confirmation Modal */} setIsClearConfirmModalVisible(false)} okText="Yes, Clear" @@ -536,7 +536,7 @@ const SSOModals: React.FC = ({ = ({ }, search_tool_info: formValues.description ? { - description: formValues.description, - } + description: formValues.description, + } : undefined, }; @@ -130,7 +130,7 @@ const CreateSearchTool: React.FC = ({ try { // Validate required fields for testing await form.validateFields(["search_provider", "api_key"]); - + setIsTestingConnection(true); // Generate a new test ID (using timestamp for uniqueness) setConnectionTestId(`test-${Date.now()}`); @@ -225,8 +225,8 @@ const CreateSearchTool: React.FC = ({ optionLabelProp="label" > {availableProviders.map((provider) => ( - void, + onEdit: (searchToolId: string) => void, + onDelete: (searchToolId: string) => void, + availableProviders: Array<{ provider_name: string; ui_friendly_name: string }>, +): ColumnsType => [ + { + title: "Search Tool ID", + dataIndex: "search_tool_id", + key: "search_tool_id", + render: (_, tool) => { + const isFromConfig = tool.is_from_config; + + if (isFromConfig) { + return -; + } + + return ( + + ); + }, + }, + { + title: "Name", + dataIndex: "search_tool_name", + key: "search_tool_name", + render: (name: string) => {name}, + }, + { + title: "Provider", + key: "provider", + render: (_, tool) => { + const provider = tool.litellm_params.search_provider; + const providerInfo = availableProviders.find((p) => p.provider_name === provider); + const displayName = providerInfo?.ui_friendly_name || provider; + + return {displayName}; + }, + }, + { + title: "Created At", + dataIndex: "created_at", + key: "created_at", + render: (_, tool) => { + return {tool.created_at ? new Date(tool.created_at).toLocaleDateString() : "-"}; + }, + }, + { + title: "Updated At", + dataIndex: "updated_at", + key: "updated_at", + render: (_, tool) => { + return {tool.updated_at ? new Date(tool.updated_at).toLocaleDateString() : "-"}; + }, + }, + { + title: "Source", + key: "source", + render: (_, tool) => { + const isFromConfig = tool.is_from_config ?? false; + + return ( + + {isFromConfig ? "Config" : "DB"} + + ); + }, + }, + { + title: "Actions", + key: "actions", + render: (_, tool) => { + const toolId = tool.search_tool_id; + const isFromConfig = tool.is_from_config ?? false; + + return ( +
+ { + if (toolId && !isFromConfig) { + onEdit(toolId); + } + }} + /> + { + if (toolId && !isFromConfig) { + onDelete(toolId); + } + }} + /> +
+ ); + }, + }, + ]; diff --git a/ui/litellm-dashboard/src/components/search_tools/search_tool_tester.tsx b/ui/litellm-dashboard/src/components/SearchTools/SearchToolTester.tsx similarity index 100% rename from ui/litellm-dashboard/src/components/search_tools/search_tool_tester.tsx rename to ui/litellm-dashboard/src/components/SearchTools/SearchToolTester.tsx diff --git a/ui/litellm-dashboard/src/components/SearchTools/SearchToolView.test.tsx b/ui/litellm-dashboard/src/components/SearchTools/SearchToolView.test.tsx new file mode 100644 index 00000000000..bb04a04a992 --- /dev/null +++ b/ui/litellm-dashboard/src/components/SearchTools/SearchToolView.test.tsx @@ -0,0 +1,278 @@ +import { render, screen, waitFor, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import { SearchToolView } from "./SearchToolView"; +import { AvailableSearchProvider, SearchTool } from "./types"; + +vi.mock("@/utils/dataUtils", () => ({ + copyToClipboard: vi.fn().mockResolvedValue(true), +})); + +vi.mock("./SearchToolTester", () => ({ + SearchToolTester: ({ searchToolName, accessToken }: { searchToolName: string; accessToken: string }) => ( +
+ Search Tool Tester for {searchToolName} + Access Token: {accessToken} +
+ ), +})); + +describe("SearchToolView", () => { + const mockSearchTool: SearchTool = { + search_tool_id: "test-tool-id-123", + search_tool_name: "Test Search Tool", + litellm_params: { + search_provider: "perplexity", + api_key: "sk-test-key", + }, + search_tool_info: { + description: "Test description", + }, + created_at: "2024-01-15T10:30:00Z", + }; + + const mockAvailableProviders: AvailableSearchProvider[] = [ + { + provider_name: "perplexity", + ui_friendly_name: "Perplexity AI", + }, + { + provider_name: "tavily", + ui_friendly_name: "Tavily Search", + }, + ]; + + const defaultProps = { + searchTool: mockSearchTool, + onBack: vi.fn(), + isEditing: false, + accessToken: "test-token", + availableProviders: mockAvailableProviders, + }; + + beforeEach(async () => { + vi.clearAllMocks(); + const { copyToClipboard } = await import("@/utils/dataUtils"); + vi.mocked(copyToClipboard).mockResolvedValue(true); + }); + + it("should render", () => { + render(); + expect(screen.getByText("Test Search Tool")).toBeInTheDocument(); + }); + + it("should display search tool name", () => { + render(); + expect(screen.getByText("Test Search Tool")).toBeInTheDocument(); + }); + + it("should display search tool ID", () => { + render(); + expect(screen.getByText("test-tool-id-123")).toBeInTheDocument(); + }); + + it("should display provider name using UI-friendly name when available", () => { + render(); + expect(screen.getByText("Perplexity AI")).toBeInTheDocument(); + }); + + it("should display provider name using provider_name when UI-friendly name is not available", () => { + const searchToolWithoutProvider: SearchTool = { + ...mockSearchTool, + litellm_params: { + search_provider: "unknown-provider", + }, + }; + + render( + , + ); + expect(screen.getByText("unknown-provider")).toBeInTheDocument(); + }); + + it("should display masked API key when API key is set", () => { + render(); + expect(screen.getByText("****")).toBeInTheDocument(); + }); + + it("should display 'Not set' when API key is not set", () => { + const searchToolWithoutApiKey: SearchTool = { + ...mockSearchTool, + litellm_params: { + search_provider: "perplexity", + }, + }; + + render( + , + ); + expect(screen.getByText("Not set")).toBeInTheDocument(); + }); + + it("should display formatted created_at date", () => { + render(); + const dateText = screen.getByText(/2024-01-15/); + expect(dateText).toBeInTheDocument(); + }); + + it("should display 'Unknown' when created_at is not set", () => { + const searchToolWithoutDate: SearchTool = { + ...mockSearchTool, + created_at: undefined, + }; + + render( + , + ); + expect(screen.getByText("Unknown")).toBeInTheDocument(); + }); + + it("should display description when search_tool_info.description is provided", () => { + render(); + expect(screen.getByText("Test description")).toBeInTheDocument(); + }); + + it("should not display description card when search_tool_info.description is not provided", () => { + const searchToolWithoutDescription: SearchTool = { + ...mockSearchTool, + search_tool_info: {}, + }; + + render( + , + ); + expect(screen.queryByText("Description")).not.toBeInTheDocument(); + }); + + it("should call onBack when back button is clicked", async () => { + const user = userEvent.setup({ delay: null }); + const onBack = vi.fn(); + render(); + + const backButton = screen.getByRole("button", { name: /back to all search tools/i }); + await user.click(backButton); + + expect(onBack).toHaveBeenCalledTimes(1); + }); + + it("should copy search tool name to clipboard when copy button is clicked", async () => { + const user = userEvent.setup({ delay: null }); + const { copyToClipboard } = await import("@/utils/dataUtils"); + render(); + + const toolNameContainer = screen.getByText("Test Search Tool").closest("div"); + expect(toolNameContainer).toBeInTheDocument(); + + const copyButtons = within(toolNameContainer!).getAllByRole("button"); + const nameCopyButton = copyButtons.find((button) => { + return button.querySelector("svg") !== null; + }); + + expect(nameCopyButton).toBeInTheDocument(); + await user.click(nameCopyButton!); + + await waitFor(() => { + expect(copyToClipboard).toHaveBeenCalledWith("Test Search Tool"); + }); + }); + + it("should copy search tool ID to clipboard when copy button is clicked", async () => { + const user = userEvent.setup({ delay: null }); + const { copyToClipboard } = await import("@/utils/dataUtils"); + render(); + + const toolIdContainer = screen.getByText("test-tool-id-123").closest("div"); + expect(toolIdContainer).toBeInTheDocument(); + + const copyButtons = within(toolIdContainer!).getAllByRole("button"); + const idCopyButton = copyButtons.find((button) => { + return button.querySelector("svg") !== null; + }); + + expect(idCopyButton).toBeInTheDocument(); + await user.click(idCopyButton!); + + await waitFor(() => { + expect(copyToClipboard).toHaveBeenCalledWith("test-tool-id-123"); + }); + }); + + it("should show check icon after copying search tool name", async () => { + const user = userEvent.setup({ delay: null }); + const { copyToClipboard } = await import("@/utils/dataUtils"); + vi.mocked(copyToClipboard).mockResolvedValue(true); + + render(); + + const toolNameContainer = screen.getByText("Test Search Tool").closest("div"); + const copyButtons = within(toolNameContainer!).getAllByRole("button"); + const nameCopyButton = copyButtons.find((button) => { + return button.querySelector("svg") !== null; + }); + + expect(nameCopyButton).toBeInTheDocument(); + + const initialSvg = nameCopyButton!.querySelector("svg"); + expect(initialSvg).toBeInTheDocument(); + + await user.click(nameCopyButton!); + + await waitFor(() => { + const updatedSvg = nameCopyButton!.querySelector("svg"); + expect(updatedSvg).toBeInTheDocument(); + expect(nameCopyButton).toHaveClass("text-green-600"); + }); + }); + + + it("should not show check icon when copy fails", async () => { + const user = userEvent.setup({ delay: null }); + const { copyToClipboard } = await import("@/utils/dataUtils"); + vi.mocked(copyToClipboard).mockResolvedValue(false); + + render(); + + const toolNameContainer = screen.getByText("Test Search Tool").closest("div"); + const copyButtons = within(toolNameContainer!).getAllByRole("button"); + const nameCopyButton = copyButtons.find((button) => { + return button.querySelector("svg") !== null; + }); + + expect(nameCopyButton).toBeInTheDocument(); + await user.click(nameCopyButton!); + + await waitFor(() => { + expect(copyToClipboard).toHaveBeenCalledWith("Test Search Tool"); + }, { timeout: 3000 }); + + expect(nameCopyButton).not.toHaveClass("text-green-600"); + }); + + it("should render SearchToolTester when accessToken is provided", () => { + render(); + expect(screen.getByTestId("search-tool-tester")).toBeInTheDocument(); + expect(screen.getByText(/Search Tool Tester for Test Search Tool/)).toBeInTheDocument(); + }); + + it("should not render SearchToolTester when accessToken is null", () => { + render(); + expect(screen.queryByTestId("search-tool-tester")).not.toBeInTheDocument(); + }); + + it("should pass correct props to SearchToolTester", () => { + render(); + expect(screen.getByText("Access Token: test-token")).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/search_tools/search_tool_view.tsx b/ui/litellm-dashboard/src/components/SearchTools/SearchToolView.tsx similarity index 86% rename from ui/litellm-dashboard/src/components/search_tools/search_tool_view.tsx rename to ui/litellm-dashboard/src/components/SearchTools/SearchToolView.tsx index a5cad4e8370..ad88acd127f 100644 --- a/ui/litellm-dashboard/src/components/search_tools/search_tool_view.tsx +++ b/ui/litellm-dashboard/src/components/SearchTools/SearchToolView.tsx @@ -1,11 +1,11 @@ -import React, { useState } from "react"; -import { ArrowLeftIcon } from "@heroicons/react/outline"; -import { Title, Card, Button, Text, Grid } from "@tremor/react"; -import { SearchTool, AvailableSearchProvider } from "./types"; import { copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils"; -import { CheckIcon, CopyIcon } from "lucide-react"; +import { ArrowLeftIcon } from "@heroicons/react/outline"; +import { Button, Card, Grid, Text, Title } from "@tremor/react"; import { Button as AntdButton } from "antd"; -import { SearchToolTester } from "./search_tool_tester"; +import { CheckIcon, CopyIcon } from "lucide-react"; +import React, { useState } from "react"; +import { SearchToolTester } from "./SearchToolTester"; +import { AvailableSearchProvider, SearchTool } from "./types"; interface SearchToolViewProps { searchTool: SearchTool; @@ -53,11 +53,10 @@ export const SearchToolView: React.FC = ({ size="small" icon={copiedStates["search-tool-name"] ? : } onClick={() => copyToClipboard(searchTool.search_tool_name, "search-tool-name")} - className={`left-2 z-10 transition-all duration-200 ${ - copiedStates["search-tool-name"] - ? "text-green-600 bg-green-50 border-green-200" - : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" - }`} + className={`left-2 z-10 transition-all duration-200 ${copiedStates["search-tool-name"] + ? "text-green-600 bg-green-50 border-green-200" + : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" + }`} />
@@ -67,11 +66,10 @@ export const SearchToolView: React.FC = ({ size="small" icon={copiedStates["search-tool-id"] ? : } onClick={() => copyToClipboard(searchTool.search_tool_id, "search-tool-id")} - className={`left-2 z-10 transition-all duration-200 ${ - copiedStates["search-tool-id"] - ? "text-green-600 bg-green-50 border-green-200" - : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" - }`} + className={`left-2 z-10 transition-all duration-200 ${copiedStates["search-tool-id"] + ? "text-green-600 bg-green-50 border-green-200" + : "text-gray-500 hover:text-gray-700 hover:bg-gray-100" + }`} />
diff --git a/ui/litellm-dashboard/src/components/SearchTools/SearchTools.test.tsx b/ui/litellm-dashboard/src/components/SearchTools/SearchTools.test.tsx new file mode 100644 index 00000000000..f1f7bf8ab42 --- /dev/null +++ b/ui/litellm-dashboard/src/components/SearchTools/SearchTools.test.tsx @@ -0,0 +1,234 @@ +import * as roles from "@/utils/roles"; +import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; +import { render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { beforeEach, describe, expect, it, vi } from "vitest"; +import * as networking from "../networking"; +import SearchTools from "./SearchTools"; +import { AvailableSearchProvider, SearchTool } from "./types"; + +vi.mock("../networking", () => ({ + fetchSearchTools: vi.fn(), + updateSearchTool: vi.fn(), + deleteSearchTool: vi.fn(), + fetchAvailableSearchProviders: vi.fn(), +})); + +vi.mock("@/utils/roles", () => ({ + isAdminRole: vi.fn(), +})); + +vi.mock("./SearchToolView", () => ({ + SearchToolView: ({ searchTool, onBack }: { searchTool: SearchTool; onBack: () => void }) => ( +
+
Search Tool View: {searchTool.search_tool_name}
+ +
+ ), +})); + +vi.mock("./CreateSearchTools", () => ({ + default: ({ + isModalVisible, + setModalVisible, + }: { + isModalVisible: boolean; + setModalVisible: (visible: boolean) => void; + }) => + isModalVisible ? ( +
+ +
+ ) : null, +})); + +vi.mock("../common_components/DeleteResourceModal", () => ({ + default: ({ + isOpen, + onOk, + onCancel, + }: { + isOpen: boolean; + onOk: () => void; + onCancel: () => void; + }) => + isOpen ? ( +
+ + +
+ ) : null, +})); + +const mockSearchTools: SearchTool[] = [ + { + search_tool_id: "tool-1", + search_tool_name: "Perplexity Search", + litellm_params: { + search_provider: "perplexity", + api_key: "sk-test-key", + }, + search_tool_info: { + description: "Test description", + }, + created_at: "2024-01-15T10:30:00Z", + }, + { + search_tool_id: "tool-2", + search_tool_name: "Tavily Search", + litellm_params: { + search_provider: "tavily", + }, + created_at: "2024-01-16T10:30:00Z", + }, +]; + +const mockAvailableProviders: AvailableSearchProvider[] = [ + { + provider_name: "perplexity", + ui_friendly_name: "Perplexity AI", + }, + { + provider_name: "tavily", + ui_friendly_name: "Tavily Search", + }, +]; + +const createWrapper = () => { + const queryClient = new QueryClient({ + defaultOptions: { + queries: { + retry: false, + }, + }, + }); + return ({ children }: { children: React.ReactNode }) => ( + {children} + ); +}; + +describe("SearchTools", () => { + const defaultProps = { + accessToken: "test-token", + userRole: "Admin", + userID: "user-1", + }; + + beforeEach(() => { + vi.clearAllMocks(); + vi.mocked(networking.fetchSearchTools).mockResolvedValue({ search_tools: mockSearchTools }); + vi.mocked(networking.fetchAvailableSearchProviders).mockResolvedValue({ providers: mockAvailableProviders }); + vi.mocked(roles.isAdminRole).mockReturnValue(true); + }); + + it("should render", async () => { + render(, { wrapper: createWrapper() }); + await waitFor(() => { + expect(screen.getByText("Search Tools")).toBeInTheDocument(); + }); + }); + + it("should display missing authentication parameters message when accessToken is missing", () => { + render(, { wrapper: createWrapper() }); + expect(screen.getByText("Missing required authentication parameters.")).toBeInTheDocument(); + }); + + it("should display missing authentication parameters message when userRole is missing", () => { + render(, { wrapper: createWrapper() }); + expect(screen.getByText("Missing required authentication parameters.")).toBeInTheDocument(); + }); + + it("should display missing authentication parameters message when userID is missing", () => { + render(, { wrapper: createWrapper() }); + expect(screen.getByText("Missing required authentication parameters.")).toBeInTheDocument(); + }); + + it("should display search tools table with tools", async () => { + render(, { wrapper: createWrapper() }); + await waitFor(() => { + expect(screen.getByText("Perplexity Search")).toBeInTheDocument(); + }); + expect(screen.getAllByText("Tavily Search").length).toBeGreaterThan(0); + }); + + it("should display empty state when no search tools are available", async () => { + vi.mocked(networking.fetchSearchTools).mockResolvedValue({ search_tools: [] }); + + render(, { wrapper: createWrapper() }); + await waitFor(() => { + expect(screen.getByText("No search tools configured")).toBeInTheDocument(); + }); + }); + + it("should show Add New Search Tool button when user is admin", async () => { + render(, { wrapper: createWrapper() }); + await waitFor(() => { + expect(screen.getByRole("button", { name: /add new search tool/i })).toBeInTheDocument(); + }); + }); + + it("should not show Add New Search Tool button when user is not admin", async () => { + vi.mocked(roles.isAdminRole).mockReturnValue(false); + + render(, { wrapper: createWrapper() }); + await waitFor(() => { + expect(screen.getByText("Search Tools")).toBeInTheDocument(); + }); + expect(screen.queryByRole("button", { name: /add new search tool/i })).not.toBeInTheDocument(); + }); + + it("should open create modal when Add New Search Tool button is clicked", async () => { + const user = userEvent.setup({ delay: null }); + render(, { wrapper: createWrapper() }); + + await waitFor(() => { + expect(screen.getByRole("button", { name: /add new search tool/i })).toBeInTheDocument(); + }); + + const addButton = screen.getByRole("button", { name: /add new search tool/i }); + await user.click(addButton); + + expect(screen.getByTestId("create-search-tool-modal")).toBeInTheDocument(); + }); + + it("should navigate to tool view when tool ID is clicked", async () => { + const user = userEvent.setup({ delay: null }); + render(, { wrapper: createWrapper() }); + + await waitFor(() => { + expect(screen.getByText("Perplexity Search")).toBeInTheDocument(); + }); + + const toolIdButton = screen.getByRole("button", { name: /tool-1/i }); + await user.click(toolIdButton); + + await waitFor(() => { + expect(screen.getByTestId("search-tool-view")).toBeInTheDocument(); + }); + expect(screen.getByText(/Search Tool View: Perplexity Search/i)).toBeInTheDocument(); + }); + + it("should navigate back from tool view to table", async () => { + const user = userEvent.setup({ delay: null }); + render(, { wrapper: createWrapper() }); + + await waitFor(() => { + expect(screen.getByText("Perplexity Search")).toBeInTheDocument(); + }); + + const toolIdButton = screen.getByRole("button", { name: /tool-1/i }); + await user.click(toolIdButton); + + await waitFor(() => { + expect(screen.getByTestId("search-tool-view")).toBeInTheDocument(); + }); + + const backButton = screen.getByRole("button", { name: /back/i }); + await user.click(backButton); + + await waitFor(() => { + expect(screen.queryByTestId("search-tool-view")).not.toBeInTheDocument(); + expect(screen.getByText("Perplexity Search")).toBeInTheDocument(); + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/search_tools/search_tools.tsx b/ui/litellm-dashboard/src/components/SearchTools/SearchTools.tsx similarity index 78% rename from ui/litellm-dashboard/src/components/search_tools/search_tools.tsx rename to ui/litellm-dashboard/src/components/SearchTools/SearchTools.tsx index 2fbdd4d27c6..dd2033fc18f 100644 --- a/ui/litellm-dashboard/src/components/search_tools/search_tools.tsx +++ b/ui/litellm-dashboard/src/components/SearchTools/SearchTools.tsx @@ -1,20 +1,21 @@ -import React, { useState } from "react"; +import { isAdminRole } from "@/utils/roles"; +import { LoadingOutlined } from "@ant-design/icons"; import { useQuery } from "@tanstack/react-query"; -import { Modal, Form, Input, Select } from "antd"; -import { Button, Title, Text, Grid, Col } from "@tremor/react"; -import { DataTable } from "../view_logs/table"; -import { searchToolColumns } from "./search_tool_columns"; +import { Button, Text, Title } from "@tremor/react"; +import { Form, Input, Modal, Select, Spin, Table } from "antd"; +import React, { useState } from "react"; +import DeleteResourceModal from "../common_components/DeleteResourceModal"; +import NotificationsManager from "../molecules/notifications_manager"; import { - fetchSearchTools, - updateSearchTool, deleteSearchTool, fetchAvailableSearchProviders, + fetchSearchTools, + updateSearchTool, } from "../networking"; -import { SearchTool, AvailableSearchProvider } from "./types"; -import { isAdminRole } from "@/utils/roles"; -import NotificationsManager from "../molecules/notifications_manager"; -import { SearchToolView } from "./search_tool_view"; -import CreateSearchTool from "./create_search_tool"; +import CreateSearchTool from "./CreateSearchTools"; +import { searchToolColumns } from "./SearchToolColumn"; +import { SearchToolView } from "./SearchToolView"; +import { AvailableSearchProvider, SearchTool } from "./types"; interface SearchToolsProps { accessToken: string | null; @@ -22,24 +23,6 @@ interface SearchToolsProps { userID: string | null; } -const DeleteModal: React.FC<{ - isModalOpen: boolean; - title: string; - confirmDelete: () => void; - cancelDelete: () => void; -}> = ({ isModalOpen, title, confirmDelete, cancelDelete }) => { - if (!isModalOpen) return null; - return ( - - - {title} - -

Are you sure you want to delete this search tool?

- -
-
- ); -}; const SearchTools: React.FC = ({ accessToken, userRole, userID }) => { const { @@ -72,6 +55,7 @@ const SearchTools: React.FC = ({ accessToken, userRole, userID // State const [toolIdToDelete, setToolToDelete] = useState(null); const [isDeleteModalOpen, setIsDeleteModalOpen] = useState(false); + const [isDeleting, setIsDeleting] = useState(false); const [selectedToolId, setSelectedToolId] = useState(null); const [editTool, setEditTool] = useState(false); const [isCreateModalVisible, setCreateModalVisible] = useState(false); @@ -116,16 +100,19 @@ const SearchTools: React.FC = ({ accessToken, userRole, userID if (toolIdToDelete == null || accessToken == null) { return; } + setIsDeleting(true); try { await deleteSearchTool(accessToken, toolIdToDelete); NotificationsManager.success("Deleted search tool successfully"); + setIsDeleteModalOpen(false); + setToolToDelete(null); refetch(); } catch (error) { console.error("Error deleting the search tool:", error); NotificationsManager.error("Failed to delete search tool"); + } finally { + setIsDeleting(false); } - setIsDeleteModalOpen(false); - setToolToDelete(null); }; const cancelDelete = () => { @@ -133,6 +120,11 @@ const SearchTools: React.FC = ({ accessToken, userRole, userID setToolToDelete(null); }; + const toolToDelete = searchTools?.find((t) => t.search_tool_id === toolIdToDelete); + const providerInfo = toolToDelete + ? availableProviders.find((p) => p.provider_name === toolToDelete.litellm_params.search_provider) + : null; + const handleCreateSuccess = (newSearchTool: SearchTool) => { setCreateModalVisible(false); refetch(); @@ -231,26 +223,46 @@ const SearchTools: React.FC = ({ accessToken, userRole, userID /> ) : (
-
- } size="large"> +
} - getRowCanExpand={() => false} - isLoading={isLoadingTools} - noDataMessage="No search tools configured" + rowKey={(record) => record.search_tool_id || record.search_tool_name} + pagination={false} + locale={{ + emptyText: "No search tools configured", + }} + size="small" /> - + + ); return (
- ([]); + const [loadingModels, setLoadingModels] = useState(true); + + // Test section state + const [testQuery, setTestQuery] = useState(""); + const [testModel, setTestModel] = useState("gpt-4o"); + const [testResult, setTestResult] = useState(null); + const [isTesting, setIsTesting] = useState(false); + + const schema = data?.field_schema; + const values = data?.values ?? {}; + + useEffect(() => { + const loadEmbeddingModels = async () => { + if (!accessToken) return; + try { + setLoadingModels(true); + const models = await fetchAvailableModels(accessToken); + const embeddingOnly = models.filter((model) => model.mode === "embedding"); + setEmbeddingModels(embeddingOnly); + } catch (error) { + console.error("Error fetching embedding models:", error); + } finally { + setLoadingModels(false); + } + }; + + loadEmbeddingModels(); + }, [accessToken]); + + useEffect(() => { + if (values) { + form.setFieldsValue({ + enabled: values.enabled ?? false, + embedding_model: values.embedding_model ?? "text-embedding-3-small", + top_k: values.top_k ?? 10, + similarity_threshold: values.similarity_threshold ?? 0.3, + }); + setIsDirty(false); + } + }, [values, form]); + + const handleSave = async () => { + try { + const formValues = await form.validateFields(); + updateSettings(formValues, { + onSuccess: () => { + setIsDirty(false); + setSaveSuccess(true); + setTimeout(() => setSaveSuccess(false), 3000); + NotificationManager.success( + "Settings updated successfully. Changes will be applied across all pods within 10 seconds." + ); + }, + onError: (error) => { + NotificationManager.fromBackend(error); + }, + }); + } catch (error) { + console.error("Form validation failed:", error); + } + }; + + const handleTest = async () => { + if (!accessToken) { + return; + } + + await runSemanticFilterTest({ + accessToken, + testModel, + testQuery, + setIsTesting, + setTestResult, + }); + }; + + if (!accessToken) { + return ( +
+ Please log in to configure semantic filter settings. +
+ ); + } + + return ( +
+ {isLoading ? ( + + ) : isError ? ( + + ) : ( + <> + + + {saveSuccess && ( + } + showIcon + closable + style={{ marginBottom: 16 }} + /> + )} + + {updateError && ( + + )} + + + {/* Left Column - Settings */} +
+ { + setIsDirty(true); + }} + > + + + Enable Semantic Filtering + + + + + } + valuePropName="checked" + > + + + + + {schema?.properties?.enabled?.description} + + + + + + Embedding Model + + + + + } + > +