mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge branch 'main' into ttl-prompt-caching-bedrock
This commit is contained in:
commit
e4b45dc32e
183 changed files with 13543 additions and 862 deletions
|
|
@ -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
|
||||
|
|
|
|||
114
cookbook/livekit_agent_sdk/README.md
Normal file
114
cookbook/livekit_agent_sdk/README.md
Normal file
|
|
@ -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)
|
||||
21
cookbook/livekit_agent_sdk/config.example.yaml
Normal file
21
cookbook/livekit_agent_sdk/config.example.yaml
Normal file
|
|
@ -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
|
||||
112
cookbook/livekit_agent_sdk/main.py
Normal file
112
cookbook/livekit_agent_sdk/main.py
Normal file
|
|
@ -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()
|
||||
2
cookbook/livekit_agent_sdk/requirements.txt
Normal file
2
cookbook/livekit_agent_sdk/requirements.txt
Normal file
|
|
@ -0,0 +1,2 @@
|
|||
livekit-agents[xai]>=1.3.12
|
||||
websockets>=15.0.1
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
280
docs/my-website/docs/a2a_invoking_agents.md
Normal file
280
docs/my-website/docs/a2a_invoking_agents.md
Normal file
|
|
@ -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
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python" default>
|
||||
|
||||
```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)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="typescript" label="TypeScript">
|
||||
|
||||
```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);
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="curl" label="cURL">
|
||||
|
||||
```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?"}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Streaming
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python" default>
|
||||
|
||||
```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)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="typescript" label="TypeScript">
|
||||
|
||||
```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);
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="curl" label="cURL">
|
||||
|
||||
```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
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Key Differences
|
||||
|
||||
| Method | Use Case | Advantages |
|
||||
|--------|----------|------------|
|
||||
| **A2A SDK** | Native A2A protocol integration | • Full A2A protocol support<br/>• Access to task states and artifacts<br/>• Context management |
|
||||
| **OpenAI SDK** | Familiar OpenAI-style interface | • Drop-in replacement for OpenAI calls<br/>• Easier migration from LLM to agent workflows<br/>• 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.
|
||||
|
||||
:::
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
308
docs/my-website/docs/providers/xai_realtime.md
Normal file
308
docs/my-website/docs/providers/xai_realtime.md
Normal file
|
|
@ -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)
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
278
docs/my-website/docs/proxy/guardrails/custom_code_guardrail.md
Normal file
278
docs/my-website/docs/proxy/guardrails/custom_code_guardrail.md
Normal file
|
|
@ -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):
|
||||
...
|
||||
```
|
||||
|
|
@ -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.
|
||||
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="monitor" label="Monitor Only">
|
||||
---
|
||||
|
||||
```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:
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="block-input" label="Block Input">
|
||||
```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:
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="full-coverage" label="Full Coverage">
|
||||
```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:
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="passthrough" label="Passthrough Mode">
|
||||
|
||||
```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.
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
- 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. |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="xai" label="xAI Grok Voice Agent">
|
||||
|
||||
```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)**
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
|
|
|
|||
99
docs/my-website/docs/tutorials/copilotkit_sdk.md
Normal file
99
docs/my-website/docs/tutorials/copilotkit_sdk.md
Normal file
|
|
@ -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)
|
||||
190
docs/my-website/docs/tutorials/livekit_xai_realtime.md
Normal file
190
docs/my-website/docs/tutorials/livekit_xai_realtime.md
Normal file
|
|
@ -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:
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python Client">
|
||||
|
||||
```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())
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="livekit" label="LiveKit 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,
|
||||
)
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## 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)
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
BIN
enterprise/dist/litellm_enterprise-0.1.29-py3-none-any.whl
vendored
Normal file
BIN
enterprise/dist/litellm_enterprise-0.1.29-py3-none-any.whl
vendored
Normal file
Binary file not shown.
BIN
enterprise/dist/litellm_enterprise-0.1.29.tar.gz
vendored
Normal file
BIN
enterprise/dist/litellm_enterprise-0.1.29.tar.gz
vendored
Normal file
Binary file not shown.
|
|
@ -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",
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
[tool.poetry]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.28"
|
||||
version = "0.1.29"
|
||||
description = "Package for LiteLLM Enterprise features"
|
||||
authors = ["BerriAI"]
|
||||
readme = "README.md"
|
||||
|
|
@ -22,7 +22,7 @@ requires = ["poetry-core"]
|
|||
build-backend = "poetry.core.masonry.api"
|
||||
|
||||
[tool.commitizen]
|
||||
version = "0.1.28"
|
||||
version = "0.1.29"
|
||||
version_files = [
|
||||
"pyproject.toml:version",
|
||||
"../requirements.txt:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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 ###
|
||||
|
|
@ -1378,6 +1380,7 @@ if TYPE_CHECKING:
|
|||
from .llms.topaz.image_variations.transformation import TopazImageVariationConfig as TopazImageVariationConfig
|
||||
from litellm.llms.openai.completion.transformation import OpenAITextCompletionConfig as OpenAITextCompletionConfig
|
||||
from .llms.groq.chat.transformation import GroqChatConfig as GroqChatConfig
|
||||
from .llms.a2a.chat.transformation import A2AConfig as A2AConfig
|
||||
from .llms.voyage.embedding.transformation import VoyageEmbeddingConfig as VoyageEmbeddingConfig
|
||||
from .llms.voyage.embedding.transformation_contextual import VoyageContextualEmbeddingConfig as VoyageContextualEmbeddingConfig
|
||||
from .llms.infinity.embedding.transformation import InfinityEmbeddingConfig as InfinityEmbeddingConfig
|
||||
|
|
|
|||
|
|
@ -213,6 +213,7 @@ LLM_CONFIG_NAMES = (
|
|||
"TopazImageVariationConfig",
|
||||
"OpenAITextCompletionConfig",
|
||||
"GroqChatConfig",
|
||||
"A2AConfig",
|
||||
"GenAIHubOrchestrationConfig",
|
||||
"VoyageEmbeddingConfig",
|
||||
"VoyageContextualEmbeddingConfig",
|
||||
|
|
@ -850,6 +851,7 @@ _LLM_CONFIGS_IMPORT_MAP = {
|
|||
"OpenAITextCompletionConfig",
|
||||
),
|
||||
"GroqChatConfig": (".llms.groq.chat.transformation", "GroqChatConfig"),
|
||||
"A2AConfig": (".llms.a2a.chat.transformation", "A2AConfig"),
|
||||
"GenAIHubOrchestrationConfig": (
|
||||
".llms.sap.chat.transformation",
|
||||
"GenAIHubOrchestrationConfig",
|
||||
|
|
|
|||
|
|
@ -329,6 +329,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
else:
|
||||
request_data[key] = value
|
||||
|
||||
if headers:
|
||||
request_data["extra_headers"] = headers
|
||||
|
||||
return request_data
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -99,6 +99,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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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 {}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
6
litellm/llms/a2a/__init__.py
Normal file
6
litellm/llms/a2a/__init__.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
"""
|
||||
A2A (Agent-to-Agent) Protocol Provider for LiteLLM
|
||||
"""
|
||||
from .chat.transformation import A2AConfig
|
||||
|
||||
__all__ = ["A2AConfig"]
|
||||
6
litellm/llms/a2a/chat/__init__.py
Normal file
6
litellm/llms/a2a/chat/__init__.py
Normal file
|
|
@ -0,0 +1,6 @@
|
|||
"""
|
||||
A2A Chat Completion Implementation
|
||||
"""
|
||||
from .transformation import A2AConfig
|
||||
|
||||
__all__ = ["A2AConfig"]
|
||||
103
litellm/llms/a2a/chat/streaming_iterator.py
Normal file
103
litellm/llms/a2a/chat/streaming_iterator.py
Normal file
|
|
@ -0,0 +1,103 @@
|
|||
"""
|
||||
A2A Streaming Response Iterator
|
||||
"""
|
||||
from typing import Optional, Union
|
||||
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
|
||||
|
||||
from ..common_utils import extract_text_from_a2a_response
|
||||
|
||||
|
||||
class A2AModelResponseIterator(BaseModelResponseIterator):
|
||||
"""
|
||||
Iterator for parsing A2A streaming responses.
|
||||
|
||||
Converts A2A JSON-RPC streaming chunks to OpenAI-compatible format.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
streaming_response,
|
||||
sync_stream: bool,
|
||||
json_mode: Optional[bool] = False,
|
||||
model: str = "a2a/agent",
|
||||
):
|
||||
super().__init__(
|
||||
streaming_response=streaming_response,
|
||||
sync_stream=sync_stream,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
self.model = model
|
||||
|
||||
def chunk_parser(self, chunk: dict) -> Union[GenericStreamingChunk, ModelResponseStream]:
|
||||
"""
|
||||
Parse A2A streaming chunk to OpenAI format.
|
||||
|
||||
A2A chunk format:
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": "request-id",
|
||||
"result": {
|
||||
"message": {
|
||||
"parts": [{"kind": "text", "text": "content"}]
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
Or for tasks:
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"result": {
|
||||
"kind": "task",
|
||||
"status": {"state": "running"},
|
||||
"artifacts": [{"parts": [{"kind": "text", "text": "content"}]}]
|
||||
}
|
||||
}
|
||||
"""
|
||||
try:
|
||||
# Extract text from A2A response
|
||||
text = extract_text_from_a2a_response(chunk)
|
||||
|
||||
# Determine finish reason
|
||||
finish_reason = self._get_finish_reason(chunk)
|
||||
|
||||
# Return generic streaming chunk
|
||||
return GenericStreamingChunk(
|
||||
text=text,
|
||||
is_finished=bool(finish_reason),
|
||||
finish_reason=finish_reason or "",
|
||||
usage=None,
|
||||
index=0,
|
||||
tool_use=None,
|
||||
)
|
||||
except Exception:
|
||||
# Return empty chunk on parse error
|
||||
return GenericStreamingChunk(
|
||||
text="",
|
||||
is_finished=False,
|
||||
finish_reason="",
|
||||
usage=None,
|
||||
index=0,
|
||||
tool_use=None,
|
||||
)
|
||||
|
||||
def _get_finish_reason(self, chunk: dict) -> Optional[str]:
|
||||
"""Extract finish reason from A2A chunk"""
|
||||
result = chunk.get("result", {})
|
||||
|
||||
# Check for task completion
|
||||
if isinstance(result, dict):
|
||||
status = result.get("status", {})
|
||||
if isinstance(status, dict):
|
||||
state = status.get("state")
|
||||
if state == "completed":
|
||||
return "stop"
|
||||
elif state == "failed":
|
||||
return "stop" # Map failed state to 'stop' (valid finish_reason)
|
||||
|
||||
# Check for [DONE] marker
|
||||
if chunk.get("done") is True:
|
||||
return "stop"
|
||||
|
||||
return None
|
||||
370
litellm/llms/a2a/chat/transformation.py
Normal file
370
litellm/llms/a2a/chat/transformation.py
Normal file
|
|
@ -0,0 +1,370 @@
|
|||
"""
|
||||
A2A Protocol Transformation for LiteLLM
|
||||
"""
|
||||
import uuid
|
||||
from typing import Any, Dict, Iterator, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import Choices, Message, ModelResponse
|
||||
|
||||
from ..common_utils import (
|
||||
A2AError,
|
||||
convert_messages_to_prompt,
|
||||
extract_text_from_a2a_response,
|
||||
)
|
||||
from .streaming_iterator import A2AModelResponseIterator
|
||||
|
||||
|
||||
class A2AConfig(BaseConfig):
|
||||
"""
|
||||
Configuration for A2A (Agent-to-Agent) Protocol.
|
||||
|
||||
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/<agent-name>".
|
||||
|
||||
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 [
|
||||
"stream",
|
||||
"temperature",
|
||||
"max_tokens",
|
||||
"top_p",
|
||||
]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OpenAI parameters to A2A parameters.
|
||||
|
||||
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(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Validate environment and set headers for A2A requests.
|
||||
|
||||
Args:
|
||||
headers: Request headers dict
|
||||
model: Model name
|
||||
messages: Messages list
|
||||
optional_params: Optional parameters
|
||||
litellm_params: LiteLLM parameters
|
||||
api_key: API key (optional for A2A)
|
||||
api_base: API base URL
|
||||
|
||||
Returns:
|
||||
Updated headers dict
|
||||
"""
|
||||
# Ensure Content-Type is set to application/json for JSON-RPC 2.0
|
||||
if "content-type" not in headers and "Content-Type" not in headers:
|
||||
headers["Content-Type"] = "application/json"
|
||||
|
||||
# Add Authorization header if API key is provided
|
||||
if api_key is not None:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Get the complete A2A agent endpoint URL.
|
||||
|
||||
A2A agents use JSON-RPC 2.0 at the base URL, not specific paths.
|
||||
The method (message/send or message/stream) is specified in the
|
||||
JSON-RPC request body, not in the URL.
|
||||
|
||||
Args:
|
||||
api_base: Base URL of the A2A agent (e.g., "http://0.0.0.0:9999")
|
||||
api_key: API key (not used for URL construction)
|
||||
model: Model name (not used for A2A, agent determined by api_base)
|
||||
optional_params: Optional parameters
|
||||
litellm_params: LiteLLM parameters
|
||||
stream: Whether this is a streaming request (affects JSON-RPC method)
|
||||
|
||||
Returns:
|
||||
Complete URL for the A2A endpoint (base URL)
|
||||
"""
|
||||
if api_base is None:
|
||||
raise ValueError("api_base is required for A2A provider")
|
||||
|
||||
# A2A uses JSON-RPC 2.0 at the base URL
|
||||
# Remove trailing slash for consistency
|
||||
return api_base.rstrip("/")
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform OpenAI request to A2A JSON-RPC 2.0 format.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
messages: List of OpenAI messages
|
||||
optional_params: Optional parameters
|
||||
litellm_params: LiteLLM parameters
|
||||
headers: Request headers
|
||||
|
||||
Returns:
|
||||
A2A JSON-RPC 2.0 request dict
|
||||
"""
|
||||
# Generate request ID
|
||||
request_id = str(uuid.uuid4())
|
||||
|
||||
if not messages:
|
||||
raise ValueError("At least one message is required for A2A completion")
|
||||
|
||||
# Convert all messages to maintain conversation history
|
||||
# Use helper to format conversation with role prefixes
|
||||
full_context = convert_messages_to_prompt(messages)
|
||||
|
||||
# Create single A2A message with full conversation context
|
||||
a2a_message = {
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": full_context}],
|
||||
"messageId": str(uuid.uuid4()),
|
||||
}
|
||||
|
||||
# Build JSON-RPC 2.0 request
|
||||
# For A2A protocol, the method is "message/send" for non-streaming
|
||||
# and "message/stream" for streaming
|
||||
stream = optional_params.get("stream", False)
|
||||
method = "message/stream" if stream else "message/send"
|
||||
|
||||
request_data = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"method": method,
|
||||
"params": {
|
||||
"message": a2a_message
|
||||
}
|
||||
}
|
||||
|
||||
return request_data
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: Any,
|
||||
request_data: dict,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> ModelResponse:
|
||||
"""
|
||||
Transform A2A JSON-RPC 2.0 response to OpenAI format.
|
||||
|
||||
Args:
|
||||
model: Model name
|
||||
raw_response: HTTP response from A2A agent
|
||||
model_response: Model response object to populate
|
||||
logging_obj: Logging object
|
||||
request_data: Original request data
|
||||
messages: Original messages
|
||||
optional_params: Optional parameters
|
||||
litellm_params: LiteLLM parameters
|
||||
encoding: Encoding object
|
||||
api_key: API key
|
||||
json_mode: JSON mode flag
|
||||
|
||||
Returns:
|
||||
Populated ModelResponse object
|
||||
"""
|
||||
try:
|
||||
response_json = raw_response.json()
|
||||
except Exception as e:
|
||||
raise A2AError(
|
||||
status_code=raw_response.status_code,
|
||||
message=f"Failed to parse A2A response: {str(e)}",
|
||||
headers=dict(raw_response.headers),
|
||||
)
|
||||
|
||||
# Check for JSON-RPC error
|
||||
if "error" in response_json:
|
||||
error = response_json["error"]
|
||||
raise A2AError(
|
||||
status_code=raw_response.status_code,
|
||||
message=f"A2A error: {error.get('message', 'Unknown error')}",
|
||||
headers=dict(raw_response.headers),
|
||||
)
|
||||
|
||||
# Extract text from A2A response
|
||||
text = extract_text_from_a2a_response(response_json)
|
||||
|
||||
# Populate model response
|
||||
model_response.choices = [
|
||||
Choices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=Message(
|
||||
content=text,
|
||||
role="assistant",
|
||||
),
|
||||
)
|
||||
]
|
||||
|
||||
# Set model
|
||||
model_response.model = model
|
||||
|
||||
# Set ID from response
|
||||
model_response.id = response_json.get("id", str(uuid.uuid4()))
|
||||
|
||||
return model_response
|
||||
|
||||
def get_model_response_iterator(
|
||||
self,
|
||||
streaming_response: Union[Iterator, Any],
|
||||
sync_stream: bool,
|
||||
json_mode: Optional[bool] = False,
|
||||
) -> BaseModelResponseIterator:
|
||||
"""
|
||||
Get streaming iterator for A2A responses.
|
||||
|
||||
Args:
|
||||
streaming_response: Streaming response iterator
|
||||
sync_stream: Whether this is a sync stream
|
||||
json_mode: JSON mode flag
|
||||
|
||||
Returns:
|
||||
A2A streaming iterator
|
||||
"""
|
||||
return A2AModelResponseIterator(
|
||||
streaming_response=streaming_response,
|
||||
sync_stream=sync_stream,
|
||||
json_mode=json_mode,
|
||||
)
|
||||
|
||||
def _openai_message_to_a2a_message(self, message: Dict[str, Any]) -> Dict[str, Any]:
|
||||
"""
|
||||
Convert OpenAI message to A2A message format.
|
||||
|
||||
Args:
|
||||
message: OpenAI message dict
|
||||
|
||||
Returns:
|
||||
A2A message dict
|
||||
"""
|
||||
content = message.get("content", "")
|
||||
role = message.get("role", "user")
|
||||
|
||||
return {
|
||||
"role": role,
|
||||
"parts": [{"kind": "text", "text": str(content)}],
|
||||
"messageId": str(uuid.uuid4()),
|
||||
}
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
"""Return appropriate error class for A2A errors"""
|
||||
# Convert headers to dict if needed
|
||||
headers_dict = dict(headers) if isinstance(headers, httpx.Headers) else headers
|
||||
return A2AError(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers_dict,
|
||||
)
|
||||
152
litellm/llms/a2a/common_utils.py
Normal file
152
litellm/llms/a2a/common_utils.py
Normal file
|
|
@ -0,0 +1,152 @@
|
|||
"""
|
||||
Common utilities for A2A (Agent-to-Agent) Protocol
|
||||
"""
|
||||
from typing import Any, Dict, List
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
|
||||
class A2AError(BaseLLMException):
|
||||
"""Base exception for A2A protocol errors"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
status_code: int,
|
||||
message: str,
|
||||
headers: Dict[str, Any] = {},
|
||||
):
|
||||
super().__init__(
|
||||
status_code=status_code,
|
||||
message=message,
|
||||
headers=headers,
|
||||
)
|
||||
|
||||
|
||||
def convert_messages_to_prompt(messages: List[AllMessageValues]) -> str:
|
||||
"""
|
||||
Convert OpenAI messages to a single prompt string for A2A agent.
|
||||
|
||||
Formats each message as "{role}: {content}" and joins with newlines
|
||||
to preserve conversation history. Handles both string and list content.
|
||||
|
||||
Args:
|
||||
messages: List of OpenAI-format messages
|
||||
|
||||
Returns:
|
||||
Formatted prompt string with full conversation context
|
||||
"""
|
||||
conversation_parts = []
|
||||
for msg in messages:
|
||||
# Use LiteLLM's helper to extract text from content (handles both str and list)
|
||||
content_text = convert_content_list_to_str(message=msg)
|
||||
|
||||
# Get role
|
||||
if isinstance(msg, BaseModel):
|
||||
role = msg.model_dump().get("role", "user")
|
||||
elif isinstance(msg, dict):
|
||||
role = msg.get("role", "user")
|
||||
else:
|
||||
role = dict(msg).get("role", "user") # type: ignore
|
||||
|
||||
if content_text:
|
||||
conversation_parts.append(f"{role}: {content_text}")
|
||||
|
||||
return "\n".join(conversation_parts)
|
||||
|
||||
|
||||
def extract_text_from_a2a_message(
|
||||
message: Dict[str, Any], depth: int = 0, max_depth: int = 10
|
||||
) -> str:
|
||||
"""
|
||||
Extract text content from A2A message parts.
|
||||
|
||||
Args:
|
||||
message: A2A message dict with 'parts' containing text parts
|
||||
depth: Current recursion depth (internal use)
|
||||
max_depth: Maximum recursion depth to prevent infinite loops
|
||||
|
||||
Returns:
|
||||
Concatenated text from all text parts
|
||||
"""
|
||||
if message is None or depth >= max_depth:
|
||||
return ""
|
||||
|
||||
parts = message.get("parts", [])
|
||||
text_parts: List[str] = []
|
||||
|
||||
for part in parts:
|
||||
if part.get("kind") == "text":
|
||||
text_parts.append(part.get("text", ""))
|
||||
# Handle nested parts if they exist
|
||||
elif "parts" in part:
|
||||
nested_text = extract_text_from_a2a_message(part, depth + 1, max_depth)
|
||||
if nested_text:
|
||||
text_parts.append(nested_text)
|
||||
|
||||
return " ".join(text_parts)
|
||||
|
||||
|
||||
def extract_text_from_a2a_response(
|
||||
response_dict: Dict[str, Any], max_depth: int = 10
|
||||
) -> str:
|
||||
"""
|
||||
Extract text content from A2A response result.
|
||||
|
||||
Args:
|
||||
response_dict: A2A response dict with 'result' containing message
|
||||
max_depth: Maximum recursion depth to prevent infinite loops
|
||||
|
||||
Returns:
|
||||
Text from response message parts
|
||||
"""
|
||||
result = response_dict.get("result", {})
|
||||
if not isinstance(result, dict):
|
||||
return ""
|
||||
|
||||
# A2A response can have different formats:
|
||||
# 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:
|
||||
return extract_text_from_a2a_message(result, depth=0, max_depth=max_depth)
|
||||
|
||||
# Check for nested message
|
||||
message = result.get("message")
|
||||
if message:
|
||||
return extract_text_from_a2a_message(message, depth=0, max_depth=max_depth)
|
||||
|
||||
# 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]
|
||||
return extract_text_from_a2a_message(
|
||||
first_artifact, depth=0, max_depth=max_depth
|
||||
)
|
||||
|
||||
return ""
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
message.pop("provider_specific_fields", None)
|
||||
|
||||
return messages
|
||||
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
|
|
|
|||
|
|
@ -1,11 +1,19 @@
|
|||
<<<<<<< ttl-prompt-caching-bedrock
|
||||
from typing import List, Optional, Tuple
|
||||
=======
|
||||
from typing import Any, List, Optional, Tuple, cast
|
||||
>>>>>>> main
|
||||
|
||||
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):
|
||||
|
|
@ -43,6 +51,7 @@ class GithubCopilotConfig(OpenAIConfig):
|
|||
):
|
||||
import litellm
|
||||
|
||||
<<<<<<< ttl-prompt-caching-bedrock
|
||||
disable_copilot_system_to_assistant = (
|
||||
litellm.disable_copilot_system_to_assistant
|
||||
)
|
||||
|
|
@ -51,6 +60,26 @@ class GithubCopilotConfig(OpenAIConfig):
|
|||
if "role" in message and message["role"] == "system":
|
||||
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
|
||||
>>>>>>> main
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
|
|
@ -67,6 +96,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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
5
litellm/llms/xai/realtime/__init__.py
Normal file
5
litellm/llms/xai/realtime/__init__.py
Normal file
|
|
@ -0,0 +1,5 @@
|
|||
"""xAI Realtime API handler."""
|
||||
|
||||
from .handler import XAIRealtime
|
||||
|
||||
__all__ = ["XAIRealtime"]
|
||||
38
litellm/llms/xai/realtime/handler.py
Normal file
38
litellm/llms/xai/realtime/handler.py
Normal file
|
|
@ -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}",
|
||||
}
|
||||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
@ -2199,6 +2206,48 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
logging_obj=logging, # model call logging done inside the class as we make need to modify I/O to fit aleph alpha's requirements
|
||||
client=client,
|
||||
)
|
||||
elif custom_llm_provider == "a2a":
|
||||
# A2A (Agent-to-Agent) Protocol
|
||||
# Resolve agent configuration from registry if model format is "a2a/<agent-name>"
|
||||
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. "
|
||||
"Either provide api_base parameter, set A2A_API_BASE environment variable, "
|
||||
"or register the agent in the proxy with model='a2a/<agent-name>'."
|
||||
)
|
||||
|
||||
headers = headers or litellm.headers
|
||||
|
||||
response = base_llm_http_handler.completion(
|
||||
model=model,
|
||||
stream=stream,
|
||||
messages=messages,
|
||||
acompletion=acompletion,
|
||||
api_base=api_base,
|
||||
model_response=model_response,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
shared_session=shared_session,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
timeout=timeout,
|
||||
headers=headers,
|
||||
encoding=_get_encoding(),
|
||||
api_key=api_key,
|
||||
logging_obj=logging,
|
||||
client=client,
|
||||
provider_config=provider_config,
|
||||
)
|
||||
elif custom_llm_provider == "gigachat":
|
||||
# GigaChat - Sber AI's LLM (Russia)
|
||||
api_key = (
|
||||
|
|
@ -2455,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
|
||||
|
||||
|
|
@ -3113,8 +3176,8 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
api_key
|
||||
or litellm.api_key
|
||||
or litellm.openrouter_key
|
||||
or get_secret("OPENROUTER_API_KEY")
|
||||
or get_secret("OR_API_KEY")
|
||||
or get_secret_str("OPENROUTER_API_KEY")
|
||||
or get_secret_str("OR_API_KEY")
|
||||
)
|
||||
|
||||
openrouter_site_url = get_secret("OR_SITE_URL") or "https://litellm.ai"
|
||||
|
|
@ -4555,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)
|
||||
|
|
@ -4884,8 +4954,8 @@ def embedding( # noqa: PLR0915
|
|||
api_key
|
||||
or litellm.api_key
|
||||
or litellm.openrouter_key
|
||||
or get_secret("OPENROUTER_API_KEY")
|
||||
or get_secret("OR_API_KEY")
|
||||
or get_secret_str("OPENROUTER_API_KEY")
|
||||
or get_secret_str("OR_API_KEY")
|
||||
)
|
||||
|
||||
openrouter_site_url = get_secret("OR_SITE_URL") or "https://litellm.ai"
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
@ -27113,6 +27189,34 @@
|
|||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"together_ai/zai-org/GLM-4.7": {
|
||||
"input_cost_per_token": 4.5e-07,
|
||||
"litellm_provider": "together_ai",
|
||||
"max_input_tokens": 200000,
|
||||
"max_output_tokens": 200000,
|
||||
"max_tokens": 200000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2e-06,
|
||||
"source": "https://www.together.ai/models/glm-4-7",
|
||||
"supports_function_calling": true,
|
||||
"supports_parallel_function_calling": true,
|
||||
"supports_reasoning": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"together_ai/moonshotai/Kimi-K2.5": {
|
||||
"input_cost_per_token": 5e-07,
|
||||
"litellm_provider": "together_ai",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2.8e-06,
|
||||
"source": "https://www.together.ai/models/kimi-k2-5",
|
||||
"supports_function_calling": true,
|
||||
"supports_tool_choice": true,
|
||||
"supports_vision": true,
|
||||
"supports_reasoning": true
|
||||
},
|
||||
"together_ai/moonshotai/Kimi-K2-Instruct-0905": {
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "together_ai",
|
||||
|
|
@ -27829,7 +27933,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,
|
||||
|
|
@ -27838,7 +27944,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,
|
||||
|
|
@ -27847,7 +27955,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,
|
||||
|
|
@ -27856,7 +27967,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,
|
||||
|
|
@ -27865,7 +27978,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,
|
||||
|
|
@ -27885,7 +28001,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,
|
||||
|
|
@ -27896,7 +28016,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,
|
||||
|
|
@ -27907,7 +28031,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,
|
||||
|
|
@ -27918,7 +28046,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,
|
||||
|
|
@ -27929,7 +28061,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,
|
||||
|
|
@ -27940,7 +28076,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,
|
||||
|
|
@ -27951,7 +28091,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,
|
||||
|
|
@ -27960,7 +28102,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,
|
||||
|
|
@ -27969,7 +28113,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,
|
||||
|
|
@ -27978,7 +28124,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,
|
||||
|
|
@ -27996,7 +28144,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,
|
||||
|
|
@ -28005,7 +28154,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,
|
||||
|
|
@ -28014,7 +28166,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",
|
||||
|
|
@ -28024,7 +28177,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",
|
||||
|
|
@ -28034,7 +28191,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,
|
||||
|
|
@ -28043,7 +28204,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,
|
||||
|
|
@ -28052,7 +28217,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,
|
||||
|
|
@ -28070,7 +28239,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,
|
||||
|
|
@ -28106,7 +28278,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,
|
||||
|
|
@ -28115,7 +28288,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,
|
||||
|
|
@ -28124,7 +28298,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,
|
||||
|
|
@ -28133,7 +28308,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,
|
||||
|
|
@ -28142,7 +28319,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,
|
||||
|
|
@ -28160,7 +28340,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,
|
||||
|
|
@ -28169,7 +28351,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,
|
||||
|
|
@ -28178,7 +28363,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,
|
||||
|
|
@ -28187,7 +28374,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,
|
||||
|
|
@ -28196,7 +28384,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,
|
||||
|
|
@ -28205,7 +28396,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,
|
||||
|
|
@ -28223,7 +28416,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,
|
||||
|
|
@ -28232,7 +28428,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,
|
||||
|
|
@ -28241,7 +28440,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,
|
||||
|
|
@ -28250,7 +28450,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,
|
||||
|
|
@ -28259,7 +28461,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,
|
||||
|
|
@ -28277,7 +28482,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,
|
||||
|
|
@ -28295,7 +28502,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,
|
||||
|
|
@ -28304,7 +28514,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,
|
||||
|
|
@ -28313,7 +28524,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,
|
||||
|
|
@ -28322,7 +28537,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,
|
||||
|
|
@ -28331,7 +28550,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,
|
||||
|
|
@ -28358,7 +28579,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,
|
||||
|
|
@ -28376,7 +28599,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,
|
||||
|
|
@ -28387,7 +28613,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,
|
||||
|
|
@ -28398,7 +28628,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,
|
||||
|
|
@ -28409,7 +28643,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,
|
||||
|
|
@ -28420,7 +28658,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,
|
||||
|
|
@ -28431,7 +28673,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,
|
||||
|
|
@ -28442,7 +28688,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,
|
||||
|
|
@ -28453,7 +28703,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,
|
||||
|
|
@ -28464,7 +28718,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,
|
||||
|
|
@ -28475,7 +28732,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,
|
||||
|
|
@ -28547,7 +28808,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,
|
||||
|
|
@ -28556,7 +28820,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,
|
||||
|
|
@ -28565,7 +28832,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,
|
||||
|
|
@ -28574,7 +28843,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,
|
||||
|
|
@ -28583,7 +28855,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,
|
||||
|
|
@ -28592,7 +28866,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,
|
||||
|
|
@ -28601,7 +28876,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,
|
||||
|
|
@ -28610,7 +28887,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,
|
||||
|
|
@ -28619,7 +28898,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,
|
||||
|
|
@ -28628,7 +28909,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,
|
||||
|
|
@ -28637,7 +28920,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",
|
||||
|
|
@ -28705,7 +28990,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,
|
||||
|
|
@ -28976,7 +29263,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,
|
||||
|
|
@ -29028,7 +29316,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,
|
||||
|
|
@ -29310,6 +29599,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",
|
||||
|
|
@ -29799,7 +30103,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
|
||||
},
|
||||
|
|
@ -29812,7 +30118,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
|
||||
},
|
||||
|
|
@ -29825,7 +30133,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
|
||||
},
|
||||
|
|
@ -29838,7 +30148,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
|
||||
},
|
||||
|
|
@ -34787,4 +35099,4 @@
|
|||
"output_cost_per_token": 0,
|
||||
"supports_reasoning": true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
53
litellm/proxy/agent_endpoints/a2a_routing.py
Normal file
53
litellm/proxy/agent_endpoints/a2a_routing.py
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
"""
|
||||
A2A Agent Routing
|
||||
|
||||
Handles routing for A2A agents (models with "a2a/<agent-name>" 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
|
||||
|
||||
|
||||
async 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)
|
||||
96
litellm/proxy/agent_endpoints/model_list_helpers.py
Normal file
96
litellm/proxy/agent_endpoints/model_list_helpers.py
Normal file
|
|
@ -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/<agent-name>" 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/<agent-name>" 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
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 <your_api_key>" \\
|
||||
-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 (<guardrail>, 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, "<guardrail>", "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(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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, "<guardrail>", "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
|
||||
|
|
@ -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"</\s*(html|head|body|div|span|p|a|script|style)\s*>",
|
||||
],
|
||||
"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,
|
||||
}
|
||||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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}"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -239,6 +239,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,
|
||||
)
|
||||
|
|
@ -8616,6 +8620,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 +9469,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}
|
||||
|
||||
|
|
|
|||
|
|
@ -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,12 @@ 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,
|
||||
)
|
||||
|
||||
return await route_a2a_agent_request(data, route_type)
|
||||
|
||||
elif user_model is not None:
|
||||
return getattr(litellm, f"{route_type}")(**data)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue