mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
Merge remote-tracking branch 'origin' into litellm_ui_callback_fix
This commit is contained in:
commit
4a0893ca22
492 changed files with 17905 additions and 3075 deletions
|
|
@ -3496,8 +3496,13 @@ jobs:
|
|||
command: |
|
||||
npx playwright test e2e_ui_tests/ --reporter=html --output=test-results
|
||||
no_output_timeout: 120m
|
||||
- store_test_results:
|
||||
- store_artifacts:
|
||||
path: test-results
|
||||
destination: playwright-results
|
||||
|
||||
- store_artifacts:
|
||||
path: playwright-report
|
||||
destination: playwright-report
|
||||
|
||||
test_nonroot_image:
|
||||
machine:
|
||||
|
|
|
|||
|
|
@ -348,7 +348,7 @@ curl 'http://0.0.0.0:4000/key/generate' \
|
|||
| [Fireworks AI (`fireworks_ai`)](https://docs.litellm.ai/docs/providers/fireworks_ai) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [FriendliAI (`friendliai`)](https://docs.litellm.ai/docs/providers/friendliai) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Galadriel (`galadriel`)](https://docs.litellm.ai/docs/providers/galadriel) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [GitHub Copilot (`github_copilot`)](https://docs.litellm.ai/docs/providers/github_copilot) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [GitHub Copilot (`github_copilot`)](https://docs.litellm.ai/docs/providers/github_copilot) | ✅ | ✅ | ✅ | ✅ | | | | | | |
|
||||
| [GitHub Models (`github`)](https://docs.litellm.ai/docs/providers/github) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Google - PaLM](https://docs.litellm.ai/docs/providers/palm) | ✅ | ✅ | ✅ | | | | | | | |
|
||||
| [Google - Vertex AI (`vertex_ai`)](https://docs.litellm.ai/docs/providers/vertex) | ✅ | ✅ | ✅ | ✅ | ✅ | | | | | |
|
||||
|
|
|
|||
|
|
@ -361,41 +361,6 @@ async def health():
|
|||
return {"status": "healthy"}
|
||||
|
||||
|
||||
@app.post(
|
||||
"/guardrail/{guardrailIdentifier}/version/{guardrailVersion}/apply",
|
||||
response_model=BedrockGuardrailResponse,
|
||||
)
|
||||
async def apply_guardrail(
|
||||
guardrailIdentifier: str,
|
||||
guardrailVersion: str,
|
||||
request: BedrockRequest,
|
||||
token: str = Depends(verify_bearer_token),
|
||||
) -> BedrockGuardrailResponse:
|
||||
"""
|
||||
Apply guardrail to input or output content.
|
||||
|
||||
This endpoint mimics the AWS Bedrock ApplyGuardrail API.
|
||||
|
||||
Args:
|
||||
guardrailIdentifier: The guardrail ID
|
||||
guardrailVersion: The guardrail version
|
||||
request: The guardrail request containing content to analyze
|
||||
token: Bearer token (verified by dependency)
|
||||
|
||||
Returns:
|
||||
BedrockGuardrailResponse with analysis results
|
||||
"""
|
||||
# Process the request
|
||||
response, output_texts = process_guardrail_request(request)
|
||||
|
||||
# Log the request (optional, for debugging)
|
||||
print(f"Guardrail applied: {guardrailIdentifier} v{guardrailVersion}")
|
||||
print(f"Source: {request.source}")
|
||||
print(f"Action: {response.action}")
|
||||
|
||||
return response
|
||||
|
||||
|
||||
"""
|
||||
LiteLLM exposes a basic guardrail API with the text extracted from the request and sent to the guardrail API, as well as the received request body for any further processing.
|
||||
|
||||
|
|
@ -426,9 +391,14 @@ This is a beta API. Please help us improve it.
|
|||
class LitellmBasicGuardrailRequest(BaseModel):
|
||||
texts: List[str]
|
||||
images: Optional[List[str]] = None
|
||||
tools: Optional[List[dict]] = None
|
||||
tool_calls: Optional[List[dict]] = None
|
||||
request_data: Dict[str, Any] = Field(default_factory=dict)
|
||||
additional_provider_specific_params: Dict[str, Any] = Field(default_factory=dict)
|
||||
input_type: Literal["request", "response"]
|
||||
litellm_call_id: Optional[str] = None
|
||||
litellm_trace_id: Optional[str] = None
|
||||
structured_messages: Optional[List[Dict[str, Any]]] = None
|
||||
|
||||
|
||||
class LitellmBasicGuardrailResponse(BaseModel):
|
||||
|
|
|
|||
232
docs/my-website/docs/a2a.md
Normal file
232
docs/my-website/docs/a2a.md
Normal file
|
|
@ -0,0 +1,232 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
import Image from '@theme/IdealImage';
|
||||
|
||||
# Agent Gateway (A2A Protocol) - Overview
|
||||
|
||||
Add A2A Agents on LiteLLM AI Gateway, Invoke agents in A2A Protocol, track request/response logs in LiteLLM Logs. Manage which Teams, Keys can access which Agents onboarded.
|
||||
|
||||
<Image
|
||||
img={require('../img/a2a_gateway.png')}
|
||||
style={{width: '80%', display: 'block', margin: '0', borderRadius: '8px'}}
|
||||
/>
|
||||
|
||||
<br />
|
||||
<br />
|
||||
|
||||
| Feature | Supported |
|
||||
|---------|-----------|
|
||||
| Logging | ✅ |
|
||||
| Load Balancing | ✅ |
|
||||
| Streaming | ✅ |
|
||||
|
||||
:::tip
|
||||
|
||||
LiteLLM follows the [A2A (Agent-to-Agent) Protocol](https://github.com/google/A2A) for invoking agents.
|
||||
|
||||
:::
|
||||
|
||||
## Adding your Agent
|
||||
|
||||
You can add A2A-compatible agents through the LiteLLM Admin UI.
|
||||
|
||||
1. Navigate to the **Agents** tab
|
||||
2. Click **Add Agent**
|
||||
3. Enter the agent name (e.g., `ij-local`) and the URL of your A2A agent
|
||||
|
||||
<Image
|
||||
img={require('../img/add_agent_1.png')}
|
||||
style={{width: '80%', display: 'block', margin: '0'}}
|
||||
/>
|
||||
|
||||
The URL should be the invocation URL for your A2A agent (e.g., `http://localhost:10001`).
|
||||
|
||||
## Invoking your Agents
|
||||
|
||||
Use the [A2A Python SDK](https://pypi.org/project/a2a/) 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())
|
||||
```
|
||||
|
||||
## Tracking Agent Logs
|
||||
|
||||
After invoking an agent, you can view the request logs in the LiteLLM **Logs** tab.
|
||||
|
||||
The logs show:
|
||||
- **Request/Response content** sent to and received from the agent
|
||||
- **User, Key, Team** information for tracking who made the request
|
||||
- **Latency and cost** metrics
|
||||
|
||||
<Image
|
||||
img={require('../img/agent2.png')}
|
||||
style={{width: '100%', display: 'block', margin: '2rem auto'}}
|
||||
/>
|
||||
|
||||
## API Reference
|
||||
|
||||
### Endpoint
|
||||
|
||||
```
|
||||
POST /a2a/{agent_name}/message/send
|
||||
```
|
||||
|
||||
### Authentication
|
||||
|
||||
Include your LiteLLM Virtual Key in the `Authorization` header:
|
||||
|
||||
```
|
||||
Authorization: Bearer sk-your-litellm-key
|
||||
```
|
||||
|
||||
### Request Format
|
||||
|
||||
LiteLLM follows the [A2A JSON-RPC 2.0 specification](https://github.com/google/A2A):
|
||||
|
||||
```json title="Request Body"
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": "unique-request-id",
|
||||
"method": "message/send",
|
||||
"params": {
|
||||
"message": {
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Your message here"}],
|
||||
"messageId": "unique-message-id"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Response Format
|
||||
|
||||
```json title="Response"
|
||||
{
|
||||
"jsonrpc": "2.0",
|
||||
"id": "unique-request-id",
|
||||
"result": {
|
||||
"kind": "task",
|
||||
"id": "task-id",
|
||||
"contextId": "context-id",
|
||||
"status": {"state": "completed", "timestamp": "2025-01-01T00:00:00Z"},
|
||||
"artifacts": [
|
||||
{
|
||||
"artifactId": "artifact-id",
|
||||
"name": "response",
|
||||
"parts": [{"kind": "text", "text": "Agent response here"}]
|
||||
}
|
||||
]
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Agent Registry
|
||||
|
||||
Want to create a central registry so your team can discover what agents are available within your company?
|
||||
|
||||
Use the [AI Hub](./proxy/ai_hub) to make agents public and discoverable across your organization. This allows developers to browse available agents without needing to rebuild them.
|
||||
259
docs/my-website/docs/a2a_agent_permissions.md
Normal file
259
docs/my-website/docs/a2a_agent_permissions.md
Normal file
|
|
@ -0,0 +1,259 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
import Image from '@theme/IdealImage';
|
||||
|
||||
# Agent Permission Management
|
||||
|
||||
Control which A2A agents can be accessed by specific keys or teams in LiteLLM.
|
||||
|
||||
## Overview
|
||||
|
||||
Agent Permission Management lets you restrict which agents a LiteLLM Virtual Key or Team can access. This is useful for:
|
||||
|
||||
- **Multi-tenant environments**: Give different teams access to different agents
|
||||
- **Security**: Prevent keys from invoking agents they shouldn't have access to
|
||||
- **Compliance**: Enforce access policies for sensitive agent workflows
|
||||
|
||||
When permissions are configured:
|
||||
- `GET /v1/agents` only returns agents the key/team can access
|
||||
- `POST /a2a/{agent_id}` (Invoking an agent) returns `403 Forbidden` if access is denied
|
||||
|
||||
## Setting Permissions on a Key
|
||||
|
||||
This example shows how to create a key with agent permissions and test access.
|
||||
|
||||
### 1. Get Your Agent ID
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="ui" label="UI">
|
||||
|
||||
1. Go to **Agents** in the sidebar
|
||||
2. Click into the agent you want
|
||||
3. Copy the **Agent ID**
|
||||
|
||||
<Image
|
||||
img={require('../img/agent_id.png')}
|
||||
style={{width: '80%', display: 'block', margin: '0', borderRadius: '8px'}}
|
||||
/>
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="api" label="API">
|
||||
|
||||
```bash title="List all agents" showLineNumbers
|
||||
curl "http://localhost:4000/v1/agents" \
|
||||
-H "Authorization: Bearer sk-master-key"
|
||||
```
|
||||
|
||||
Response:
|
||||
```json title="Response" showLineNumbers
|
||||
{
|
||||
"agents": [
|
||||
{"agent_id": "agent-123", "name": "Support Agent"},
|
||||
{"agent_id": "agent-456", "name": "Sales Agent"}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### 2. Create a Key with Agent Permissions
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="ui" label="UI">
|
||||
|
||||
1. Go to **Keys** → **Create Key**
|
||||
2. Expand **Agent Settings**
|
||||
3. Select the agents you want to allow
|
||||
|
||||
<Image
|
||||
img={require('../img/agent_key.png')}
|
||||
style={{width: '80%', display: 'block', margin: '0', borderRadius: '8px'}}
|
||||
/>
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="api" label="API">
|
||||
|
||||
```bash title="Create key with agent permissions" showLineNumbers
|
||||
curl -X POST "http://localhost:4000/key/generate" \
|
||||
-H "Authorization: Bearer sk-master-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"object_permission": {
|
||||
"agents": ["agent-123"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### 3. Test Access
|
||||
|
||||
**Allowed agent (succeeds):**
|
||||
```bash title="Invoke allowed agent" showLineNumbers
|
||||
curl -X POST "http://localhost:4000/a2a/agent-123" \
|
||||
-H "Authorization: Bearer sk-your-new-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"message": {"role": "user", "parts": [{"type": "text", "text": "Hello"}]}}'
|
||||
```
|
||||
|
||||
**Blocked agent (fails with 403):**
|
||||
```bash title="Invoke blocked agent" showLineNumbers
|
||||
curl -X POST "http://localhost:4000/a2a/agent-456" \
|
||||
-H "Authorization: Bearer sk-your-new-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"message": {"role": "user", "parts": [{"type": "text", "text": "Hello"}]}}'
|
||||
```
|
||||
|
||||
Response:
|
||||
```json title="403 Forbidden Response" showLineNumbers
|
||||
{
|
||||
"error": {
|
||||
"message": "Access denied to agent: agent-456",
|
||||
"code": 403
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Setting Permissions on a Team
|
||||
|
||||
Restrict all keys belonging to a team to only access specific agents.
|
||||
|
||||
### 1. Create a Team with Agent Permissions
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="ui" label="UI">
|
||||
|
||||
1. Go to **Teams** → **Create Team**
|
||||
2. Expand **Agent Settings**
|
||||
3. Select the agents you want to allow for this team
|
||||
|
||||
<Image
|
||||
img={require('../img/agent_key.png')}
|
||||
style={{width: '80%', display: 'block', margin: '0', borderRadius: '8px'}}
|
||||
/>
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="api" label="API">
|
||||
|
||||
```bash title="Create team with agent permissions" showLineNumbers
|
||||
curl -X POST "http://localhost:4000/team/new" \
|
||||
-H "Authorization: Bearer sk-master-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"team_alias": "support-team",
|
||||
"object_permission": {
|
||||
"agents": ["agent-123"]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
Response:
|
||||
```json title="Response" showLineNumbers
|
||||
{
|
||||
"team_id": "team-abc-123",
|
||||
"team_alias": "support-team"
|
||||
}
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### 2. Create a Key for the Team
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="ui" label="UI">
|
||||
|
||||
1. Go to **Keys** → **Create Key**
|
||||
2. Select the **Team** from the dropdown
|
||||
|
||||
<Image
|
||||
img={require('../img/agent_team.png')}
|
||||
style={{width: '80%', display: 'block', margin: '0', borderRadius: '8px'}}
|
||||
/>
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="api" label="API">
|
||||
|
||||
```bash title="Create key for team" showLineNumbers
|
||||
curl -X POST "http://localhost:4000/key/generate" \
|
||||
-H "Authorization: Bearer sk-master-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"team_id": "team-abc-123"
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### 3. Test Access
|
||||
|
||||
The key inherits agent permissions from the team.
|
||||
|
||||
**Allowed agent (succeeds):**
|
||||
```bash title="Invoke allowed agent" showLineNumbers
|
||||
curl -X POST "http://localhost:4000/a2a/agent-123" \
|
||||
-H "Authorization: Bearer sk-team-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"message": {"role": "user", "parts": [{"type": "text", "text": "Hello"}]}}'
|
||||
```
|
||||
|
||||
**Blocked agent (fails with 403):**
|
||||
```bash title="Invoke blocked agent" showLineNumbers
|
||||
curl -X POST "http://localhost:4000/a2a/agent-456" \
|
||||
-H "Authorization: Bearer sk-team-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{"message": {"role": "user", "parts": [{"type": "text", "text": "Hello"}]}}'
|
||||
```
|
||||
|
||||
## How It Works
|
||||
|
||||
```mermaid
|
||||
flowchart TD
|
||||
A[Request to invoke agent] --> B{LiteLLM Virtual Key has agent restrictions?}
|
||||
B -->|Yes| C{LiteLLM Team has agent restrictions?}
|
||||
B -->|No| D{LiteLLM Team has agent restrictions?}
|
||||
|
||||
C -->|Yes| E[Use intersection of key + team permissions]
|
||||
C -->|No| F[Use key permissions only]
|
||||
|
||||
D -->|Yes| G[Inherit team permissions]
|
||||
D -->|No| H[Allow ALL agents]
|
||||
|
||||
E --> I{Agent in allowed list?}
|
||||
F --> I
|
||||
G --> I
|
||||
H --> J[Allow request]
|
||||
|
||||
I -->|Yes| J
|
||||
I -->|No| K[Return 403 Forbidden]
|
||||
```
|
||||
|
||||
| Key Permissions | Team Permissions | Result | Notes |
|
||||
|-----------------|------------------|--------|-------|
|
||||
| None | None | Key can access **all** agents | Open access by default when no restrictions are set |
|
||||
| `["agent-1", "agent-2"]` | None | Key can access `agent-1` and `agent-2` | Key uses its own permissions |
|
||||
| None | `["agent-1", "agent-3"]` | Key can access `agent-1` and `agent-3` | Key inherits team's permissions |
|
||||
| `["agent-1", "agent-2"]` | `["agent-1", "agent-3"]` | Key can access `agent-1` only | Intersection of both lists (most restrictive wins) |
|
||||
|
||||
## Viewing Permissions
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="ui" label="UI">
|
||||
|
||||
1. Go to **Keys** or **Teams**
|
||||
2. Click into the key/team you want to view
|
||||
3. Agent permissions are displayed in the info view
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="api" label="API">
|
||||
|
||||
```bash title="Get key info" showLineNumbers
|
||||
curl "http://localhost:4000/key/info?key=sk-your-key" \
|
||||
-H "Authorization: Bearer sk-master-key"
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
|
@ -21,6 +21,20 @@ The **Generic Guardrail API** lets you integrate with LiteLLM **instantly** by i
|
|||
5. **Custom Parameters** - Pass provider-specific params via config
|
||||
6. **Full Control** - You own and maintain your guardrail API
|
||||
|
||||
## Supported Endpoints
|
||||
|
||||
The Generic Guardrail API works with the following LiteLLM endpoints:
|
||||
|
||||
- `/v1/chat/completions` - OpenAI Chat Completions
|
||||
- `/v1/completions` - OpenAI Text Completions
|
||||
- `/v1/responses` - OpenAI Responses API
|
||||
- `/v1/images/generations` - OpenAI Image Generation
|
||||
- `/v1/audio/transcriptions` - OpenAI Audio Transcriptions
|
||||
- `/v1/audio/speech` - OpenAI Text-to-Speech
|
||||
- `/v1/messages` - Anthropic Messages
|
||||
- `/v1/rerank` - Cohere Rerank
|
||||
- Pass-through endpoints
|
||||
|
||||
## How It Works
|
||||
|
||||
1. LiteLLM extracts text and images from any request (chat messages, embeddings, image prompts, etc.)
|
||||
|
|
@ -40,6 +54,35 @@ Implement `POST /beta/litellm_basic_guardrail_api`
|
|||
{
|
||||
"texts": ["extracted text from the request"], // array of text strings
|
||||
"images": ["base64_encoded_image_data"], // optional array of images
|
||||
"tools": [ // tool calls sent to the LLM (in the OpenAI Chat Completions spec)
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get the current weather",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": {"type": "string"}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
],
|
||||
"tool_calls": [ // tool calls received from the LLM (in the OpenAI Chat Completions spec)
|
||||
{
|
||||
"id": "call_abc123",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": "{\"location\": \"San Francisco\"}"
|
||||
}
|
||||
}
|
||||
],
|
||||
"structured_messages": [ // optional, full messages in OpenAI format (for chat endpoints)
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "Hello"}
|
||||
],
|
||||
"request_data": {
|
||||
"user_api_key_hash": "hash of the litellm virtual key used",
|
||||
"user_api_key_alias": "alias of the litellm virtual key used",
|
||||
|
|
@ -75,6 +118,106 @@ Implement `POST /beta/litellm_basic_guardrail_api`
|
|||
- `NONE` - Request proceeds unchanged
|
||||
- `GUARDRAIL_INTERVENED` - Request proceeds with modified texts/images (provide `texts` and/or `images` fields)
|
||||
|
||||
## Parameters
|
||||
|
||||
### `tools` Parameter
|
||||
|
||||
The `tools` parameter provides information about available function/tool definitions in the request.
|
||||
|
||||
**Format:** OpenAI `ChatCompletionToolParam` format (see [OpenAI API reference](https://platform.openai.com/docs/api-reference/chat/create#chat-create-tools))
|
||||
|
||||
**Example:**
|
||||
```json
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"description": "Get the current weather in a location",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"location": {
|
||||
"type": "string",
|
||||
"description": "City and state, e.g. San Francisco, CA"
|
||||
},
|
||||
"unit": {
|
||||
"type": "string",
|
||||
"enum": ["celsius", "fahrenheit"]
|
||||
}
|
||||
},
|
||||
"required": ["location"]
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Availability:**
|
||||
- **Input only:** Tools are only passed for `input_type="request"` (pre-call guardrails). Output/response guardrails do not currently receive tool definitions.
|
||||
- **Supported endpoints:** The `tools` parameter is supported on: `/v1/chat/completions`, `/v1/responses`, and `/v1/messages`. Other endpoints do not have tool support.
|
||||
|
||||
**Use cases:**
|
||||
- Enforce tool permission policies (e.g., only allow certain users/teams to access specific tools)
|
||||
- Validate tool schemas before sending to LLM
|
||||
- Log tool usage for audit purposes
|
||||
- Block sensitive tools based on user context
|
||||
|
||||
### `tool_calls` Parameter
|
||||
|
||||
The `tool_calls` parameter contains actual function/tool invocations being made in the request or response.
|
||||
|
||||
**Format:** OpenAI `ChatCompletionMessageToolCall` format (see [OpenAI API reference](https://platform.openai.com/docs/api-reference/chat/object#chat/object-tool_calls))
|
||||
|
||||
**Example:**
|
||||
```json
|
||||
{
|
||||
"id": "call_abc123",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "get_weather",
|
||||
"arguments": "{\"location\": \"San Francisco\", \"unit\": \"celsius\"}"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
**Key Difference from `tools`:**
|
||||
- **`tools`** = Tool definitions/schemas (what tools are *available*)
|
||||
- **`tool_calls`** = Tool invocations/executions (what tools are *being called* with what arguments)
|
||||
|
||||
**Availability:**
|
||||
- **Both input and output:** Tool calls can be present in both `input_type="request"` (assistant messages requesting tool calls) and `input_type="response"` (LLM responses with tool calls).
|
||||
- **Supported endpoints:** The `tool_calls` parameter is supported on: `/v1/chat/completions`, `/v1/responses`, and `/v1/messages`.
|
||||
|
||||
**Use cases:**
|
||||
- Validate tool call arguments before execution
|
||||
- Redact sensitive data from tool call arguments (e.g., PII)
|
||||
- Log tool invocations for audit/debugging
|
||||
- Block tool calls with dangerous parameters
|
||||
- Modify tool call arguments (e.g., enforce constraints, sanitize inputs)
|
||||
- Monitor tool usage patterns across users/teams
|
||||
|
||||
### `structured_messages` Parameter
|
||||
|
||||
The `structured_messages` parameter provides the full input in OpenAI chat completion spec format, useful for distinguishing between system and user messages.
|
||||
|
||||
**Format:** Array of OpenAI chat completion messages (see [OpenAI API reference](https://platform.openai.com/docs/api-reference/chat/create#chat-create-messages))
|
||||
|
||||
**Example:**
|
||||
```json
|
||||
[
|
||||
{"role": "system", "content": "You are a helpful assistant"},
|
||||
{"role": "user", "content": "Hello"}
|
||||
]
|
||||
```
|
||||
|
||||
**Availability:**
|
||||
- **Supported endpoints:** `/v1/chat/completions`, `/v1/messages`, `/v1/responses`
|
||||
- **Input only:** Only passed for `input_type="request"` (pre-call guardrails)
|
||||
|
||||
**Use cases:**
|
||||
- Apply different policies for system vs user messages
|
||||
- Enforce role-based content restrictions
|
||||
- Log structured conversation context
|
||||
|
||||
## LiteLLM Configuration
|
||||
|
||||
Add to `config.yaml`:
|
||||
|
|
@ -138,6 +281,9 @@ app = FastAPI()
|
|||
class GuardrailRequest(BaseModel):
|
||||
texts: List[str]
|
||||
images: Optional[List[str]] = None
|
||||
tools: Optional[List[Dict[str, Any]]] = None # OpenAI ChatCompletionToolParam format (tool definitions)
|
||||
tool_calls: Optional[List[Dict[str, Any]]] = None # OpenAI ChatCompletionMessageToolCall format (tool invocations)
|
||||
structured_messages: Optional[List[Dict[str, Any]]] = None # OpenAI messages format (for chat endpoints)
|
||||
request_data: Dict[str, Any]
|
||||
input_type: str # "request" or "response"
|
||||
litellm_call_id: Optional[str] = None
|
||||
|
|
@ -153,6 +299,8 @@ class GuardrailResponse(BaseModel):
|
|||
@app.post("/beta/litellm_basic_guardrail_api")
|
||||
async def apply_guardrail(request: GuardrailRequest):
|
||||
# Your guardrail logic here
|
||||
|
||||
# Example: Check text content
|
||||
for text in request.texts:
|
||||
if "badword" in text.lower():
|
||||
return GuardrailResponse(
|
||||
|
|
@ -160,6 +308,49 @@ async def apply_guardrail(request: GuardrailRequest):
|
|||
blocked_reason="Content contains prohibited terms"
|
||||
)
|
||||
|
||||
# Example: Check tool definitions (if present in request)
|
||||
if request.tools:
|
||||
for tool in request.tools:
|
||||
if tool.get("type") == "function":
|
||||
function_name = tool.get("function", {}).get("name", "")
|
||||
# Block sensitive tool definitions
|
||||
if function_name in ["delete_data", "access_admin_panel"]:
|
||||
return GuardrailResponse(
|
||||
action="BLOCKED",
|
||||
blocked_reason=f"Tool '{function_name}' is not allowed"
|
||||
)
|
||||
|
||||
# Example: Check tool calls (if present in request or response)
|
||||
if request.tool_calls:
|
||||
for tool_call in request.tool_calls:
|
||||
if tool_call.get("type") == "function":
|
||||
function_name = tool_call.get("function", {}).get("name", "")
|
||||
arguments_str = tool_call.get("function", {}).get("arguments", "{}")
|
||||
|
||||
# Parse arguments and validate
|
||||
import json
|
||||
try:
|
||||
arguments = json.loads(arguments_str)
|
||||
# Block dangerous arguments
|
||||
if "file_path" in arguments and ".." in str(arguments["file_path"]):
|
||||
return GuardrailResponse(
|
||||
action="BLOCKED",
|
||||
blocked_reason="Tool call contains path traversal attempt"
|
||||
)
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
# Example: Check structured messages (if present in request)
|
||||
if request.structured_messages:
|
||||
for message in request.structured_messages:
|
||||
if message.get("role") == "system":
|
||||
# Apply stricter policies to system messages
|
||||
if "admin" in message.get("content", "").lower():
|
||||
return GuardrailResponse(
|
||||
action="BLOCKED",
|
||||
blocked_reason="System message contains restricted terms"
|
||||
)
|
||||
|
||||
return GuardrailResponse(action="NONE")
|
||||
```
|
||||
|
||||
|
|
|
|||
|
|
@ -371,6 +371,22 @@ model_list:
|
|||
web_search_options: {} # Enables web search with default settings
|
||||
```
|
||||
|
||||
### Advanced
|
||||
You can configure LiteLLM's router to optionally drop models that do not support WebSearch, for example
|
||||
```yaml
|
||||
- model_name: gpt-4.1
|
||||
litellm_params:
|
||||
model: openai/gpt-4.1
|
||||
- model_name: gpt-4.1
|
||||
litellm_params:
|
||||
model: azure/gpt-4.1
|
||||
api_base: "x.openai.azure.com/"
|
||||
api_version: 2025-03-01-preview
|
||||
model_info:
|
||||
supports_web_search: False <---- KEY CHANGE!
|
||||
```
|
||||
In this example, LiteLLM will still route LLM requests to both deployments, but for WebSearch, will solely route to OpenAI.
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="custom" label="Custom Search Context">
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,130 @@
|
|||
# Adding OpenAI-Compatible Providers
|
||||
|
||||
For simple OpenAI-compatible providers (like Hyperbolic, Nscale, etc.), you can add support by editing a single JSON file.
|
||||
|
||||
## Quick Start
|
||||
|
||||
1. Edit `litellm/llms/openai_like/providers.json`
|
||||
2. Add your provider configuration
|
||||
3. Test with: `litellm.completion(model="your_provider/model-name", ...)`
|
||||
|
||||
## Basic Configuration
|
||||
|
||||
For a fully OpenAI-compatible provider:
|
||||
|
||||
```json
|
||||
{
|
||||
"your_provider": {
|
||||
"base_url": "https://api.yourprovider.com/v1",
|
||||
"api_key_env": "YOUR_PROVIDER_API_KEY"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
That's it! The provider is now available.
|
||||
|
||||
## Configuration Options
|
||||
|
||||
### Required Fields
|
||||
|
||||
- `base_url` - API endpoint (e.g., `https://api.provider.com/v1`)
|
||||
- `api_key_env` - Environment variable name for API key (e.g., `PROVIDER_API_KEY`)
|
||||
|
||||
### Optional Fields
|
||||
|
||||
- `api_base_env` - Environment variable to override `base_url`
|
||||
- `base_class` - Use `"openai_gpt"` (default) or `"openai_like"`
|
||||
- `param_mappings` - Map OpenAI parameter names to provider-specific names
|
||||
- `constraints` - Parameter value constraints (min/max)
|
||||
- `special_handling` - Special behaviors like content format conversion
|
||||
|
||||
## Examples
|
||||
|
||||
### Simple Provider (Fully Compatible)
|
||||
|
||||
```json
|
||||
{
|
||||
"hyperbolic": {
|
||||
"base_url": "https://api.hyperbolic.xyz/v1",
|
||||
"api_key_env": "HYPERBOLIC_API_KEY"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Provider with Parameter Mapping
|
||||
|
||||
```json
|
||||
{
|
||||
"publicai": {
|
||||
"base_url": "https://api.publicai.co/v1",
|
||||
"api_key_env": "PUBLICAI_API_KEY",
|
||||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Provider with Constraints
|
||||
|
||||
```json
|
||||
{
|
||||
"custom_provider": {
|
||||
"base_url": "https://api.custom.com/v1",
|
||||
"api_key_env": "CUSTOM_API_KEY",
|
||||
"constraints": {
|
||||
"temperature_max": 1.0,
|
||||
"temperature_min": 0.0
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
import litellm
|
||||
import os
|
||||
|
||||
# Set your API key
|
||||
os.environ["YOUR_PROVIDER_API_KEY"] = "your-key-here"
|
||||
|
||||
# Use the provider
|
||||
response = litellm.completion(
|
||||
model="your_provider/model-name",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
)
|
||||
```
|
||||
|
||||
## When to Use Python Instead
|
||||
|
||||
Use a Python config class if you need:
|
||||
|
||||
- Custom authentication flows (OAuth, JWT, etc.)
|
||||
- Complex request/response transformations
|
||||
- Provider-specific streaming logic
|
||||
- Advanced tool calling modifications
|
||||
|
||||
For these cases, create a config class in `litellm/llms/your_provider/chat/transformation.py` that inherits from `OpenAIGPTConfig` or `OpenAILikeChatConfig`.
|
||||
|
||||
## Testing
|
||||
|
||||
Test your provider:
|
||||
|
||||
```bash
|
||||
# Quick test
|
||||
python -c "
|
||||
import litellm
|
||||
import os
|
||||
os.environ['PROVIDER_API_KEY'] = 'your-key'
|
||||
response = litellm.completion(
|
||||
model='provider/model-name',
|
||||
messages=[{'role': 'user', 'content': 'test'}]
|
||||
)
|
||||
print(response.choices[0].message.content)
|
||||
"
|
||||
```
|
||||
|
||||
## Reference
|
||||
|
||||
See existing providers in `litellm/llms/openai_like/providers.json` for examples.
|
||||
|
|
@ -301,6 +301,17 @@ content = await litellm.afile_content(
|
|||
print("file content=", content)
|
||||
```
|
||||
|
||||
**Get File Content (Bedrock)**
|
||||
```python
|
||||
# For Bedrock batch output files stored in S3
|
||||
content = await litellm.afile_content(
|
||||
file_id="s3://bucket-name/path/to/file.jsonl", # S3 URI or unified file ID
|
||||
custom_llm_provider="bedrock",
|
||||
aws_region_name="us-west-2"
|
||||
)
|
||||
print("file content=", content.text)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
|
|
@ -313,4 +324,6 @@ print("file content=", content)
|
|||
|
||||
### [Vertex AI](./providers/vertex#batch-apis)
|
||||
|
||||
### [Bedrock](./providers/bedrock_batches#4-retrieve-batch-results)
|
||||
|
||||
## [Swagger API Reference](https://litellm-api.up.railway.app/#/files)
|
||||
|
|
|
|||
|
|
@ -2,6 +2,12 @@
|
|||
title: "Integrate as a Model Provider"
|
||||
---
|
||||
|
||||
## Quick Start for OpenAI-Compatible Providers
|
||||
|
||||
If your API is OpenAI-compatible, you can add support by editing a single JSON file. See [Adding OpenAI-Compatible Providers](/docs/contributing/adding_openai_compatible_providers) for the simple approach.
|
||||
|
||||
---
|
||||
|
||||
This guide focuses on how to setup the classes and configuration necessary to act as a chat provider.
|
||||
|
||||
Please see this guide first and look at the existing code in the codebase to understand how to act as a different provider, e.g. handling embeddings or image-generation.
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ ALL Bedrock models (Anthropic, Meta, Deepseek, Mistral, Amazon, etc.) are Suppor
|
|||
| Property | Details |
|
||||
|-------|-------|
|
||||
| Description | Amazon Bedrock is a fully managed service that offers a choice of high-performing foundation models (FMs). |
|
||||
| Provider Route on LiteLLM | `bedrock/`, [`bedrock/converse/`](#set-converse--invoke-route), [`bedrock/invoke/`](#set-invoke-route), [`bedrock/converse_like/`](#calling-via-internal-proxy), [`bedrock/llama/`](#deepseek-not-r1), [`bedrock/deepseek_r1/`](#deepseek-r1), [`bedrock/qwen3/`](#qwen3-imported-models), [`bedrock/openai/`](./bedrock_imported.md#openai-compatible-imported-models-qwen-25-vl-etc) |
|
||||
| Provider Route on LiteLLM | `bedrock/`, [`bedrock/converse/`](#set-converse--invoke-route), [`bedrock/invoke/`](#set-invoke-route), [`bedrock/converse_like/`](#calling-via-internal-proxy), [`bedrock/llama/`](#deepseek-not-r1), [`bedrock/deepseek_r1/`](#deepseek-r1), [`bedrock/qwen3/`](#qwen3-imported-models), [`bedrock/qwen2/`](./bedrock_imported.md#qwen2-imported-models), [`bedrock/openai/`](./bedrock_imported.md#openai-compatible-imported-models-qwen-25-vl-etc) |
|
||||
| Provider Doc | [Amazon Bedrock ↗](https://docs.aws.amazon.com/bedrock/latest/userguide/what-is-bedrock.html) |
|
||||
| Supported OpenAI Endpoints | `/chat/completions`, `/completions`, `/embeddings`, `/images/generations` |
|
||||
| Rerank Endpoint | `/rerank` |
|
||||
|
|
|
|||
|
|
@ -172,6 +172,97 @@ curl http://localhost:4000/v1/batches \
|
|||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### 4. Retrieve batch results
|
||||
|
||||
Once the batch job is completed, download the results from S3:
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python">
|
||||
|
||||
```python showLineNumbers title="bedrock_batch.py"
|
||||
...
|
||||
# Wait for batch completion (check status periodically)
|
||||
batch_status = client.batches.retrieve(batch_id=batch.id)
|
||||
|
||||
if batch_status.status == "completed":
|
||||
# Download the output file
|
||||
result = client.files.content(
|
||||
file_id=batch_status.output_file_id,
|
||||
extra_headers={"custom-llm-provider": "bedrock"}
|
||||
)
|
||||
|
||||
# Save or process the results
|
||||
with open("batch_output.jsonl", "wb") as f:
|
||||
f.write(result.content)
|
||||
|
||||
# Parse JSONL results
|
||||
for line in result.text.strip().split('\n'):
|
||||
record = json.loads(line)
|
||||
print(f"Record ID: {record['recordId']}")
|
||||
print(f"Output: {record.get('modelOutput', {})}")
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="curl" label="Curl">
|
||||
|
||||
```bash showLineNumbers title="Download Batch Results"
|
||||
# First retrieve batch to get output_file_id
|
||||
curl http://localhost:4000/v1/batches/batch_abc123 \
|
||||
-H "Authorization: Bearer sk-1234"
|
||||
|
||||
# Then download the output file
|
||||
curl http://localhost:4000/v1/files/{output_file_id}/content \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "custom-llm-provider: bedrock" \
|
||||
-o batch_output.jsonl
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="litellm-direct" label="LiteLLM Direct">
|
||||
|
||||
```python showLineNumbers title="bedrock_batch.py"
|
||||
import litellm
|
||||
from litellm import file_content
|
||||
|
||||
# Download using litellm directly (bypasses proxy managed files)
|
||||
result = file_content(
|
||||
file_id=batch_status.output_file_id, # Can be S3 URI or unified file ID
|
||||
custom_llm_provider="bedrock",
|
||||
aws_region_name="us-west-2",
|
||||
)
|
||||
|
||||
# Process results
|
||||
print(result.text)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
**Output Format:**
|
||||
|
||||
The batch output file is in JSONL format with each line containing:
|
||||
|
||||
```json
|
||||
{
|
||||
"recordId": "request-1",
|
||||
"modelInput": {
|
||||
"messages": [...],
|
||||
"max_tokens": 1000
|
||||
},
|
||||
"modelOutput": {
|
||||
"content": [...],
|
||||
"id": "msg_abc123",
|
||||
"model": "claude-3-5-sonnet-20240620-v1:0",
|
||||
"role": "assistant",
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {
|
||||
"input_tokens": 15,
|
||||
"output_tokens": 10
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## FAQ
|
||||
|
||||
### Where are my files written?
|
||||
|
|
|
|||
|
|
@ -203,6 +203,71 @@ curl --location 'http://0.0.0.0:4000/chat/completions' \
|
|||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Qwen2 Imported Models
|
||||
|
||||
| Property | Details |
|
||||
|----------|---------|
|
||||
| Provider Route | `bedrock/qwen2/{model_arn}` |
|
||||
| Provider Documentation | [Bedrock Imported Models](https://docs.aws.amazon.com/bedrock/latest/userguide/model-customization-import-model.html) |
|
||||
| Note | Qwen2 and Qwen3 architectures are mostly similar. The main difference is in the response format: Qwen2 uses "text" field while Qwen3 uses "generation" field. |
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
import os
|
||||
|
||||
response = completion(
|
||||
model="bedrock/qwen2/arn:aws:bedrock:us-east-1:086734376398:imported-model/your-qwen2-model", # bedrock/qwen2/{your-model-arn}
|
||||
messages=[{"role": "user", "content": "Tell me a joke"}],
|
||||
max_tokens=100,
|
||||
temperature=0.7
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="Proxy">
|
||||
|
||||
**1. Add to config**
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: Qwen2-72B
|
||||
litellm_params:
|
||||
model: bedrock/qwen2/arn:aws:bedrock:us-east-1:086734376398:imported-model/your-qwen2-model
|
||||
|
||||
```
|
||||
|
||||
**2. Start proxy**
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
|
||||
# RUNNING at http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
**3. Test it!**
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/chat/completions' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"model": "Qwen2-72B", # 👈 the 'model_name' in config
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": "what llm are you"
|
||||
}
|
||||
],
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### OpenAI-Compatible Imported Models (Qwen 2.5 VL, etc.)
|
||||
|
||||
Use this route for Bedrock imported models that follow the **OpenAI Chat Completions API spec**. This includes models like Qwen 2.5 VL that accept OpenAI-formatted messages with support for vision (images), tool calling, and other OpenAI features.
|
||||
|
|
|
|||
|
|
@ -15,7 +15,7 @@ https://docs.github.com/en/copilot
|
|||
|-------|-------|
|
||||
| Description | GitHub Copilot Chat API provides access to GitHub's AI-powered coding assistant. |
|
||||
| Provider Route on LiteLLM | `github_copilot/` |
|
||||
| Supported Endpoints | `/chat/completions` |
|
||||
| Supported Endpoints | `/chat/completions`, `/embeddings` |
|
||||
| API Reference | [GitHub Copilot docs](https://docs.github.com/en/copilot) |
|
||||
|
||||
## Authentication
|
||||
|
|
@ -62,6 +62,34 @@ for chunk in stream:
|
|||
print(chunk.choices[0].delta.content, end="")
|
||||
```
|
||||
|
||||
### Responses
|
||||
|
||||
For GPT Codex models, only responses API is supported.
|
||||
|
||||
```python showLineNumbers title="GitHub Copilot Responses"
|
||||
import litellm
|
||||
|
||||
response = await litellm.aresponses(
|
||||
model="github_copilot/gpt-5.1-codex",
|
||||
input="Write a Python hello world",
|
||||
max_output_tokens=500
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
### Embedding
|
||||
|
||||
```python showLineNumbers title="GitHub Copilot Embedding"
|
||||
import litellm
|
||||
|
||||
response = litellm.embedding(
|
||||
model="github_copilot/text-embedding-3-small",
|
||||
input=["good morning from litellm"]
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
|
||||
## Usage - LiteLLM Proxy
|
||||
|
||||
Add the following to your LiteLLM Proxy configuration file:
|
||||
|
|
@ -71,6 +99,16 @@ model_list:
|
|||
- model_name: github_copilot/gpt-4
|
||||
litellm_params:
|
||||
model: github_copilot/gpt-4
|
||||
- model_name: github_copilot/gpt-5.1-codex
|
||||
model_info:
|
||||
mode: responses
|
||||
litellm_params:
|
||||
model: github_copilot/gpt-5.1-codex
|
||||
- model_name: github_copilot/text-embedding-ada-002
|
||||
model_info:
|
||||
mode: embedding
|
||||
litellm_params:
|
||||
model: github_copilot/text-embedding-ada-002
|
||||
```
|
||||
|
||||
Start your LiteLLM Proxy server:
|
||||
|
|
@ -180,7 +218,7 @@ extra_headers = {
|
|||
"editor-version": "vscode/1.85.1", # Editor version
|
||||
"editor-plugin-version": "copilot/1.155.0", # Plugin version
|
||||
"Copilot-Integration-Id": "vscode-chat", # Integration ID
|
||||
"user-agent": "GithubCopilot/1.155.0" # User agent
|
||||
"user-agent": "GithubCopilot/1.155.0" # User agent
|
||||
}
|
||||
```
|
||||
|
||||
|
|
|
|||
|
|
@ -11,7 +11,7 @@ Selecting `openai` as the provider routes your request to an OpenAI-compatible e
|
|||
This library **requires** an API key for all requests, either through the `api_key` parameter
|
||||
or the `OPENAI_API_KEY` environment variable.
|
||||
|
||||
If you don’t want to provide a fake API key in each request, consider using a provider that directly matches your
|
||||
If you don't want to provide a fake API key in each request, consider using a provider that directly matches your
|
||||
OpenAI-compatible endpoint, such as [`hosted_vllm`](/docs/providers/vllm) or [`llamafile`](/docs/providers/llamafile).
|
||||
|
||||
:::
|
||||
|
|
@ -150,4 +150,4 @@ model_list:
|
|||
api_base: http://my-custom-base
|
||||
api_key: ""
|
||||
supports_system_message: False # 👈 KEY CHANGE
|
||||
```
|
||||
```
|
||||
|
|
|
|||
|
|
@ -1604,6 +1604,53 @@ litellm.vertex_location = "us-central1 # Your Location
|
|||
| gemini-2.5-flash-preview-09-2025 | `completion('gemini-2.5-flash-preview-09-2025', messages)`, `completion('vertex_ai/gemini-2.5-flash-preview-09-2025', messages)` |
|
||||
| gemini-2.5-flash-lite-preview-09-2025 | `completion('gemini-2.5-flash-lite-preview-09-2025', messages)`, `completion('vertex_ai/gemini-2.5-flash-lite-preview-09-2025', messages)` |
|
||||
|
||||
## Private Service Connect (PSC) Endpoints
|
||||
|
||||
LiteLLM supports Vertex AI models deployed to Private Service Connect (PSC) endpoints, allowing you to use custom `api_base` URLs for private deployments.
|
||||
|
||||
### Usage
|
||||
|
||||
```python
|
||||
from litellm import completion
|
||||
|
||||
# Use PSC endpoint with custom api_base
|
||||
response = completion(
|
||||
model="vertex_ai/1234567890", # Numeric endpoint ID
|
||||
messages=[{"role": "user", "content": "Hello!"}],
|
||||
api_base="http://10.96.32.8", # Your PSC endpoint
|
||||
vertex_project="my-project-id",
|
||||
vertex_location="us-central1"
|
||||
)
|
||||
```
|
||||
|
||||
**Key Features:**
|
||||
- Supports both numeric endpoint IDs and custom model names
|
||||
- Works with both completion and embedding endpoints
|
||||
- Automatically constructs full PSC URL: `{api_base}/v1/projects/{project}/locations/{location}/endpoints/{model}:{endpoint}`
|
||||
- Compatible with streaming requests
|
||||
|
||||
### Configuration
|
||||
|
||||
Add PSC endpoints to your `config.yaml`:
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: psc-gemini
|
||||
litellm_params:
|
||||
model: vertex_ai/1234567890 # Numeric endpoint ID
|
||||
api_base: "http://10.96.32.8" # Your PSC endpoint
|
||||
vertex_project: "my-project-id"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: "/path/to/service_account.json"
|
||||
- model_name: psc-embedding
|
||||
litellm_params:
|
||||
model: vertex_ai/text-embedding-004
|
||||
api_base: "http://10.96.32.8" # Your PSC endpoint
|
||||
vertex_project: "my-project-id"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: "/path/to/service_account.json"
|
||||
```
|
||||
|
||||
## Fine-tuned Models
|
||||
|
||||
You can call fine-tuned Vertex AI Gemini models through LiteLLM
|
||||
|
|
|
|||
587
docs/my-website/docs/providers/vertex_embedding.md
Normal file
587
docs/my-website/docs/providers/vertex_embedding.md
Normal file
|
|
@ -0,0 +1,587 @@
|
|||
import Image from '@theme/IdealImage';
|
||||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Vertex AI Embedding
|
||||
|
||||
## Usage - Embedding
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
import litellm
|
||||
from litellm import embedding
|
||||
litellm.vertex_project = "hardy-device-38811" # Your Project ID
|
||||
litellm.vertex_location = "us-central1" # proj location
|
||||
|
||||
response = embedding(
|
||||
model="vertex_ai/textembedding-gecko",
|
||||
input=["good morning from litellm"],
|
||||
)
|
||||
print(response)
|
||||
```
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="LiteLLM PROXY">
|
||||
|
||||
|
||||
1. Add model to config.yaml
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: snowflake-arctic-embed-m-long-1731622468876
|
||||
litellm_params:
|
||||
model: vertex_ai/<your-model-id>
|
||||
vertex_project: "adroit-crow-413218"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: adroit-crow-413218-a956eef1a2a8.json
|
||||
|
||||
litellm_settings:
|
||||
drop_params: True
|
||||
```
|
||||
|
||||
2. Start Proxy
|
||||
|
||||
```
|
||||
$ litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
3. Make Request using OpenAI Python SDK, Langchain Python SDK
|
||||
|
||||
```python
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(api_key="sk-1234", base_url="http://0.0.0.0:4000")
|
||||
|
||||
response = client.embeddings.create(
|
||||
model="snowflake-arctic-embed-m-long-1731622468876",
|
||||
input = ["good morning from litellm", "this is another item"],
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
#### Supported Embedding Models
|
||||
All models listed [here](https://github.com/BerriAI/litellm/blob/57f37f743886a0249f630a6792d49dffc2c5d9b7/model_prices_and_context_window.json#L835) are supported
|
||||
|
||||
| Model Name | Function Call |
|
||||
|--------------------------|------------------------------------------------------------------------------------------------------------------------------------------------------------------|
|
||||
| text-embedding-004 | `embedding(model="vertex_ai/text-embedding-004", input)` |
|
||||
| text-multilingual-embedding-002 | `embedding(model="vertex_ai/text-multilingual-embedding-002", input)` |
|
||||
| textembedding-gecko | `embedding(model="vertex_ai/textembedding-gecko", input)` |
|
||||
| textembedding-gecko-multilingual | `embedding(model="vertex_ai/textembedding-gecko-multilingual", input)` |
|
||||
| textembedding-gecko-multilingual@001 | `embedding(model="vertex_ai/textembedding-gecko-multilingual@001", input)` |
|
||||
| textembedding-gecko@001 | `embedding(model="vertex_ai/textembedding-gecko@001", input)` |
|
||||
| textembedding-gecko@003 | `embedding(model="vertex_ai/textembedding-gecko@003", input)` |
|
||||
| text-embedding-preview-0409 | `embedding(model="vertex_ai/text-embedding-preview-0409", input)` |
|
||||
| text-multilingual-embedding-preview-0409 | `embedding(model="vertex_ai/text-multilingual-embedding-preview-0409", input)` |
|
||||
| Fine-tuned OR Custom Embedding models | `embedding(model="vertex_ai/<your-model-id>", input)` |
|
||||
|
||||
### Supported OpenAI (Unified) Params
|
||||
|
||||
| [param](../embedding/supported_embedding.md#input-params-for-litellmembedding) | type | [vertex equivalent](https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/text-embeddings-api) |
|
||||
|-------|-------------|--------------------|
|
||||
| `input` | **string or List[string]** | `instances` |
|
||||
| `dimensions` | **int** | `output_dimensionality` |
|
||||
| `input_type` | **Literal["RETRIEVAL_QUERY","RETRIEVAL_DOCUMENT", "SEMANTIC_SIMILARITY", "CLASSIFICATION", "CLUSTERING", "QUESTION_ANSWERING", "FACT_VERIFICATION"]** | `task_type` |
|
||||
|
||||
#### Usage with OpenAI (Unified) Params
|
||||
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
response = litellm.embedding(
|
||||
model="vertex_ai/text-embedding-004",
|
||||
input=["good morning from litellm", "gm"]
|
||||
input_type = "RETRIEVAL_DOCUMENT",
|
||||
dimensions=1,
|
||||
)
|
||||
```
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="LiteLLM PROXY">
|
||||
|
||||
|
||||
```python
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(api_key="sk-1234", base_url="http://0.0.0.0:4000")
|
||||
|
||||
response = client.embeddings.create(
|
||||
model="text-embedding-004",
|
||||
input = ["good morning from litellm", "gm"],
|
||||
dimensions=1,
|
||||
extra_body = {
|
||||
"input_type": "RETRIEVAL_QUERY",
|
||||
}
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
|
||||
### Supported Vertex Specific Params
|
||||
|
||||
| param | type |
|
||||
|-------|-------------|
|
||||
| `auto_truncate` | **bool** |
|
||||
| `task_type` | **Literal["RETRIEVAL_QUERY","RETRIEVAL_DOCUMENT", "SEMANTIC_SIMILARITY", "CLASSIFICATION", "CLUSTERING", "QUESTION_ANSWERING", "FACT_VERIFICATION"]** |
|
||||
| `title` | **str** |
|
||||
|
||||
#### Usage with Vertex Specific Params (Use `task_type` and `title`)
|
||||
|
||||
You can pass any vertex specific params to the embedding model. Just pass them to the embedding function like this:
|
||||
|
||||
[Relevant Vertex AI doc with all embedding params](https://cloud.google.com/vertex-ai/generative-ai/docs/model-reference/text-embeddings-api#request_body)
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python
|
||||
response = litellm.embedding(
|
||||
model="vertex_ai/text-embedding-004",
|
||||
input=["good morning from litellm", "gm"]
|
||||
task_type = "RETRIEVAL_DOCUMENT",
|
||||
title = "test",
|
||||
dimensions=1,
|
||||
auto_truncate=True,
|
||||
)
|
||||
```
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="LiteLLM PROXY">
|
||||
|
||||
|
||||
```python
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(api_key="sk-1234", base_url="http://0.0.0.0:4000")
|
||||
|
||||
response = client.embeddings.create(
|
||||
model="text-embedding-004",
|
||||
input = ["good morning from litellm", "gm"],
|
||||
dimensions=1,
|
||||
extra_body = {
|
||||
"task_type": "RETRIEVAL_QUERY",
|
||||
"auto_truncate": True,
|
||||
"title": "test",
|
||||
}
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## **BGE Embeddings**
|
||||
|
||||
Use BGE (Baidu General Embedding) models deployed on Vertex AI.
|
||||
|
||||
### Usage
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
```python showLineNumbers title="Using BGE on Vertex AI"
|
||||
import litellm
|
||||
|
||||
response = litellm.embedding(
|
||||
model="vertex_ai/bge/<your-endpoint-id>",
|
||||
input=["Hello", "World"],
|
||||
vertex_project="your-project-id",
|
||||
vertex_location="your-location"
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="proxy" label="LiteLLM PROXY">
|
||||
|
||||
1. Add model to config.yaml
|
||||
```yaml showLineNumbers title="config.yaml"
|
||||
model_list:
|
||||
- model_name: bge-embedding
|
||||
litellm_params:
|
||||
model: vertex_ai/bge/<your-endpoint-id>
|
||||
vertex_project: "your-project-id"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: your-credentials.json
|
||||
|
||||
litellm_settings:
|
||||
drop_params: True
|
||||
```
|
||||
|
||||
2. Start Proxy
|
||||
|
||||
```bash
|
||||
$ litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
3. Make Request using OpenAI Python SDK
|
||||
|
||||
```python showLineNumbers title="Making requests to BGE"
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(api_key="sk-1234", base_url="http://0.0.0.0:4000")
|
||||
|
||||
response = client.embeddings.create(
|
||||
model="bge-embedding",
|
||||
input=["good morning from litellm", "this is another item"]
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
Using a Private Service Connect (PSC) endpoint
|
||||
|
||||
```yaml showLineNumbers title="config.yaml (PSC)"
|
||||
model_list:
|
||||
- model_name: bge-small-en-v1.5
|
||||
litellm_params:
|
||||
model: vertex_ai/bge/1234567890
|
||||
api_base: http://10.96.32.8 # Your PSC IP
|
||||
vertex_project: my-project-id #optional
|
||||
vertex_location: us-central1 #optional
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## **Multi-Modal Embeddings**
|
||||
|
||||
|
||||
Known Limitations:
|
||||
- Only supports 1 image / video / image per request
|
||||
- Only supports GCS or base64 encoded images / videos
|
||||
|
||||
### Usage
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
Using GCS Images
|
||||
|
||||
```python
|
||||
response = await litellm.aembedding(
|
||||
model="vertex_ai/multimodalembedding@001",
|
||||
input="gs://cloud-samples-data/vertex-ai/llm/prompts/landmark1.png" # will be sent as a gcs image
|
||||
)
|
||||
```
|
||||
|
||||
Using base 64 encoded images
|
||||
|
||||
```python
|
||||
response = await litellm.aembedding(
|
||||
model="vertex_ai/multimodalembedding@001",
|
||||
input="data:image/jpeg;base64,..." # will be sent as a base64 encoded image
|
||||
)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="LiteLLM PROXY (Unified Endpoint)">
|
||||
|
||||
1. Add model to config.yaml
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: multimodalembedding@001
|
||||
litellm_params:
|
||||
model: vertex_ai/multimodalembedding@001
|
||||
vertex_project: "adroit-crow-413218"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: adroit-crow-413218-a956eef1a2a8.json
|
||||
|
||||
litellm_settings:
|
||||
drop_params: True
|
||||
```
|
||||
|
||||
2. Start Proxy
|
||||
|
||||
```
|
||||
$ litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
3. Make Request use OpenAI Python SDK, Langchain Python SDK
|
||||
|
||||
|
||||
<Tabs>
|
||||
|
||||
<TabItem value="OpenAI SDK" label="OpenAI SDK">
|
||||
|
||||
Requests with GCS Image / Video URI
|
||||
|
||||
```python
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(api_key="sk-1234", base_url="http://0.0.0.0:4000")
|
||||
|
||||
# # request sent to model set on litellm proxy, `litellm --model`
|
||||
response = client.embeddings.create(
|
||||
model="multimodalembedding@001",
|
||||
input = "gs://cloud-samples-data/vertex-ai/llm/prompts/landmark1.png",
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
Requests with base64 encoded images
|
||||
|
||||
```python
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(api_key="sk-1234", base_url="http://0.0.0.0:4000")
|
||||
|
||||
# # request sent to model set on litellm proxy, `litellm --model`
|
||||
response = client.embeddings.create(
|
||||
model="multimodalembedding@001",
|
||||
input = "data:image/jpeg;base64,...",
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="langchain" label="Langchain">
|
||||
|
||||
Requests with GCS Image / Video URI
|
||||
```python
|
||||
from langchain_openai import OpenAIEmbeddings
|
||||
|
||||
embeddings_models = "multimodalembedding@001"
|
||||
|
||||
embeddings = OpenAIEmbeddings(
|
||||
model="multimodalembedding@001",
|
||||
base_url="http://0.0.0.0:4000",
|
||||
api_key="sk-1234", # type: ignore
|
||||
)
|
||||
|
||||
|
||||
query_result = embeddings.embed_query(
|
||||
"gs://cloud-samples-data/vertex-ai/llm/prompts/landmark1.png"
|
||||
)
|
||||
print(query_result)
|
||||
|
||||
```
|
||||
|
||||
Requests with base64 encoded images
|
||||
|
||||
```python
|
||||
from langchain_openai import OpenAIEmbeddings
|
||||
|
||||
embeddings_models = "multimodalembedding@001"
|
||||
|
||||
embeddings = OpenAIEmbeddings(
|
||||
model="multimodalembedding@001",
|
||||
base_url="http://0.0.0.0:4000",
|
||||
api_key="sk-1234", # type: ignore
|
||||
)
|
||||
|
||||
|
||||
query_result = embeddings.embed_query(
|
||||
"data:image/jpeg;base64,..."
|
||||
)
|
||||
print(query_result)
|
||||
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
</Tabs>
|
||||
</TabItem>
|
||||
|
||||
|
||||
<TabItem value="proxy-vtx" label="LiteLLM PROXY (Vertex SDK)">
|
||||
|
||||
1. Add model to config.yaml
|
||||
```yaml
|
||||
default_vertex_config:
|
||||
vertex_project: "adroit-crow-413218"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: adroit-crow-413218-a956eef1a2a8.json
|
||||
```
|
||||
|
||||
2. Start Proxy
|
||||
|
||||
```
|
||||
$ litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
3. Make Request use OpenAI Python SDK
|
||||
|
||||
```python
|
||||
import vertexai
|
||||
|
||||
from vertexai.vision_models import Image, MultiModalEmbeddingModel, Video
|
||||
from vertexai.vision_models import VideoSegmentConfig
|
||||
from google.auth.credentials import Credentials
|
||||
|
||||
|
||||
LITELLM_PROXY_API_KEY = "sk-1234"
|
||||
LITELLM_PROXY_BASE = "http://0.0.0.0:4000/vertex-ai"
|
||||
|
||||
import datetime
|
||||
|
||||
class CredentialsWrapper(Credentials):
|
||||
def __init__(self, token=None):
|
||||
super().__init__()
|
||||
self.token = token
|
||||
self.expiry = None # or set to a future date if needed
|
||||
|
||||
def refresh(self, request):
|
||||
pass
|
||||
|
||||
def apply(self, headers, token=None):
|
||||
headers['Authorization'] = f'Bearer {self.token}'
|
||||
|
||||
@property
|
||||
def expired(self):
|
||||
return False # Always consider the token as non-expired
|
||||
|
||||
@property
|
||||
def valid(self):
|
||||
return True # Always consider the credentials as valid
|
||||
|
||||
credentials = CredentialsWrapper(token=LITELLM_PROXY_API_KEY)
|
||||
|
||||
vertexai.init(
|
||||
project="adroit-crow-413218",
|
||||
location="us-central1",
|
||||
api_endpoint=LITELLM_PROXY_BASE,
|
||||
credentials = credentials,
|
||||
api_transport="rest",
|
||||
|
||||
)
|
||||
|
||||
model = MultiModalEmbeddingModel.from_pretrained("multimodalembedding")
|
||||
image = Image.load_from_file(
|
||||
"gs://cloud-samples-data/vertex-ai/llm/prompts/landmark1.png"
|
||||
)
|
||||
|
||||
embeddings = model.get_embeddings(
|
||||
image=image,
|
||||
contextual_text="Colosseum",
|
||||
dimension=1408,
|
||||
)
|
||||
print(f"Image Embedding: {embeddings.image_embedding}")
|
||||
print(f"Text Embedding: {embeddings.text_embedding}")
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
|
||||
### Text + Image + Video Embeddings
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="sdk" label="SDK">
|
||||
|
||||
Text + Image
|
||||
|
||||
```python
|
||||
response = await litellm.aembedding(
|
||||
model="vertex_ai/multimodalembedding@001",
|
||||
input=["hey", "gs://cloud-samples-data/vertex-ai/llm/prompts/landmark1.png"] # will be sent as a gcs image
|
||||
)
|
||||
```
|
||||
|
||||
Text + Video
|
||||
|
||||
```python
|
||||
response = await litellm.aembedding(
|
||||
model="vertex_ai/multimodalembedding@001",
|
||||
input=["hey", "gs://my-bucket/embeddings/supermarket-video.mp4"] # will be sent as a gcs image
|
||||
)
|
||||
```
|
||||
|
||||
Image + Video
|
||||
|
||||
```python
|
||||
response = await litellm.aembedding(
|
||||
model="vertex_ai/multimodalembedding@001",
|
||||
input=["gs://cloud-samples-data/vertex-ai/llm/prompts/landmark1.png", "gs://my-bucket/embeddings/supermarket-video.mp4"] # will be sent as a gcs image
|
||||
)
|
||||
```
|
||||
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="proxy" label="LiteLLM PROXY (Unified Endpoint)">
|
||||
|
||||
1. Add model to config.yaml
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: multimodalembedding@001
|
||||
litellm_params:
|
||||
model: vertex_ai/multimodalembedding@001
|
||||
vertex_project: "adroit-crow-413218"
|
||||
vertex_location: "us-central1"
|
||||
vertex_credentials: adroit-crow-413218-a956eef1a2a8.json
|
||||
|
||||
litellm_settings:
|
||||
drop_params: True
|
||||
```
|
||||
|
||||
2. Start Proxy
|
||||
|
||||
```
|
||||
$ litellm --config /path/to/config.yaml
|
||||
```
|
||||
|
||||
3. Make Request use OpenAI Python SDK, Langchain Python SDK
|
||||
|
||||
|
||||
Text + Image
|
||||
|
||||
```python
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(api_key="sk-1234", base_url="http://0.0.0.0:4000")
|
||||
|
||||
# # request sent to model set on litellm proxy, `litellm --model`
|
||||
response = client.embeddings.create(
|
||||
model="multimodalembedding@001",
|
||||
input = ["hey", "gs://cloud-samples-data/vertex-ai/llm/prompts/landmark1.png"],
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
Text + Video
|
||||
```python
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(api_key="sk-1234", base_url="http://0.0.0.0:4000")
|
||||
|
||||
# # request sent to model set on litellm proxy, `litellm --model`
|
||||
response = client.embeddings.create(
|
||||
model="multimodalembedding@001",
|
||||
input = ["hey", "gs://my-bucket/embeddings/supermarket-video.mp4"],
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
Image + Video
|
||||
```python
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(api_key="sk-1234", base_url="http://0.0.0.0:4000")
|
||||
|
||||
# # request sent to model set on litellm proxy, `litellm --model`
|
||||
response = client.embeddings.create(
|
||||
model="multimodalembedding@001",
|
||||
input = ["gs://cloud-samples-data/vertex-ai/llm/prompts/landmark1.png", "gs://my-bucket/embeddings/supermarket-video.mp4"],
|
||||
)
|
||||
|
||||
print(response)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
|
@ -29,7 +29,8 @@ litellm_settings:
|
|||
request_timeout: 10 # (int) llm requesttimeout in seconds. Raise Timeout error if call takes longer than 10s. Sets litellm.request_timeout
|
||||
force_ipv4: boolean # If true, litellm will force ipv4 for all LLM requests. Some users have seen httpx ConnectionError when using ipv6 + Anthropic API
|
||||
|
||||
set_verbose: boolean # sets litellm.set_verbose=True to view verbose debug logs. DO NOT LEAVE THIS ON IN PRODUCTION
|
||||
# Debugging - see debugging docs for more options
|
||||
# Use `--debug` or `--detailed_debug` CLI flags, or set LITELLM_LOG env var to "INFO", "DEBUG", or "ERROR"
|
||||
json_logs: boolean # if true, logs will be in json format
|
||||
|
||||
# Fallbacks, reliability
|
||||
|
|
@ -171,7 +172,7 @@ router_settings:
|
|||
| redact_user_api_key_info | boolean | If true, redacts information about the user api key from logs [Proxy Logging](logging#redacting-userapikeyinfo) |
|
||||
| mcp_aliases | object | Maps friendly aliases to MCP server names for easier tool access. Only the first alias for each server is used. [MCP Aliases](../mcp#mcp-aliases) |
|
||||
| langfuse_default_tags | array of strings | Default tags for Langfuse Logging. Use this if you want to control which LiteLLM-specific fields are logged as tags by the LiteLLM proxy. By default LiteLLM Proxy logs no LiteLLM-specific fields as tags. [Further docs](./logging#litellm-specific-tags-on-langfuse---cache_hit-cache_key) |
|
||||
| set_verbose | boolean | If true, sets litellm.set_verbose=True to view verbose debug logs. DO NOT LEAVE THIS ON IN PRODUCTION |
|
||||
| set_verbose | boolean | [DEPRECATED - see debugging docs](./debugging) Use `--debug` or `--detailed_debug` CLI flags, or set `LITELLM_LOG` env var to "INFO", "DEBUG", or "ERROR" instead. |
|
||||
| json_logs | boolean | If true, logs will be in json format. If you need to store the logs as JSON, just set the `litellm.json_logs = True`. We currently just log the raw POST request from litellm as a JSON [Further docs](./debugging) |
|
||||
| default_fallbacks | array of strings | List of fallback models to use if a specific model group is misconfigured / bad. [Further docs](./reliability#default-fallbacks) |
|
||||
| request_timeout | integer | The timeout for requests in seconds. If not set, the default value is `6000 seconds`. [For reference OpenAI Python SDK defaults to `600 seconds`.](https://github.com/openai/openai-python/blob/main/src/openai/_constants.py) |
|
||||
|
|
@ -333,7 +334,7 @@ router_settings:
|
|||
| caching_groups | Optional[List[tuple]] | List of model groups for caching across model groups. Defaults to None. - e.g. caching_groups=[("openai-gpt-3.5-turbo", "azure-gpt-3.5-turbo")]|
|
||||
| alerting_config | AlertingConfig | [SDK-only arg] Slack alerting configuration. Defaults to None. [Further Docs](../routing.md#alerting-) |
|
||||
| assistants_config | AssistantsConfig | Set on proxy via `assistant_settings`. [Further docs](../assistants.md) |
|
||||
| set_verbose | boolean | [DEPRECATED PARAM - see debug docs](./debugging.md) If true, sets the logging level to verbose. |
|
||||
| set_verbose | boolean | [DEPRECATED PARAM - see debug docs](./debugging) If true, sets the logging level to verbose. |
|
||||
| retry_after | int | Time to wait before retrying a request in seconds. Defaults to 0. If `x-retry-after` is received from LLM API, this value is overridden. |
|
||||
| provider_budget_config | ProviderBudgetConfig | Provider budget configuration. Use this to set llm_provider budget limits. example $100/day to OpenAI, $100/day to Azure, etc. Defaults to None. [Further Docs](./provider_budget_routing.md) |
|
||||
| enable_pre_call_checks | boolean | If true, checks if a call is within the model's context window before making the call. [More information here](reliability) |
|
||||
|
|
@ -359,6 +360,7 @@ router_settings:
|
|||
| AISPEND_ACCOUNT_ID | Account ID for AI Spend
|
||||
| AISPEND_API_KEY | API Key for AI Spend
|
||||
| AIOHTTP_CONNECTOR_LIMIT | Connection limit for aiohttp connector. When set to 0, no limit is applied. **Default is 0**
|
||||
| AIOHTTP_CONNECTOR_LIMIT_PER_HOST | Connection limit per host for aiohttp connector. When set to 0, no limit is applied. **Default is 0**
|
||||
| AIOHTTP_KEEPALIVE_TIMEOUT | Keep-alive timeout for aiohttp connections in seconds. **Default is 120**
|
||||
| AIOHTTP_TRUST_ENV | Flag to enable aiohttp trust environment. When this is set to True, aiohttp will respect HTTP(S)_PROXY env vars. **Default is False**
|
||||
| AIOHTTP_TTL_DNS_CACHE | DNS cache time-to-live for aiohttp in seconds. **Default is 300**
|
||||
|
|
@ -377,6 +379,8 @@ router_settings:
|
|||
| ATHINA_API_KEY | API key for Athina service
|
||||
| ATHINA_BASE_URL | Base URL for Athina service (defaults to `https://log.athina.ai`)
|
||||
| AUTH_STRATEGY | Strategy used for authentication (e.g., OAuth, API key)
|
||||
| AUTO_REDIRECT_UI_LOGIN_TO_SSO | Flag to enable automatic redirect of UI login page to SSO when SSO is configured. Default is **true**
|
||||
| AUDIO_SPEECH_CHUNK_SIZE | Chunk size for audio speech processing. Default is 1024
|
||||
| ANTHROPIC_API_KEY | API key for Anthropic service
|
||||
| ANTHROPIC_API_BASE | Base URL for Anthropic API. Default is https://api.anthropic.com
|
||||
| AWS_ACCESS_KEY_ID | Access Key ID for AWS services
|
||||
|
|
@ -439,6 +443,7 @@ router_settings:
|
|||
| CYBERARK_CLIENT_CERT | Path to client certificate for CyberArk authentication
|
||||
| CYBERARK_CLIENT_KEY | Path to client key for CyberArk authentication
|
||||
| CYBERARK_USERNAME | Username for CyberArk authentication
|
||||
| CYBERARK_SSL_VERIFY | Flag to enable or disable SSL certificate verification for CyberArk. Default is True
|
||||
| CONFIDENT_API_KEY | API key for DeepEval integration
|
||||
| CUSTOM_TIKTOKEN_CACHE_DIR | Custom directory for Tiktoken cache
|
||||
| CONFIDENT_API_KEY | API key for Confident AI (Deepeval) Logging service
|
||||
|
|
@ -653,6 +658,8 @@ router_settings:
|
|||
| LITERAL_API_URL | API URL for Literal service
|
||||
| LITERAL_BATCH_SIZE | Batch size for Literal operations
|
||||
| LITELLM_ANTHROPIC_DISABLE_URL_SUFFIX | Disable automatic URL suffix appending for Anthropic API base URLs. When set to `true`, prevents LiteLLM from automatically adding `/v1/messages` or `/v1/complete` to custom Anthropic API endpoints
|
||||
| LITELLM_DD_AGENT_HOST | Hostname or IP of DataDog agent for LiteLLM-specific logging. When set, logs are sent to agent instead of direct API
|
||||
| LITELLM_DD_AGENT_PORT | Port of DataDog agent for LiteLLM-specific log intake. Default is 10518
|
||||
| LITELLM_DONT_SHOW_FEEDBACK_BOX | Flag to hide feedback box in LiteLLM UI
|
||||
| LITELLM_DROP_PARAMS | Parameters to drop in LiteLLM requests
|
||||
| LITELLM_MODIFY_PARAMS | Parameters to modify in LiteLLM requests
|
||||
|
|
@ -797,7 +804,7 @@ router_settings:
|
|||
| SEND_USER_API_KEY_ALIAS | Flag to send user API key alias to Zscaler AI Guard. Default is False
|
||||
| SEND_USER_API_KEY_TEAM_ID | Flag to send user API key team ID to Zscaler AI Guard. Default is False
|
||||
| SEND_USER_API_KEY_USER_ID | Flag to send user API key user ID to Zscaler AI Guard. Default is False
|
||||
| SET_VERBOSE | Flag to enable verbose logging
|
||||
| SET_VERBOSE | [DEPRECATED] Use `LITELLM_LOG` instead with values "INFO", "DEBUG", or "ERROR". See [debugging docs](./debugging)
|
||||
| SINGLE_DEPLOYMENT_TRAFFIC_FAILURE_THRESHOLD | Minimum number of requests to consider "reasonable traffic" for single-deployment cooldown logic. Default is 1000
|
||||
| SLACK_DAILY_REPORT_FREQUENCY | Frequency of daily Slack reports (e.g., daily, weekly)
|
||||
| SLACK_WEBHOOK_URL | Webhook URL for Slack integration
|
||||
|
|
@ -839,6 +846,9 @@ router_settings:
|
|||
| UPSTREAM_LANGFUSE_SECRET_KEY | Secret key for upstream Langfuse authentication
|
||||
| USE_AWS_KMS | Flag to enable AWS Key Management Service for encryption
|
||||
| USE_PRISMA_MIGRATE | Flag to use prisma migrate instead of prisma db push. Recommended for production environments.
|
||||
| WANDB_API_KEY | API key for Weights & Biases (W&B) logging integration
|
||||
| WANDB_HOST | Host URL for Weights & Biases (W&B) service
|
||||
| WANDB_PROJECT_ID | Project ID for Weights & Biases (W&B) logging integration
|
||||
| WEBHOOK_URL | URL for receiving webhooks from external services
|
||||
| SPEND_LOG_RUN_LOOPS | Constant for setting how many runs of 1000 batch deletes should spend_log_cleanup task run
|
||||
| SPEND_LOG_CLEANUP_BATCH_SIZE | Number of logs deleted per batch during cleanup. Default is 1000
|
||||
|
|
|
|||
108
docs/my-website/docs/proxy/cursor.md
Normal file
108
docs/my-website/docs/proxy/cursor.md
Normal file
|
|
@ -0,0 +1,108 @@
|
|||
---
|
||||
id: cursor
|
||||
title: /cursor/chat/completions - Cursor Endpoint
|
||||
description: Accept Responses API input from Cursor and return OpenAI Chat Completions output
|
||||
---
|
||||
|
||||
LiteLLM provides a Cursor-specific endpoint to make Cursor IDE work seamlessly with the LiteLLM Proxy when using BYOK + custom `base_url`.
|
||||
|
||||
- Accepts Requests in OpenAI Responses API input format (Cursor sends this)
|
||||
- Returns Responses in OpenAI Chat Completions format (Cursor expects this)
|
||||
- Supports streaming and non‑streaming
|
||||
|
||||
## Endpoint
|
||||
|
||||
- Path: `/cursor/chat/completions`
|
||||
- Auth: Standard LiteLLM Proxy auth (`Authorization: Bearer <key>`)
|
||||
- Behavior: Internally routes to LiteLLM `/responses` flow and transforms output to Chat Completions
|
||||
|
||||
## Why this exists
|
||||
|
||||
When setting up Cursor with BYOK against a custom `base_url`, Cursor sends requests to the Chat Completions endpoint but in the OpenAI Responses API input shape. Without translation, Cursor won’t display streamed output. This endpoint bridges the formats:
|
||||
|
||||
- Input: Responses API (`input`, tool calls, etc.)
|
||||
- Output: Chat Completions (`choices`, `delta`, `finish_reason`, etc.)
|
||||
|
||||
## Usage
|
||||
|
||||
### Non-streaming
|
||||
|
||||
```bash
|
||||
curl -X POST https://litellm-internal/cursor/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "gpt-4o",
|
||||
"input": [{"role": "user", "content": "Hello"}]
|
||||
}'
|
||||
```
|
||||
|
||||
Example response (shape):
|
||||
|
||||
```json
|
||||
{
|
||||
"id": "chatcmpl-123",
|
||||
"object": "chat.completion",
|
||||
"created": 1733333333,
|
||||
"model": "gpt-4o",
|
||||
"choices": [
|
||||
{
|
||||
"index": 0,
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"content": "Hello! How can I help you?"
|
||||
},
|
||||
"finish_reason": "stop"
|
||||
}
|
||||
],
|
||||
"usage": {
|
||||
"prompt_tokens": 10,
|
||||
"completion_tokens": 8,
|
||||
"total_tokens": 18
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
### Streaming
|
||||
|
||||
```bash
|
||||
curl -N -X POST https://litellm-internal/cursor/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "gpt-4o",
|
||||
"input": [{"role": "user", "content": "Hello"}],
|
||||
"stream": true
|
||||
}'
|
||||
```
|
||||
|
||||
- Server-Sent Events (SSE)
|
||||
- Emits `chat.completion.chunk` deltas (`choices[].delta`) and ends with `data: [DONE]`
|
||||
|
||||
## Configuration
|
||||
|
||||
### Base URL Setup
|
||||
|
||||
**Important**: When configuring Cursor IDE to use this endpoint, you must include `/cursor` in the base URL.
|
||||
|
||||
Cursor automatically appends `/chat/completions` to the base URL you provide. To ensure requests go to `/cursor/chat/completions`, configure your base URL in Cursor as:
|
||||
|
||||
```
|
||||
Base URL: https://litellm-internal/cursor
|
||||
```
|
||||
|
||||
This way, when Cursor appends `/chat/completions`, the full path becomes `/cursor/chat/completions`, which is the correct endpoint.
|
||||
|
||||
**Example**: If your LiteLLM Proxy is running at `https://litellm-internal`, set the base URL in Cursor to `https://litellm-internal/cursor` (not just `https://litellm-internal`).
|
||||
|
||||
### General Setup
|
||||
|
||||
No special configuration is required beyond your normal LiteLLM Proxy setup. Ensure that:
|
||||
|
||||
- Your `config.yaml` includes the models you want to call via this endpoint
|
||||
- Your Cursor project uses your LiteLLM Proxy `base_url` (with `/cursor` included) and a valid API key
|
||||
|
||||
## Notes
|
||||
- This endpoint is intended specifically for Cursor’s request/response expectations. Other clients should continue to use `/v1/chat/completions` or `/v1/responses` as appropriate.
|
||||
|
||||
|
||||
710
docs/my-website/docs/proxy/multi_tenant_architecture.md
Normal file
710
docs/my-website/docs/proxy/multi_tenant_architecture.md
Normal file
|
|
@ -0,0 +1,710 @@
|
|||
import Image from '@theme/IdealImage';
|
||||
|
||||
# Multi-Tenant Architecture with LiteLLM
|
||||
|
||||
## Overview
|
||||
|
||||
LiteLLM provides a centralized solution that scales across multiple tenants, enabling organizations to:
|
||||
|
||||
- **Centrally manage** LLM access for multiple tenants (organizations, teams, departments)
|
||||
- **Isolate spend and usage** across different organizational units
|
||||
- **Delegate administration** without compromising security
|
||||
- **Track costs** at granular levels (organization → team → user → key)
|
||||
- **Scale seamlessly** as new teams and users are added
|
||||
|
||||
:::info Open Source vs. Enterprise
|
||||
- **Teams + Virtual Keys**: ✅ Available in open source
|
||||
- **Organizations + Org Admins**: ✨ Enterprise feature ([Get a 7 day trial](https://www.litellm.ai/#trial))
|
||||
|
||||
You can implement multi-tenancy using **Teams** alone in the open source version, or add **Organizations** on top for additional hierarchy in the enterprise version.
|
||||
:::
|
||||
|
||||
## The Multi-Tenant Challenge
|
||||
|
||||
Organizations with multi-tenant architectures face several challenges when deploying LLM solutions:
|
||||
|
||||
1. **Centralized vs. Decentralized**: Need a single unified gateway while maintaining tenant isolation
|
||||
2. **Cost Attribution**: Tracking spend across different business units, departments, or customers
|
||||
3. **Access Control**: Different teams need different models, budgets, and rate limits
|
||||
4. **Delegation**: Team leads should manage their teams without platform-wide admin access
|
||||
5. **Scalability**: Solution must scale from 10 to 10,000+ users without architectural changes
|
||||
|
||||
## How LiteLLM Solves Multi-Tenancy
|
||||
|
||||
<Image img={require('../../img/litellm_user_heirarchy.png')} style={{ width: '100%', maxWidth: '4000px' }} />
|
||||
|
||||
LiteLLM implements a hierarchical multi-tenant architecture with four levels:
|
||||
|
||||
### 1. Organizations (Top-Level Tenants) ✨ Enterprise Feature
|
||||
|
||||
**Organizations** represent the highest level of tenant isolation - typically different business units, departments, or customers.
|
||||
|
||||
- Each organization has its own:
|
||||
- Budget limits
|
||||
- Allowed models
|
||||
- Admin users (org admins)
|
||||
- Teams
|
||||
- Spend tracking
|
||||
|
||||
**Use Cases:**
|
||||
- **Enterprise Departments**: Separate organizations for Engineering, Marketing, Sales
|
||||
- **Multi-Customer SaaS**: Each customer is an organization with full isolation
|
||||
- **Geographic Regions**: EMEA, APAC, Americas as separate organizations
|
||||
|
||||
**Key Features:**
|
||||
- Organizations cannot see each other's data
|
||||
- Each organization can have multiple teams
|
||||
- Organization admins manage teams within their organization only
|
||||
- Spend and usage tracked at organization level
|
||||
|
||||
[API Reference for Organizations](https://litellm-api.up.railway.app/#/organization%20management)
|
||||
|
||||
---
|
||||
|
||||
### 2. Teams (Mid-Level Grouping) ✅ Open Source
|
||||
|
||||
**Teams** can work independently or sit within organizations, representing logical groupings of users working together.
|
||||
|
||||
:::tip
|
||||
Teams are available in **open source** and can be used as your primary multi-tenant boundary without needing Organizations. Organizations provide an additional layer of hierarchy for enterprise deployments.
|
||||
:::
|
||||
|
||||
- Each team has:
|
||||
- Team-specific budgets and rate limits
|
||||
- Team admins who manage members
|
||||
- Service account keys for shared resources
|
||||
- Model access controls
|
||||
- Granular team member permissions
|
||||
|
||||
**Use Cases:**
|
||||
- **Project Teams**: ML Research team, Product team, Data Science team
|
||||
- **Customer Sub-Groups**: Different divisions within a customer organization
|
||||
- **Environment Separation**: Development, Staging, Production teams
|
||||
|
||||
**Key Features:**
|
||||
- Teams inherit organization constraints (can't exceed org budget/models)
|
||||
- Team admins can manage their team without affecting others
|
||||
- Service account keys survive team member changes
|
||||
- Per-team spend tracking and billing
|
||||
|
||||
[API Reference for Teams](https://litellm-api.up.railway.app/#/team%20management)
|
||||
|
||||
---
|
||||
|
||||
### 3. Users (Individual Members) ✅ Open Source
|
||||
|
||||
**Users** are individuals who belong to teams and create/use API keys.
|
||||
|
||||
- Each user can:
|
||||
- Belong to multiple teams
|
||||
- Have their own budget limits
|
||||
- Create personal API keys
|
||||
- Track individual spend
|
||||
|
||||
**User Types:**
|
||||
- **Internal Users**: Employees, developers, data scientists
|
||||
- **Team Admins**: Lead their teams, manage members
|
||||
- **Org Admins**: Manage multiple teams within their organization
|
||||
- **Proxy Admins**: Platform-wide administrators
|
||||
|
||||
**Key Features:**
|
||||
- User spend tracked individually
|
||||
- Users can be on multiple teams simultaneously
|
||||
- Role-based permissions control what users can do
|
||||
- User keys deleted when user is removed
|
||||
|
||||
[API Reference for Users](https://litellm-api.up.railway.app/#/user%20management)
|
||||
|
||||
---
|
||||
|
||||
### 4. Virtual Keys (Authentication Layer) ✅ Open Source
|
||||
|
||||
**Virtual Keys** are the API keys used to authenticate requests and track spend.
|
||||
|
||||
Each key can be one of three types:
|
||||
|
||||
| Key Type | Configuration | Use Case | Spend Tracking | Lifecycle |
|
||||
|----------|---------------|----------|----------------|-----------|
|
||||
| **User-only** | `user_id` only | Developer personal keys | User level | Deleted with user |
|
||||
| **Team Service Account** | `team_id` only | Production apps, CI/CD | Team level | Survives member changes |
|
||||
| **User + Team** | Both `user_id` and `team_id` | User within team context | User AND Team | Deleted with user |
|
||||
|
||||
**Example Scenarios:**
|
||||
- Use **user-only keys** for developers testing locally
|
||||
- Use **team service account keys** for your production application that shouldn't break when employees leave
|
||||
- Use **user + team keys** when you want individual accountability within a team budget
|
||||
|
||||
[API Reference for Keys](https://litellm-api.up.railway.app/#/key%20management)
|
||||
|
||||
---
|
||||
|
||||
## Role-Based Access Control (RBAC)
|
||||
|
||||
LiteLLM provides granular RBAC across the hierarchy:
|
||||
|
||||
### Global Proxy Roles (Platform-Wide)
|
||||
|
||||
| Role | Scope | Permissions |
|
||||
|------|-------|-------------|
|
||||
| **Proxy Admin** | Entire platform | Create orgs, teams, users. View all spend. Full control. |
|
||||
| **Proxy Admin Viewer** | Entire platform | View-only access to all data. Cannot make changes. |
|
||||
| **Internal User** | Own resources | Create/delete own keys. View own spend. |
|
||||
|
||||
### Organization/Team Roles (Scoped)
|
||||
|
||||
| Role | Scope | Permissions |
|
||||
|------|-------|-------------|
|
||||
| **Org Admin** ✨ | Specific organization | Create teams, add users, view org spend within their org only. |
|
||||
| **Team Admin** ✨ | Specific team | Manage team members, budgets, keys within their team only. |
|
||||
|
||||
✨ = Premium Feature
|
||||
|
||||
### Team Member Permissions
|
||||
|
||||
Team admins can configure granular permissions for regular team members:
|
||||
|
||||
**Read-only** (default):
|
||||
```json
|
||||
["/key/info", "/key/health"]
|
||||
```
|
||||
|
||||
**Allow key creation**:
|
||||
```json
|
||||
["/key/info", "/key/health", "/key/generate", "/key/update"]
|
||||
```
|
||||
|
||||
**Full key management**:
|
||||
```json
|
||||
["/key/info", "/key/health", "/key/generate", "/key/update", "/key/delete", "/key/regenerate", "/key/block", "/key/unblock"]
|
||||
```
|
||||
|
||||
[Learn more about RBAC](./access_control)
|
||||
|
||||
---
|
||||
|
||||
## Spend Tracking & Cost Attribution
|
||||
|
||||
LiteLLM provides multi-level spend tracking that flows through the hierarchy:
|
||||
|
||||
### Hierarchical Spend Flow
|
||||
|
||||
```
|
||||
Organization Spend
|
||||
├── Team 1 Spend
|
||||
│ ├── User A Spend
|
||||
│ │ ├── Key 1 Spend
|
||||
│ │ └── Key 2 Spend
|
||||
│ └── Service Account Spend
|
||||
│ └── Key 3 Spend
|
||||
└── Team 2 Spend
|
||||
└── User B Spend
|
||||
└── Key 4 Spend
|
||||
```
|
||||
|
||||
### Budget Enforcement
|
||||
|
||||
Budgets can be set at every level with inheritance:
|
||||
|
||||
1. **Organization Budget**: `$10,000/month`
|
||||
- Team 1: `$6,000/month` (within org limit)
|
||||
- User A: `$3,000/month` (within team limit)
|
||||
- User B: `$3,000/month` (within team limit)
|
||||
- Team 2: `$4,000/month` (within org limit)
|
||||
|
||||
**Enforcement Rules:**
|
||||
- Team budgets cannot exceed organization budget
|
||||
- User budgets cannot exceed team budget
|
||||
- Requests blocked when any level exceeds budget
|
||||
- Real-time tracking prevents overruns
|
||||
|
||||
[Learn more about Budgets](./team_budgets)
|
||||
|
||||
---
|
||||
|
||||
## Common Multi-Tenant Patterns
|
||||
|
||||
### Pattern 1: Enterprise Departments
|
||||
|
||||
**Scenario**: Large enterprise with multiple departments needing centralized LLM access
|
||||
|
||||
**Enterprise Setup** (with Organizations):
|
||||
```
|
||||
Platform (LiteLLM Instance)
|
||||
├── Engineering Organization ✨
|
||||
│ ├── Backend Team
|
||||
│ ├── Frontend Team
|
||||
│ └── ML Team
|
||||
├── Marketing Organization ✨
|
||||
│ ├── Content Team
|
||||
│ └── Analytics Team
|
||||
└── Sales Organization ✨
|
||||
├── Sales Ops Team
|
||||
└── Customer Success Team
|
||||
```
|
||||
|
||||
**Open Source Alternative** (Teams only):
|
||||
```
|
||||
Platform (LiteLLM Instance)
|
||||
├── Engineering Backend Team
|
||||
├── Engineering Frontend Team
|
||||
├── Engineering ML Team
|
||||
├── Marketing Content Team
|
||||
├── Marketing Analytics Team
|
||||
├── Sales Ops Team
|
||||
└── Customer Success Team
|
||||
```
|
||||
|
||||
**Benefits:**
|
||||
- Each department/team manages their own budget
|
||||
- Department leads (org/team admins) control their teams
|
||||
- Centralized billing and model access
|
||||
- Cross-department cost visibility for finance
|
||||
|
||||
---
|
||||
|
||||
### Pattern 2: Multi-Customer SaaS
|
||||
|
||||
**Scenario**: SaaS provider offering LLM-powered features to multiple customers
|
||||
|
||||
**Enterprise Setup** (with Organizations):
|
||||
```
|
||||
Platform (LiteLLM Instance)
|
||||
├── Customer A Organization ✨
|
||||
│ ├── Production Team (Service Accounts)
|
||||
│ ├── Development Team
|
||||
│ └── QA Team
|
||||
├── Customer B Organization ✨
|
||||
│ ├── Production Team (Service Accounts)
|
||||
│ └── Development Team
|
||||
└── Customer C Organization ✨
|
||||
└── Production Team (Service Accounts)
|
||||
```
|
||||
|
||||
**Open Source Alternative** (Teams only):
|
||||
```
|
||||
Platform (LiteLLM Instance)
|
||||
├── Customer A Production Team (Service Accounts)
|
||||
├── Customer A Development Team
|
||||
├── Customer A QA Team
|
||||
├── Customer B Production Team (Service Accounts)
|
||||
├── Customer B Development Team
|
||||
└── Customer C Production Team (Service Accounts)
|
||||
```
|
||||
|
||||
**Benefits:**
|
||||
- Complete isolation between customers/teams
|
||||
- Per-customer/team billing and usage tracking
|
||||
- Customer/team admins can self-serve
|
||||
- Production service account keys survive employee turnover
|
||||
|
||||
---
|
||||
|
||||
### Pattern 3: Environment Separation
|
||||
|
||||
**Scenario**: Single organization with multiple environments
|
||||
|
||||
```
|
||||
Platform (LiteLLM Instance)
|
||||
└── Company Organization
|
||||
├── Production Team
|
||||
│ └── Service Account Keys (strict rate limits)
|
||||
├── Staging Team
|
||||
│ └── Service Account Keys (moderate limits)
|
||||
└── Development Team
|
||||
└── User Keys (generous limits for testing)
|
||||
```
|
||||
|
||||
**Benefits:**
|
||||
- Separate budgets for each environment
|
||||
- Different model access (production vs. development)
|
||||
- Prevent development usage from affecting production budget
|
||||
- Easy cost attribution by environment
|
||||
|
||||
---
|
||||
|
||||
## Delegation & Self-Service
|
||||
|
||||
One of LiteLLM's key advantages is delegated administration:
|
||||
|
||||
### Without LiteLLM
|
||||
```
|
||||
Every team → Requests platform admin → Admin makes changes
|
||||
```
|
||||
❌ Bottleneck on platform team
|
||||
❌ Slow onboarding
|
||||
❌ Poor scalability
|
||||
|
||||
### With LiteLLM
|
||||
```
|
||||
Proxy Admin → Creates org + org admin
|
||||
Org Admin → Creates teams + team admins
|
||||
Team Admin → Manages their team independently
|
||||
```
|
||||
✅ Decentralized management
|
||||
✅ Fast onboarding
|
||||
✅ Scales to thousands of users
|
||||
|
||||
### Self-Service Capabilities
|
||||
|
||||
**Team Admins Can:**
|
||||
- Add/remove team members
|
||||
- Create API keys for team members
|
||||
- Update team budgets (within org limits)
|
||||
- Configure team member permissions
|
||||
- View team usage and spend
|
||||
|
||||
**Org Admins Can:**
|
||||
- Create new teams within their organization
|
||||
- Assign team admins
|
||||
- View organization-wide spend
|
||||
- Manage users across their teams
|
||||
|
||||
**Platform Admins Can:**
|
||||
- Create organizations
|
||||
- Assign org admins
|
||||
- Set organization-level policies
|
||||
- View platform-wide analytics
|
||||
|
||||
---
|
||||
|
||||
## Scalability
|
||||
|
||||
LiteLLM's architecture scales from small teams to enterprise deployments:
|
||||
|
||||
### Small Team (10-100 users)
|
||||
- Single organization
|
||||
- Few teams (5-10)
|
||||
- Proxy admins manage everything
|
||||
|
||||
### Mid-Size (100-1,000 users)
|
||||
- Multiple organizations
|
||||
- Many teams (50+)
|
||||
- Org admins delegate to team admins
|
||||
|
||||
### Enterprise (1,000+ users)
|
||||
- Many organizations (departments/regions)
|
||||
- Hundreds of teams
|
||||
- Fully delegated admin structure
|
||||
- Centralized observability and billing
|
||||
|
||||
**Key Scalability Features:**
|
||||
- No architectural changes needed as you grow
|
||||
- Database-backed (PostgreSQL) for reliability
|
||||
- Horizontal scaling support
|
||||
- Efficient spend tracking and logging
|
||||
|
||||
---
|
||||
|
||||
## Security & Isolation
|
||||
|
||||
### Tenant Isolation
|
||||
|
||||
Each tenant (organization) is isolated:
|
||||
- ✅ Cannot view other organizations' data
|
||||
- ✅ Cannot access other organizations' keys
|
||||
- ✅ Cannot exceed their budget limits
|
||||
- ✅ Cannot access models not in their allowed list
|
||||
|
||||
### Authentication Security
|
||||
|
||||
- Master key for platform admins
|
||||
- Virtual keys with scoped permissions
|
||||
- SSO integration support
|
||||
- JWT authentication
|
||||
- IP allowlisting
|
||||
|
||||
### Audit & Compliance
|
||||
|
||||
- All API calls logged with user/team/org context
|
||||
- Spend tracking for chargeback/showback
|
||||
- Admin actions audited
|
||||
- Integration with observability tools
|
||||
|
||||
[Learn more about Security](../data_security)
|
||||
|
||||
---
|
||||
|
||||
## Getting Started
|
||||
|
||||
:::info Enterprise vs. Open Source Setup
|
||||
The steps below show the **full enterprise hierarchy** with Organizations.
|
||||
|
||||
For **open source**, skip Steps 1-2 and start directly with **Step 3** (creating teams). Teams can function as your top-level tenant boundary without Organizations.
|
||||
:::
|
||||
|
||||
### Step 1: Set Up Organizations ✨ Enterprise
|
||||
|
||||
Create your first organization:
|
||||
|
||||
```bash
|
||||
curl --location 'http://0.0.0.0:4000/organization/new' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"organization_alias": "engineering_department",
|
||||
"models": ["gpt-4", "gpt-4o", "claude-3-5-sonnet"],
|
||||
"max_budget": 10000
|
||||
}'
|
||||
```
|
||||
|
||||
### Step 2: Add an Organization Admin ✨ Enterprise
|
||||
|
||||
```bash
|
||||
curl -X POST 'http://0.0.0.0:4000/organization/member_add' \
|
||||
-H 'Authorization: Bearer sk-1234' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"organization_id": "org-123",
|
||||
"member": {
|
||||
"role": "org_admin",
|
||||
"user_id": "admin@company.com"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
### Step 3: Create Teams ✅ Open Source
|
||||
|
||||
**For Enterprise:** Organization admin creates team within their organization
|
||||
**For Open Source:** Proxy admin creates team directly (no `organization_id` needed)
|
||||
|
||||
```bash
|
||||
# Enterprise: Org admin creates team in their organization
|
||||
curl --location 'http://0.0.0.0:4000/team/new' \
|
||||
--header 'Authorization: Bearer sk-org-admin-key' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"team_alias": "ml_team",
|
||||
"organization_id": "org-123",
|
||||
"max_budget": 5000
|
||||
}'
|
||||
|
||||
# Open Source: Proxy admin creates team directly
|
||||
curl --location 'http://0.0.0.0:4000/team/new' \
|
||||
--header 'Authorization: Bearer sk-1234' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"team_alias": "ml_team",
|
||||
"max_budget": 5000
|
||||
}'
|
||||
```
|
||||
|
||||
### Step 4: Add Team Admin
|
||||
|
||||
```bash
|
||||
curl -X POST 'http://0.0.0.0:4000/team/member_add' \
|
||||
-H 'Authorization: Bearer sk-org-admin-key' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"team_id": "team-456",
|
||||
"member": {
|
||||
"role": "admin",
|
||||
"user_id": "team-lead@company.com"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
### Step 5: Team Admin Manages Their Team
|
||||
|
||||
```bash
|
||||
# Team admin adds members
|
||||
curl -X POST 'http://0.0.0.0:4000/team/member_add' \
|
||||
-H 'Authorization: Bearer sk-team-admin-key' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"team_id": "team-456",
|
||||
"member": {
|
||||
"role": "user",
|
||||
"user_id": "developer@company.com"
|
||||
}
|
||||
}'
|
||||
|
||||
# Team admin creates keys for members
|
||||
curl --location 'http://0.0.0.0:4000/key/generate' \
|
||||
--header 'Authorization: Bearer sk-team-admin-key' \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"user_id": "developer@company.com",
|
||||
"team_id": "team-456"
|
||||
}'
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Use Case Examples
|
||||
|
||||
### Example 1: Chargeback Model
|
||||
|
||||
**Goal**: Each business unit pays for their own LLM usage
|
||||
|
||||
**Setup:**
|
||||
1. Create organization per business unit
|
||||
2. Set budgets based on allocated budgets
|
||||
3. Track spend per organization
|
||||
4. Generate monthly reports for finance
|
||||
|
||||
**Result**: Finance can charge back costs to respective departments with accurate attribution.
|
||||
|
||||
---
|
||||
|
||||
### Example 2: Customer-Facing AI Product
|
||||
|
||||
**Goal**: Provide LLM capabilities to customers with isolation and cost tracking
|
||||
|
||||
**Setup:**
|
||||
1. Create organization per customer
|
||||
2. Use service account keys for production workloads
|
||||
3. Track spend per customer organization
|
||||
4. Set rate limits per customer tier
|
||||
|
||||
**Result**: Bill customers accurately, prevent noisy neighbors, maintain isolation.
|
||||
|
||||
---
|
||||
|
||||
### Example 3: Development vs. Production
|
||||
|
||||
**Goal**: Separate development and production environments with different policies
|
||||
|
||||
**Setup:**
|
||||
1. Create "Development" and "Production" teams
|
||||
2. Development: Generous budgets, all models, user keys
|
||||
3. Production: Strict budgets, approved models only, service account keys
|
||||
4. Different rate limits per environment
|
||||
|
||||
**Result**: Developers can experiment freely without impacting production budget or reliability.
|
||||
|
||||
---
|
||||
|
||||
## Best Practices
|
||||
|
||||
### 1. Organization Design
|
||||
|
||||
- ✅ Map organizations to cost centers or customers
|
||||
- ✅ Set realistic budgets with buffer for growth
|
||||
- ✅ Assign 1-2 org admins per organization
|
||||
- ❌ Don't create too many organizations (adds management overhead)
|
||||
|
||||
### 2. Team Structure
|
||||
|
||||
- ✅ Keep teams aligned with actual working groups
|
||||
- ✅ Use service account keys for production
|
||||
- ✅ Give team admins enough permissions to self-serve
|
||||
- ❌ Don't create single-user teams (use user-only keys instead)
|
||||
|
||||
### 3. Key Management
|
||||
|
||||
- ✅ Use descriptive key names
|
||||
- ✅ Rotate keys regularly
|
||||
- ✅ Delete unused keys
|
||||
- ✅ Use appropriate key type for use case
|
||||
- ❌ Don't share keys across users/teams
|
||||
|
||||
### 4. Budget Management
|
||||
|
||||
- ✅ Set budgets at multiple levels (org → team → user)
|
||||
- ✅ Monitor spend regularly
|
||||
- ✅ Alert before budget exhaustion
|
||||
- ❌ Don't set budgets too tight (may block legitimate usage)
|
||||
|
||||
### 5. Delegation
|
||||
|
||||
- ✅ Assign org admins for large organizations
|
||||
- ✅ Assign team admins for active teams
|
||||
- ✅ Configure team member permissions appropriately
|
||||
- ❌ Don't make everyone a proxy admin
|
||||
|
||||
---
|
||||
|
||||
## Monitoring & Observability
|
||||
|
||||
LiteLLM provides comprehensive monitoring:
|
||||
|
||||
- **Spend Tracking**: Real-time spend by org/team/user/key
|
||||
- **Usage Analytics**: Request counts, token usage, model usage
|
||||
- **Admin UI**: Visual dashboard for all metrics
|
||||
- **Logging**: Detailed logs with tenant context
|
||||
- **Alerting**: Budget alerts, rate limit alerts, error alerts
|
||||
|
||||
[Learn more about Logging](./logging)
|
||||
|
||||
---
|
||||
|
||||
## Comparison with Other Approaches
|
||||
|
||||
| Approach | Pros | Cons | LiteLLM Advantage |
|
||||
|----------|------|------|-------------------|
|
||||
| **Separate instances per tenant** | Strong isolation | High operational overhead, cost inefficient | Single instance, same isolation, 90% cost reduction |
|
||||
| **Single shared pool** | Simple setup | No cost attribution, no access control | Full attribution, granular access control |
|
||||
| **API key prefixes** | Basic separation | Manual tracking, no hierarchy, no RBAC | Automatic tracking, hierarchical, full RBAC |
|
||||
| **External auth layer** | Flexible | Complex integration, no built-in budgets | Native integration, built-in budgets |
|
||||
|
||||
---
|
||||
|
||||
## FAQ
|
||||
|
||||
**Q: Can users belong to multiple teams?**
|
||||
A: Yes, users can be members of multiple teams and have different keys for each team.
|
||||
|
||||
**Q: What happens when a user leaves?**
|
||||
A: User-specific keys are deleted, but team service account keys remain active.
|
||||
|
||||
**Q: Can team budgets exceed organization budget?**
|
||||
A: No, the system enforces that team budgets cannot exceed their organization's budget.
|
||||
|
||||
**Q: How granular is the cost tracking?**
|
||||
A: Every API call is tracked with organization, team, user, and key context.
|
||||
|
||||
**Q: Can I have teams without organizations?**
|
||||
A: Yes! Teams work independently in **open source** without needing Organizations. Organizations are an **enterprise feature** that adds an additional hierarchy layer on top of teams.
|
||||
|
||||
**Q: Is there a limit to hierarchy depth?**
|
||||
A: The hierarchy is: Organization → Team → User → Key (4 levels). This covers most use cases.
|
||||
|
||||
**Q: How do I migrate from flat structure to hierarchical?**
|
||||
A: You can gradually create organizations and teams, then move existing users/keys into them.
|
||||
|
||||
---
|
||||
|
||||
## Related Documentation
|
||||
|
||||
- [User Management Hierarchy](./user_management_heirarchy) - Visual hierarchy overview
|
||||
- [Access Control (RBAC)](./access_control) - Detailed role permissions
|
||||
- [Team Budgets](./team_budgets) - Budget management guide
|
||||
- [Virtual Keys](./virtual_keys) - API key management
|
||||
- [Admin UI](./ui) - Visual dashboard for management
|
||||
|
||||
---
|
||||
|
||||
## Summary
|
||||
|
||||
LiteLLM solves multi-tenant architecture challenges through:
|
||||
|
||||
1. **Hierarchical Structure**: Organizations → Teams → Users → Keys
|
||||
2. **Granular RBAC**: Platform-wide and tenant-scoped roles
|
||||
3. **Cost Attribution**: Spend tracking at every level
|
||||
4. **Delegation**: Org admins and team admins self-manage
|
||||
5. **Isolation**: Strong tenant boundaries
|
||||
6. **Scalability**: Handles 10 to 10,000+ users with same architecture
|
||||
|
||||
### Open Source vs. Enterprise
|
||||
|
||||
**Open Source** (Teams + Users + Keys):
|
||||
- ✅ Teams as primary tenant boundary
|
||||
- ✅ Team admins manage their teams
|
||||
- ✅ Virtual keys with team/user tracking
|
||||
- ✅ Budget and rate limits per team
|
||||
- ✅ Spend tracking and logging
|
||||
|
||||
**Enterprise** (Adds Organizations layer):
|
||||
- ✨ Organizations for top-level tenant isolation
|
||||
- ✨ Organization admins manage multiple teams
|
||||
- ✨ Organization-level budgets and model access
|
||||
- ✨ Hierarchical delegation and reporting
|
||||
|
||||
This makes LiteLLM ideal for:
|
||||
- ✅ Enterprises with multiple departments
|
||||
- ✅ SaaS providers with multiple customers
|
||||
- ✅ Organizations needing cost chargeback/showback
|
||||
- ✅ Teams requiring self-service LLM access
|
||||
- ✅ Any multi-tenant LLM deployment
|
||||
|
||||
[Start with LiteLLM Proxy →](./quick_start)
|
||||
|
|
@ -41,6 +41,7 @@ CYBERARK_CLIENT_KEY="path/to/client.key"
|
|||
|
||||
# OPTIONAL
|
||||
CYBERARK_REFRESH_INTERVAL="300" # defaults to 300 seconds (5 minutes), frequency of token refresh
|
||||
CYBERARK_SSL_VERIFY="true" # defaults to true, set to "false" to disable SSL verification (for self-signed certificates)
|
||||
```
|
||||
|
||||
**Step 2.** Add to proxy config.yaml
|
||||
|
|
@ -172,6 +173,24 @@ If these commands work successfully against your CyberArk instance, then CyberAr
|
|||
- The `CYBERARK_API_BASE` URL is accessible from your LiteLLM instance
|
||||
- Your API key or certificates have the necessary permissions in CyberArk
|
||||
|
||||
### SSL Certificate Errors
|
||||
|
||||
If you encounter SSL certificate verification errors like:
|
||||
|
||||
```
|
||||
RuntimeError: Could not authenticate to CyberArk Conjur: [SSL: CERTIFICATE_VERIFY_FAILED] certificate verify failed: self-signed certificate in certificate chain
|
||||
```
|
||||
|
||||
This typically occurs when your CyberArk Conjur instance uses a self-signed certificate. You can disable SSL verification by setting:
|
||||
|
||||
```bash
|
||||
CYBERARK_SSL_VERIFY="false"
|
||||
```
|
||||
|
||||
:::warning
|
||||
Disabling SSL verification is insecure and should only be used for testing or development environments with self-signed certificates. For production, configure your certificate chain properly or use certificate-based authentication with `CYBERARK_CLIENT_CERT` and `CYBERARK_CLIENT_KEY`.
|
||||
:::
|
||||
|
||||
## Video Walkthrough
|
||||
|
||||
This video walks through using CyberArk Conjur as a secret manager with LiteLLM. We create a virtual key in the LiteLLM Admin UI and verify it exists in CyberArk. Then we rotate the secret key and verify it exists in CyberArk.
|
||||
|
|
|
|||
226
docs/my-website/docs/tutorials/cursor_integration.md
Normal file
226
docs/my-website/docs/tutorials/cursor_integration.md
Normal file
|
|
@ -0,0 +1,226 @@
|
|||
---
|
||||
sidebar_label: "Cursor IDE"
|
||||
---
|
||||
|
||||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Cursor IDE Integration with LiteLLM
|
||||
|
||||
This tutorial shows you how to integrate Cursor IDE with LiteLLM Proxy, allowing you to use any LiteLLM-supported model through Cursor's interface with BYOK (Bring Your Own Key) and custom base URL.
|
||||
|
||||
## Benefits of using Cursor with LiteLLM
|
||||
|
||||
When you use Cursor IDE with LiteLLM you get the following benefits:
|
||||
|
||||
**Developer Benefits:**
|
||||
- Universal Model Access: Use any LiteLLM supported model (Anthropic, OpenAI, Vertex AI, Bedrock, etc.) through the Cursor IDE interface.
|
||||
- Higher Rate Limits & Reliability: Load balance across multiple models and providers to avoid hitting individual provider limits, with fallbacks to ensure you get responses even if one provider fails.
|
||||
- Streaming Support: Full streaming support with proper response transformation for Cursor's expected format.
|
||||
|
||||
**Proxy Admin Benefits:**
|
||||
- Centralized Management: Control access to all models through a single LiteLLM proxy instance without giving your developers API Keys to each provider.
|
||||
- Budget Controls: Set spending limits and track costs across all Cursor usage.
|
||||
- Request Logging: Track all requests made through Cursor for debugging and monitoring.
|
||||
|
||||
## Prerequisites
|
||||
|
||||
Before you begin, ensure you have:
|
||||
- Cursor IDE installed
|
||||
- A running LiteLLM Proxy instance with **HTTPS enabled** (HTTP is not supported)
|
||||
- A valid LiteLLM Proxy API key
|
||||
- An HTTPS domain for your LiteLLM Proxy (required by Cursor)
|
||||
|
||||
## Quick Start Guide
|
||||
|
||||
### Step 1: Install LiteLLM
|
||||
|
||||
Install LiteLLM with proxy support:
|
||||
|
||||
```bash
|
||||
pip install litellm[proxy]
|
||||
```
|
||||
|
||||
### Step 2: Configure LiteLLM Proxy
|
||||
|
||||
Create a `config.yaml` file with your model configurations:
|
||||
|
||||
```yaml showLineNumbers title="config.yaml"
|
||||
model_list:
|
||||
- model_name: gpt-4o
|
||||
litellm_params:
|
||||
model: gpt-4o
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
- model_name: claude-3-5-sonnet
|
||||
litellm_params:
|
||||
model: anthropic/claude-3-5-sonnet-20241022
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
|
||||
general_settings:
|
||||
master_key: sk-1234567890 # Change this to a secure key
|
||||
```
|
||||
|
||||
### Step 3: Start LiteLLM Proxy
|
||||
|
||||
Start the proxy server with HTTPS enabled:
|
||||
|
||||
```bash
|
||||
litellm --config config.yaml --port 4000
|
||||
```
|
||||
|
||||
:::warning HTTPS Required
|
||||
|
||||
**Important**: Cursor IDE requires HTTPS connections. HTTP (`http://`) will not work. You must:
|
||||
- Deploy your LiteLLM Proxy with HTTPS enabled
|
||||
- Use a valid SSL certificate
|
||||
- Access the proxy via an HTTPS domain (e.g., `https://your-proxy-domain.com`)
|
||||
|
||||
For local development, you'll need to set up HTTPS (e.g., using a reverse proxy like nginx with SSL, or deploying to a cloud service with HTTPS).
|
||||
|
||||
:::
|
||||
|
||||
### Step 4: Configure Cursor IDE
|
||||
|
||||
Configure Cursor IDE to use your LiteLLM proxy with the `/cursor/chat/completions` endpoint:
|
||||
|
||||
1. Open Cursor IDE
|
||||
2. Go to **Settings** → **Features** → **AI**
|
||||
3. Enable **"Use Custom API"** or **"Bring Your Own Key"**
|
||||
4. Set the following:
|
||||
- **Base URL**: `https://your-proxy-domain.com/cursor` (⚠️ **Important**: Must use HTTPS and include `/cursor`)
|
||||
- **API Key**: Your LiteLLM Proxy API key (e.g., `sk-1234567890`)
|
||||
|
||||
:::warning HTTPS Required
|
||||
|
||||
Cursor IDE **requires HTTPS** connections. HTTP (`http://`) will not work. You must:
|
||||
- Use an HTTPS URL for your base URL (e.g., `https://your-proxy-domain.com/cursor`)
|
||||
- Ensure your LiteLLM Proxy is accessible via HTTPS
|
||||
- Have a valid SSL certificate configured
|
||||
|
||||
:::
|
||||
|
||||
**Example Configuration:**
|
||||
|
||||
```
|
||||
Base URL: https://your-proxy-domain.com/cursor
|
||||
API Key: sk-1234567890
|
||||
```
|
||||
|
||||
Replace `your-proxy-domain.com` with your actual HTTPS domain where LiteLLM Proxy is running.
|
||||
|
||||
:::info Why `/cursor` in the base URL?
|
||||
|
||||
Cursor automatically appends `/chat/completions` to the base URL you provide. By setting the base URL to `https://your-proxy-domain.com/cursor`, Cursor will send requests to `/cursor/chat/completions`, which is the special endpoint that handles Cursor's Responses API input format and transforms it to Chat Completions output format.
|
||||
|
||||
If you set the base URL to just `https://your-proxy-domain.com`, Cursor would send requests to `/chat/completions`, which won't work correctly with Cursor's request format.
|
||||
|
||||
|
||||
:::
|
||||
|
||||
### Step 5: Test the Integration
|
||||
|
||||
1. Restart Cursor IDE to apply the settings
|
||||
2. Open a code file and try using Cursor's AI features (completions, chat, etc.)
|
||||
3. Your requests will now be routed through LiteLLM Proxy
|
||||
|
||||
You can verify it's working by:
|
||||
- Checking the LiteLLM Proxy logs for incoming requests
|
||||
- Using Cursor's chat feature and seeing responses stream correctly
|
||||
- Checking your LiteLLM dashboard for request logs and cost tracking
|
||||
|
||||
## How It Works
|
||||
|
||||
The `/cursor/chat/completions` endpoint is specifically designed to handle Cursor's unique request format:
|
||||
|
||||
1. **Input**: Cursor sends requests in OpenAI Responses API format (with `input` field)
|
||||
2. **Processing**: LiteLLM processes the request through its internal `/responses` flow
|
||||
3. **Output**: The response is transformed to OpenAI Chat Completions format (with `choices` field) that Cursor expects
|
||||
|
||||
This transformation happens automatically for both streaming and non-streaming responses.
|
||||
|
||||
## Advanced Configuration
|
||||
|
||||
### Using Different Models
|
||||
|
||||
You can configure Cursor to use different models by updating your `config.yaml`:
|
||||
|
||||
```yaml showLineNumbers title="config.yaml"
|
||||
model_list:
|
||||
- model_name: gpt-4o
|
||||
litellm_params:
|
||||
model: gpt-4o
|
||||
api_key: os.environ/OPENAI_API_KEY
|
||||
|
||||
- model_name: claude-3-5-sonnet
|
||||
litellm_params:
|
||||
model: anthropic/claude-3-5-sonnet-20241022
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
|
||||
- model_name: gemini-pro
|
||||
litellm_params:
|
||||
model: gemini/gemini-1.5-pro
|
||||
api_key: os.environ/GEMINI_API_KEY
|
||||
```
|
||||
|
||||
Then in Cursor, you can specify which model to use in your requests.
|
||||
|
||||
### Rate Limiting and Budgets
|
||||
|
||||
Set up rate limits and budgets in your `config.yaml`:
|
||||
|
||||
```yaml showLineNumbers title="config.yaml"
|
||||
general_settings:
|
||||
master_key: sk-1234567890
|
||||
|
||||
litellm_settings:
|
||||
# Set max budget per user
|
||||
max_budget: 100.0
|
||||
|
||||
# Set rate limits
|
||||
rate_limit: 100 # requests per minute
|
||||
```
|
||||
|
||||
### Request Logging
|
||||
|
||||
All requests from Cursor will be logged by LiteLLM Proxy. You can:
|
||||
- View logs in the LiteLLM Admin UI
|
||||
- Export logs to your preferred logging service
|
||||
- Track costs per user/team
|
||||
|
||||
## Troubleshooting
|
||||
|
||||
### Cursor shows no output
|
||||
|
||||
- **Check base URL**: Ensure it uses HTTPS and includes `/cursor` (e.g., `https://your-proxy-domain.com/cursor`, not `http://` or without `/cursor`)
|
||||
- **Verify HTTPS**: Cursor requires HTTPS - HTTP connections will not work
|
||||
- **Check API key**: Verify your LiteLLM Proxy API key is correct
|
||||
- **Check proxy logs**: Look for errors in the LiteLLM Proxy logs
|
||||
|
||||
### Requests failing
|
||||
|
||||
- **Verify HTTPS is enabled**: Cursor requires HTTPS connections. Ensure your LiteLLM Proxy is accessible via HTTPS with a valid SSL certificate
|
||||
- **Verify proxy is running**: Check that LiteLLM Proxy is accessible at your HTTPS base URL
|
||||
- **Check SSL certificate**: Ensure your SSL certificate is valid and not expired
|
||||
- **Check model configuration**: Ensure the model you're trying to use is configured in `config.yaml`
|
||||
- **Check API keys**: Verify provider API keys are set correctly in environment variables
|
||||
|
||||
### HTTP not working
|
||||
|
||||
If you're trying to use HTTP (`http://`) and it's not working:
|
||||
- **This is expected**: Cursor IDE requires HTTPS connections
|
||||
- **Solution**: Deploy your LiteLLM Proxy with HTTPS enabled (use a reverse proxy like nginx, or deploy to a cloud service that provides HTTPS)
|
||||
|
||||
### Streaming not working
|
||||
|
||||
The `/cursor/chat/completions` endpoint automatically handles streaming. If streaming isn't working:
|
||||
- Check that your model supports streaming
|
||||
- Verify the proxy logs for any transformation errors
|
||||
- Ensure Cursor IDE is up to date
|
||||
|
||||
## Related Documentation
|
||||
|
||||
- [Cursor Endpoint Documentation](/docs/proxy/cursor) - Detailed endpoint documentation
|
||||
- [LiteLLM Proxy Setup](/docs/proxy/quick_start) - General proxy setup guide
|
||||
- [Model Configuration](/docs/proxy/configs) - How to configure models
|
||||
|
||||
BIN
docs/my-website/img/a2a_gateway.png
Normal file
BIN
docs/my-website/img/a2a_gateway.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 1 MiB |
0
docs/my-website/img/add_agent1.png
Normal file
0
docs/my-website/img/add_agent1.png
Normal file
BIN
docs/my-website/img/add_agent_1.png
Normal file
BIN
docs/my-website/img/add_agent_1.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 288 KiB |
BIN
docs/my-website/img/agent2.png
Normal file
BIN
docs/my-website/img/agent2.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 806 KiB |
BIN
docs/my-website/img/agent_id.png
Normal file
BIN
docs/my-website/img/agent_id.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 230 KiB |
BIN
docs/my-website/img/agent_key.png
Normal file
BIN
docs/my-website/img/agent_key.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 176 KiB |
BIN
docs/my-website/img/agent_team.png
Normal file
BIN
docs/my-website/img/agent_team.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 352 KiB |
|
|
@ -1,5 +1,5 @@
|
|||
---
|
||||
title: "[PREVIEW] v1.80.5.rc.2 - Gemini 3.0 Support"
|
||||
title: "v1.80.5-stable - Gemini 3.0 Support"
|
||||
slug: "v1-80-5"
|
||||
date: 2025-11-22T10:00:00
|
||||
authors:
|
||||
|
|
@ -27,7 +27,7 @@ import TabItem from '@theme/TabItem';
|
|||
docker run \
|
||||
-e STORE_MODEL_IN_DB=True \
|
||||
-p 4000:4000 \
|
||||
ghcr.io/berriai/litellm:v1.80.5.rc.2
|
||||
ghcr.io/berriai/litellm:v1.80.5-stable
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
|
|
|||
|
|
@ -105,6 +105,7 @@ const sidebars = {
|
|||
items: [
|
||||
"tutorials/claude_responses_api",
|
||||
"tutorials/cost_tracking_coding",
|
||||
"tutorials/cursor_integration",
|
||||
"tutorials/github_copilot_integration",
|
||||
"tutorials/litellm_gemini_cli",
|
||||
"tutorials/litellm_qwen_code_cli",
|
||||
|
|
@ -520,6 +521,7 @@ const sidebars = {
|
|||
"providers/vertex_ai/videos",
|
||||
"providers/vertex_partner",
|
||||
"providers/vertex_self_deployed",
|
||||
"providers/vertex_embedding",
|
||||
"providers/vertex_image",
|
||||
"providers/vertex_speech",
|
||||
"providers/vertex_batch",
|
||||
|
|
|
|||
BIN
enterprise/dist/litellm_enterprise-0.1.23-py3-none-any.whl
vendored
Normal file
BIN
enterprise/dist/litellm_enterprise-0.1.23-py3-none-any.whl
vendored
Normal file
Binary file not shown.
BIN
enterprise/dist/litellm_enterprise-0.1.23.tar.gz
vendored
Normal file
BIN
enterprise/dist/litellm_enterprise-0.1.23.tar.gz
vendored
Normal file
Binary file not shown.
345
enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py
Normal file
345
enterprise/litellm_enterprise/proxy/vector_stores/endpoints.py
Normal file
|
|
@ -0,0 +1,345 @@
|
|||
"""
|
||||
VECTOR STORE MANAGEMENT
|
||||
|
||||
All /vector_store management endpoints
|
||||
|
||||
/vector_store/new
|
||||
/vector_store/delete
|
||||
/vector_store/list
|
||||
"""
|
||||
|
||||
import copy
|
||||
import json
|
||||
from typing import List, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.proxy._types import (
|
||||
LiteLLM_ManagedVectorStoresTable,
|
||||
ResponseLiteLLM_ManagedVectorStore,
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.vector_stores import (
|
||||
LiteLLM_ManagedVectorStore,
|
||||
LiteLLM_ManagedVectorStoreListResponse,
|
||||
VectorStoreDeleteRequest,
|
||||
VectorStoreInfoRequest,
|
||||
VectorStoreUpdateRequest,
|
||||
)
|
||||
from litellm.vector_stores.vector_store_registry import VectorStoreRegistry
|
||||
|
||||
router = APIRouter()
|
||||
|
||||
|
||||
########################################################
|
||||
# Management Endpoints
|
||||
########################################################
|
||||
@router.post(
|
||||
"/vector_store/new",
|
||||
tags=["vector store management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def new_vector_store(
|
||||
vector_store: LiteLLM_ManagedVectorStore,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Create a new vector store.
|
||||
|
||||
Parameters:
|
||||
- vector_store_id: str - Unique identifier for the vector store
|
||||
- custom_llm_provider: str - Provider of the vector store
|
||||
- vector_store_name: Optional[str] - Name of the vector store
|
||||
- vector_store_description: Optional[str] - Description of the vector store
|
||||
- vector_store_metadata: Optional[Dict] - Additional metadata for the vector store
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
# Check if vector store already exists
|
||||
existing_vector_store = (
|
||||
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
|
||||
where={"vector_store_id": vector_store.get("vector_store_id")}
|
||||
)
|
||||
)
|
||||
if existing_vector_store is not None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"Vector store with ID {vector_store.get('vector_store_id')} already exists",
|
||||
)
|
||||
|
||||
if vector_store.get("vector_store_metadata") is not None:
|
||||
vector_store["vector_store_metadata"] = safe_dumps(
|
||||
vector_store.get("vector_store_metadata")
|
||||
)
|
||||
|
||||
# Safely handle JSON serialization of litellm_params
|
||||
litellm_params_json: Optional[str] = None
|
||||
_input_litellm_params: dict = vector_store.get("litellm_params", {}) or {}
|
||||
if _input_litellm_params is not None:
|
||||
litellm_params_dict = GenericLiteLLMParams(
|
||||
**_input_litellm_params
|
||||
).model_dump(exclude_none=True)
|
||||
litellm_params_json = safe_dumps(litellm_params_dict)
|
||||
del vector_store["litellm_params"]
|
||||
|
||||
_new_vector_store = (
|
||||
await prisma_client.db.litellm_managedvectorstorestable.create(
|
||||
data={
|
||||
**vector_store,
|
||||
"litellm_params": litellm_params_json,
|
||||
}
|
||||
)
|
||||
)
|
||||
|
||||
new_vector_store: LiteLLM_ManagedVectorStore = LiteLLM_ManagedVectorStore(
|
||||
**_new_vector_store.model_dump()
|
||||
)
|
||||
|
||||
# Add vector store to registry
|
||||
if litellm.vector_store_registry is not None:
|
||||
litellm.vector_store_registry.add_vector_store_to_registry(
|
||||
vector_store=new_vector_store
|
||||
)
|
||||
|
||||
return {
|
||||
"status": "success",
|
||||
"message": f"Vector store {vector_store.get('vector_store_id')} created successfully",
|
||||
"vector_store": new_vector_store,
|
||||
}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error creating vector store: {str(e)}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.get(
|
||||
"/vector_store/list",
|
||||
tags=["vector store management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=LiteLLM_ManagedVectorStoreListResponse,
|
||||
)
|
||||
async def list_vector_stores(
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
page: int = 1,
|
||||
page_size: int = 100,
|
||||
):
|
||||
"""
|
||||
List all available vector stores with optional filtering and pagination.
|
||||
Combines both in-memory vector stores and those stored in the database.
|
||||
|
||||
Parameters:
|
||||
- page: int - Page number for pagination (default: 1)
|
||||
- page_size: int - Number of items per page (default: 100)
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
try:
|
||||
# Get vector stores from database (source of truth)
|
||||
# Only return what's in the database to ensure consistency across instances
|
||||
vector_stores_from_db = await VectorStoreRegistry._get_vector_stores_from_db(
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
|
||||
# Also clean up in-memory registry to remove any deleted vector stores
|
||||
if litellm.vector_store_registry is not None:
|
||||
db_vector_store_ids = {
|
||||
vs.get("vector_store_id")
|
||||
for vs in vector_stores_from_db
|
||||
if vs.get("vector_store_id")
|
||||
}
|
||||
# Remove any in-memory vector stores that no longer exist in database
|
||||
vector_stores_to_remove = []
|
||||
for vs in litellm.vector_store_registry.vector_stores:
|
||||
vs_id = vs.get("vector_store_id")
|
||||
if vs_id and vs_id not in db_vector_store_ids:
|
||||
vector_stores_to_remove.append(vs_id)
|
||||
for vs_id in vector_stores_to_remove:
|
||||
litellm.vector_store_registry.delete_vector_store_from_registry(
|
||||
vector_store_id=vs_id
|
||||
)
|
||||
verbose_proxy_logger.debug(
|
||||
f"Removed deleted vector store {vs_id} from in-memory registry"
|
||||
)
|
||||
|
||||
# Use database as single source of truth for listing
|
||||
combined_vector_stores: List[LiteLLM_ManagedVectorStore] = vector_stores_from_db
|
||||
|
||||
total_count = len(combined_vector_stores)
|
||||
total_pages = (total_count + page_size - 1) // page_size
|
||||
|
||||
# Format response using LiteLLM_ManagedVectorStoreListResponse
|
||||
response = LiteLLM_ManagedVectorStoreListResponse(
|
||||
object="list",
|
||||
data=combined_vector_stores,
|
||||
total_count=total_count,
|
||||
current_page=page,
|
||||
total_pages=total_pages,
|
||||
)
|
||||
|
||||
return response
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error listing vector stores: {str(e)}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/vector_store/delete",
|
||||
tags=["vector store management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def delete_vector_store(
|
||||
data: VectorStoreDeleteRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Delete a vector store.
|
||||
|
||||
Parameters:
|
||||
- vector_store_id: str - ID of the vector store to delete
|
||||
"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
# Check if vector store exists
|
||||
existing_vector_store = (
|
||||
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
|
||||
where={"vector_store_id": data.vector_store_id}
|
||||
)
|
||||
)
|
||||
if existing_vector_store is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"Vector store with ID {data.vector_store_id} not found",
|
||||
)
|
||||
|
||||
# Delete vector store
|
||||
await prisma_client.db.litellm_managedvectorstorestable.delete(
|
||||
where={"vector_store_id": data.vector_store_id}
|
||||
)
|
||||
|
||||
# Delete vector store from registry
|
||||
if litellm.vector_store_registry is not None:
|
||||
litellm.vector_store_registry.delete_vector_store_from_registry(
|
||||
vector_store_id=data.vector_store_id
|
||||
)
|
||||
|
||||
return {"message": f"Vector store {data.vector_store_id} deleted successfully"}
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/vector_store/info",
|
||||
tags=["vector store management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
response_model=ResponseLiteLLM_ManagedVectorStore,
|
||||
)
|
||||
async def get_vector_store_info(
|
||||
data: VectorStoreInfoRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Return a single vector store's details"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
if litellm.vector_store_registry is not None:
|
||||
vector_store = litellm.vector_store_registry.get_litellm_managed_vector_store_from_registry(
|
||||
vector_store_id=data.vector_store_id
|
||||
)
|
||||
if vector_store is not None:
|
||||
vector_store_metadata = vector_store.get("vector_store_metadata")
|
||||
# Parse metadata if it's a JSON string
|
||||
parsed_metadata: Optional[dict] = None
|
||||
if isinstance(vector_store_metadata, str):
|
||||
parsed_metadata = json.loads(vector_store_metadata)
|
||||
elif isinstance(vector_store_metadata, dict):
|
||||
parsed_metadata = vector_store_metadata
|
||||
|
||||
vector_store_pydantic_obj = LiteLLM_ManagedVectorStoresTable(
|
||||
vector_store_id=vector_store.get("vector_store_id") or "",
|
||||
custom_llm_provider=vector_store.get("custom_llm_provider") or "",
|
||||
vector_store_name=vector_store.get("vector_store_name") or None,
|
||||
vector_store_description=vector_store.get(
|
||||
"vector_store_description"
|
||||
)
|
||||
or None,
|
||||
vector_store_metadata=parsed_metadata,
|
||||
created_at=vector_store.get("created_at") or None,
|
||||
updated_at=vector_store.get("updated_at") or None,
|
||||
litellm_credential_name=vector_store.get("litellm_credential_name"),
|
||||
litellm_params=vector_store.get("litellm_params") or None,
|
||||
)
|
||||
return {"vector_store": vector_store_pydantic_obj}
|
||||
|
||||
vector_store = (
|
||||
await prisma_client.db.litellm_managedvectorstorestable.find_unique(
|
||||
where={"vector_store_id": data.vector_store_id}
|
||||
)
|
||||
)
|
||||
if vector_store is None:
|
||||
raise HTTPException(
|
||||
status_code=404,
|
||||
detail=f"Vector store with ID {data.vector_store_id} not found",
|
||||
)
|
||||
|
||||
vector_store_dict = vector_store.model_dump() # type: ignore[attr-defined]
|
||||
return {"vector_store": vector_store_dict}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error getting vector store info: {str(e)}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
||||
|
||||
@router.post(
|
||||
"/vector_store/update",
|
||||
tags=["vector store management"],
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
async def update_vector_store(
|
||||
data: VectorStoreUpdateRequest,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""Update vector store details"""
|
||||
from litellm.proxy.proxy_server import prisma_client
|
||||
|
||||
if prisma_client is None:
|
||||
raise HTTPException(status_code=500, detail="Database not connected")
|
||||
|
||||
try:
|
||||
update_data = data.model_dump(exclude_unset=True)
|
||||
vector_store_id = update_data.pop("vector_store_id")
|
||||
if update_data.get("vector_store_metadata") is not None:
|
||||
update_data["vector_store_metadata"] = safe_dumps(
|
||||
update_data["vector_store_metadata"]
|
||||
)
|
||||
|
||||
updated = await prisma_client.db.litellm_managedvectorstorestable.update(
|
||||
where={"vector_store_id": vector_store_id},
|
||||
data=update_data,
|
||||
)
|
||||
|
||||
updated_vs = LiteLLM_ManagedVectorStore(**updated.model_dump())
|
||||
|
||||
if litellm.vector_store_registry is not None:
|
||||
litellm.vector_store_registry.update_vector_store_in_registry(
|
||||
vector_store_id=vector_store_id,
|
||||
updated_data=updated_vs,
|
||||
)
|
||||
|
||||
return {"vector_store": updated_vs}
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error updating vector store: {str(e)}")
|
||||
raise HTTPException(status_code=500, detail=str(e))
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
[tool.poetry]
|
||||
name = "litellm-enterprise"
|
||||
version = "0.1.22"
|
||||
version = "0.1.23"
|
||||
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.22"
|
||||
version = "0.1.23"
|
||||
version_files = [
|
||||
"pyproject.toml:version",
|
||||
"../requirements.txt:litellm-enterprise==",
|
||||
|
|
|
|||
|
|
@ -0,0 +1,7 @@
|
|||
-- Add agent permission fields to LiteLLM_ObjectPermissionTable
|
||||
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN IF NOT EXISTS "agents" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
ALTER TABLE "LiteLLM_ObjectPermissionTable" ADD COLUMN IF NOT EXISTS "agent_access_groups" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
|
||||
-- Add agent_access_groups field to LiteLLM_AgentsTable
|
||||
ALTER TABLE "LiteLLM_AgentsTable" ADD COLUMN IF NOT EXISTS "agent_access_groups" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
|
||||
|
|
@ -61,6 +61,7 @@ model LiteLLM_AgentsTable {
|
|||
agent_name String @unique
|
||||
litellm_params Json?
|
||||
agent_card_params Json
|
||||
agent_access_groups String[] @default([])
|
||||
created_at DateTime @default(now()) @map("created_at")
|
||||
created_by String
|
||||
updated_at DateTime @default(now()) @updatedAt @map("updated_at")
|
||||
|
|
@ -172,6 +173,8 @@ model LiteLLM_ObjectPermissionTable {
|
|||
mcp_access_groups String[] @default([])
|
||||
mcp_tool_permissions Json? // Tool-level permissions for MCP servers. Format: {"server_id": ["tool_name_1", "tool_name_2"]}
|
||||
vector_stores String[] @default([])
|
||||
agents String[] @default([])
|
||||
agent_access_groups String[] @default([])
|
||||
teams LiteLLM_TeamTable[]
|
||||
verification_tokens LiteLLM_VerificationToken[]
|
||||
organizations LiteLLM_OrganizationTable[]
|
||||
|
|
|
|||
|
|
@ -151,6 +151,7 @@ _custom_logger_compatible_callbacks_literal = Literal[
|
|||
"mlflow",
|
||||
"langfuse",
|
||||
"langfuse_otel",
|
||||
"weave_otel",
|
||||
"pagerduty",
|
||||
"humanloop",
|
||||
"gcs_pubsub",
|
||||
|
|
@ -1056,57 +1057,10 @@ from .timeout import timeout
|
|||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.litellm_core_utils.core_helpers import remove_index_from_tool_calls
|
||||
from litellm.litellm_core_utils.token_counter import get_modified_max_tokens
|
||||
from .utils import (
|
||||
client,
|
||||
exception_type,
|
||||
get_optional_params,
|
||||
get_response_string,
|
||||
token_counter,
|
||||
create_pretrained_tokenizer,
|
||||
create_tokenizer,
|
||||
supports_function_calling,
|
||||
supports_web_search,
|
||||
supports_url_context,
|
||||
supports_response_schema,
|
||||
supports_parallel_function_calling,
|
||||
supports_vision,
|
||||
supports_audio_input,
|
||||
supports_audio_output,
|
||||
supports_system_messages,
|
||||
supports_reasoning,
|
||||
get_litellm_params,
|
||||
acreate,
|
||||
get_max_tokens,
|
||||
get_model_info,
|
||||
register_prompt_template,
|
||||
validate_environment,
|
||||
check_valid_key,
|
||||
register_model,
|
||||
encode,
|
||||
decode,
|
||||
_calculate_retry_after,
|
||||
_should_retry,
|
||||
get_supported_openai_params,
|
||||
get_api_base,
|
||||
get_first_chars_messages,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
EmbeddingResponse,
|
||||
ImageResponse,
|
||||
TranscriptionResponse,
|
||||
TextCompletionResponse,
|
||||
get_provider_fields,
|
||||
ModelResponseListIterator,
|
||||
get_valid_models,
|
||||
)
|
||||
|
||||
ALL_LITELLM_RESPONSE_TYPES = [
|
||||
ModelResponse,
|
||||
EmbeddingResponse,
|
||||
ImageResponse,
|
||||
TranscriptionResponse,
|
||||
TextCompletionResponse,
|
||||
]
|
||||
# client must be imported immediately as it's used as a decorator at function definition time
|
||||
from .utils import client
|
||||
# Note: Most other utils imports are lazy-loaded via __getattr__ to avoid loading utils.py
|
||||
# (which imports tiktoken) at import time
|
||||
|
||||
from .llms.bytez.chat.transformation import BytezChatConfig
|
||||
from .llms.custom_llm import CustomLLM
|
||||
|
|
@ -1210,6 +1164,9 @@ from .llms.bedrock.chat.invoke_transformations.amazon_ai21_transformation import
|
|||
from .llms.bedrock.chat.invoke_transformations.amazon_nova_transformation import (
|
||||
AmazonInvokeNovaConfig,
|
||||
)
|
||||
from .llms.bedrock.chat.invoke_transformations.amazon_qwen2_transformation import (
|
||||
AmazonQwen2Config,
|
||||
)
|
||||
from .llms.bedrock.chat.invoke_transformations.amazon_qwen3_transformation import (
|
||||
AmazonQwen3Config,
|
||||
)
|
||||
|
|
@ -1382,7 +1339,7 @@ from .llms.nebius.chat.transformation import NebiusConfig
|
|||
from .llms.wandb.chat.transformation import WandbConfig
|
||||
from .llms.dashscope.chat.transformation import DashScopeChatConfig
|
||||
from .llms.moonshot.chat.transformation import MoonshotChatConfig
|
||||
from .llms.publicai.chat.transformation import PublicAIChatConfig
|
||||
# PublicAI now uses JSON-based configuration (see litellm/llms/openai_like/providers.json)
|
||||
from .llms.docker_model_runner.chat.transformation import DockerModelRunnerChatConfig
|
||||
from .llms.v0.chat.transformation import V0ChatConfig
|
||||
from .llms.oci.chat.transformation import OCIChatConfig
|
||||
|
|
@ -1538,56 +1495,6 @@ def set_global_gitlab_config(config: Dict[str, Any]) -> None:
|
|||
|
||||
|
||||
# Lazy loading system for heavy modules to reduce initial import time and memory usage
|
||||
def _lazy_import_cost_calculator(name: str) -> Any:
|
||||
"""Lazy import for cost_calculator functions."""
|
||||
from .cost_calculator import (
|
||||
completion_cost as _completion_cost,
|
||||
cost_per_token as _cost_per_token,
|
||||
response_cost_calculator as _response_cost_calculator,
|
||||
)
|
||||
|
||||
_cost_functions = {
|
||||
"completion_cost": _completion_cost,
|
||||
"cost_per_token": _cost_per_token,
|
||||
"response_cost_calculator": _response_cost_calculator,
|
||||
}
|
||||
|
||||
func = _cost_functions[name]
|
||||
globals()[name] = func
|
||||
return func
|
||||
|
||||
|
||||
def _lazy_import_litellm_logging(name: str) -> Any:
|
||||
"""Lazy import for litellm_logging module."""
|
||||
try:
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
Logging as _Logging,
|
||||
modify_integration as _modify_integration,
|
||||
)
|
||||
|
||||
_logging_objects = {
|
||||
"Logging": _Logging,
|
||||
"modify_integration": _modify_integration,
|
||||
}
|
||||
|
||||
obj = _logging_objects[name]
|
||||
globals()[name] = obj
|
||||
return obj
|
||||
except Exception as e:
|
||||
raise AttributeError(
|
||||
f"module {__name__!r} has no attribute {name!r}. "
|
||||
f"Lazy import failed: {e}"
|
||||
) from e
|
||||
|
||||
|
||||
_LAZY_LOAD_REGISTRY: Dict[str, Callable[[str], Any]] = {
|
||||
"completion_cost": _lazy_import_cost_calculator,
|
||||
"cost_per_token": _lazy_import_cost_calculator,
|
||||
"response_cost_calculator": _lazy_import_cost_calculator,
|
||||
"Logging": _lazy_import_litellm_logging,
|
||||
"modify_integration": _lazy_import_litellm_logging,
|
||||
}
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
cost_per_token: Callable[..., Tuple[float, float]]
|
||||
|
|
@ -1598,7 +1505,45 @@ if TYPE_CHECKING:
|
|||
|
||||
def __getattr__(name: str) -> Any:
|
||||
"""Lazy import handler for cost_calculator and litellm_logging functions."""
|
||||
if name in _LAZY_LOAD_REGISTRY:
|
||||
return _LAZY_LOAD_REGISTRY[name](name)
|
||||
# Lazy load cost_calculator functions
|
||||
_cost_calculator_names = (
|
||||
"completion_cost",
|
||||
"cost_per_token",
|
||||
"response_cost_calculator",
|
||||
)
|
||||
if name in _cost_calculator_names:
|
||||
from ._lazy_imports import _lazy_import_cost_calculator
|
||||
return _lazy_import_cost_calculator(name)
|
||||
|
||||
# Lazy load litellm_logging functions
|
||||
_litellm_logging_names = (
|
||||
"Logging",
|
||||
"modify_integration",
|
||||
)
|
||||
if name in _litellm_logging_names:
|
||||
from ._lazy_imports import _lazy_import_litellm_logging
|
||||
return _lazy_import_litellm_logging(name)
|
||||
|
||||
# Lazy load utils functions
|
||||
_utils_names = (
|
||||
"exception_type", "get_optional_params", "get_response_string", "token_counter",
|
||||
"create_pretrained_tokenizer", "create_tokenizer", "supports_function_calling",
|
||||
"supports_web_search", "supports_url_context", "supports_response_schema",
|
||||
"supports_parallel_function_calling", "supports_vision", "supports_audio_input",
|
||||
"supports_audio_output", "supports_system_messages", "supports_reasoning",
|
||||
"get_litellm_params", "acreate", "get_max_tokens", "get_model_info",
|
||||
"register_prompt_template", "validate_environment", "check_valid_key",
|
||||
"register_model", "encode", "decode", "_calculate_retry_after", "_should_retry",
|
||||
"get_supported_openai_params", "get_api_base", "get_first_chars_messages",
|
||||
"ModelResponse", "ModelResponseStream", "EmbeddingResponse", "ImageResponse",
|
||||
"TranscriptionResponse", "TextCompletionResponse", "get_provider_fields",
|
||||
"ModelResponseListIterator", "get_valid_models",
|
||||
)
|
||||
if name in _utils_names:
|
||||
from ._lazy_imports import _lazy_import_utils
|
||||
return _lazy_import_utils(name)
|
||||
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}")
|
||||
|
||||
|
||||
# ALL_LITELLM_RESPONSE_TYPES is lazy-loaded via __getattr__ to avoid loading utils at import time
|
||||
|
|
|
|||
259
litellm/_lazy_imports.py
Normal file
259
litellm/_lazy_imports.py
Normal file
|
|
@ -0,0 +1,259 @@
|
|||
from typing import Any
|
||||
import sys
|
||||
|
||||
def _get_litellm_globals() -> dict:
|
||||
"""Helper to get the globals dictionary of the litellm module."""
|
||||
return sys.modules["litellm"].__dict__
|
||||
|
||||
# Lazy import for utils module - imports only the requested item by name.
|
||||
# Note: PLR0915 (too many statements) is suppressed because the many if statements
|
||||
# are intentional - each attribute is imported individually only when requested,
|
||||
# ensuring true lazy imports rather than importing the entire utils module.
|
||||
def _lazy_import_utils(name: str) -> Any: # noqa: PLR0915
|
||||
"""Lazy import for utils module - imports only the requested item by name."""
|
||||
_globals = _get_litellm_globals()
|
||||
if name == "exception_type":
|
||||
from .utils import exception_type as _exception_type
|
||||
_globals["exception_type"] = _exception_type
|
||||
return _exception_type
|
||||
|
||||
if name == "get_optional_params":
|
||||
from .utils import get_optional_params as _get_optional_params
|
||||
_globals["get_optional_params"] = _get_optional_params
|
||||
return _get_optional_params
|
||||
|
||||
if name == "get_response_string":
|
||||
from .utils import get_response_string as _get_response_string
|
||||
_globals["get_response_string"] = _get_response_string
|
||||
return _get_response_string
|
||||
|
||||
if name == "token_counter":
|
||||
from .utils import token_counter as _token_counter
|
||||
_globals["token_counter"] = _token_counter
|
||||
return _token_counter
|
||||
|
||||
if name == "create_pretrained_tokenizer":
|
||||
from .utils import create_pretrained_tokenizer as _create_pretrained_tokenizer
|
||||
_globals["create_pretrained_tokenizer"] = _create_pretrained_tokenizer
|
||||
return _create_pretrained_tokenizer
|
||||
|
||||
if name == "create_tokenizer":
|
||||
from .utils import create_tokenizer as _create_tokenizer
|
||||
_globals["create_tokenizer"] = _create_tokenizer
|
||||
return _create_tokenizer
|
||||
|
||||
if name == "supports_function_calling":
|
||||
from .utils import supports_function_calling as _supports_function_calling
|
||||
_globals["supports_function_calling"] = _supports_function_calling
|
||||
return _supports_function_calling
|
||||
|
||||
if name == "supports_web_search":
|
||||
from .utils import supports_web_search as _supports_web_search
|
||||
_globals["supports_web_search"] = _supports_web_search
|
||||
return _supports_web_search
|
||||
|
||||
if name == "supports_url_context":
|
||||
from .utils import supports_url_context as _supports_url_context
|
||||
_globals["supports_url_context"] = _supports_url_context
|
||||
return _supports_url_context
|
||||
|
||||
if name == "supports_response_schema":
|
||||
from .utils import supports_response_schema as _supports_response_schema
|
||||
_globals["supports_response_schema"] = _supports_response_schema
|
||||
return _supports_response_schema
|
||||
|
||||
if name == "supports_parallel_function_calling":
|
||||
from .utils import supports_parallel_function_calling as _supports_parallel_function_calling
|
||||
_globals["supports_parallel_function_calling"] = _supports_parallel_function_calling
|
||||
return _supports_parallel_function_calling
|
||||
|
||||
if name == "supports_vision":
|
||||
from .utils import supports_vision as _supports_vision
|
||||
_globals["supports_vision"] = _supports_vision
|
||||
return _supports_vision
|
||||
|
||||
if name == "supports_audio_input":
|
||||
from .utils import supports_audio_input as _supports_audio_input
|
||||
_globals["supports_audio_input"] = _supports_audio_input
|
||||
return _supports_audio_input
|
||||
|
||||
if name == "supports_audio_output":
|
||||
from .utils import supports_audio_output as _supports_audio_output
|
||||
_globals["supports_audio_output"] = _supports_audio_output
|
||||
return _supports_audio_output
|
||||
|
||||
if name == "supports_system_messages":
|
||||
from .utils import supports_system_messages as _supports_system_messages
|
||||
_globals["supports_system_messages"] = _supports_system_messages
|
||||
return _supports_system_messages
|
||||
|
||||
if name == "supports_reasoning":
|
||||
from .utils import supports_reasoning as _supports_reasoning
|
||||
_globals["supports_reasoning"] = _supports_reasoning
|
||||
return _supports_reasoning
|
||||
|
||||
if name == "get_litellm_params":
|
||||
from .utils import get_litellm_params as _get_litellm_params
|
||||
_globals["get_litellm_params"] = _get_litellm_params
|
||||
return _get_litellm_params
|
||||
|
||||
if name == "acreate":
|
||||
from .utils import acreate as _acreate
|
||||
_globals["acreate"] = _acreate
|
||||
return _acreate
|
||||
|
||||
if name == "get_max_tokens":
|
||||
from .utils import get_max_tokens as _get_max_tokens
|
||||
_globals["get_max_tokens"] = _get_max_tokens
|
||||
return _get_max_tokens
|
||||
|
||||
if name == "get_model_info":
|
||||
from .utils import get_model_info as _get_model_info
|
||||
_globals["get_model_info"] = _get_model_info
|
||||
return _get_model_info
|
||||
|
||||
if name == "register_prompt_template":
|
||||
from .utils import register_prompt_template as _register_prompt_template
|
||||
_globals["register_prompt_template"] = _register_prompt_template
|
||||
return _register_prompt_template
|
||||
|
||||
if name == "validate_environment":
|
||||
from .utils import validate_environment as _validate_environment
|
||||
_globals["validate_environment"] = _validate_environment
|
||||
return _validate_environment
|
||||
|
||||
if name == "check_valid_key":
|
||||
from .utils import check_valid_key as _check_valid_key
|
||||
_globals["check_valid_key"] = _check_valid_key
|
||||
return _check_valid_key
|
||||
|
||||
if name == "register_model":
|
||||
from .utils import register_model as _register_model
|
||||
_globals["register_model"] = _register_model
|
||||
return _register_model
|
||||
|
||||
if name == "encode":
|
||||
from .utils import encode as _encode
|
||||
_globals["encode"] = _encode
|
||||
return _encode
|
||||
|
||||
if name == "decode":
|
||||
from .utils import decode as _decode
|
||||
_globals["decode"] = _decode
|
||||
return _decode
|
||||
|
||||
if name == "_calculate_retry_after":
|
||||
from .utils import _calculate_retry_after as __calculate_retry_after
|
||||
_globals["_calculate_retry_after"] = __calculate_retry_after
|
||||
return __calculate_retry_after
|
||||
|
||||
if name == "_should_retry":
|
||||
from .utils import _should_retry as __should_retry
|
||||
_globals["_should_retry"] = __should_retry
|
||||
return __should_retry
|
||||
|
||||
if name == "get_supported_openai_params":
|
||||
from .utils import get_supported_openai_params as _get_supported_openai_params
|
||||
_globals["get_supported_openai_params"] = _get_supported_openai_params
|
||||
return _get_supported_openai_params
|
||||
|
||||
if name == "get_api_base":
|
||||
from .utils import get_api_base as _get_api_base
|
||||
_globals["get_api_base"] = _get_api_base
|
||||
return _get_api_base
|
||||
|
||||
if name == "get_first_chars_messages":
|
||||
from .utils import get_first_chars_messages as _get_first_chars_messages
|
||||
_globals["get_first_chars_messages"] = _get_first_chars_messages
|
||||
return _get_first_chars_messages
|
||||
|
||||
if name == "ModelResponse":
|
||||
from .utils import ModelResponse as _ModelResponse
|
||||
_globals["ModelResponse"] = _ModelResponse
|
||||
return _ModelResponse
|
||||
|
||||
if name == "ModelResponseStream":
|
||||
from .utils import ModelResponseStream as _ModelResponseStream
|
||||
_globals["ModelResponseStream"] = _ModelResponseStream
|
||||
return _ModelResponseStream
|
||||
|
||||
if name == "EmbeddingResponse":
|
||||
from .utils import EmbeddingResponse as _EmbeddingResponse
|
||||
_globals["EmbeddingResponse"] = _EmbeddingResponse
|
||||
return _EmbeddingResponse
|
||||
|
||||
if name == "ImageResponse":
|
||||
from .utils import ImageResponse as _ImageResponse
|
||||
_globals["ImageResponse"] = _ImageResponse
|
||||
return _ImageResponse
|
||||
|
||||
if name == "TranscriptionResponse":
|
||||
from .utils import TranscriptionResponse as _TranscriptionResponse
|
||||
_globals["TranscriptionResponse"] = _TranscriptionResponse
|
||||
return _TranscriptionResponse
|
||||
|
||||
if name == "TextCompletionResponse":
|
||||
from .utils import TextCompletionResponse as _TextCompletionResponse
|
||||
_globals["TextCompletionResponse"] = _TextCompletionResponse
|
||||
return _TextCompletionResponse
|
||||
|
||||
if name == "get_provider_fields":
|
||||
from .utils import get_provider_fields as _get_provider_fields
|
||||
_globals["get_provider_fields"] = _get_provider_fields
|
||||
return _get_provider_fields
|
||||
|
||||
if name == "ModelResponseListIterator":
|
||||
from .utils import ModelResponseListIterator as _ModelResponseListIterator
|
||||
_globals["ModelResponseListIterator"] = _ModelResponseListIterator
|
||||
return _ModelResponseListIterator
|
||||
|
||||
if name == "get_valid_models":
|
||||
from .utils import get_valid_models as _get_valid_models
|
||||
_globals["get_valid_models"] = _get_valid_models
|
||||
return _get_valid_models
|
||||
|
||||
raise AttributeError(f"Utils lazy import: unknown attribute {name!r}")
|
||||
|
||||
|
||||
def _lazy_import_cost_calculator(name: str) -> Any:
|
||||
"""Lazy import for cost_calculator functions."""
|
||||
_globals = _get_litellm_globals()
|
||||
from .cost_calculator import (
|
||||
completion_cost as _completion_cost,
|
||||
cost_per_token as _cost_per_token,
|
||||
response_cost_calculator as _response_cost_calculator,
|
||||
)
|
||||
|
||||
_cost_functions = {
|
||||
"completion_cost": _completion_cost,
|
||||
"cost_per_token": _cost_per_token,
|
||||
"response_cost_calculator": _response_cost_calculator,
|
||||
}
|
||||
|
||||
func = _cost_functions[name]
|
||||
_globals[name] = func
|
||||
return func
|
||||
|
||||
|
||||
def _lazy_import_litellm_logging(name: str) -> Any:
|
||||
"""Lazy import for litellm_logging module."""
|
||||
_globals = _get_litellm_globals()
|
||||
try:
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
Logging as _Logging,
|
||||
modify_integration as _modify_integration,
|
||||
)
|
||||
|
||||
_logging_objects = {
|
||||
"Logging": _Logging,
|
||||
"modify_integration": _modify_integration,
|
||||
}
|
||||
|
||||
obj = _logging_objects[name]
|
||||
_globals[name] = obj
|
||||
return obj
|
||||
except Exception as e:
|
||||
raise AttributeError(
|
||||
f"module 'litellm' has no attribute {name!r}. "
|
||||
f"Lazy import failed: {e}"
|
||||
) from e
|
||||
59
litellm/a2a_protocol/__init__.py
Normal file
59
litellm/a2a_protocol/__init__.py
Normal file
|
|
@ -0,0 +1,59 @@
|
|||
"""
|
||||
LiteLLM A2A - Wrapper for invoking A2A protocol agents.
|
||||
|
||||
This module provides a thin wrapper around the official `a2a` SDK that:
|
||||
- Handles httpx client creation and agent card resolution
|
||||
- Adds LiteLLM logging via @client decorator
|
||||
- Matches the A2A SDK interface (SendMessageRequest, SendMessageResponse, etc.)
|
||||
|
||||
Example usage (standalone functions with @client decorator):
|
||||
```python
|
||||
from litellm.a2a_protocol import asend_message
|
||||
from a2a.types import SendMessageRequest, MessageSendParams
|
||||
from uuid import uuid4
|
||||
|
||||
request = SendMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(
|
||||
message={
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hello!"}],
|
||||
"messageId": uuid4().hex,
|
||||
}
|
||||
)
|
||||
)
|
||||
response = await asend_message(
|
||||
base_url="http://localhost:10001",
|
||||
request=request,
|
||||
)
|
||||
print(response.model_dump(mode='json', exclude_none=True))
|
||||
```
|
||||
|
||||
Example usage (class-based):
|
||||
```python
|
||||
from litellm.a2a_protocol import A2AClient
|
||||
|
||||
client = A2AClient(base_url="http://localhost:10001")
|
||||
response = await client.send_message(request)
|
||||
```
|
||||
"""
|
||||
|
||||
from litellm.a2a_protocol.client import A2AClient
|
||||
from litellm.a2a_protocol.main import (
|
||||
aget_agent_card,
|
||||
asend_message,
|
||||
asend_message_streaming,
|
||||
create_a2a_client,
|
||||
send_message,
|
||||
)
|
||||
from litellm.types.agents import LiteLLMSendMessageResponse
|
||||
|
||||
__all__ = [
|
||||
"A2AClient",
|
||||
"asend_message",
|
||||
"send_message",
|
||||
"asend_message_streaming",
|
||||
"aget_agent_card",
|
||||
"create_a2a_client",
|
||||
"LiteLLMSendMessageResponse",
|
||||
]
|
||||
107
litellm/a2a_protocol/client.py
Normal file
107
litellm/a2a_protocol/client.py
Normal file
|
|
@ -0,0 +1,107 @@
|
|||
"""
|
||||
LiteLLM A2A Client class.
|
||||
|
||||
Provides a class-based interface for A2A agent invocation.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, AsyncIterator, Dict, Optional
|
||||
|
||||
from litellm.types.agents import LiteLLMSendMessageResponse
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from a2a.client import A2AClient as A2AClientType
|
||||
from a2a.types import (
|
||||
AgentCard,
|
||||
SendMessageRequest,
|
||||
SendStreamingMessageRequest,
|
||||
SendStreamingMessageResponse,
|
||||
)
|
||||
|
||||
|
||||
class A2AClient:
|
||||
"""
|
||||
LiteLLM wrapper for A2A agent invocation.
|
||||
|
||||
Creates the underlying A2A client once on first use and reuses it.
|
||||
|
||||
Example:
|
||||
```python
|
||||
from litellm.a2a_protocol import A2AClient
|
||||
from a2a.types import SendMessageRequest, MessageSendParams
|
||||
from uuid import uuid4
|
||||
|
||||
client = A2AClient(base_url="http://localhost:10001")
|
||||
|
||||
request = SendMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(
|
||||
message={
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hello!"}],
|
||||
"messageId": uuid4().hex,
|
||||
}
|
||||
)
|
||||
)
|
||||
response = await client.send_message(request)
|
||||
```
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
base_url: str,
|
||||
timeout: float = 60.0,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
):
|
||||
"""
|
||||
Initialize the A2A client wrapper.
|
||||
|
||||
Args:
|
||||
base_url: The base URL of the A2A agent (e.g., "http://localhost:10001")
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
extra_headers: Optional additional headers to include in requests
|
||||
"""
|
||||
self.base_url = base_url
|
||||
self.timeout = timeout
|
||||
self.extra_headers = extra_headers
|
||||
self._a2a_client: Optional["A2AClientType"] = None
|
||||
|
||||
async def _get_client(self) -> "A2AClientType":
|
||||
"""Get or create the underlying A2A client."""
|
||||
if self._a2a_client is None:
|
||||
from litellm.a2a_protocol.main import create_a2a_client
|
||||
|
||||
self._a2a_client = await create_a2a_client(
|
||||
base_url=self.base_url,
|
||||
timeout=self.timeout,
|
||||
extra_headers=self.extra_headers,
|
||||
)
|
||||
return self._a2a_client
|
||||
|
||||
async def get_agent_card(self) -> "AgentCard":
|
||||
"""Fetch the agent card from the server."""
|
||||
from litellm.a2a_protocol.main import aget_agent_card
|
||||
|
||||
return await aget_agent_card(
|
||||
base_url=self.base_url,
|
||||
timeout=self.timeout,
|
||||
extra_headers=self.extra_headers,
|
||||
)
|
||||
|
||||
async def send_message(
|
||||
self, request: "SendMessageRequest"
|
||||
) -> LiteLLMSendMessageResponse:
|
||||
"""Send a message to the A2A agent."""
|
||||
from litellm.a2a_protocol.main import asend_message
|
||||
|
||||
a2a_client = await self._get_client()
|
||||
return await asend_message(a2a_client=a2a_client, request=request)
|
||||
|
||||
async def send_message_streaming(
|
||||
self, request: "SendStreamingMessageRequest"
|
||||
) -> AsyncIterator["SendStreamingMessageResponse"]:
|
||||
"""Send a streaming message to the A2A agent."""
|
||||
from litellm.a2a_protocol.main import asend_message_streaming
|
||||
|
||||
a2a_client = await self._get_client()
|
||||
async for chunk in asend_message_streaming(a2a_client=a2a_client, request=request):
|
||||
yield chunk
|
||||
36
litellm/a2a_protocol/cost_calculator.py
Normal file
36
litellm/a2a_protocol/cost_calculator.py
Normal file
|
|
@ -0,0 +1,36 @@
|
|||
"""
|
||||
Cost calculator for A2A (Agent-to-Agent) calls.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
Logging as LitellmLoggingObject,
|
||||
)
|
||||
else:
|
||||
LitellmLoggingObject = Any
|
||||
|
||||
|
||||
class A2ACostCalculator:
|
||||
@staticmethod
|
||||
def calculate_a2a_cost(
|
||||
litellm_logging_obj: Optional[LitellmLoggingObject],
|
||||
) -> float:
|
||||
"""
|
||||
Calculate the cost of an A2A send_message call.
|
||||
|
||||
Default is 0.0. In the future, users can configure cost per agent call.
|
||||
"""
|
||||
if litellm_logging_obj is None:
|
||||
return 0.0
|
||||
|
||||
# Check if user set a custom response cost
|
||||
response_cost = litellm_logging_obj.model_call_details.get(
|
||||
"response_cost", None
|
||||
)
|
||||
if response_cost is not None:
|
||||
return response_cost
|
||||
|
||||
# Default to 0.0 for A2A calls
|
||||
return 0.0
|
||||
298
litellm/a2a_protocol/main.py
Normal file
298
litellm/a2a_protocol/main.py
Normal file
|
|
@ -0,0 +1,298 @@
|
|||
"""
|
||||
LiteLLM A2A SDK functions.
|
||||
|
||||
Provides standalone functions with @client decorator for LiteLLM logging integration.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import TYPE_CHECKING, Any, AsyncIterator, Coroutine, Dict, Optional, Union
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
)
|
||||
from litellm.types.agents import LiteLLMSendMessageResponse
|
||||
from litellm.utils import client
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from a2a.client import A2AClient as A2AClientType
|
||||
from a2a.types import (
|
||||
AgentCard,
|
||||
SendMessageRequest,
|
||||
SendStreamingMessageRequest,
|
||||
SendStreamingMessageResponse,
|
||||
)
|
||||
|
||||
# Runtime imports with availability check
|
||||
A2A_SDK_AVAILABLE = False
|
||||
A2ACardResolver: Any = None
|
||||
_A2AClient: Any = None
|
||||
|
||||
try:
|
||||
from a2a.client import A2ACardResolver # type: ignore[no-redef]
|
||||
from a2a.client import A2AClient as _A2AClient # type: ignore[no-redef]
|
||||
|
||||
A2A_SDK_AVAILABLE = True
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
|
||||
def _get_a2a_model_info(a2a_client: Any, kwargs: Dict[str, Any]) -> str:
|
||||
"""
|
||||
Extract agent info and set model/custom_llm_provider for cost tracking.
|
||||
|
||||
Sets model info on the litellm_logging_obj if available.
|
||||
Returns the agent name for logging.
|
||||
"""
|
||||
agent_name = "unknown"
|
||||
|
||||
# Try to get agent card from our stored attribute first, then fallback to SDK attribute
|
||||
agent_card = getattr(a2a_client, "_litellm_agent_card", None)
|
||||
if agent_card is None:
|
||||
agent_card = getattr(a2a_client, "agent_card", None)
|
||||
|
||||
if agent_card is not None:
|
||||
agent_name = getattr(agent_card, "name", "unknown") or "unknown"
|
||||
|
||||
# Build model string
|
||||
model = f"a2a_agent/{agent_name}"
|
||||
custom_llm_provider = "a2a_agent"
|
||||
|
||||
# Set on litellm_logging_obj if available (for standard logging payload)
|
||||
litellm_logging_obj = kwargs.get("litellm_logging_obj")
|
||||
if litellm_logging_obj is not None:
|
||||
litellm_logging_obj.model = model
|
||||
litellm_logging_obj.custom_llm_provider = custom_llm_provider
|
||||
litellm_logging_obj.model_call_details["model"] = model
|
||||
litellm_logging_obj.model_call_details["custom_llm_provider"] = custom_llm_provider
|
||||
|
||||
return agent_name
|
||||
|
||||
|
||||
@client
|
||||
async def asend_message(
|
||||
a2a_client: "A2AClientType",
|
||||
request: "SendMessageRequest",
|
||||
**kwargs: Any,
|
||||
) -> LiteLLMSendMessageResponse:
|
||||
"""
|
||||
Async: Send a message to an A2A agent.
|
||||
|
||||
Uses the @client decorator for LiteLLM logging and tracking.
|
||||
|
||||
Args:
|
||||
a2a_client: An initialized a2a.client.A2AClient instance
|
||||
request: SendMessageRequest from a2a.types
|
||||
**kwargs: Additional arguments passed to the client decorator
|
||||
|
||||
Returns:
|
||||
LiteLLMSendMessageResponse (wraps a2a SendMessageResponse with _hidden_params)
|
||||
|
||||
Example:
|
||||
```python
|
||||
from litellm.a2a_protocol import asend_message, create_a2a_client
|
||||
from a2a.types import SendMessageRequest, MessageSendParams
|
||||
from uuid import uuid4
|
||||
|
||||
# Create client once
|
||||
a2a_client = await create_a2a_client(base_url="http://localhost:10001")
|
||||
|
||||
# Use it for multiple requests
|
||||
request = SendMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(
|
||||
message={
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hello!"}],
|
||||
"messageId": uuid4().hex,
|
||||
}
|
||||
)
|
||||
)
|
||||
response = await asend_message(a2a_client=a2a_client, request=request)
|
||||
```
|
||||
"""
|
||||
agent_name = _get_a2a_model_info(a2a_client, kwargs)
|
||||
|
||||
verbose_logger.info(f"A2A send_message request_id={request.id}, agent={agent_name}")
|
||||
|
||||
a2a_response = await a2a_client.send_message(request)
|
||||
|
||||
verbose_logger.info(f"A2A send_message completed, request_id={request.id}")
|
||||
|
||||
# Wrap in LiteLLM response type for _hidden_params support
|
||||
response = LiteLLMSendMessageResponse.from_a2a_response(a2a_response)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
@client
|
||||
def send_message(
|
||||
a2a_client: "A2AClientType",
|
||||
request: "SendMessageRequest",
|
||||
**kwargs: Any,
|
||||
) -> Union[LiteLLMSendMessageResponse, Coroutine[Any, Any, LiteLLMSendMessageResponse]]:
|
||||
"""
|
||||
Sync: Send a message to an A2A agent.
|
||||
|
||||
Uses the @client decorator for LiteLLM logging and tracking.
|
||||
|
||||
Args:
|
||||
a2a_client: An initialized a2a.client.A2AClient instance
|
||||
request: SendMessageRequest from a2a.types
|
||||
**kwargs: Additional arguments passed to the client decorator
|
||||
|
||||
Returns:
|
||||
LiteLLMSendMessageResponse (wraps a2a SendMessageResponse with _hidden_params)
|
||||
"""
|
||||
try:
|
||||
loop = asyncio.get_running_loop()
|
||||
except RuntimeError:
|
||||
loop = None
|
||||
|
||||
if loop is not None:
|
||||
return asend_message(a2a_client=a2a_client, request=request, **kwargs)
|
||||
else:
|
||||
return asyncio.run(asend_message(a2a_client=a2a_client, request=request, **kwargs))
|
||||
|
||||
|
||||
async def asend_message_streaming(
|
||||
a2a_client: "A2AClientType",
|
||||
request: "SendStreamingMessageRequest",
|
||||
) -> AsyncIterator["SendStreamingMessageResponse"]:
|
||||
"""
|
||||
Async: Send a streaming message to an A2A agent.
|
||||
|
||||
Args:
|
||||
a2a_client: An initialized a2a.client.A2AClient instance
|
||||
request: SendStreamingMessageRequest from a2a.types
|
||||
|
||||
Yields:
|
||||
SendStreamingMessageResponse chunks from the agent
|
||||
"""
|
||||
verbose_logger.info(f"A2A send_message_streaming request_id={request.id}")
|
||||
|
||||
stream = a2a_client.send_message_streaming(request)
|
||||
|
||||
chunk_count = 0
|
||||
async for chunk in stream:
|
||||
chunk_count += 1
|
||||
yield chunk
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A send_message_streaming completed, request_id={request.id}, chunks={chunk_count}"
|
||||
)
|
||||
|
||||
|
||||
async def create_a2a_client(
|
||||
base_url: str,
|
||||
timeout: float = 60.0,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> "A2AClientType":
|
||||
"""
|
||||
Create an A2A client for the given agent URL.
|
||||
|
||||
This resolves the agent card and returns a ready-to-use A2A client.
|
||||
The client can be reused for multiple requests.
|
||||
|
||||
Args:
|
||||
base_url: The base URL of the A2A agent (e.g., "http://localhost:10001")
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
extra_headers: Optional additional headers to include in requests
|
||||
|
||||
Returns:
|
||||
An initialized a2a.client.A2AClient instance
|
||||
|
||||
Example:
|
||||
```python
|
||||
from litellm.a2a_protocol import create_a2a_client, asend_message
|
||||
|
||||
# Create client once
|
||||
client = await create_a2a_client(base_url="http://localhost:10001")
|
||||
|
||||
# Reuse for multiple requests
|
||||
response1 = await asend_message(a2a_client=client, request=request1)
|
||||
response2 = await asend_message(a2a_client=client, request=request2)
|
||||
```
|
||||
"""
|
||||
if not A2A_SDK_AVAILABLE:
|
||||
raise ImportError(
|
||||
"The 'a2a' package is required for A2A agent invocation. "
|
||||
"Install it with: pip install a2a"
|
||||
)
|
||||
|
||||
verbose_logger.info(f"Creating A2A client for {base_url}")
|
||||
|
||||
# Use LiteLLM's cached httpx client
|
||||
http_handler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.A2A,
|
||||
params={"timeout": timeout},
|
||||
)
|
||||
httpx_client = http_handler.client
|
||||
|
||||
# Resolve agent card
|
||||
resolver = A2ACardResolver(
|
||||
httpx_client=httpx_client,
|
||||
base_url=base_url,
|
||||
)
|
||||
agent_card = await resolver.get_agent_card()
|
||||
|
||||
verbose_logger.debug(
|
||||
f"Resolved agent card: {agent_card.name if hasattr(agent_card, 'name') else 'unknown'}"
|
||||
)
|
||||
|
||||
# Create A2A client
|
||||
a2a_client = _A2AClient(
|
||||
httpx_client=httpx_client,
|
||||
agent_card=agent_card,
|
||||
)
|
||||
|
||||
# Store agent_card on client for later retrieval (SDK doesn't expose it)
|
||||
a2a_client._litellm_agent_card = agent_card # type: ignore[attr-defined]
|
||||
|
||||
verbose_logger.info(f"A2A client created for {base_url}")
|
||||
|
||||
return a2a_client
|
||||
|
||||
|
||||
async def aget_agent_card(
|
||||
base_url: str,
|
||||
timeout: float = 60.0,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> "AgentCard":
|
||||
"""
|
||||
Fetch the agent card from an A2A agent.
|
||||
|
||||
Args:
|
||||
base_url: The base URL of the A2A agent (e.g., "http://localhost:10001")
|
||||
timeout: Request timeout in seconds (default: 60.0)
|
||||
extra_headers: Optional additional headers to include in requests
|
||||
|
||||
Returns:
|
||||
AgentCard from the A2A agent
|
||||
"""
|
||||
if not A2A_SDK_AVAILABLE:
|
||||
raise ImportError(
|
||||
"The 'a2a' package is required for A2A agent invocation. "
|
||||
"Install it with: pip install a2a"
|
||||
)
|
||||
|
||||
verbose_logger.info(f"Fetching agent card from {base_url}")
|
||||
|
||||
# Use LiteLLM's cached httpx client
|
||||
http_handler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.A2A,
|
||||
params={"timeout": timeout},
|
||||
)
|
||||
httpx_client = http_handler.client
|
||||
|
||||
resolver = A2ACardResolver(
|
||||
httpx_client=httpx_client,
|
||||
base_url=base_url,
|
||||
)
|
||||
agent_card = await resolver.get_agent_card()
|
||||
|
||||
verbose_logger.info(
|
||||
f"Fetched agent card: {agent_card.name if hasattr(agent_card, 'name') else 'unknown'}"
|
||||
)
|
||||
return agent_card
|
||||
|
|
@ -15,6 +15,7 @@ from litellm.utils import token_counter
|
|||
async def calculate_batch_cost_and_usage(
|
||||
file_content_dictionary: List[dict],
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm"],
|
||||
model_name: Optional[str] = None,
|
||||
) -> Tuple[float, Usage, List[str]]:
|
||||
"""
|
||||
Calculate the cost and usage of a batch
|
||||
|
|
@ -37,6 +38,7 @@ async def calculate_batch_cost_and_usage(
|
|||
async def _handle_completed_batch(
|
||||
batch: Batch,
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm"],
|
||||
model_name: Optional[str] = None,
|
||||
) -> Tuple[float, Usage, List[str]]:
|
||||
"""Helper function to process a completed batch and handle logging"""
|
||||
# Get batch results
|
||||
|
|
@ -83,6 +85,7 @@ def _get_batch_models_from_file_content(
|
|||
def _batch_cost_calculator(
|
||||
file_content_dictionary: List[dict],
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm"] = "openai",
|
||||
model_name: Optional[str] = None,
|
||||
) -> float:
|
||||
"""
|
||||
Calculate the cost of a batch based on the output file id
|
||||
|
|
@ -251,6 +254,7 @@ def _get_batch_job_cost_from_file_content(
|
|||
def _get_batch_job_total_usage_from_file_content(
|
||||
file_content_dictionary: List[dict],
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm"] = "openai",
|
||||
model_name: Optional[str] = None,
|
||||
) -> Usage:
|
||||
"""
|
||||
Get the tokens of a batch job from the file content
|
||||
|
|
|
|||
|
|
@ -367,49 +367,14 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
reasoning_content = None # flush reasoning content
|
||||
index += 1
|
||||
elif isinstance(item, ResponseFunctionToolCall):
|
||||
|
||||
provider_specific_fields = getattr(
|
||||
item, "provider_specific_fields", None
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
if provider_specific_fields and not isinstance(
|
||||
provider_specific_fields, dict
|
||||
):
|
||||
provider_specific_fields = (
|
||||
dict(provider_specific_fields)
|
||||
if hasattr(provider_specific_fields, "__dict__")
|
||||
else {}
|
||||
)
|
||||
elif hasattr(item, "get") and callable(item.get): # type: ignore
|
||||
provider_fields = item.get("provider_specific_fields") # type: ignore
|
||||
if provider_fields:
|
||||
provider_specific_fields = (
|
||||
provider_fields
|
||||
if isinstance(provider_fields, dict)
|
||||
else (
|
||||
dict(provider_fields) # type: ignore
|
||||
if hasattr(provider_fields, "__dict__")
|
||||
else {}
|
||||
)
|
||||
)
|
||||
|
||||
function_dict: Dict[str, Any] = {
|
||||
"name": item.name,
|
||||
"arguments": item.arguments,
|
||||
}
|
||||
|
||||
if provider_specific_fields:
|
||||
function_dict["provider_specific_fields"] = provider_specific_fields
|
||||
|
||||
tool_call_dict: Dict[str, Any] = {
|
||||
"id": item.call_id,
|
||||
"function": function_dict,
|
||||
"type": "function",
|
||||
}
|
||||
|
||||
if provider_specific_fields:
|
||||
tool_call_dict["provider_specific_fields"] = (
|
||||
provider_specific_fields
|
||||
)
|
||||
tool_call_dict = LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
|
||||
tool_call_item=item,
|
||||
index=index,
|
||||
)
|
||||
|
||||
msg = Message(
|
||||
content=None,
|
||||
|
|
@ -718,17 +683,9 @@ class LiteLLMResponsesTransformationHandler(CompletionTransformationBridge):
|
|||
}
|
||||
}
|
||||
elif format_type == "json_object":
|
||||
return {
|
||||
"format": {
|
||||
"type": "json_object"
|
||||
}
|
||||
}
|
||||
return {"format": {"type": "json_object"}}
|
||||
elif format_type == "text":
|
||||
return {
|
||||
"format": {
|
||||
"type": "text"
|
||||
}
|
||||
}
|
||||
return {"format": {"type": "text"}}
|
||||
|
||||
return None
|
||||
|
||||
|
|
|
|||
|
|
@ -100,8 +100,10 @@ RUNWAYML_POLLING_TIMEOUT = int(
|
|||
########## Networking constants ##############################################################
|
||||
_DEFAULT_TTL_FOR_HTTPX_CLIENTS = 3600 # 1 hour, re-use the same httpx client for 1 hour
|
||||
|
||||
# Aiohttp connection pooling constants
|
||||
AIOHTTP_CONNECTOR_LIMIT = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT", 0))
|
||||
# Aiohttp connection pooling - prevents memory leaks from unbounded connection growth
|
||||
# Set to 0 for unlimited (not recommended for production)
|
||||
AIOHTTP_CONNECTOR_LIMIT = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT", 300))
|
||||
AIOHTTP_CONNECTOR_LIMIT_PER_HOST = int(os.getenv("AIOHTTP_CONNECTOR_LIMIT_PER_HOST", 50))
|
||||
AIOHTTP_KEEPALIVE_TIMEOUT = int(os.getenv("AIOHTTP_KEEPALIVE_TIMEOUT", 120))
|
||||
AIOHTTP_TTL_DNS_CACHE = int(os.getenv("AIOHTTP_TTL_DNS_CACHE", 300))
|
||||
# enable_cleanup_closed is only needed for Python versions with the SSL leak bug
|
||||
|
|
@ -543,7 +545,7 @@ openai_compatible_endpoints: List = [
|
|||
"api.studio.nebius.ai/v1",
|
||||
"https://dashscope-intl.aliyuncs.com/compatible-mode/v1",
|
||||
"https://api.moonshot.ai/v1",
|
||||
"https://platform.publicai.co/v1",
|
||||
"https://api.publicai.co/v1",
|
||||
"https://api.v0.dev/v1",
|
||||
"https://api.morphllm.com/v1",
|
||||
"https://api.lambda.ai/v1",
|
||||
|
|
@ -585,6 +587,7 @@ openai_compatible_providers: List = [
|
|||
"github_copilot", # GitHub Copilot Chat API
|
||||
"novita",
|
||||
"meta_llama",
|
||||
"publicai", # PublicAI - JSON-configured provider
|
||||
"featherless_ai",
|
||||
"nscale",
|
||||
"nebius",
|
||||
|
|
@ -874,6 +877,7 @@ BEDROCK_INVOKE_PROVIDERS_LITERAL = Literal[
|
|||
"nova",
|
||||
"deepseek_r1",
|
||||
"qwen3",
|
||||
"qwen2",
|
||||
"twelvelabs",
|
||||
"openai",
|
||||
]
|
||||
|
|
|
|||
|
|
@ -95,6 +95,7 @@ from litellm.utils import (
|
|||
EmbeddingResponse,
|
||||
ImageResponse,
|
||||
ModelResponse,
|
||||
ModelResponseStream,
|
||||
ProviderConfigManager,
|
||||
TextCompletionResponse,
|
||||
TranscriptionResponse,
|
||||
|
|
@ -654,7 +655,9 @@ def _infer_call_type(
|
|||
if completion_response is None:
|
||||
return None
|
||||
|
||||
if isinstance(completion_response, ModelResponse):
|
||||
if isinstance(completion_response, ModelResponse) or isinstance(
|
||||
completion_response, ModelResponseStream
|
||||
):
|
||||
return "completion"
|
||||
elif isinstance(completion_response, EmbeddingResponse):
|
||||
return "embedding"
|
||||
|
|
@ -934,6 +937,17 @@ def completion_cost( # noqa: PLR0915
|
|||
prompt_tokens = token_counter(model=model, text=prompt)
|
||||
completion_tokens = token_counter(model=model, text=completion)
|
||||
|
||||
# Handle A2A calls before model check - A2A doesn't require a model
|
||||
if call_type in (
|
||||
CallTypes.asend_message.value,
|
||||
CallTypes.send_message.value,
|
||||
):
|
||||
from litellm.a2a_protocol.cost_calculator import A2ACostCalculator
|
||||
|
||||
return A2ACostCalculator.calculate_a2a_cost(
|
||||
litellm_logging_obj=litellm_logging_obj
|
||||
)
|
||||
|
||||
if model is None:
|
||||
raise ValueError(
|
||||
f"Model is None and does not exist in passed completion_response. Passed completion_response={completion_response}, model={model}"
|
||||
|
|
@ -1046,29 +1060,33 @@ def completion_cost( # noqa: PLR0915
|
|||
number_of_queries = len(query)
|
||||
elif query is not None:
|
||||
number_of_queries = 1
|
||||
|
||||
|
||||
search_model = model or ""
|
||||
if custom_llm_provider and "/" not in search_model:
|
||||
# If model is like "tavily-search", construct "tavily/search" for cost lookup
|
||||
search_model = f"{custom_llm_provider}/search"
|
||||
|
||||
prompt_cost, completion_cost_result = search_provider_cost_per_query(
|
||||
model=search_model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
number_of_queries=number_of_queries,
|
||||
optional_params=optional_params,
|
||||
|
||||
prompt_cost, completion_cost_result = (
|
||||
search_provider_cost_per_query(
|
||||
model=search_model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
number_of_queries=number_of_queries,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# Return the total cost (prompt_cost + completion_cost, but for search it's just prompt_cost)
|
||||
_final_cost = prompt_cost + completion_cost_result
|
||||
|
||||
|
||||
# Apply discount
|
||||
original_cost = _final_cost
|
||||
_final_cost, discount_percent, discount_amount = _apply_cost_discount(
|
||||
base_cost=_final_cost,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
_final_cost, discount_percent, discount_amount = (
|
||||
_apply_cost_discount(
|
||||
base_cost=_final_cost,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
# Store cost breakdown in logging object if available
|
||||
_store_cost_breakdown_in_logging_obj(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
|
|
@ -1080,7 +1098,7 @@ def completion_cost( # noqa: PLR0915
|
|||
discount_percent=discount_percent,
|
||||
discount_amount=discount_amount,
|
||||
)
|
||||
|
||||
|
||||
return _final_cost
|
||||
elif call_type == CallTypes.arealtime.value and isinstance(
|
||||
completion_response, LiteLLMRealtimeStreamLoggingObject
|
||||
|
|
|
|||
|
|
@ -18,6 +18,7 @@ from litellm import get_secret_str
|
|||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.azure.files.handler import AzureOpenAIFilesAPI
|
||||
from litellm.llms.bedrock.files.handler import BedrockFilesHandler
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.llms.openai.openai import FileDeleted, FileObject, OpenAIFilesAPI
|
||||
|
|
@ -47,6 +48,7 @@ base_llm_http_handler = BaseLLMHTTPHandler()
|
|||
openai_files_instance = OpenAIFilesAPI()
|
||||
azure_files_instance = AzureOpenAIFilesAPI()
|
||||
vertex_ai_files_instance = VertexAIFilesHandler()
|
||||
bedrock_files_instance = BedrockFilesHandler()
|
||||
#################################################
|
||||
|
||||
|
||||
|
|
@ -755,7 +757,7 @@ def file_list(
|
|||
@client
|
||||
async def afile_content(
|
||||
file_id: str,
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "hosted_vllm"] = "openai",
|
||||
custom_llm_provider: Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"] = "openai",
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_body: Optional[Dict[str, str]] = None,
|
||||
**kwargs,
|
||||
|
|
@ -800,7 +802,7 @@ def file_content(
|
|||
file_id: str,
|
||||
model: Optional[str] = None,
|
||||
custom_llm_provider: Optional[
|
||||
Union[Literal["openai", "azure", "vertex_ai", "hosted_vllm"], str]
|
||||
Union[Literal["openai", "azure", "vertex_ai", "bedrock", "hosted_vllm"], str]
|
||||
] = None,
|
||||
extra_headers: Optional[Dict[str, str]] = None,
|
||||
extra_body: Optional[Dict[str, str]] = None,
|
||||
|
|
@ -938,9 +940,18 @@ def file_content(
|
|||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
)
|
||||
elif custom_llm_provider == "bedrock":
|
||||
response = bedrock_files_instance.file_content(
|
||||
_is_async=_is_async,
|
||||
file_content_request=_file_content_request,
|
||||
api_base=optional_params.api_base,
|
||||
optional_params=litellm_params_dict,
|
||||
timeout=timeout,
|
||||
max_retries=optional_params.max_retries,
|
||||
)
|
||||
else:
|
||||
raise litellm.exceptions.BadRequestError(
|
||||
message="LiteLLM doesn't support {} for 'custom_llm_provider'. Supported providers are 'openai', 'azure', 'vertex_ai'.".format(
|
||||
message="LiteLLM doesn't support {} for 'custom_llm_provider'. Supported providers are 'openai', 'azure', 'vertex_ai', 'bedrock'.".format(
|
||||
custom_llm_provider
|
||||
),
|
||||
model="n/a",
|
||||
|
|
|
|||
|
|
@ -6,7 +6,9 @@ from typing import Any, Coroutine, Dict, List, Literal, Optional, Union, cast, o
|
|||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm import client, exception_type, get_litellm_params
|
||||
from litellm.utils import exception_type, get_litellm_params
|
||||
# client is imported from litellm as it's a decorator
|
||||
from litellm import client
|
||||
from litellm.constants import DEFAULT_IMAGE_ENDPOINT_MODEL
|
||||
from litellm.constants import request_timeout as DEFAULT_REQUEST_TIMEOUT
|
||||
from litellm.exceptions import LiteLLMUnknownProvider
|
||||
|
|
|
|||
|
|
@ -1,18 +1,20 @@
|
|||
import os
|
||||
from typing import TYPE_CHECKING, Any, Union
|
||||
from typing import TYPE_CHECKING, Any, Optional, Union
|
||||
from datetime import datetime
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations.arize import _utils
|
||||
from litellm.integrations.arize._utils import ArizeOTELAttributes
|
||||
from litellm.types.integrations.arize_phoenix import ArizePhoenixConfig
|
||||
from litellm.types.services import ServiceLoggerPayload
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
||||
from litellm.integrations.opentelemetry import OpenTelemetryConfig as _OpenTelemetryConfig
|
||||
from litellm.types.integrations.arize import Protocol as _Protocol
|
||||
|
||||
from .opentelemetry import OpenTelemetryConfig as _OpenTelemetryConfig
|
||||
|
||||
Protocol = _Protocol
|
||||
OpenTelemetryConfig = _OpenTelemetryConfig
|
||||
Span = Union[_Span, Any]
|
||||
|
|
@ -25,7 +27,11 @@ else:
|
|||
ARIZE_HOSTED_PHOENIX_ENDPOINT = "https://otlp.arize.com/v1/traces"
|
||||
|
||||
|
||||
class ArizePhoenixLogger:
|
||||
class ArizePhoenixLogger(OpenTelemetry):
|
||||
def set_attributes(self, span: Span, kwargs, response_obj: Optional[Any]):
|
||||
ArizePhoenixLogger.set_arize_phoenix_attributes(span, kwargs, response_obj)
|
||||
return
|
||||
|
||||
@staticmethod
|
||||
def set_arize_phoenix_attributes(span: Span, kwargs, response_obj):
|
||||
_utils.set_attributes(span, kwargs, response_obj, ArizeOTELAttributes)
|
||||
|
|
@ -97,3 +103,46 @@ class ArizePhoenixLogger:
|
|||
endpoint=endpoint,
|
||||
project_name=project_name,
|
||||
)
|
||||
|
||||
async def async_service_success_hook(
|
||||
self,
|
||||
payload: ServiceLoggerPayload,
|
||||
parent_otel_span: Optional[Span] = None,
|
||||
start_time: Optional[Union[datetime, float]] = None,
|
||||
end_time: Optional[Union[datetime, float]] = None,
|
||||
event_metadata: Optional[dict] = None,
|
||||
):
|
||||
pass # suppress additional spans
|
||||
|
||||
async def async_service_failure_hook(
|
||||
self,
|
||||
payload: ServiceLoggerPayload,
|
||||
error: Optional[str] = "",
|
||||
parent_otel_span: Optional[Span] = None,
|
||||
start_time: Optional[Union[datetime, float]] = None,
|
||||
end_time: Optional[Union[float, datetime]] = None,
|
||||
event_metadata: Optional[dict] = None,
|
||||
):
|
||||
pass # suppress additional spans
|
||||
|
||||
def create_litellm_proxy_request_started_span(
|
||||
self,
|
||||
start_time: datetime,
|
||||
headers: dict,
|
||||
):
|
||||
pass # suppress additional spans
|
||||
|
||||
async def async_health_check(self):
|
||||
|
||||
config = self.get_arize_phoenix_config()
|
||||
|
||||
if not config.otlp_auth_headers:
|
||||
return {
|
||||
"status": "unhealthy",
|
||||
"error_message": "PHOENIX_API_KEY environment variable not set",
|
||||
}
|
||||
|
||||
return {
|
||||
"status": "healthy",
|
||||
"message": "Arize-Phoenix credentials are configured properly",
|
||||
}
|
||||
|
|
@ -6,7 +6,6 @@ from typing import (
|
|||
List,
|
||||
Literal,
|
||||
Optional,
|
||||
Tuple,
|
||||
Type,
|
||||
Union,
|
||||
get_args,
|
||||
|
|
@ -17,6 +16,7 @@ from litellm.caching import DualCache
|
|||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.types.guardrails import (
|
||||
DynamicGuardrailParams,
|
||||
GenericGuardrailAPIInputs,
|
||||
GuardrailEventHooks,
|
||||
LitellmParams,
|
||||
Mode,
|
||||
|
|
@ -449,20 +449,22 @@ class CustomGuardrail(CustomLogger):
|
|||
|
||||
async def apply_guardrail(
|
||||
self,
|
||||
texts: List[str],
|
||||
inputs: GenericGuardrailAPIInputs,
|
||||
request_data: dict,
|
||||
input_type: Literal["request", "response"],
|
||||
logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
images: Optional[List[str]] = None,
|
||||
) -> Tuple[List[str], Optional[List[str]]]:
|
||||
) -> GenericGuardrailAPIInputs:
|
||||
"""
|
||||
Apply your guardrail logic to the given text
|
||||
Apply your guardrail logic to the given inputs
|
||||
|
||||
Args:
|
||||
texts: The texts to apply the guardrail to
|
||||
images: The images to apply the guardrail to
|
||||
inputs: Dictionary containing:
|
||||
- texts: List of texts to apply the guardrail to
|
||||
- images: Optional list of images to apply the guardrail to
|
||||
- tool_calls: Optional list of tool calls to apply the guardrail to
|
||||
request_data: The request data dictionary - containing user api key metadata (e.g. user_id, team_id, etc.)
|
||||
input_type: The type of input to apply the guardrail to - "request" or "response"
|
||||
logging_obj: Optional logging object for tracking the guardrail execution
|
||||
|
||||
Any of the custom guardrails can override this method to provide custom guardrail logic
|
||||
|
||||
|
|
@ -473,7 +475,7 @@ class CustomGuardrail(CustomLogger):
|
|||
- If the guardrail raises an exception
|
||||
|
||||
"""
|
||||
return texts, images
|
||||
return inputs
|
||||
|
||||
def _process_response(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -1065,14 +1065,7 @@ class OpenTelemetry(CustomLogger):
|
|||
self, span: Span, kwargs, response_obj: Optional[Any]
|
||||
):
|
||||
try:
|
||||
if self.callback_name == "arize_phoenix":
|
||||
from litellm.integrations.arize.arize_phoenix import ArizePhoenixLogger
|
||||
|
||||
ArizePhoenixLogger.set_arize_phoenix_attributes(
|
||||
span, kwargs, response_obj
|
||||
)
|
||||
return
|
||||
elif self.callback_name == "langtrace":
|
||||
if self.callback_name == "langtrace":
|
||||
from litellm.integrations.langtrace import LangtraceAttributes
|
||||
|
||||
LangtraceAttributes().set_langtrace_attributes(
|
||||
|
|
@ -1088,6 +1081,11 @@ class OpenTelemetry(CustomLogger):
|
|||
span, kwargs, response_obj
|
||||
)
|
||||
return
|
||||
elif self.callback_name == "weave_otel":
|
||||
from litellm.integrations.weave.weave_otel import set_weave_otel_attributes
|
||||
|
||||
set_weave_otel_attributes(span, kwargs, response_obj)
|
||||
return
|
||||
from litellm.proxy._types import SpanAttributes
|
||||
|
||||
optional_params = kwargs.get("optional_params", {})
|
||||
|
|
|
|||
|
|
@ -74,9 +74,20 @@ class VectorStorePreCallHook(CustomLogger):
|
|||
if litellm.vector_store_registry is None:
|
||||
return model, messages, non_default_params
|
||||
|
||||
# Get prisma_client for database fallback
|
||||
prisma_client = None
|
||||
try:
|
||||
from litellm.proxy.proxy_server import prisma_client as _prisma_client
|
||||
prisma_client = _prisma_client
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
# Use database fallback to ensure synchronization across instances
|
||||
vector_stores_to_run: List[LiteLLM_ManagedVectorStore] = (
|
||||
litellm.vector_store_registry.pop_vector_stores_to_run(
|
||||
non_default_params=non_default_params, tools=tools
|
||||
await litellm.vector_store_registry.pop_vector_stores_to_run_with_db_fallback(
|
||||
non_default_params=non_default_params,
|
||||
tools=tools,
|
||||
prisma_client=prisma_client
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
7
litellm/integrations/weave/__init__.py
Normal file
7
litellm/integrations/weave/__init__.py
Normal file
|
|
@ -0,0 +1,7 @@
|
|||
"""
|
||||
Weave (W&B) integration for LiteLLM via OpenTelemetry.
|
||||
"""
|
||||
|
||||
from litellm.integrations.weave.weave_otel import WeaveOtelLogger
|
||||
|
||||
__all__ = ["WeaveOtelLogger"]
|
||||
329
litellm/integrations/weave/weave_otel.py
Normal file
329
litellm/integrations/weave/weave_otel.py
Normal file
|
|
@ -0,0 +1,329 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import json
|
||||
import os
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
||||
from opentelemetry.trace import Status, StatusCode
|
||||
from typing_extensions import override
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.integrations._types.open_inference import SpanAttributes as OpenInferenceSpanAttributes
|
||||
from litellm.integrations.arize import _utils
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry, OpenTelemetryConfig
|
||||
from litellm.integrations.opentelemetry_utils.base_otel_llm_obs_attributes import (
|
||||
BaseLLMObsOTELAttributes,
|
||||
safe_set_attribute,
|
||||
)
|
||||
from litellm.litellm_core_utils.safe_json_dumps import safe_dumps
|
||||
from litellm.types.integrations.weave_otel import WeaveOtelConfig, WeaveSpanAttributes
|
||||
from litellm.types.utils import StandardCallbackDynamicParams
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span
|
||||
|
||||
|
||||
# Weave OTEL endpoint
|
||||
# Multi-tenant cloud: https://trace.wandb.ai/otel/v1/traces
|
||||
# Dedicated cloud: https://<your-subdomain>.wandb.io/traces/otel/v1/traces
|
||||
WEAVE_BASE_URL = "https://trace.wandb.ai"
|
||||
WEAVE_OTEL_ENDPOINT = "/otel/v1/traces"
|
||||
|
||||
|
||||
class WeaveLLMObsOTELAttributes(BaseLLMObsOTELAttributes):
|
||||
"""
|
||||
Weave-specific LLM observability OTEL attributes.
|
||||
|
||||
Weave automatically maps attributes from multiple frameworks including
|
||||
GenAI, OpenInference, Langfuse, and others.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
@override
|
||||
def set_messages(span: "Span", kwargs: dict[str, Any]):
|
||||
"""Set input messages as span attributes using OpenInference conventions."""
|
||||
|
||||
messages = kwargs.get("messages") or []
|
||||
optional_params = kwargs.get("optional_params") or {}
|
||||
|
||||
prompt = {"messages": messages}
|
||||
functions = optional_params.get("functions")
|
||||
tools = optional_params.get("tools")
|
||||
if functions is not None:
|
||||
prompt["functions"] = functions
|
||||
if tools is not None:
|
||||
prompt["tools"] = tools
|
||||
safe_set_attribute(span, OpenInferenceSpanAttributes.INPUT_VALUE, json.dumps(prompt))
|
||||
|
||||
|
||||
def _set_weave_specific_attributes(span: Span, kwargs: dict[str, Any], response_obj: Any):
|
||||
"""
|
||||
Sets Weave-specific metadata attributes onto the OTEL span.
|
||||
|
||||
Based on Weave's OTEL attribute mappings from:
|
||||
https://github.com/wandb/weave/blob/master/weave/trace_server/opentelemetry/constants.py
|
||||
"""
|
||||
|
||||
# Extract all needed data upfront
|
||||
litellm_params = kwargs.get("litellm_params") or {}
|
||||
# optional_params = kwargs.get("optional_params") or {}
|
||||
metadata = kwargs.get("metadata") or {}
|
||||
model = kwargs.get("model") or ""
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider") or ""
|
||||
|
||||
# Weave supports a custom display name and will default to the model name if not provided.
|
||||
display_name = metadata.get("display_name")
|
||||
if not display_name and model:
|
||||
if custom_llm_provider:
|
||||
display_name = f"{custom_llm_provider}/{model}"
|
||||
else:
|
||||
display_name = model
|
||||
if display_name:
|
||||
display_name = display_name.replace("/", "__")
|
||||
safe_set_attribute(span, WeaveSpanAttributes.DISPLAY_NAME.value, display_name)
|
||||
|
||||
# Weave threads are OpenInference sessions.
|
||||
if (session_id := metadata.get("session_id")) is not None:
|
||||
if isinstance(session_id, (list, dict)):
|
||||
session_id = safe_dumps(session_id)
|
||||
safe_set_attribute(span, WeaveSpanAttributes.THREAD_ID.value, session_id)
|
||||
safe_set_attribute(span, WeaveSpanAttributes.IS_TURN.value, True)
|
||||
|
||||
# Response attributes are already set by _utils.set_attributes,
|
||||
# but we override them here to better match Weave's expectations
|
||||
if response_obj:
|
||||
output_dict = None
|
||||
if hasattr(response_obj, "model_dump"):
|
||||
output_dict = response_obj.model_dump()
|
||||
elif hasattr(response_obj, "get"):
|
||||
output_dict = response_obj
|
||||
|
||||
if output_dict:
|
||||
safe_set_attribute(span, OpenInferenceSpanAttributes.OUTPUT_VALUE, safe_dumps(output_dict))
|
||||
|
||||
|
||||
def _get_weave_authorization_header(api_key: str) -> str:
|
||||
"""
|
||||
Get the authorization header for Weave OpenTelemetry.
|
||||
|
||||
Weave uses Basic auth with format: api:<WANDB_API_KEY>
|
||||
"""
|
||||
auth_string = f"api:{api_key}"
|
||||
auth_header = base64.b64encode(auth_string.encode()).decode()
|
||||
return f"Basic {auth_header}"
|
||||
|
||||
|
||||
def get_weave_otel_config() -> WeaveOtelConfig:
|
||||
"""
|
||||
Retrieves the Weave OpenTelemetry configuration based on environment variables.
|
||||
|
||||
Environment Variables:
|
||||
WANDB_API_KEY: Required. W&B API key for authentication.
|
||||
WANDB_PROJECT_ID: Required. Project ID in format <entity>/<project_name>.
|
||||
WANDB_HOST: Optional. Custom Weave host URL. Defaults to cloud endpoint.
|
||||
|
||||
Returns:
|
||||
WeaveOtelConfig: A Pydantic model containing Weave OTEL configuration.
|
||||
|
||||
Raises:
|
||||
ValueError: If required environment variables are missing.
|
||||
"""
|
||||
api_key = os.getenv("WANDB_API_KEY")
|
||||
project_id = os.getenv("WANDB_PROJECT_ID")
|
||||
host = os.getenv("WANDB_HOST")
|
||||
|
||||
if not api_key:
|
||||
raise ValueError("WANDB_API_KEY must be set for Weave OpenTelemetry integration.")
|
||||
|
||||
if not project_id:
|
||||
raise ValueError(
|
||||
"WANDB_PROJECT_ID must be set for Weave OpenTelemetry integration. Format: <entity>/<project_name>"
|
||||
)
|
||||
|
||||
if host:
|
||||
if not host.startswith("http"):
|
||||
host = "https://" + host
|
||||
# Self-managed instances use a different path
|
||||
endpoint = host.rstrip("/") + WEAVE_OTEL_ENDPOINT
|
||||
verbose_logger.debug(f"Using Weave OTEL endpoint from host: {endpoint}")
|
||||
else:
|
||||
endpoint = WEAVE_BASE_URL + WEAVE_OTEL_ENDPOINT
|
||||
verbose_logger.debug(f"Using Weave cloud endpoint: {endpoint}")
|
||||
|
||||
# Weave uses Basic auth with format: api:<WANDB_API_KEY>
|
||||
auth_header = _get_weave_authorization_header(api_key=api_key)
|
||||
otlp_auth_headers = f"Authorization={auth_header},project_id={project_id}"
|
||||
|
||||
# Set standard OTEL environment variables
|
||||
os.environ["OTEL_EXPORTER_OTLP_ENDPOINT"] = endpoint
|
||||
os.environ["OTEL_EXPORTER_OTLP_HEADERS"] = otlp_auth_headers
|
||||
|
||||
return WeaveOtelConfig(
|
||||
otlp_auth_headers=otlp_auth_headers,
|
||||
endpoint=endpoint,
|
||||
project_id=project_id,
|
||||
protocol="otlp_http",
|
||||
)
|
||||
|
||||
|
||||
def set_weave_otel_attributes(span: Span, kwargs: dict[str, Any], response_obj: Any):
|
||||
"""
|
||||
Sets OpenTelemetry span attributes for Weave observability.
|
||||
Uses the same attribute setting logic as other OTEL integrations for consistency.
|
||||
"""
|
||||
_utils.set_attributes(span, kwargs, response_obj, WeaveLLMObsOTELAttributes)
|
||||
_set_weave_specific_attributes(span=span, kwargs=kwargs, response_obj=response_obj)
|
||||
|
||||
|
||||
class WeaveOtelLogger(OpenTelemetry):
|
||||
"""
|
||||
Weave (W&B) OpenTelemetry Logger for LiteLLM.
|
||||
|
||||
Sends LLM traces to Weave via the OpenTelemetry Protocol (OTLP).
|
||||
|
||||
Environment Variables:
|
||||
WANDB_API_KEY: Required. Weights & Biases API key for authentication.
|
||||
WANDB_PROJECT_ID: Required. Project ID in format <entity>/<project_name>.
|
||||
WANDB_HOST: Optional. Custom Weave host URL. Defaults to cloud endpoint.
|
||||
|
||||
Usage:
|
||||
litellm.callbacks = ["weave_otel"]
|
||||
|
||||
Or manually:
|
||||
from litellm.integrations.weave.weave_otel import WeaveOtelLogger
|
||||
weave_logger = WeaveOtelLogger(callback_name="weave_otel")
|
||||
litellm.callbacks = [weave_logger]
|
||||
|
||||
Reference:
|
||||
https://docs.wandb.ai/weave/guides/tracking/otel
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
config: Optional[OpenTelemetryConfig] = None,
|
||||
callback_name: Optional[str] = "weave_otel",
|
||||
**kwargs,
|
||||
):
|
||||
"""
|
||||
Initialize WeaveOtelLogger.
|
||||
|
||||
If config is not provided, automatically configures from environment variables
|
||||
(WANDB_API_KEY, WANDB_PROJECT_ID, WANDB_HOST) via get_weave_otel_config().
|
||||
"""
|
||||
if config is None:
|
||||
# Auto-configure from Weave environment variables
|
||||
weave_config = get_weave_otel_config()
|
||||
|
||||
config = OpenTelemetryConfig(
|
||||
exporter=weave_config.protocol,
|
||||
endpoint=weave_config.endpoint,
|
||||
headers=weave_config.otlp_auth_headers,
|
||||
)
|
||||
|
||||
super().__init__(config=config, callback_name=callback_name, **kwargs)
|
||||
|
||||
def _maybe_log_raw_request(self, kwargs, response_obj, start_time, end_time, parent_span):
|
||||
"""
|
||||
Override to skip creating the raw_gen_ai_request child span.
|
||||
|
||||
For Weave, we only want a single span per LLM call. The parent span
|
||||
already contains all the necessary attributes, so the child span
|
||||
is redundant.
|
||||
"""
|
||||
pass
|
||||
|
||||
def _start_primary_span(
|
||||
self,
|
||||
kwargs,
|
||||
response_obj,
|
||||
start_time,
|
||||
end_time,
|
||||
context,
|
||||
parent_span=None,
|
||||
):
|
||||
"""
|
||||
Override to always create a child span instead of reusing the parent span.
|
||||
|
||||
This ensures that wrapper spans (like "B", "C", "D", "E") remain separate
|
||||
from the LiteLLM LLM call spans, creating proper nesting in Weave.
|
||||
"""
|
||||
|
||||
otel_tracer = self.get_tracer_to_use_for_request(kwargs)
|
||||
# Always create a new child span, even if parent_span is provided
|
||||
# This ensures wrapper spans remain separate from LLM call spans
|
||||
span = otel_tracer.start_span(
|
||||
name=self._get_span_name(kwargs),
|
||||
start_time=self._to_ns(start_time),
|
||||
context=context,
|
||||
)
|
||||
span.set_status(Status(StatusCode.OK))
|
||||
self.set_attributes(span, kwargs, response_obj)
|
||||
span.end(end_time=self._to_ns(end_time))
|
||||
return span
|
||||
|
||||
def _handle_success(self, kwargs, response_obj, start_time, end_time):
|
||||
"""
|
||||
Override to prevent ending externally created parent spans.
|
||||
|
||||
When wrapper spans (like "B", "C", "D", "E") are provided as parent spans,
|
||||
they should be managed by the user code, not ended by LiteLLM.
|
||||
"""
|
||||
|
||||
verbose_logger.debug(
|
||||
"Weave OpenTelemetry Logger: Logging kwargs: %s, OTEL config settings=%s",
|
||||
kwargs,
|
||||
self.config,
|
||||
)
|
||||
ctx, parent_span = self._get_span_context(kwargs)
|
||||
|
||||
# Always create a child span (handled by _start_primary_span override)
|
||||
primary_span_parent = None
|
||||
|
||||
# 1. Primary span
|
||||
span = self._start_primary_span(kwargs, response_obj, start_time, end_time, ctx, primary_span_parent)
|
||||
|
||||
# 2. Raw-request sub-span (skipped for Weave via _maybe_log_raw_request override)
|
||||
self._maybe_log_raw_request(kwargs, response_obj, start_time, end_time, span)
|
||||
|
||||
# 3. Guardrail span
|
||||
self._create_guardrail_span(kwargs=kwargs, context=ctx)
|
||||
|
||||
# 4. Metrics & cost recording
|
||||
self._record_metrics(kwargs, response_obj, start_time, end_time)
|
||||
|
||||
# 5. Semantic logs.
|
||||
if self.config.enable_events:
|
||||
self._emit_semantic_logs(kwargs, response_obj, span)
|
||||
|
||||
# 6. Don't end parent span - it's managed by user code
|
||||
# Since we always create a child span (never reuse parent), the parent span
|
||||
# lifecycle is owned by the user. This prevents double-ending of wrapper spans
|
||||
# like "B", "C", "D", "E" that users create and manage themselves.
|
||||
|
||||
def construct_dynamic_otel_headers(
|
||||
self, standard_callback_dynamic_params: StandardCallbackDynamicParams
|
||||
) -> dict | None:
|
||||
"""
|
||||
Construct dynamic Weave headers from standard callback dynamic params.
|
||||
|
||||
This is used for team/key based logging.
|
||||
|
||||
Returns:
|
||||
dict: A dictionary of dynamic Weave headers
|
||||
"""
|
||||
dynamic_headers = {}
|
||||
|
||||
dynamic_wandb_api_key = standard_callback_dynamic_params.get("wandb_api_key")
|
||||
dynamic_weave_project_id = standard_callback_dynamic_params.get("weave_project_id")
|
||||
|
||||
if dynamic_wandb_api_key:
|
||||
auth_header = _get_weave_authorization_header(
|
||||
api_key=dynamic_wandb_api_key,
|
||||
)
|
||||
dynamic_headers["Authorization"] = auth_header
|
||||
|
||||
if dynamic_weave_project_id:
|
||||
dynamic_headers["project_id"] = dynamic_weave_project_id
|
||||
|
||||
return dynamic_headers if dynamic_headers else None
|
||||
|
|
@ -75,6 +75,7 @@ class CustomLoggerRegistry:
|
|||
"langfuse_otel": OpenTelemetry,
|
||||
"arize_phoenix": OpenTelemetry,
|
||||
"langtrace": OpenTelemetry,
|
||||
"weave_otel": OpenTelemetry,
|
||||
"mlflow": MlflowLogger,
|
||||
"langfuse": LangfusePromptManagement,
|
||||
"otel": OpenTelemetry,
|
||||
|
|
|
|||
|
|
@ -468,6 +468,18 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
|
|||
custom_llm_provider = model.split("/", 1)[0]
|
||||
model = model.split("/", 1)[1]
|
||||
|
||||
# Check JSON providers FIRST (before hardcoded ones)
|
||||
from litellm.llms.openai_like.json_loader import JSONProviderRegistry
|
||||
from litellm.llms.openai_like.dynamic_config import create_config_class
|
||||
|
||||
if JSONProviderRegistry.exists(custom_llm_provider):
|
||||
provider_config = JSONProviderRegistry.get(custom_llm_provider)
|
||||
config_class = create_config_class(provider_config)
|
||||
api_base, dynamic_api_key = config_class()._get_openai_compatible_provider_info(
|
||||
api_base, api_key
|
||||
)
|
||||
return model, custom_llm_provider, dynamic_api_key, api_base
|
||||
|
||||
if custom_llm_provider == "perplexity":
|
||||
# perplexity is openai compatible, we just need to set this to custom_openai and have the api_base be https://api.perplexity.ai
|
||||
(
|
||||
|
|
@ -763,13 +775,7 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
|
|||
) = litellm.MoonshotChatConfig()._get_openai_compatible_provider_info(
|
||||
api_base, api_key
|
||||
)
|
||||
elif custom_llm_provider == "publicai":
|
||||
(
|
||||
api_base,
|
||||
dynamic_api_key,
|
||||
) = litellm.PublicAIChatConfig()._get_openai_compatible_provider_info(
|
||||
api_base, api_key
|
||||
)
|
||||
# publicai is now handled by JSON config (see litellm/llms/openai_like/providers.json)
|
||||
elif custom_llm_provider == "docker_model_runner":
|
||||
(
|
||||
api_base,
|
||||
|
|
|
|||
|
|
@ -71,6 +71,7 @@ from litellm.litellm_core_utils.redact_messages import (
|
|||
from litellm.llms.base_llm.ocr.transformation import OCRResponse
|
||||
from litellm.llms.base_llm.search.transformation import SearchResponse
|
||||
from litellm.responses.utils import ResponseAPILoggingUtils
|
||||
from litellm.types.agents import LiteLLMSendMessageResponse
|
||||
from litellm.types.containers.main import ContainerObject
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
|
|
@ -1738,6 +1739,7 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
and logging_result.get("object") == "search" # Search API (dict format)
|
||||
or isinstance(logging_result, VideoObject)
|
||||
or isinstance(logging_result, ContainerObject)
|
||||
or isinstance(logging_result, LiteLLMSendMessageResponse) # A2A
|
||||
or (self.call_type == CallTypes.call_mcp_tool.value)
|
||||
):
|
||||
return True
|
||||
|
|
@ -3617,15 +3619,15 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
|
|||
|
||||
for callback in _in_memory_loggers:
|
||||
if (
|
||||
isinstance(callback, OpenTelemetry)
|
||||
isinstance(callback, ArizePhoenixLogger)
|
||||
and callback.callback_name == "arize_phoenix"
|
||||
):
|
||||
return callback # type: ignore
|
||||
_otel_logger = OpenTelemetry(
|
||||
_arize_phoenix_otel_logger = ArizePhoenixLogger(
|
||||
config=otel_config, callback_name="arize_phoenix"
|
||||
)
|
||||
_in_memory_loggers.append(_otel_logger)
|
||||
return _otel_logger # type: ignore
|
||||
_in_memory_loggers.append(_arize_phoenix_otel_logger)
|
||||
return _arize_phoenix_otel_logger # type: ignore
|
||||
elif logging_integration == "otel":
|
||||
from litellm.integrations.opentelemetry import OpenTelemetry
|
||||
|
||||
|
|
@ -3800,6 +3802,31 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
|
|||
)
|
||||
_in_memory_loggers.append(_otel_logger)
|
||||
return _otel_logger # type: ignore
|
||||
elif logging_integration == "weave_otel":
|
||||
from litellm.integrations.opentelemetry import (
|
||||
OpenTelemetryConfig,
|
||||
)
|
||||
from litellm.integrations.weave.weave_otel import WeaveOtelLogger, get_weave_otel_config
|
||||
|
||||
weave_otel_config = get_weave_otel_config()
|
||||
|
||||
otel_config = OpenTelemetryConfig(
|
||||
exporter=weave_otel_config.protocol,
|
||||
endpoint=weave_otel_config.endpoint,
|
||||
headers=weave_otel_config.otlp_auth_headers,
|
||||
)
|
||||
|
||||
for callback in _in_memory_loggers:
|
||||
if (
|
||||
isinstance(callback, WeaveOtelLogger)
|
||||
and callback.callback_name == "weave_otel"
|
||||
):
|
||||
return callback # type: ignore
|
||||
_otel_logger = WeaveOtelLogger(
|
||||
config=otel_config, callback_name="weave_otel"
|
||||
)
|
||||
_in_memory_loggers.append(_otel_logger)
|
||||
return _otel_logger # type: ignore
|
||||
elif logging_integration == "pagerduty":
|
||||
for callback in _in_memory_loggers:
|
||||
if isinstance(callback, PagerDutyAlerting):
|
||||
|
|
|
|||
|
|
@ -1071,7 +1071,7 @@ def _parse_content_for_reasoning(
|
|||
return None, message_text
|
||||
|
||||
reasoning_match = re.match(
|
||||
r"<(?:think|thinking)>(.*?)</(?:think|thinking)>(.*)", message_text, re.DOTALL
|
||||
r"<(?:think|thinking|budget:thinking)>(.*?)</(?:think|thinking|budget:thinking)>(.*)", message_text, re.DOTALL
|
||||
)
|
||||
|
||||
if reasoning_match:
|
||||
|
|
|
|||
|
|
@ -3446,8 +3446,25 @@ class BedrockConverseMessagesProcessor:
|
|||
@staticmethod
|
||||
def _initial_message_setup(
|
||||
messages: List,
|
||||
model: str,
|
||||
llm_provider: str,
|
||||
user_continue_message: Optional[ChatCompletionUserMessage] = None,
|
||||
) -> List:
|
||||
# gracefully handle base case of no messages at all
|
||||
if len(messages) == 0:
|
||||
if user_continue_message is not None:
|
||||
messages.append(user_continue_message)
|
||||
elif litellm.modify_params:
|
||||
messages.append(DEFAULT_USER_CONTINUE_MESSAGE)
|
||||
else:
|
||||
raise litellm.BadRequestError(
|
||||
message=BAD_MESSAGE_ERROR_STR
|
||||
+ "bedrock requires at least one non-system message",
|
||||
model=model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
# if initial message is assistant message
|
||||
if messages[0].get("role") is not None and messages[0]["role"] == "assistant":
|
||||
if user_continue_message is not None:
|
||||
messages.insert(0, user_continue_message)
|
||||
|
|
@ -3475,18 +3492,8 @@ class BedrockConverseMessagesProcessor:
|
|||
contents: List[BedrockMessageBlock] = []
|
||||
msg_i = 0
|
||||
|
||||
## BASE CASE ##
|
||||
if len(messages) == 0:
|
||||
raise litellm.BadRequestError(
|
||||
message=BAD_MESSAGE_ERROR_STR
|
||||
+ "bedrock requires at least one non-system message",
|
||||
model=model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
# if initial message is assistant message
|
||||
messages = BedrockConverseMessagesProcessor._initial_message_setup(
|
||||
messages, user_continue_message
|
||||
messages, model, llm_provider, user_continue_message
|
||||
)
|
||||
|
||||
while msg_i < len(messages):
|
||||
|
|
@ -3847,28 +3854,9 @@ def _bedrock_converse_messages_pt( # noqa: PLR0915
|
|||
contents: List[BedrockMessageBlock] = []
|
||||
msg_i = 0
|
||||
|
||||
## BASE CASE ##
|
||||
if len(messages) == 0:
|
||||
raise litellm.BadRequestError(
|
||||
message=BAD_MESSAGE_ERROR_STR
|
||||
+ "bedrock requires at least one non-system message",
|
||||
model=model,
|
||||
llm_provider=llm_provider,
|
||||
)
|
||||
|
||||
# if initial message is assistant message
|
||||
if messages[0].get("role") is not None and messages[0]["role"] == "assistant":
|
||||
if user_continue_message is not None:
|
||||
messages.insert(0, user_continue_message)
|
||||
elif litellm.modify_params:
|
||||
messages.insert(0, DEFAULT_USER_CONTINUE_MESSAGE)
|
||||
|
||||
# if final message is assistant message
|
||||
if messages[-1].get("role") is not None and messages[-1]["role"] == "assistant":
|
||||
if user_continue_message is not None:
|
||||
messages.append(user_continue_message)
|
||||
elif litellm.modify_params:
|
||||
messages.append(DEFAULT_USER_CONTINUE_MESSAGE)
|
||||
messages = BedrockConverseMessagesProcessor._initial_message_setup(
|
||||
messages, model, llm_provider, user_continue_message
|
||||
)
|
||||
|
||||
while msg_i < len(messages):
|
||||
user_content: List[BedrockContentBlock] = []
|
||||
|
|
|
|||
|
|
@ -12,10 +12,24 @@ Pattern Overview:
|
|||
4. Apply guardrail responses back to the original structure
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, cast
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.llms.anthropic.chat.transformation import AnthropicConfig
|
||||
from litellm.llms.anthropic.experimental_pass_through.adapters.transformation import (
|
||||
LiteLLMAnthropicMessagesAdapter,
|
||||
)
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.types.guardrails import GenericGuardrailAPIInputs
|
||||
from litellm.types.llms.anthropic import (
|
||||
AllAnthropicToolsValues,
|
||||
AnthropicMessagesRequest,
|
||||
)
|
||||
from litellm.types.llms.openai import (
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolParam,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
|
@ -36,6 +50,10 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
Methods can be overridden to customize behavior for different message formats.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.adapter = LiteLLMAnthropicMessagesAdapter()
|
||||
|
||||
async def process_input_messages(
|
||||
self,
|
||||
data: dict,
|
||||
|
|
@ -49,8 +67,19 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
if messages is None:
|
||||
return data
|
||||
|
||||
chat_completion_compatible_request = (
|
||||
LiteLLMAnthropicMessagesAdapter().translate_anthropic_to_openai(
|
||||
anthropic_message_request=cast(AnthropicMessagesRequest, data)
|
||||
)
|
||||
)
|
||||
|
||||
structured_messages = chat_completion_compatible_request.get("messages", [])
|
||||
|
||||
texts_to_check: List[str] = []
|
||||
images_to_check: List[str] = []
|
||||
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
|
||||
|
|
@ -67,16 +96,22 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
# Step 2: Apply guardrail to all texts in batch
|
||||
if texts_to_check:
|
||||
guardrailed_texts, guardrailed_images = (
|
||||
await guardrail_to_apply.apply_guardrail(
|
||||
texts=texts_to_check,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
images=images_to_check if images_to_check else None,
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if images_to_check:
|
||||
inputs["images"] = images_to_check
|
||||
if tools_to_check:
|
||||
inputs["tools"] = tools_to_check
|
||||
if structured_messages:
|
||||
inputs["structured_messages"] = structured_messages
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
|
||||
# Step 3: Map guardrail responses back to original message structure
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=messages,
|
||||
|
|
@ -104,15 +139,17 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
Override this method to customize text/image extraction logic.
|
||||
"""
|
||||
content = message.get("content", None)
|
||||
if content is None:
|
||||
tools = message.get("tools", None)
|
||||
if content is None and tools is None:
|
||||
return
|
||||
|
||||
if isinstance(content, str):
|
||||
## CHECK FOR TEXT + IMAGES
|
||||
if content is not None and isinstance(content, str):
|
||||
# Simple string content
|
||||
texts_to_check.append(content)
|
||||
task_mappings.append((msg_idx, None))
|
||||
|
||||
elif isinstance(content, list):
|
||||
elif content is not None and isinstance(content, list):
|
||||
# List content (e.g., multimodal with text and images)
|
||||
for content_idx, content_item in enumerate(content):
|
||||
# Extract text
|
||||
|
|
@ -130,6 +167,22 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
if data:
|
||||
images_to_check.append(data)
|
||||
|
||||
def _extract_input_tools(
|
||||
self,
|
||||
tools: List[Dict[str, Any]],
|
||||
tools_to_check: List[ChatCompletionToolParam],
|
||||
) -> None:
|
||||
"""
|
||||
Extract tools from a message.
|
||||
"""
|
||||
## CHECK FOR TOOLS
|
||||
if tools is not None and isinstance(tools, list):
|
||||
# TRANSFORM ANTHROPIC TOOLS TO OPENAI TOOLS
|
||||
openai_tools = self.adapter.translate_anthropic_tools_to_openai(
|
||||
tools=cast(List[AllAnthropicToolsValues], tools)
|
||||
)
|
||||
tools_to_check.extend(openai_tools)
|
||||
|
||||
async def _apply_guardrail_responses_to_input(
|
||||
self,
|
||||
messages: List[Dict[str, Any]],
|
||||
|
|
@ -168,7 +221,7 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
user_api_key_dict: Optional[Any] = None,
|
||||
) -> Any:
|
||||
"""
|
||||
Process output response by applying guardrails to text content.
|
||||
Process output response by applying guardrails to text content and tool calls.
|
||||
|
||||
Args:
|
||||
response: Anthropic MessagesResponse object
|
||||
|
|
@ -180,17 +233,15 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
Modified response with guardrail applied to content
|
||||
|
||||
Response Format Support:
|
||||
- List content: response.content = [{"type": "text", "text": "text here"}, ...]
|
||||
- List content: response.content = [
|
||||
{"type": "text", "text": "text here"},
|
||||
{"type": "tool_use", "id": "...", "name": "...", "input": {...}},
|
||||
...
|
||||
]
|
||||
"""
|
||||
# Step 0: Check if response has any text content to process
|
||||
if not self._has_text_content(response):
|
||||
verbose_proxy_logger.warning(
|
||||
"Anthropic Messages: No text content in response, skipping guardrail"
|
||||
)
|
||||
return response
|
||||
|
||||
texts_to_check: List[str] = []
|
||||
images_to_check: List[str] = []
|
||||
tool_calls_to_check: List[ChatCompletionToolCallChunk] = []
|
||||
task_mappings: List[Tuple[int, Optional[int]]] = []
|
||||
# Track (content_index, None) for each text
|
||||
|
||||
|
|
@ -198,10 +249,13 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
if not response_content:
|
||||
return response
|
||||
|
||||
# Step 1: Extract all text content from response
|
||||
# Step 1: Extract all text content and tool calls from response
|
||||
for content_idx, content_block in enumerate(response_content):
|
||||
# Check if this is a text block by checking the 'type' field
|
||||
if isinstance(content_block, dict) and content_block.get("type") == "text":
|
||||
# Check if this is a text or tool_use block by checking the 'type' field
|
||||
if isinstance(content_block, dict) and content_block.get("type") in [
|
||||
"text",
|
||||
"tool_use",
|
||||
]:
|
||||
# Cast to dict to handle the union type properly
|
||||
self._extract_output_text_and_images(
|
||||
content_block=cast(Dict[str, Any], content_block),
|
||||
|
|
@ -209,10 +263,11 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
texts_to_check=texts_to_check,
|
||||
images_to_check=images_to_check,
|
||||
task_mappings=task_mappings,
|
||||
tool_calls_to_check=tool_calls_to_check,
|
||||
)
|
||||
|
||||
# Step 2: Apply guardrail to all texts in batch
|
||||
if texts_to_check:
|
||||
if texts_to_check or tool_calls_to_check:
|
||||
# Create a request_data dict with response info and user API key metadata
|
||||
request_data: dict = {"response": response}
|
||||
|
||||
|
|
@ -223,16 +278,21 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
if user_metadata:
|
||||
request_data["litellm_metadata"] = user_metadata
|
||||
|
||||
guardrailed_texts, guardrailed_images = (
|
||||
await guardrail_to_apply.apply_guardrail(
|
||||
texts=texts_to_check,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
images=images_to_check if images_to_check else None,
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if images_to_check:
|
||||
inputs["images"] = images_to_check
|
||||
if tool_calls_to_check:
|
||||
inputs["tool_calls"] = tool_calls_to_check
|
||||
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
|
||||
# Step 3: Map guardrail responses back to original response structure
|
||||
await self._apply_guardrail_responses_to_output(
|
||||
response=response,
|
||||
|
|
@ -246,6 +306,112 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
|
||||
return response
|
||||
|
||||
async def process_output_streaming_response(
|
||||
self,
|
||||
responses_so_far: List[Any],
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
litellm_logging_obj: Optional[Any] = None,
|
||||
user_api_key_dict: Optional[Any] = None,
|
||||
) -> List[Any]:
|
||||
"""
|
||||
Process output streaming response by applying guardrails to text content.
|
||||
|
||||
Get the string so far, check the apply guardrail to the string so far, and return the list of responses so far.
|
||||
"""
|
||||
string_so_far = self.get_streaming_string_so_far(responses_so_far)
|
||||
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail( # allow rejecting the response, if invalid
|
||||
inputs={"texts": [string_so_far]},
|
||||
request_data={},
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
return responses_so_far
|
||||
|
||||
def get_streaming_string_so_far(self, responses_so_far: List[Any]) -> str:
|
||||
"""
|
||||
Parse streaming responses and extract accumulated text content.
|
||||
|
||||
Handles two formats:
|
||||
1. Raw bytes in SSE (Server-Sent Events) format from Anthropic API
|
||||
2. Parsed dict objects (for backwards compatibility)
|
||||
|
||||
SSE format example:
|
||||
b'event: content_block_delta\\ndata: {"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":" curious"}}\\n\\n'
|
||||
|
||||
Dict format example:
|
||||
{
|
||||
"type": "content_block_delta",
|
||||
"index": 0,
|
||||
"delta": {
|
||||
"type": "text_delta",
|
||||
"text": " curious"
|
||||
}
|
||||
}
|
||||
"""
|
||||
text_so_far = ""
|
||||
for response in responses_so_far:
|
||||
# Handle raw bytes in SSE format
|
||||
if isinstance(response, bytes):
|
||||
text_so_far += self._extract_text_from_sse(response)
|
||||
# Handle already-parsed dict format
|
||||
elif isinstance(response, dict):
|
||||
delta = response.get("delta") if response.get("delta") else None
|
||||
if delta and delta.get("type") == "text_delta":
|
||||
text = delta.get("text", "")
|
||||
if text:
|
||||
text_so_far += text
|
||||
return text_so_far
|
||||
|
||||
def _extract_text_from_sse(self, sse_bytes: bytes) -> str:
|
||||
"""
|
||||
Extract text content from Server-Sent Events (SSE) format.
|
||||
|
||||
Args:
|
||||
sse_bytes: Raw bytes in SSE format
|
||||
|
||||
Returns:
|
||||
Accumulated text from all content_block_delta events
|
||||
"""
|
||||
text = ""
|
||||
try:
|
||||
# Decode bytes to string
|
||||
sse_string = sse_bytes.decode("utf-8")
|
||||
|
||||
# Split by double newline to get individual events
|
||||
events = sse_string.split("\n\n")
|
||||
|
||||
for event in events:
|
||||
if not event.strip():
|
||||
continue
|
||||
|
||||
# Parse event lines
|
||||
lines = event.strip().split("\n")
|
||||
event_type = None
|
||||
data_line = None
|
||||
|
||||
for line in lines:
|
||||
if line.startswith("event:"):
|
||||
event_type = line[6:].strip()
|
||||
elif line.startswith("data:"):
|
||||
data_line = line[5:].strip()
|
||||
|
||||
# Only process content_block_delta events
|
||||
if event_type == "content_block_delta" and data_line:
|
||||
try:
|
||||
data = json.loads(data_line)
|
||||
delta = data.get("delta", {})
|
||||
if delta.get("type") == "text_delta":
|
||||
text += delta.get("text", "")
|
||||
except json.JSONDecodeError:
|
||||
verbose_proxy_logger.warning(
|
||||
f"Failed to parse JSON from SSE data: {data_line}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error extracting text from SSE: {e}")
|
||||
|
||||
return text
|
||||
|
||||
def _has_text_content(self, response: "AnthropicMessagesResponse") -> bool:
|
||||
"""
|
||||
Check if response has any text content to process.
|
||||
|
|
@ -270,17 +436,32 @@ class AnthropicMessagesHandler(BaseTranslation):
|
|||
texts_to_check: List[str],
|
||||
images_to_check: List[str],
|
||||
task_mappings: List[Tuple[int, Optional[int]]],
|
||||
tool_calls_to_check: Optional[List[ChatCompletionToolCallChunk]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Extract text content and images from a response content block.
|
||||
Extract text content, images, and tool calls from a response content block.
|
||||
|
||||
Override this method to customize text/image extraction logic.
|
||||
Override this method to customize text/image/tool extraction logic.
|
||||
"""
|
||||
content_text = content_block.get("text")
|
||||
if content_text and isinstance(content_text, str):
|
||||
# Simple string content
|
||||
texts_to_check.append(content_text)
|
||||
task_mappings.append((content_idx, None))
|
||||
content_type = content_block.get("type")
|
||||
|
||||
# Extract text content
|
||||
if content_type == "text":
|
||||
content_text = content_block.get("text")
|
||||
if content_text and isinstance(content_text, str):
|
||||
# Simple string content
|
||||
texts_to_check.append(content_text)
|
||||
task_mappings.append((content_idx, None))
|
||||
|
||||
# Extract tool calls
|
||||
elif content_type == "tool_use":
|
||||
tool_call = AnthropicConfig.convert_tool_use_to_openai_format(
|
||||
anthropic_tool_content=content_block,
|
||||
index=content_idx,
|
||||
)
|
||||
if tool_calls_to_check is None:
|
||||
tool_calls_to_check = []
|
||||
tool_calls_to_check.append(tool_call)
|
||||
|
||||
async def _apply_guardrail_responses_to_output(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -436,9 +436,7 @@ class AnthropicChatCompletion(BaseLLM):
|
|||
|
||||
else:
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
client = _get_httpx_client(
|
||||
params={"timeout": timeout}
|
||||
)
|
||||
client = _get_httpx_client(params={"timeout": timeout})
|
||||
else:
|
||||
client = client
|
||||
|
||||
|
|
@ -528,9 +526,7 @@ class ModelResponseIterator:
|
|||
usage_object=cast(dict, anthropic_usage_chunk), reasoning_content=None
|
||||
)
|
||||
|
||||
def _content_block_delta_helper(
|
||||
self, chunk: dict
|
||||
) -> Tuple[
|
||||
def _content_block_delta_helper(self, chunk: dict) -> Tuple[
|
||||
str,
|
||||
Optional[ChatCompletionToolCallChunk],
|
||||
List[Union[ChatCompletionThinkingBlock, ChatCompletionRedactedThinkingBlock]],
|
||||
|
|
|
|||
|
|
@ -54,10 +54,7 @@ from litellm.types.utils import (
|
|||
CompletionTokensDetailsWrapper,
|
||||
)
|
||||
from litellm.types.utils import Message as LitellmMessage
|
||||
from litellm.types.utils import (
|
||||
PromptTokensDetailsWrapper,
|
||||
ServerToolUse,
|
||||
)
|
||||
from litellm.types.utils import PromptTokensDetailsWrapper, ServerToolUse
|
||||
from litellm.utils import (
|
||||
ModelResponse,
|
||||
Usage,
|
||||
|
|
@ -119,6 +116,36 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
def get_config(cls):
|
||||
return super().get_config()
|
||||
|
||||
@staticmethod
|
||||
def convert_tool_use_to_openai_format(
|
||||
anthropic_tool_content: Dict[str, Any],
|
||||
index: int,
|
||||
) -> ChatCompletionToolCallChunk:
|
||||
"""
|
||||
Convert Anthropic tool_use format to OpenAI ChatCompletionToolCallChunk format.
|
||||
|
||||
Args:
|
||||
anthropic_tool_content: Anthropic tool_use content block with format:
|
||||
{"type": "tool_use", "id": "...", "name": "...", "input": {...}}
|
||||
index: The index of this tool call
|
||||
|
||||
Returns:
|
||||
ChatCompletionToolCallChunk in OpenAI format
|
||||
"""
|
||||
tool_call = ChatCompletionToolCallChunk(
|
||||
id=anthropic_tool_content["id"],
|
||||
type="function",
|
||||
function=ChatCompletionToolCallFunctionChunk(
|
||||
name=anthropic_tool_content["name"],
|
||||
arguments=json.dumps(anthropic_tool_content["input"]),
|
||||
),
|
||||
index=index,
|
||||
)
|
||||
# Include caller information if present (for programmatic tool calling)
|
||||
if "caller" in anthropic_tool_content:
|
||||
tool_call["caller"] = cast(Dict[str, Any], anthropic_tool_content["caller"]) # type: ignore[typeddict-item]
|
||||
return tool_call
|
||||
|
||||
def _is_claude_opus_4_5(self, model: str) -> bool:
|
||||
"""Check if the model is Claude Opus 4.5."""
|
||||
return "opus-4-5" in model.lower() or "opus_4_5" in model.lower()
|
||||
|
|
@ -279,7 +306,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
elif tool["type"] == "tool_search_tool_regex_20251119":
|
||||
# Tool search tool using regex
|
||||
from litellm.types.llms.anthropic import AnthropicToolSearchToolRegex
|
||||
|
||||
|
||||
tool_name_obj = tool.get("name", "tool_search_tool_regex")
|
||||
if not isinstance(tool_name_obj, str):
|
||||
raise ValueError("Tool search tool must have a valid name")
|
||||
|
|
@ -291,7 +318,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
elif tool["type"] == "tool_search_tool_bm25_20251119":
|
||||
# Tool search tool using BM25
|
||||
from litellm.types.llms.anthropic import AnthropicToolSearchToolBM25
|
||||
|
||||
|
||||
tool_name_obj = tool.get("name", "tool_search_tool_bm25")
|
||||
if not isinstance(tool_name_obj, str):
|
||||
raise ValueError("Tool search tool must have a valid name")
|
||||
|
|
@ -309,7 +336,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
if returned_tool is not None:
|
||||
# Only set cache_control on tools that support it (not tool search tools)
|
||||
tool_type = returned_tool.get("type", "")
|
||||
if tool_type not in ("tool_search_tool_regex_20251119", "tool_search_tool_bm25_20251119"):
|
||||
if tool_type not in (
|
||||
"tool_search_tool_regex_20251119",
|
||||
"tool_search_tool_bm25_20251119",
|
||||
):
|
||||
if _cache_control is not None:
|
||||
returned_tool["cache_control"] = _cache_control # type: ignore[typeddict-item]
|
||||
elif _cache_control_function is not None and isinstance(
|
||||
|
|
@ -318,14 +348,19 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
returned_tool["cache_control"] = ChatCompletionCachedContent( # type: ignore[typeddict-item]
|
||||
**_cache_control_function # type: ignore
|
||||
)
|
||||
|
||||
|
||||
## check if defer_loading is set in the tool
|
||||
_defer_loading = tool.get("defer_loading", None)
|
||||
_defer_loading_function = tool.get("function", {}).get("defer_loading", None)
|
||||
if returned_tool is not None:
|
||||
# Only set defer_loading on tools that support it (not tool search tools or computer tools)
|
||||
tool_type = returned_tool.get("type", "")
|
||||
if tool_type not in ("tool_search_tool_regex_20251119", "tool_search_tool_bm25_20251119", "computer_20241022", "computer_20250124"):
|
||||
if tool_type not in (
|
||||
"tool_search_tool_regex_20251119",
|
||||
"tool_search_tool_bm25_20251119",
|
||||
"computer_20241022",
|
||||
"computer_20250124",
|
||||
):
|
||||
if _defer_loading is not None:
|
||||
if not isinstance(_defer_loading, bool):
|
||||
raise ValueError("defer_loading must be a boolean")
|
||||
|
|
@ -334,14 +369,21 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
if not isinstance(_defer_loading_function, bool):
|
||||
raise ValueError("defer_loading must be a boolean")
|
||||
returned_tool["defer_loading"] = _defer_loading_function # type: ignore[typeddict-item]
|
||||
|
||||
|
||||
## check if allowed_callers is set in the tool
|
||||
_allowed_callers = tool.get("allowed_callers", None)
|
||||
_allowed_callers_function = tool.get("function", {}).get("allowed_callers", None)
|
||||
_allowed_callers_function = tool.get("function", {}).get(
|
||||
"allowed_callers", None
|
||||
)
|
||||
if returned_tool is not None:
|
||||
# Only set allowed_callers on tools that support it (not tool search tools or computer tools)
|
||||
tool_type = returned_tool.get("type", "")
|
||||
if tool_type not in ("tool_search_tool_regex_20251119", "tool_search_tool_bm25_20251119", "computer_20241022", "computer_20250124"):
|
||||
if tool_type not in (
|
||||
"tool_search_tool_regex_20251119",
|
||||
"tool_search_tool_bm25_20251119",
|
||||
"computer_20241022",
|
||||
"computer_20250124",
|
||||
):
|
||||
if _allowed_callers is not None:
|
||||
if not isinstance(_allowed_callers, list) or not all(
|
||||
isinstance(item, str) for item in _allowed_callers
|
||||
|
|
@ -354,7 +396,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
):
|
||||
raise ValueError("allowed_callers must be a list of strings")
|
||||
returned_tool["allowed_callers"] = _allowed_callers_function # type: ignore[typeddict-item]
|
||||
|
||||
|
||||
## check if input_examples is set in the tool
|
||||
_input_examples = tool.get("input_examples", None)
|
||||
_input_examples_function = tool.get("function", {}).get("input_examples", None)
|
||||
|
|
@ -423,31 +465,32 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
"""Check if tool search tools are present in the tools list."""
|
||||
if not tools:
|
||||
return False
|
||||
|
||||
|
||||
for tool in tools:
|
||||
tool_type = tool.get("type", "")
|
||||
if tool_type in ["tool_search_tool_regex_20251119", "tool_search_tool_bm25_20251119"]:
|
||||
if tool_type in [
|
||||
"tool_search_tool_regex_20251119",
|
||||
"tool_search_tool_bm25_20251119",
|
||||
]:
|
||||
return True
|
||||
return False
|
||||
|
||||
def _separate_deferred_tools(
|
||||
self, tools: List
|
||||
) -> Tuple[List, List]:
|
||||
def _separate_deferred_tools(self, tools: List) -> Tuple[List, List]:
|
||||
"""
|
||||
Separate tools into deferred and non-deferred lists.
|
||||
|
||||
|
||||
Returns:
|
||||
Tuple of (non_deferred_tools, deferred_tools)
|
||||
"""
|
||||
non_deferred = []
|
||||
deferred = []
|
||||
|
||||
|
||||
for tool in tools:
|
||||
if tool.get("defer_loading", False):
|
||||
deferred.append(tool)
|
||||
else:
|
||||
non_deferred.append(tool)
|
||||
|
||||
|
||||
return non_deferred, deferred
|
||||
|
||||
def _expand_tool_references(
|
||||
|
|
@ -457,28 +500,28 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
) -> List:
|
||||
"""
|
||||
Expand tool_reference blocks to full tool definitions.
|
||||
|
||||
|
||||
When Anthropic's tool search returns results, it includes tool_reference blocks
|
||||
that reference tools by name. This method expands those references to full
|
||||
tool definitions from the deferred_tools catalog.
|
||||
|
||||
|
||||
Args:
|
||||
content: Response content that may contain tool_reference blocks
|
||||
deferred_tools: List of deferred tools that can be referenced
|
||||
|
||||
|
||||
Returns:
|
||||
Content with tool_reference blocks expanded to full tool definitions
|
||||
"""
|
||||
if not deferred_tools:
|
||||
return content
|
||||
|
||||
|
||||
# Create a mapping of tool names to tool definitions
|
||||
tool_map = {}
|
||||
for tool in deferred_tools:
|
||||
tool_name = tool.get("name") or tool.get("function", {}).get("name")
|
||||
if tool_name:
|
||||
tool_map[tool_name] = tool
|
||||
|
||||
|
||||
# Expand tool references in content
|
||||
expanded_content = []
|
||||
for item in content:
|
||||
|
|
@ -492,7 +535,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
expanded_content.append(item)
|
||||
else:
|
||||
expanded_content.append(item)
|
||||
|
||||
|
||||
return expanded_content
|
||||
|
||||
def _map_stop_sequences(
|
||||
|
|
@ -786,6 +829,9 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
valid_content: bool = False
|
||||
system_message_block = ChatCompletionSystemMessage(**message)
|
||||
if isinstance(system_message_block["content"], str):
|
||||
# Skip empty text blocks - Anthropic API raises errors for empty text
|
||||
if not system_message_block["content"]:
|
||||
continue
|
||||
anthropic_system_message_content = AnthropicSystemMessageContent(
|
||||
type="text",
|
||||
text=system_message_block["content"],
|
||||
|
|
@ -800,10 +846,14 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
valid_content = True
|
||||
elif isinstance(message["content"], list):
|
||||
for _content in message["content"]:
|
||||
# Skip empty text blocks - Anthropic API raises errors for empty text
|
||||
text_value = _content.get("text")
|
||||
if _content.get("type") == "text" and not text_value:
|
||||
continue
|
||||
anthropic_system_message_content = (
|
||||
AnthropicSystemMessageContent(
|
||||
type=_content.get("type"),
|
||||
text=_content.get("text"),
|
||||
text=text_value,
|
||||
)
|
||||
)
|
||||
if "cache_control" in _content:
|
||||
|
|
@ -988,7 +1038,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
"messages": anthropic_messages,
|
||||
**optional_params,
|
||||
}
|
||||
|
||||
|
||||
## Handle output_config (Anthropic-specific parameter)
|
||||
if "output_config" in optional_params:
|
||||
output_config = optional_params.get("output_config")
|
||||
|
|
@ -1047,34 +1097,20 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
text_content += content["text"]
|
||||
## TOOL CALLING
|
||||
elif content["type"] == "tool_use":
|
||||
tool_call = ChatCompletionToolCallChunk(
|
||||
id=content["id"],
|
||||
type="function",
|
||||
function=ChatCompletionToolCallFunctionChunk(
|
||||
name=content["name"],
|
||||
arguments=json.dumps(content["input"]),
|
||||
),
|
||||
tool_call = AnthropicConfig.convert_tool_use_to_openai_format(
|
||||
anthropic_tool_content=content,
|
||||
index=idx,
|
||||
)
|
||||
# Include caller information if present (for programmatic tool calling)
|
||||
if "caller" in content:
|
||||
tool_call["caller"] = cast(Dict[str, Any], content["caller"]) # type: ignore[typeddict-item]
|
||||
tool_calls.append(tool_call)
|
||||
## SERVER TOOL USE (for tool search)
|
||||
elif content["type"] == "server_tool_use":
|
||||
# Server tool use blocks are for tool search - treat as tool calls
|
||||
tool_call = ChatCompletionToolCallChunk(
|
||||
id=content["id"],
|
||||
type="function",
|
||||
function=ChatCompletionToolCallFunctionChunk(
|
||||
name=content["name"],
|
||||
arguments=json.dumps(content.get("input", {})),
|
||||
),
|
||||
# Note: using .get("input", {}) for server_tool_use as input may not be present
|
||||
content_with_input = {**content, "input": content.get("input", {})}
|
||||
tool_call = AnthropicConfig.convert_tool_use_to_openai_format(
|
||||
anthropic_tool_content=content_with_input,
|
||||
index=idx,
|
||||
)
|
||||
# Include caller information if present (for programmatic tool calling)
|
||||
if "caller" in content:
|
||||
tool_call["caller"] = cast(Dict[str, Any], content["caller"]) # type: ignore[typeddict-item]
|
||||
tool_calls.append(tool_call)
|
||||
## TOOL SEARCH TOOL RESULT (skip - this is metadata about tool discovery)
|
||||
elif content["type"] == "tool_search_tool_result":
|
||||
|
|
@ -1115,7 +1151,10 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
return text_content, citations, thinking_blocks, reasoning_content, tool_calls
|
||||
|
||||
def calculate_usage(
|
||||
self, usage_object: dict, reasoning_content: Optional[str], completion_response: Optional[dict] = None
|
||||
self,
|
||||
usage_object: dict,
|
||||
reasoning_content: Optional[str],
|
||||
completion_response: Optional[dict] = None,
|
||||
) -> Usage:
|
||||
# NOTE: Sometimes the usage object has None set explicitly for token counts, meaning .get() & key access returns None, and we need to account for this
|
||||
prompt_tokens = usage_object.get("input_tokens", 0) or 0
|
||||
|
|
@ -1153,7 +1192,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
|
|||
tool_search_requests = cast(
|
||||
int, _usage["server_tool_use"]["tool_search_requests"]
|
||||
)
|
||||
|
||||
|
||||
# Count tool_search_requests from content blocks if not in usage
|
||||
# Anthropic doesn't always include tool_search_requests in the usage object
|
||||
if tool_search_requests is None and completion_response is not None:
|
||||
|
|
|
|||
|
|
@ -1020,7 +1020,7 @@ class AzureChatCompletion(BaseAzureLLM, BaseLLM):
|
|||
headers: dict,
|
||||
client=None,
|
||||
timeout=None,
|
||||
) -> litellm.ImageResponse:
|
||||
) -> ImageResponse:
|
||||
|
||||
response: Optional[dict] = None
|
||||
try:
|
||||
|
|
|
|||
|
|
@ -40,7 +40,11 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config):
|
|||
or optional_params.get("reasoning_effort")
|
||||
)
|
||||
|
||||
if reasoning_effort_value == "none":
|
||||
# gpt-5.1 supports reasoning_effort='none', but other gpt-5 models don't
|
||||
# See: https://learn.microsoft.com/en-us/azure/ai-foundry/openai/how-to/reasoning
|
||||
is_gpt_5_1 = self.is_model_gpt_5_1_model(model)
|
||||
|
||||
if reasoning_effort_value == "none" and not is_gpt_5_1:
|
||||
if litellm.drop_params is True or (
|
||||
drop_params is not None and drop_params is True
|
||||
):
|
||||
|
|
@ -54,7 +58,7 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config):
|
|||
raise UnsupportedParamsError(
|
||||
status_code=400,
|
||||
message=(
|
||||
"Azure OpenAI does not support reasoning_effort='none'. "
|
||||
"Azure OpenAI does not support reasoning_effort='none' for this model. "
|
||||
"Supported values are: 'low', 'medium', and 'high'. "
|
||||
"To drop this parameter, set `litellm.drop_params=True` or for proxy:\n\n"
|
||||
"`litellm_settings:\n drop_params: true`\n"
|
||||
|
|
@ -70,7 +74,8 @@ class AzureOpenAIGPT5Config(AzureOpenAIConfig, OpenAIGPT5Config):
|
|||
drop_params=drop_params,
|
||||
)
|
||||
|
||||
if result.get("reasoning_effort") == "none":
|
||||
# Only drop reasoning_effort='none' for non-gpt-5.1 models
|
||||
if result.get("reasoning_effort") == "none" and not is_gpt_5_1:
|
||||
result.pop("reasoning_effort")
|
||||
|
||||
return result
|
||||
|
|
|
|||
|
|
@ -58,7 +58,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
|
|||
data: ImageEmbeddingRequest,
|
||||
timeout: float,
|
||||
logging_obj,
|
||||
model_response: litellm.EmbeddingResponse,
|
||||
model_response: EmbeddingResponse,
|
||||
optional_params: dict,
|
||||
api_key: Optional[str],
|
||||
api_base: Optional[str],
|
||||
|
|
@ -138,7 +138,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
|
|||
input: List,
|
||||
timeout: float,
|
||||
logging_obj,
|
||||
model_response: litellm.EmbeddingResponse,
|
||||
model_response: EmbeddingResponse,
|
||||
optional_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
from abc import ABC, abstractmethod
|
||||
from typing import TYPE_CHECKING, Any, Dict, Optional
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
|
@ -87,7 +87,7 @@ class BaseTranslation(ABC):
|
|||
|
||||
async def process_output_streaming_response(
|
||||
self,
|
||||
response: Any,
|
||||
responses_so_far: List[Any],
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
litellm_logging_obj: Optional["LiteLLMLoggingObj"] = None,
|
||||
user_api_key_dict: Optional["UserAPIKeyAuth"] = None,
|
||||
|
|
@ -97,4 +97,4 @@ class BaseTranslation(ABC):
|
|||
|
||||
Optional to override in subclasses.
|
||||
"""
|
||||
return response
|
||||
return responses_so_far
|
||||
|
|
|
|||
|
|
@ -10,7 +10,6 @@ from typing import Any, List, Optional
|
|||
|
||||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.types.llms.bedrock import BedrockInvokeNovaRequest
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
|
@ -80,7 +79,7 @@ class AmazonInvokeNovaConfig(AmazonInvokeConfig, AmazonConverseConfig):
|
|||
encoding: Any,
|
||||
api_key: Optional[str] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
) -> litellm.ModelResponse:
|
||||
) -> ModelResponse:
|
||||
return AmazonConverseConfig.transform_response(
|
||||
self,
|
||||
model,
|
||||
|
|
|
|||
|
|
@ -0,0 +1,98 @@
|
|||
"""
|
||||
Handles transforming requests for `bedrock/invoke/{qwen2} models`
|
||||
|
||||
Inherits from `AmazonQwen3Config` since Qwen2 and Qwen3 architectures are mostly similar.
|
||||
The main difference is in the response format: Qwen2 uses "text" field while Qwen3 uses "generation" field.
|
||||
|
||||
Qwen2 + Invoke API Tutorial: https://docs.aws.amazon.com/bedrock/latest/userguide/invoke-imported-model.html
|
||||
"""
|
||||
|
||||
from typing import Any, List, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm.llms.bedrock.chat.invoke_transformations.amazon_qwen3_transformation import (
|
||||
AmazonQwen3Config,
|
||||
)
|
||||
from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import (
|
||||
LiteLLMLoggingObj,
|
||||
)
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
class AmazonQwen2Config(AmazonQwen3Config):
|
||||
"""
|
||||
Config for sending `qwen2` requests to `/bedrock/invoke/`
|
||||
|
||||
Inherits from AmazonQwen3Config since Qwen2 and Qwen3 architectures are mostly similar.
|
||||
The main difference is in the response format: Qwen2 uses "text" field while Qwen3 uses "generation" field.
|
||||
|
||||
Reference: https://docs.aws.amazon.com/bedrock/latest/userguide/invoke-imported-model.html
|
||||
"""
|
||||
|
||||
def transform_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
model_response: ModelResponse,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
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 Qwen2 Bedrock response to OpenAI format
|
||||
|
||||
Qwen2 uses "text" field, but we also support "generation" field for compatibility.
|
||||
"""
|
||||
try:
|
||||
if hasattr(raw_response, 'json'):
|
||||
response_data = raw_response.json()
|
||||
else:
|
||||
response_data = raw_response
|
||||
|
||||
# Extract the generated text - Qwen2 uses "text" field, but also support "generation" for compatibility
|
||||
generated_text = response_data.get("generation", "") or response_data.get("text", "")
|
||||
|
||||
# Clean up the response (remove assistant start token if present)
|
||||
if generated_text.startswith("<|im_start|>assistant\n"):
|
||||
generated_text = generated_text[len("<|im_start|>assistant\n"):]
|
||||
if generated_text.endswith("<|im_end|>"):
|
||||
generated_text = generated_text[:-len("<|im_end|>")]
|
||||
|
||||
# Set the content in the existing model_response structure
|
||||
if hasattr(model_response, 'choices') and len(model_response.choices) > 0:
|
||||
choice = model_response.choices[0]
|
||||
if hasattr(choice, 'message'):
|
||||
choice.message.content = generated_text
|
||||
choice.finish_reason = "stop"
|
||||
else:
|
||||
# Handle streaming choices
|
||||
choice.delta.content = generated_text
|
||||
choice.finish_reason = "stop"
|
||||
|
||||
# Set usage information if available in response
|
||||
if "usage" in response_data:
|
||||
usage_data = response_data["usage"]
|
||||
if hasattr(model_response, 'usage'):
|
||||
model_response.usage.prompt_tokens = usage_data.get("prompt_tokens", 0)
|
||||
model_response.usage.completion_tokens = usage_data.get("completion_tokens", 0)
|
||||
model_response.usage.total_tokens = usage_data.get("total_tokens", 0)
|
||||
|
||||
return model_response
|
||||
|
||||
except Exception as e:
|
||||
if logging_obj:
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key=api_key,
|
||||
original_response=raw_response,
|
||||
additional_args={"error": str(e)},
|
||||
)
|
||||
raise e
|
||||
|
||||
|
|
@ -134,6 +134,12 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
fake_stream=fake_stream,
|
||||
)
|
||||
|
||||
def _apply_config_to_params(self, config: dict, inference_params: dict) -> None:
|
||||
"""Apply config values to inference_params if not already set."""
|
||||
for k, v in config.items():
|
||||
if k not in inference_params:
|
||||
inference_params[k] = v
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
|
|
@ -166,11 +172,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
if model.startswith("cohere.command-r"):
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonCohereChatConfig().get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
self._apply_config_to_params(config, inference_params)
|
||||
_data = {"message": prompt, **inference_params}
|
||||
if chat_history is not None:
|
||||
_data["chat_history"] = chat_history
|
||||
|
|
@ -178,11 +180,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
else:
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonCohereConfig.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
self._apply_config_to_params(config, inference_params)
|
||||
if stream is True:
|
||||
inference_params[
|
||||
"stream"
|
||||
|
|
@ -211,32 +209,17 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
elif provider == "ai21":
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonAI21Config.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
|
||||
self._apply_config_to_params(config, inference_params)
|
||||
request_data = {"prompt": prompt, **inference_params}
|
||||
elif provider == "mistral":
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonMistralConfig.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
|
||||
self._apply_config_to_params(config, inference_params)
|
||||
request_data = {"prompt": prompt, **inference_params}
|
||||
elif provider == "amazon": # amazon titan
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonTitanConfig.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > amazon_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
|
||||
self._apply_config_to_params(config, inference_params)
|
||||
request_data = {
|
||||
"inputText": prompt,
|
||||
"textGenerationConfig": inference_params,
|
||||
|
|
@ -244,11 +227,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM):
|
|||
elif provider == "meta" or provider == "llama" or provider == "deepseek_r1":
|
||||
## LOAD CONFIG
|
||||
config = litellm.AmazonLlamaConfig.get_config()
|
||||
for k, v in config.items():
|
||||
if (
|
||||
k not in inference_params
|
||||
): # completion(top_k=3) > anthropic_config(top_k=3) <- allows for dynamic variables to be passed in
|
||||
inference_params[k] = v
|
||||
self._apply_config_to_params(config, inference_params)
|
||||
request_data = {"prompt": prompt, **inference_params}
|
||||
elif provider == "twelvelabs":
|
||||
return litellm.AmazonTwelveLabsPegasusConfig().transform_request(
|
||||
|
|
|
|||
|
|
@ -27,6 +27,25 @@ class BedrockError(BaseLLMException):
|
|||
pass
|
||||
|
||||
|
||||
# Lazy import cache to avoid circular imports and performance impact
|
||||
_get_model_info = None
|
||||
|
||||
|
||||
def get_cached_model_info():
|
||||
"""
|
||||
Lazy import and cache get_model_info to avoid circular imports.
|
||||
|
||||
This function is used by bedrock transformation classes that need get_model_info
|
||||
but cannot import it at module level due to circular import issues.
|
||||
The function is cached after first use to avoid performance impact.
|
||||
"""
|
||||
global _get_model_info
|
||||
if _get_model_info is None:
|
||||
from litellm import get_model_info
|
||||
_get_model_info = get_model_info
|
||||
return _get_model_info
|
||||
|
||||
|
||||
class AmazonBedrockGlobalConfig:
|
||||
def __init__(self):
|
||||
pass
|
||||
|
|
@ -616,6 +635,8 @@ def get_bedrock_chat_config(model: str):
|
|||
return litellm.AmazonInvokeNovaConfig()
|
||||
elif bedrock_invoke_provider == "qwen3":
|
||||
return litellm.AmazonQwen3Config()
|
||||
elif bedrock_invoke_provider == "qwen2":
|
||||
return litellm.AmazonQwen2Config()
|
||||
elif bedrock_invoke_provider == "twelvelabs":
|
||||
return litellm.AmazonTwelveLabsPegasusConfig()
|
||||
else:
|
||||
|
|
|
|||
206
litellm/llms/bedrock/files/handler.py
Normal file
206
litellm/llms/bedrock/files/handler.py
Normal file
|
|
@ -0,0 +1,206 @@
|
|||
import asyncio
|
||||
import base64
|
||||
from typing import Any, Coroutine, Optional, Tuple, Union
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm import LlmProviders
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
from litellm.types.llms.openai import (
|
||||
FileContentRequest,
|
||||
HttpxBinaryResponseContent,
|
||||
)
|
||||
from litellm.types.utils import SpecialEnums
|
||||
|
||||
from ..base_aws_llm import BaseAWSLLM
|
||||
|
||||
|
||||
class BedrockFilesHandler(BaseAWSLLM):
|
||||
"""
|
||||
Handles downloading files from S3 for Bedrock batch processing.
|
||||
|
||||
This implementation downloads files from S3 buckets where Bedrock
|
||||
stores batch output files.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.async_httpx_client = get_async_httpx_client(
|
||||
llm_provider=LlmProviders.BEDROCK,
|
||||
)
|
||||
|
||||
def _extract_s3_uri_from_file_id(self, file_id: str) -> str:
|
||||
"""
|
||||
Extract S3 URI from encoded file ID.
|
||||
|
||||
The file ID can be in two formats:
|
||||
1. Base64-encoded unified file ID containing: llm_output_file_id,s3://bucket/path
|
||||
2. Direct S3 URI: s3://bucket/path
|
||||
|
||||
Args:
|
||||
file_id: Encoded file ID or direct S3 URI
|
||||
|
||||
Returns:
|
||||
S3 URI (e.g., "s3://bucket-name/path/to/file")
|
||||
"""
|
||||
# First, try to decode if it's a base64-encoded unified file ID
|
||||
try:
|
||||
# Add padding if needed
|
||||
padded = file_id + "=" * (-len(file_id) % 4)
|
||||
decoded = base64.urlsafe_b64decode(padded).decode()
|
||||
|
||||
# Check if it's a unified file ID format
|
||||
if decoded.startswith(SpecialEnums.LITELM_MANAGED_FILE_ID_PREFIX.value):
|
||||
# Extract llm_output_file_id from the decoded string
|
||||
if "llm_output_file_id," in decoded:
|
||||
s3_uri = decoded.split("llm_output_file_id,")[1].split(";")[0]
|
||||
return s3_uri
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# If not base64 encoded or doesn't contain llm_output_file_id, assume it's already an S3 URI
|
||||
if file_id.startswith("s3://"):
|
||||
return file_id
|
||||
|
||||
# If it doesn't start with s3://, assume it's a direct S3 URI and add the prefix
|
||||
return f"s3://{file_id}"
|
||||
|
||||
def _parse_s3_uri(self, s3_uri: str) -> Tuple[str, str]:
|
||||
"""
|
||||
Parse S3 URI to extract bucket name and object key.
|
||||
|
||||
Args:
|
||||
s3_uri: S3 URI (e.g., "s3://bucket-name/path/to/file")
|
||||
|
||||
Returns:
|
||||
Tuple of (bucket_name, object_key)
|
||||
"""
|
||||
if not s3_uri.startswith("s3://"):
|
||||
raise ValueError(f"Invalid S3 URI format: {s3_uri}. Expected format: s3://bucket-name/path/to/file")
|
||||
|
||||
# Remove 's3://' prefix
|
||||
path = s3_uri[5:]
|
||||
|
||||
if "/" in path:
|
||||
bucket_name, object_key = path.split("/", 1)
|
||||
else:
|
||||
bucket_name = path
|
||||
object_key = ""
|
||||
|
||||
return bucket_name, object_key
|
||||
|
||||
async def afile_content(
|
||||
self,
|
||||
file_content_request: FileContentRequest,
|
||||
optional_params: dict,
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
max_retries: Optional[int],
|
||||
) -> HttpxBinaryResponseContent:
|
||||
"""
|
||||
Download file content from S3 bucket for Bedrock files.
|
||||
|
||||
Args:
|
||||
file_content_request: Contains file_id (encoded or S3 URI)
|
||||
optional_params: Optional parameters containing AWS credentials
|
||||
timeout: Request timeout
|
||||
max_retries: Max retry attempts
|
||||
|
||||
Returns:
|
||||
HttpxBinaryResponseContent: Binary content wrapped in compatible response format
|
||||
"""
|
||||
import boto3
|
||||
from botocore.credentials import Credentials
|
||||
|
||||
file_id = file_content_request.get("file_id")
|
||||
if not file_id:
|
||||
raise ValueError("file_id is required in file_content_request")
|
||||
|
||||
# Extract S3 URI from file ID
|
||||
s3_uri = self._extract_s3_uri_from_file_id(file_id)
|
||||
bucket_name, object_key = self._parse_s3_uri(s3_uri)
|
||||
|
||||
# Get AWS credentials
|
||||
aws_region_name = self._get_aws_region_name(
|
||||
optional_params=optional_params, model=""
|
||||
)
|
||||
credentials: Credentials = self.get_credentials(
|
||||
aws_access_key_id=optional_params.get("aws_access_key_id"),
|
||||
aws_secret_access_key=optional_params.get("aws_secret_access_key"),
|
||||
aws_session_token=optional_params.get("aws_session_token"),
|
||||
aws_region_name=aws_region_name,
|
||||
aws_session_name=optional_params.get("aws_session_name"),
|
||||
aws_profile_name=optional_params.get("aws_profile_name"),
|
||||
aws_role_name=optional_params.get("aws_role_name"),
|
||||
aws_web_identity_token=optional_params.get("aws_web_identity_token"),
|
||||
aws_sts_endpoint=optional_params.get("aws_sts_endpoint"),
|
||||
)
|
||||
|
||||
# Create S3 client
|
||||
s3_client = boto3.client(
|
||||
"s3",
|
||||
aws_access_key_id=credentials.access_key,
|
||||
aws_secret_access_key=credentials.secret_key,
|
||||
aws_session_token=credentials.token,
|
||||
region_name=aws_region_name,
|
||||
)
|
||||
|
||||
# Download file from S3
|
||||
try:
|
||||
response = s3_client.get_object(Bucket=bucket_name, Key=object_key)
|
||||
file_content = response["Body"].read()
|
||||
except Exception as e:
|
||||
raise ValueError(f"Failed to download file from S3: {s3_uri}. Error: {str(e)}")
|
||||
|
||||
# Create mock HTTP response
|
||||
mock_response = httpx.Response(
|
||||
status_code=200,
|
||||
content=file_content,
|
||||
headers={"content-type": "application/octet-stream"},
|
||||
request=httpx.Request(method="GET", url=s3_uri),
|
||||
)
|
||||
|
||||
return HttpxBinaryResponseContent(response=mock_response)
|
||||
|
||||
def file_content(
|
||||
self,
|
||||
_is_async: bool,
|
||||
file_content_request: FileContentRequest,
|
||||
api_base: Optional[str],
|
||||
optional_params: dict,
|
||||
timeout: Union[float, httpx.Timeout],
|
||||
max_retries: Optional[int],
|
||||
) -> Union[
|
||||
HttpxBinaryResponseContent, Coroutine[Any, Any, HttpxBinaryResponseContent]
|
||||
]:
|
||||
"""
|
||||
Download file content from S3 bucket for Bedrock files.
|
||||
Supports both sync and async operations.
|
||||
|
||||
Args:
|
||||
_is_async: Whether to run asynchronously
|
||||
file_content_request: Contains file_id (encoded or S3 URI)
|
||||
api_base: API base (unused for S3 operations)
|
||||
optional_params: Optional parameters containing AWS credentials
|
||||
timeout: Request timeout
|
||||
max_retries: Max retry attempts
|
||||
|
||||
Returns:
|
||||
HttpxBinaryResponseContent or Coroutine: Binary content wrapped in compatible response format
|
||||
"""
|
||||
if _is_async:
|
||||
return self.afile_content(
|
||||
file_content_request=file_content_request,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
)
|
||||
else:
|
||||
return asyncio.run(
|
||||
self.afile_content(
|
||||
file_content_request=file_content_request,
|
||||
optional_params=optional_params,
|
||||
timeout=timeout,
|
||||
max_retries=max_retries,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
@ -3,7 +3,6 @@ from typing import Any, Dict, List, Optional
|
|||
|
||||
from openai.types.image import Image
|
||||
|
||||
from litellm import get_model_info
|
||||
from litellm.types.llms.bedrock import (
|
||||
AmazonNovaCanvasColorGuidedGenerationParams,
|
||||
AmazonNovaCanvasColorGuidedRequest,
|
||||
|
|
@ -15,6 +14,7 @@ from litellm.types.llms.bedrock import (
|
|||
AmazonNovaCanvasTextToImageRequest,
|
||||
AmazonNovaCanvasTextToImageResponse,
|
||||
)
|
||||
from litellm.llms.bedrock.common_utils import get_cached_model_info
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
|
||||
|
|
@ -207,6 +207,7 @@ class AmazonNovaCanvasConfig:
|
|||
size: Optional[str] = None,
|
||||
optional_params: Optional[dict] = None,
|
||||
) -> float:
|
||||
get_model_info = get_cached_model_info()
|
||||
model_info = get_model_info(
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ from typing import List, Optional
|
|||
|
||||
from openai.types.image import Image
|
||||
|
||||
from litellm import get_model_info
|
||||
from litellm.llms.bedrock.common_utils import get_cached_model_info
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
|
||||
|
|
@ -151,6 +151,7 @@ class AmazonStabilityConfig:
|
|||
size = size or "1024-x-1024"
|
||||
model = f"{size}/{steps}/{model}"
|
||||
|
||||
get_model_info = get_cached_model_info()
|
||||
model_info = get_model_info(
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
|
|
|
|||
|
|
@ -3,12 +3,12 @@ from typing import List, Optional
|
|||
|
||||
from openai.types.image import Image
|
||||
|
||||
from litellm import get_model_info
|
||||
from litellm.llms.bedrock.common_utils import BedrockError
|
||||
from litellm.types.llms.bedrock import (
|
||||
AmazonStability3TextToImageRequest,
|
||||
AmazonStability3TextToImageResponse,
|
||||
)
|
||||
from litellm.llms.bedrock.common_utils import get_cached_model_info
|
||||
from litellm.types.utils import ImageResponse
|
||||
|
||||
|
||||
|
|
@ -115,6 +115,7 @@ class AmazonStability3Config:
|
|||
size: Optional[str] = None,
|
||||
optional_params: Optional[dict] = None,
|
||||
) -> float:
|
||||
get_model_info = get_cached_model_info()
|
||||
model_info = get_model_info(
|
||||
model=model,
|
||||
custom_llm_provider="bedrock",
|
||||
|
|
|
|||
|
|
@ -7,7 +7,7 @@ from typing import List, Optional
|
|||
|
||||
from openai.types.image import Image
|
||||
|
||||
from litellm import get_model_info
|
||||
from litellm.utils import get_model_info
|
||||
from litellm.types.llms.bedrock import (
|
||||
AmazonNovaCanvasImageGenerationConfig,
|
||||
AmazonTitanImageGenerationRequestBody,
|
||||
|
|
|
|||
|
|
@ -49,12 +49,13 @@ class CohereRerankHandler(BaseTranslation):
|
|||
# Process query only
|
||||
query = data.get("query")
|
||||
if query is not None and isinstance(query, str):
|
||||
guardrailed_texts, _ = await guardrail_to_apply.apply_guardrail(
|
||||
texts=[query],
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": [query]},
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
data["query"] = guardrailed_texts[0] if guardrailed_texts else query
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
|
|||
|
|
@ -16,8 +16,8 @@ from litellm._logging import verbose_logger
|
|||
from litellm.constants import (
|
||||
_DEFAULT_TTL_FOR_HTTPX_CLIENTS,
|
||||
AIOHTTP_CONNECTOR_LIMIT,
|
||||
AIOHTTP_CONNECTOR_LIMIT_PER_HOST,
|
||||
AIOHTTP_KEEPALIVE_TIMEOUT,
|
||||
AIOHTTP_NEEDS_CLEANUP_CLOSED,
|
||||
AIOHTTP_TTL_DNS_CACHE,
|
||||
DEFAULT_SSL_CIPHERS,
|
||||
)
|
||||
|
|
@ -793,15 +793,20 @@ class AsyncHTTPHandler:
|
|||
verbose_logger.debug(
|
||||
"NEW SESSION: Creating new ClientSession (no shared session provided)"
|
||||
)
|
||||
transport_connector_kwargs = {
|
||||
"keepalive_timeout": AIOHTTP_KEEPALIVE_TIMEOUT,
|
||||
"ttl_dns_cache": AIOHTTP_TTL_DNS_CACHE,
|
||||
"enable_cleanup_closed": True,
|
||||
**connector_kwargs,
|
||||
}
|
||||
if AIOHTTP_CONNECTOR_LIMIT > 0:
|
||||
transport_connector_kwargs["limit"] = AIOHTTP_CONNECTOR_LIMIT
|
||||
if AIOHTTP_CONNECTOR_LIMIT_PER_HOST > 0:
|
||||
transport_connector_kwargs["limit_per_host"] = AIOHTTP_CONNECTOR_LIMIT_PER_HOST
|
||||
|
||||
return LiteLLMAiohttpTransport(
|
||||
client=lambda: ClientSession(
|
||||
connector=TCPConnector(
|
||||
limit=AIOHTTP_CONNECTOR_LIMIT,
|
||||
keepalive_timeout=AIOHTTP_KEEPALIVE_TIMEOUT,
|
||||
ttl_dns_cache=AIOHTTP_TTL_DNS_CACHE,
|
||||
enable_cleanup_closed=AIOHTTP_NEEDS_CLEANUP_CLOSED,
|
||||
**connector_kwargs,
|
||||
),
|
||||
connector=TCPConnector(**transport_connector_kwargs),
|
||||
trust_env=trust_env,
|
||||
),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -114,20 +114,27 @@ class GoogleAIStudioGeminiConfig(VertexGeminiConfig):
|
|||
img_element = element
|
||||
_image_url: Optional[str] = None
|
||||
format: Optional[str] = None
|
||||
detail: Optional[str] = None
|
||||
if isinstance(img_element.get("image_url"), dict):
|
||||
_image_url = img_element["image_url"].get("url") # type: ignore
|
||||
format = img_element["image_url"].get("format") # type: ignore
|
||||
detail = img_element["image_url"].get("detail") # type: ignore
|
||||
else:
|
||||
_image_url = img_element.get("image_url") # type: ignore
|
||||
if _image_url and "https://" in _image_url:
|
||||
image_obj = convert_to_anthropic_image_obj(
|
||||
_image_url, format=format
|
||||
)
|
||||
img_element["image_url"] = ( # type: ignore
|
||||
convert_generic_image_chunk_to_openai_image_obj(
|
||||
image_obj
|
||||
)
|
||||
converted_image_url = convert_generic_image_chunk_to_openai_image_obj(
|
||||
image_obj
|
||||
)
|
||||
if detail is not None:
|
||||
img_element["image_url"] = { # type: ignore
|
||||
"url": converted_image_url,
|
||||
"detail": detail
|
||||
}
|
||||
else:
|
||||
img_element["image_url"] = converted_image_url # type: ignore
|
||||
elif element.get("type") == "file":
|
||||
file_element = cast(ChatCompletionFileObject, element)
|
||||
file_id = file_element["file"].get("file_id")
|
||||
|
|
|
|||
|
|
@ -218,22 +218,47 @@ class GroqChatConfig(OpenAILikeChatConfig):
|
|||
When using tools in this way: - https://docs.anthropic.com/en/docs/build-with-claude/tool-use#json-mode
|
||||
- You usually want to provide a single tool
|
||||
- You should set tool_choice (see Forcing tool use) to instruct the model to explicitly use that tool
|
||||
- Remember that the model will pass the input to the tool, so the name of the tool and description should be from the model’s perspective.
|
||||
- Remember that the model will pass the input to the tool, so the name of the tool and description should be from the model's perspective.
|
||||
|
||||
Note: This workaround is only for models that don't support native json_schema.
|
||||
Models like gpt-oss-120b, llama-4, kimi-k2 support native json_schema and should
|
||||
pass response_format directly to Groq.
|
||||
See: https://console.groq.com/docs/structured-outputs#supported-models
|
||||
"""
|
||||
if json_schema is not None:
|
||||
_tool_choice = {
|
||||
"type": "function",
|
||||
"function": {"name": "json_tool_call"},
|
||||
}
|
||||
_tool = self._create_json_tool_call_for_response_format(
|
||||
json_schema=json_schema,
|
||||
)
|
||||
optional_params["tools"] = [_tool]
|
||||
optional_params["tool_choice"] = _tool_choice
|
||||
optional_params["json_mode"] = True
|
||||
non_default_params.pop(
|
||||
"response_format", None
|
||||
) # only remove if it's a json_schema - handled via using groq's tool calling params.
|
||||
# Check if model supports native response_schema
|
||||
if not litellm.supports_response_schema(
|
||||
model=model, custom_llm_provider="groq"
|
||||
):
|
||||
# Check if user is also passing tools - this combination won't work
|
||||
# See: https://console.groq.com/docs/structured-outputs
|
||||
# "Streaming and tool use are not currently supported with Structured Outputs"
|
||||
if "tools" in non_default_params:
|
||||
raise litellm.BadRequestError(
|
||||
message=f"Groq model '{model}' does not support native structured outputs. "
|
||||
"LiteLLM uses a tool-calling workaround for structured outputs on this model, "
|
||||
"which is incompatible with user-provided tools. "
|
||||
"Either use a model that supports native structured outputs "
|
||||
"(e.g., gpt-oss-120b, llama-4, kimi-k2), or remove the tools parameter. "
|
||||
"See: https://console.groq.com/docs/structured-outputs#supported-models",
|
||||
model=model,
|
||||
llm_provider="groq",
|
||||
)
|
||||
# Use workaround only for models without native support
|
||||
_tool_choice = {
|
||||
"type": "function",
|
||||
"function": {"name": "json_tool_call"},
|
||||
}
|
||||
_tool = self._create_json_tool_call_for_response_format(
|
||||
json_schema=json_schema,
|
||||
)
|
||||
optional_params["tools"] = [_tool]
|
||||
optional_params["tool_choice"] = _tool_choice
|
||||
optional_params["json_mode"] = True
|
||||
non_default_params.pop(
|
||||
"response_format", None
|
||||
) # only remove if it's a json_schema - handled via using groq's tool calling params.
|
||||
# else: model supports native json_schema, let response_format pass through
|
||||
optional_params = super().map_openai_params(
|
||||
non_default_params, optional_params, model, drop_params
|
||||
)
|
||||
|
|
|
|||
|
|
@ -19,6 +19,8 @@ from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
|||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.types.guardrails import GenericGuardrailAPIInputs
|
||||
from litellm.types.llms.openai import ChatCompletionToolParam
|
||||
from litellm.types.utils import Choices, StreamingChoices
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -52,38 +54,64 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
|
||||
texts_to_check: List[str] = []
|
||||
images_to_check: List[str] = []
|
||||
task_mappings: List[Tuple[int, Optional[int]]] = []
|
||||
# Track (message_index, content_index) for each text
|
||||
tool_calls_to_check: List[ChatCompletionToolParam] = []
|
||||
text_task_mappings: List[Tuple[int, Optional[int]]] = []
|
||||
tool_call_task_mappings: List[Tuple[int, int]] = []
|
||||
# text_task_mappings: Track (message_index, content_index) for each text
|
||||
# content_index is None for string content, int for list content
|
||||
# tool_call_task_mappings: Track (message_index, tool_call_index) for each tool call
|
||||
|
||||
# Step 1: Extract all text content and images
|
||||
# Step 1: Extract all text content, images, and tool calls
|
||||
for msg_idx, message in enumerate(messages):
|
||||
self._extract_input_text_and_images(
|
||||
self._extract_inputs(
|
||||
message=message,
|
||||
msg_idx=msg_idx,
|
||||
texts_to_check=texts_to_check,
|
||||
images_to_check=images_to_check,
|
||||
task_mappings=task_mappings,
|
||||
tool_calls_to_check=tool_calls_to_check,
|
||||
text_task_mappings=text_task_mappings,
|
||||
tool_call_task_mappings=tool_call_task_mappings,
|
||||
)
|
||||
|
||||
# Step 2: Apply guardrail to all texts in batch
|
||||
if texts_to_check:
|
||||
guardrailed_texts, guardrailed_images = (
|
||||
await guardrail_to_apply.apply_guardrail(
|
||||
texts=texts_to_check,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
images=images_to_check if images_to_check else None,
|
||||
logging_obj=litellm_logging_obj,
|
||||
# Step 2: Apply guardrail to all texts and tool calls in batch
|
||||
if texts_to_check or tool_calls_to_check:
|
||||
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if images_to_check:
|
||||
inputs["images"] = images_to_check
|
||||
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
|
||||
)
|
||||
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
guardrailed_tool_calls = guardrailed_inputs.get("tool_calls", [])
|
||||
|
||||
# Step 3: Map guardrail responses back to original message structure
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=messages,
|
||||
responses=guardrailed_texts,
|
||||
task_mappings=task_mappings,
|
||||
)
|
||||
if guardrailed_texts and texts_to_check:
|
||||
await self._apply_guardrail_responses_to_input_texts(
|
||||
messages=messages,
|
||||
responses=guardrailed_texts,
|
||||
task_mappings=text_task_mappings,
|
||||
)
|
||||
|
||||
# Step 4: Apply guardrailed tool calls back to messages
|
||||
if guardrailed_tool_calls:
|
||||
# Note: The guardrail may modify tool_calls_to_check in place
|
||||
# or we may need to handle returned tool calls differently
|
||||
await self._apply_guardrail_responses_to_input_tool_calls(
|
||||
messages=messages,
|
||||
tool_calls=guardrailed_tool_calls, # type: ignore
|
||||
task_mappings=tool_call_task_mappings,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"OpenAI Chat Completions: Processed input messages: %s", messages
|
||||
|
|
@ -91,61 +119,71 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
def _extract_input_text_and_images(
|
||||
def _extract_inputs(
|
||||
self,
|
||||
message: Dict[str, Any],
|
||||
msg_idx: int,
|
||||
texts_to_check: List[str],
|
||||
images_to_check: List[str],
|
||||
task_mappings: List[Tuple[int, Optional[int]]],
|
||||
tool_calls_to_check: List[ChatCompletionToolParam],
|
||||
text_task_mappings: List[Tuple[int, Optional[int]]],
|
||||
tool_call_task_mappings: List[Tuple[int, int]],
|
||||
) -> None:
|
||||
"""
|
||||
Extract text content and images from a message.
|
||||
Extract text content, images, and tool calls from a message.
|
||||
|
||||
Override this method to customize text/image extraction logic.
|
||||
Override this method to customize text/image/tool call extraction logic.
|
||||
"""
|
||||
content = message.get("content", None)
|
||||
if content is None:
|
||||
return
|
||||
if content is not None:
|
||||
if isinstance(content, str):
|
||||
# Simple string content
|
||||
texts_to_check.append(content)
|
||||
text_task_mappings.append((msg_idx, None))
|
||||
|
||||
if isinstance(content, str):
|
||||
# Simple string content
|
||||
texts_to_check.append(content)
|
||||
task_mappings.append((msg_idx, None))
|
||||
elif isinstance(content, list):
|
||||
# List content (e.g., multimodal with text and images)
|
||||
for content_idx, content_item in enumerate(content):
|
||||
# Extract text
|
||||
text_str = content_item.get("text", None)
|
||||
if text_str is not None:
|
||||
texts_to_check.append(text_str)
|
||||
text_task_mappings.append((msg_idx, int(content_idx)))
|
||||
|
||||
elif isinstance(content, list):
|
||||
# List content (e.g., multimodal with text and images)
|
||||
for content_idx, content_item in enumerate(content):
|
||||
# Extract text
|
||||
text_str = content_item.get("text", None)
|
||||
if text_str is not None:
|
||||
texts_to_check.append(text_str)
|
||||
task_mappings.append((msg_idx, int(content_idx)))
|
||||
# Extract images (image_url)
|
||||
if content_item.get("type") == "image_url":
|
||||
image_url = content_item.get("image_url", {})
|
||||
if isinstance(image_url, dict):
|
||||
url = image_url.get("url")
|
||||
if url:
|
||||
images_to_check.append(url)
|
||||
|
||||
# Extract images (image_url)
|
||||
if content_item.get("type") == "image_url":
|
||||
image_url = content_item.get("image_url", {})
|
||||
if isinstance(image_url, dict):
|
||||
url = image_url.get("url")
|
||||
if url:
|
||||
images_to_check.append(url)
|
||||
# Extract tool calls (typically in assistant messages)
|
||||
tool_calls = message.get("tool_calls", None)
|
||||
if tool_calls is not None and isinstance(tool_calls, list):
|
||||
for tool_call_idx, tool_call in enumerate(tool_calls):
|
||||
if isinstance(tool_call, dict):
|
||||
# Add the full tool call object to the list
|
||||
tool_calls_to_check.append(ChatCompletionToolParam(**tool_call))
|
||||
tool_call_task_mappings.append((msg_idx, int(tool_call_idx)))
|
||||
|
||||
async def _apply_guardrail_responses_to_input(
|
||||
async def _apply_guardrail_responses_to_input_texts(
|
||||
self,
|
||||
messages: List[Dict[str, Any]],
|
||||
responses: List[str],
|
||||
task_mappings: List[Tuple[int, Optional[int]]],
|
||||
) -> None:
|
||||
"""
|
||||
Apply guardrail responses back to input messages.
|
||||
Apply guardrail responses back to input message text content.
|
||||
|
||||
Override this method to customize how responses are applied.
|
||||
Override this method to customize how text responses are applied.
|
||||
"""
|
||||
for task_idx, guardrail_response in enumerate(responses):
|
||||
mapping = task_mappings[task_idx]
|
||||
msg_idx = cast(int, mapping[0])
|
||||
content_idx_optional = cast(Optional[int], mapping[1])
|
||||
|
||||
# Handle content
|
||||
content = messages[msg_idx].get("content", None)
|
||||
if content is None:
|
||||
continue
|
||||
|
|
@ -160,6 +198,31 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
"text"
|
||||
] = guardrail_response
|
||||
|
||||
async def _apply_guardrail_responses_to_input_tool_calls(
|
||||
self,
|
||||
messages: List[Dict[str, Any]],
|
||||
tool_calls: List[Dict[str, Any]],
|
||||
task_mappings: List[Tuple[int, int]],
|
||||
) -> None:
|
||||
"""
|
||||
Apply guardrailed tool calls back to input messages.
|
||||
|
||||
The guardrail may have modified the tool_calls list in place,
|
||||
so we apply the modified tool calls back to the original messages.
|
||||
|
||||
Override this method to customize how tool call responses are applied.
|
||||
"""
|
||||
for task_idx, (msg_idx, tool_call_idx) in enumerate(task_mappings):
|
||||
if task_idx < len(tool_calls):
|
||||
guardrailed_tool_call = tool_calls[task_idx]
|
||||
message_tool_calls = messages[msg_idx].get("tool_calls", None)
|
||||
if message_tool_calls is not None and isinstance(
|
||||
message_tool_calls, list
|
||||
):
|
||||
if tool_call_idx < len(message_tool_calls):
|
||||
# Replace the tool call with the guardrailed version
|
||||
message_tool_calls[tool_call_idx] = guardrailed_tool_call
|
||||
|
||||
async def process_output_response(
|
||||
self,
|
||||
response: "ModelResponse",
|
||||
|
|
@ -193,21 +256,27 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
|
||||
texts_to_check: List[str] = []
|
||||
images_to_check: List[str] = []
|
||||
task_mappings: List[Tuple[int, Optional[int]]] = []
|
||||
# Track (choice_index, content_index) for each text
|
||||
tool_calls_to_check: List[Dict[str, Any]] = []
|
||||
text_task_mappings: List[Tuple[int, Optional[int]]] = []
|
||||
tool_call_task_mappings: List[Tuple[int, int]] = []
|
||||
# text_task_mappings: Track (choice_index, content_index) for each text
|
||||
# content_index is None for string content, int for list content
|
||||
# tool_call_task_mappings: Track (choice_index, tool_call_index) for each tool call
|
||||
|
||||
# Step 1: Extract all text content and images from response choices
|
||||
# Step 1: Extract all text content, images, and tool calls from response choices
|
||||
for choice_idx, choice in enumerate(response.choices):
|
||||
self._extract_output_text_and_images(
|
||||
self._extract_output_text_images_and_tool_calls(
|
||||
choice=choice,
|
||||
choice_idx=choice_idx,
|
||||
texts_to_check=texts_to_check,
|
||||
images_to_check=images_to_check,
|
||||
task_mappings=task_mappings,
|
||||
tool_calls_to_check=tool_calls_to_check,
|
||||
text_task_mappings=text_task_mappings,
|
||||
tool_call_task_mappings=tool_call_task_mappings,
|
||||
)
|
||||
|
||||
# Step 2: Apply guardrail to all texts in batch
|
||||
if texts_to_check:
|
||||
# Step 2: Apply guardrail to all texts and tool calls in batch
|
||||
if texts_to_check or tool_calls_to_check:
|
||||
# Create a request_data dict with response info and user API key metadata
|
||||
request_data: dict = {"response": response}
|
||||
|
||||
|
|
@ -218,22 +287,36 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
if user_metadata:
|
||||
request_data["litellm_metadata"] = user_metadata
|
||||
|
||||
guardrailed_texts, guardrailed_images = (
|
||||
await guardrail_to_apply.apply_guardrail(
|
||||
texts=texts_to_check,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
images=images_to_check if images_to_check else None,
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if images_to_check:
|
||||
inputs["images"] = images_to_check
|
||||
if tool_calls_to_check:
|
||||
inputs["tool_calls"] = tool_calls_to_check # type: ignore
|
||||
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
|
||||
# Step 3: Map guardrail responses back to original response structure
|
||||
await self._apply_guardrail_responses_to_output(
|
||||
response=response,
|
||||
responses=guardrailed_texts,
|
||||
task_mappings=task_mappings,
|
||||
)
|
||||
if guardrailed_texts and texts_to_check:
|
||||
await self._apply_guardrail_responses_to_output_texts(
|
||||
response=response,
|
||||
responses=guardrailed_texts,
|
||||
task_mappings=text_task_mappings,
|
||||
)
|
||||
|
||||
# Step 4: Apply guardrailed tool calls back to response
|
||||
if tool_calls_to_check:
|
||||
await self._apply_guardrail_responses_to_output_tool_calls(
|
||||
response=response,
|
||||
tool_calls=tool_calls_to_check,
|
||||
task_mappings=tool_call_task_mappings,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"OpenAI Chat Completions: Processed output response: %s", response
|
||||
|
|
@ -243,53 +326,129 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
|
||||
async def process_output_streaming_response(
|
||||
self,
|
||||
response: "ModelResponseStream",
|
||||
responses_so_far: List["ModelResponseStream"],
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
litellm_logging_obj: Optional[Any] = None,
|
||||
user_api_key_dict: Optional[Any] = None,
|
||||
) -> Any:
|
||||
) -> List["ModelResponseStream"]:
|
||||
"""
|
||||
Process output streaming response by applying guardrails to text content.
|
||||
Process output streaming responses by applying guardrails to text content.
|
||||
|
||||
Args:
|
||||
response: LiteLLM ModelResponseStream object
|
||||
responses_so_far: List of LiteLLM ModelResponseStream objects
|
||||
guardrail_to_apply: The guardrail instance to apply
|
||||
litellm_logging_obj: Optional logging object
|
||||
user_api_key_dict: User API key metadata to pass to guardrails
|
||||
|
||||
Returns:
|
||||
Modified response with guardrail applied to content
|
||||
Modified list of responses with guardrail applied to content
|
||||
|
||||
Response Format Support:
|
||||
- String content: choice.message.content = "text here"
|
||||
- List content: choice.message.content = [{"type": "text", "text": "text here"}, ...]
|
||||
"""
|
||||
|
||||
# Step 0: Check if response has any text content to process
|
||||
if not self._has_text_content(response):
|
||||
return response
|
||||
# Step 0: Check if any response has text content to process
|
||||
has_any_text_content = False
|
||||
for response in responses_so_far:
|
||||
if self._has_text_content(response):
|
||||
has_any_text_content = True
|
||||
break
|
||||
|
||||
if not has_any_text_content:
|
||||
verbose_proxy_logger.warning(
|
||||
"OpenAI Chat Completions: No text content in streaming responses, skipping guardrail"
|
||||
)
|
||||
return responses_so_far
|
||||
|
||||
# Step 1: Combine all streaming chunks into complete text per choice
|
||||
# For streaming, we need to concatenate all delta.content across all chunks
|
||||
# Key: (choice_idx, content_idx), Value: combined text
|
||||
combined_texts: Dict[Tuple[int, Optional[int]], str] = {}
|
||||
|
||||
for response_idx, response in enumerate(responses_so_far):
|
||||
for choice_idx, choice in enumerate(response.choices):
|
||||
if isinstance(choice, litellm.StreamingChoices):
|
||||
content = choice.delta.content
|
||||
elif isinstance(choice, litellm.Choices):
|
||||
content = choice.message.content
|
||||
else:
|
||||
continue
|
||||
|
||||
if content is None:
|
||||
continue
|
||||
|
||||
if isinstance(content, str):
|
||||
# String content - accumulate for this choice
|
||||
key = (choice_idx, None)
|
||||
if key not in combined_texts:
|
||||
combined_texts[key] = ""
|
||||
combined_texts[key] += content
|
||||
|
||||
elif isinstance(content, list):
|
||||
# List content - accumulate for each content item
|
||||
for content_idx, content_item in enumerate(content):
|
||||
text_str = content_item.get("text")
|
||||
if text_str:
|
||||
key = (choice_idx, content_idx)
|
||||
if key not in combined_texts:
|
||||
combined_texts[key] = ""
|
||||
combined_texts[key] += text_str
|
||||
|
||||
# Step 2: Create lists for guardrail processing
|
||||
texts_to_check: List[str] = []
|
||||
images_to_check: List[str] = []
|
||||
task_mappings: List[Tuple[int, Optional[int]]] = []
|
||||
# Track (choice_index, content_index) for each text
|
||||
# Track (choice_index, content_index) for each combined text
|
||||
|
||||
# Step 1: Extract all text content and images from response choices
|
||||
for choice_idx, choice in enumerate(response.choices):
|
||||
for (choice_idx, content_idx), combined_text in combined_texts.items():
|
||||
texts_to_check.append(combined_text)
|
||||
task_mappings.append((choice_idx, content_idx))
|
||||
|
||||
self._extract_output_text_and_images(
|
||||
choice=choice,
|
||||
choice_idx=choice_idx,
|
||||
texts_to_check=texts_to_check,
|
||||
images_to_check=images_to_check,
|
||||
# Step 3: Apply guardrail to all combined texts in batch
|
||||
if texts_to_check:
|
||||
# Create a request_data dict with response info and user API key metadata
|
||||
request_data: dict = {"responses": responses_so_far}
|
||||
|
||||
# Add user API key metadata with prefixed keys
|
||||
user_metadata = self.transform_user_api_key_dict_to_metadata(
|
||||
user_api_key_dict
|
||||
)
|
||||
if user_metadata:
|
||||
request_data["litellm_metadata"] = user_metadata
|
||||
|
||||
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if images_to_check:
|
||||
inputs["images"] = images_to_check
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
|
||||
# Step 4: Apply guardrailed text back to all streaming chunks
|
||||
# For each choice, replace the combined text across all chunks
|
||||
await self._apply_guardrail_responses_to_output_streaming(
|
||||
responses=responses_so_far,
|
||||
guardrailed_texts=guardrailed_texts,
|
||||
task_mappings=task_mappings,
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
"OpenAI Chat Completions: Processed output streaming responses: %s",
|
||||
responses_so_far,
|
||||
)
|
||||
|
||||
return responses_so_far
|
||||
|
||||
def _has_text_content(
|
||||
self, response: Union["ModelResponse", "ModelResponseStream"]
|
||||
) -> bool:
|
||||
"""
|
||||
Check if response has any text content to process.
|
||||
Check if response has any text content or tool calls to process.
|
||||
|
||||
Override this method to customize text content detection.
|
||||
"""
|
||||
|
|
@ -298,42 +457,59 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
if isinstance(response, ModelResponse):
|
||||
for choice in response.choices:
|
||||
if isinstance(choice, litellm.Choices):
|
||||
# Check for text content
|
||||
if choice.message.content and isinstance(
|
||||
choice.message.content, str
|
||||
):
|
||||
return True
|
||||
# Check for tool calls
|
||||
if choice.message.tool_calls and isinstance(
|
||||
choice.message.tool_calls, list
|
||||
):
|
||||
if len(choice.message.tool_calls) > 0:
|
||||
return True
|
||||
elif isinstance(response, ModelResponseStream):
|
||||
for choice in response.choices:
|
||||
if isinstance(choice, litellm.Choices):
|
||||
if choice.message.content and isinstance(
|
||||
choice.message.content, str
|
||||
):
|
||||
if isinstance(choice, litellm.StreamingChoices):
|
||||
# Check for text content
|
||||
if choice.delta.content and isinstance(choice.delta.content, str):
|
||||
return True
|
||||
# Check for tool calls
|
||||
if choice.delta.tool_calls and isinstance(
|
||||
choice.delta.tool_calls, list
|
||||
):
|
||||
if len(choice.delta.tool_calls) > 0:
|
||||
return True
|
||||
return False
|
||||
|
||||
def _extract_output_text_and_images(
|
||||
def _extract_output_text_images_and_tool_calls(
|
||||
self,
|
||||
choice: Union[Choices, StreamingChoices],
|
||||
choice_idx: int,
|
||||
texts_to_check: List[str],
|
||||
images_to_check: List[str],
|
||||
task_mappings: List[Tuple[int, Optional[int]]],
|
||||
tool_calls_to_check: List[Dict[str, Any]],
|
||||
text_task_mappings: List[Tuple[int, Optional[int]]],
|
||||
tool_call_task_mappings: List[Tuple[int, int]],
|
||||
) -> None:
|
||||
"""
|
||||
Extract text content and images from a response choice.
|
||||
Extract text content, images, and tool calls from a response choice.
|
||||
|
||||
Override this method to customize text/image extraction logic.
|
||||
Override this method to customize text/image/tool call extraction logic.
|
||||
"""
|
||||
verbose_proxy_logger.debug(
|
||||
"OpenAI Chat Completions: Processing choice: %s", choice
|
||||
)
|
||||
|
||||
# Determine content source based on choice type
|
||||
# Determine content source and tool calls based on choice type
|
||||
content = None
|
||||
tool_calls = None
|
||||
if isinstance(choice, litellm.Choices):
|
||||
content = choice.message.content
|
||||
tool_calls = choice.message.tool_calls
|
||||
elif isinstance(choice, litellm.StreamingChoices):
|
||||
content = choice.delta.content
|
||||
tool_calls = choice.delta.tool_calls
|
||||
else:
|
||||
# Unknown choice type, skip processing
|
||||
return
|
||||
|
|
@ -342,7 +518,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
if content and isinstance(content, str):
|
||||
# Simple string content
|
||||
texts_to_check.append(content)
|
||||
task_mappings.append((choice_idx, None))
|
||||
text_task_mappings.append((choice_idx, None))
|
||||
|
||||
elif content and isinstance(content, list):
|
||||
# List content (e.g., multimodal response)
|
||||
|
|
@ -351,7 +527,7 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
content_text = content_item.get("text")
|
||||
if content_text:
|
||||
texts_to_check.append(content_text)
|
||||
task_mappings.append((choice_idx, int(content_idx)))
|
||||
text_task_mappings.append((choice_idx, int(content_idx)))
|
||||
|
||||
# Extract images
|
||||
if content_item.get("type") == "image_url":
|
||||
|
|
@ -361,36 +537,181 @@ class OpenAIChatCompletionsHandler(BaseTranslation):
|
|||
if url:
|
||||
images_to_check.append(url)
|
||||
|
||||
async def _apply_guardrail_responses_to_output(
|
||||
# Process tool calls if they exist
|
||||
if tool_calls is not None and isinstance(tool_calls, list):
|
||||
for tool_call_idx, tool_call in enumerate(tool_calls):
|
||||
# Convert tool call to dict format for guardrail processing
|
||||
tool_call_dict = self._convert_tool_call_to_dict(tool_call)
|
||||
if tool_call_dict:
|
||||
tool_calls_to_check.append(tool_call_dict)
|
||||
tool_call_task_mappings.append((choice_idx, int(tool_call_idx)))
|
||||
|
||||
def _convert_tool_call_to_dict(
|
||||
self, tool_call: Union[Dict[str, Any], Any]
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Convert a tool call object to dictionary format.
|
||||
|
||||
Tool calls can be either dict or object depending on the type.
|
||||
"""
|
||||
if isinstance(tool_call, dict):
|
||||
return tool_call
|
||||
elif hasattr(tool_call, "id") and hasattr(tool_call, "function"):
|
||||
# Convert object to dict
|
||||
function = tool_call.function
|
||||
function_dict = {}
|
||||
if hasattr(function, "name"):
|
||||
function_dict["name"] = function.name
|
||||
if hasattr(function, "arguments"):
|
||||
function_dict["arguments"] = function.arguments
|
||||
|
||||
tool_call_dict = {
|
||||
"id": tool_call.id if hasattr(tool_call, "id") else None,
|
||||
"type": tool_call.type if hasattr(tool_call, "type") else "function",
|
||||
"function": function_dict,
|
||||
}
|
||||
return tool_call_dict
|
||||
return None
|
||||
|
||||
async def _apply_guardrail_responses_to_output_texts(
|
||||
self,
|
||||
response: "ModelResponse",
|
||||
responses: List[str],
|
||||
task_mappings: List[Tuple[int, Optional[int]]],
|
||||
) -> None:
|
||||
"""
|
||||
Apply guardrail responses back to output response.
|
||||
Apply guardrail text responses back to output response.
|
||||
|
||||
Override this method to customize how responses are applied.
|
||||
Override this method to customize how text responses are applied.
|
||||
"""
|
||||
for task_idx, guardrail_response in enumerate(responses):
|
||||
mapping = task_mappings[task_idx]
|
||||
choice_idx = cast(int, mapping[0])
|
||||
content_idx_optional = cast(Optional[int], mapping[1])
|
||||
|
||||
content = cast(Choices, response.choices[choice_idx]).message.content
|
||||
choice = cast(Choices, response.choices[choice_idx])
|
||||
|
||||
# Handle content
|
||||
content = choice.message.content
|
||||
if content is None:
|
||||
continue
|
||||
|
||||
if isinstance(content, str) and content_idx_optional is None:
|
||||
# Replace string content with guardrail response
|
||||
cast(Choices, response.choices[choice_idx]).message.content = (
|
||||
guardrail_response
|
||||
)
|
||||
choice.message.content = guardrail_response
|
||||
|
||||
elif isinstance(content, list) and content_idx_optional is not None:
|
||||
# Replace specific text item in list content
|
||||
cast(Choices, response.choices[choice_idx]).message.content[ # type: ignore
|
||||
content_idx_optional
|
||||
][
|
||||
"text"
|
||||
] = guardrail_response
|
||||
choice.message.content[content_idx_optional]["text"] = guardrail_response # type: ignore
|
||||
|
||||
async def _apply_guardrail_responses_to_output_tool_calls(
|
||||
self,
|
||||
response: "ModelResponse",
|
||||
tool_calls: List[Dict[str, Any]],
|
||||
task_mappings: List[Tuple[int, int]],
|
||||
) -> None:
|
||||
"""
|
||||
Apply guardrailed tool calls back to output response.
|
||||
|
||||
The guardrail may have modified the tool_calls list in place,
|
||||
so we apply the modified tool calls back to the original response.
|
||||
|
||||
Override this method to customize how tool call responses are applied.
|
||||
"""
|
||||
for task_idx, (choice_idx, tool_call_idx) in enumerate(task_mappings):
|
||||
if task_idx < len(tool_calls):
|
||||
guardrailed_tool_call = tool_calls[task_idx]
|
||||
choice = cast(Choices, response.choices[choice_idx])
|
||||
choice_tool_calls = choice.message.tool_calls
|
||||
|
||||
if choice_tool_calls is not None and isinstance(
|
||||
choice_tool_calls, list
|
||||
):
|
||||
if tool_call_idx < len(choice_tool_calls):
|
||||
# Update the tool call with guardrailed version
|
||||
existing_tool_call = choice_tool_calls[tool_call_idx]
|
||||
# Update object attributes (output responses always have typed objects)
|
||||
if "function" in guardrailed_tool_call:
|
||||
func_dict = guardrailed_tool_call["function"]
|
||||
if "arguments" in func_dict:
|
||||
existing_tool_call.function.arguments = func_dict[
|
||||
"arguments"
|
||||
]
|
||||
if "name" in func_dict:
|
||||
existing_tool_call.function.name = func_dict["name"]
|
||||
|
||||
async def _apply_guardrail_responses_to_output_streaming(
|
||||
self,
|
||||
responses: List["ModelResponseStream"],
|
||||
guardrailed_texts: List[str],
|
||||
task_mappings: List[Tuple[int, Optional[int]]],
|
||||
) -> None:
|
||||
"""
|
||||
Apply guardrail responses back to output streaming responses.
|
||||
|
||||
For streaming responses, the guardrailed text (which is the combined text from all chunks)
|
||||
is placed in the first chunk, and subsequent chunks are cleared.
|
||||
|
||||
Args:
|
||||
responses: List of ModelResponseStream objects to modify
|
||||
guardrailed_texts: List of guardrailed text responses (combined from all chunks)
|
||||
task_mappings: List of tuples (choice_idx, content_idx)
|
||||
|
||||
Override this method to customize how responses are applied to streaming responses.
|
||||
"""
|
||||
# Build a mapping of what guardrailed text to use for each (choice_idx, content_idx)
|
||||
guardrail_map: Dict[Tuple[int, Optional[int]], str] = {}
|
||||
for task_idx, guardrail_response in enumerate(guardrailed_texts):
|
||||
mapping = task_mappings[task_idx]
|
||||
choice_idx = cast(int, mapping[0])
|
||||
content_idx_optional = cast(Optional[int], mapping[1])
|
||||
guardrail_map[(choice_idx, content_idx_optional)] = guardrail_response
|
||||
|
||||
# Track which choices we've already set the guardrailed text for
|
||||
# Key: (choice_idx, content_idx), Value: boolean (True if already set)
|
||||
already_set: Dict[Tuple[int, Optional[int]], bool] = {}
|
||||
|
||||
# Iterate through all responses and update content
|
||||
for response_idx, response in enumerate(responses):
|
||||
for choice_idx_in_response, choice in enumerate(response.choices):
|
||||
if isinstance(choice, litellm.StreamingChoices):
|
||||
content = choice.delta.content
|
||||
elif isinstance(choice, litellm.Choices):
|
||||
content = choice.message.content
|
||||
else:
|
||||
continue
|
||||
|
||||
if content is None:
|
||||
continue
|
||||
|
||||
if isinstance(content, str):
|
||||
# String content
|
||||
key = (choice_idx_in_response, None)
|
||||
if key in guardrail_map:
|
||||
if key not in already_set:
|
||||
# First chunk - set the complete guardrailed text
|
||||
if isinstance(choice, litellm.StreamingChoices):
|
||||
choice.delta.content = guardrail_map[key]
|
||||
elif isinstance(choice, litellm.Choices):
|
||||
choice.message.content = guardrail_map[key]
|
||||
already_set[key] = True
|
||||
else:
|
||||
# Subsequent chunks - clear the content
|
||||
if isinstance(choice, litellm.StreamingChoices):
|
||||
choice.delta.content = ""
|
||||
elif isinstance(choice, litellm.Choices):
|
||||
choice.message.content = ""
|
||||
|
||||
elif isinstance(content, list):
|
||||
# List content - handle each content item
|
||||
for content_idx, content_item in enumerate(content):
|
||||
if "text" in content_item:
|
||||
key = (choice_idx_in_response, content_idx)
|
||||
if key in guardrail_map:
|
||||
if key not in already_set:
|
||||
# First chunk - set the complete guardrailed text
|
||||
content_item["text"] = guardrail_map[key]
|
||||
already_set[key] = True
|
||||
else:
|
||||
# Subsequent chunks - clear the text
|
||||
content_item["text"] = ""
|
||||
|
|
|
|||
|
|
@ -53,12 +53,13 @@ class OpenAITextCompletionHandler(BaseTranslation):
|
|||
|
||||
if isinstance(prompt, str):
|
||||
# Single string prompt
|
||||
guardrailed_texts, _ = await guardrail_to_apply.apply_guardrail(
|
||||
texts=[prompt],
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": [prompt]},
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
data["prompt"] = guardrailed_texts[0] if guardrailed_texts else prompt
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
@ -79,12 +80,13 @@ class OpenAITextCompletionHandler(BaseTranslation):
|
|||
text_indices.append(idx)
|
||||
|
||||
if texts_to_check:
|
||||
guardrailed_texts, _ = await guardrail_to_apply.apply_guardrail(
|
||||
texts=texts_to_check,
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": texts_to_check},
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
|
||||
# Replace guardrailed texts back
|
||||
for guardrail_idx, prompt_idx in enumerate(text_indices):
|
||||
|
|
@ -152,12 +154,13 @@ class OpenAITextCompletionHandler(BaseTranslation):
|
|||
if user_metadata:
|
||||
request_data["litellm_metadata"] = user_metadata
|
||||
|
||||
guardrailed_texts, _ = await guardrail_to_apply.apply_guardrail(
|
||||
texts=texts_to_check,
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": texts_to_check},
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
|
||||
# Apply guardrailed texts back to choices
|
||||
for guardrail_idx, choice_idx in enumerate(choice_indices):
|
||||
|
|
|
|||
|
|
@ -52,12 +52,13 @@ class OpenAIImageGenerationHandler(BaseTranslation):
|
|||
|
||||
# Apply guardrail to the prompt
|
||||
if isinstance(prompt, str):
|
||||
guardrailed_texts, _ = await guardrail_to_apply.apply_guardrail(
|
||||
texts=[prompt],
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": [prompt]},
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
data["prompt"] = guardrailed_texts[0] if guardrailed_texts else prompt
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
|
|||
|
|
@ -444,6 +444,11 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
else:
|
||||
headers = {}
|
||||
response = raw_response.parse()
|
||||
if not data.get("stream") and not hasattr(response, "model_dump"):
|
||||
raise OpenAIError(
|
||||
status_code=500,
|
||||
message=f"Empty or invalid response from LLM endpoint. Received: {response!r}. Check the reverse proxy or model server configuration.",
|
||||
)
|
||||
return headers, response
|
||||
except openai.APITimeoutError as e:
|
||||
end_time = time.time()
|
||||
|
|
@ -477,7 +482,14 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
|
|||
else:
|
||||
headers = {}
|
||||
response = raw_response.parse()
|
||||
if not data.get("stream") and not hasattr(response, "model_dump"):
|
||||
raise OpenAIError(
|
||||
status_code=500,
|
||||
message=f"Empty or invalid response from LLM endpoint. Received: {response!r}. Check the reverse proxy or model server configuration.",
|
||||
)
|
||||
return headers, response
|
||||
except OpenAIError:
|
||||
raise
|
||||
except Exception as e:
|
||||
if raw_response is not None:
|
||||
raise Exception(
|
||||
|
|
|
|||
|
|
@ -28,11 +28,25 @@ Output: response.output is List[GenericResponseOutputItem] where each has:
|
|||
- text: str
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, List, Optional, Tuple, Union, cast
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
from openai import BaseModel
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.llms.base_llm.guardrail_translation.base_translation import BaseTranslation
|
||||
from litellm.types.responses.main import GenericResponseOutputItem, OutputText
|
||||
from litellm.responses.litellm_completion_transformation.transformation import (
|
||||
LiteLLMCompletionResponsesConfig,
|
||||
)
|
||||
from litellm.types.guardrails import GenericGuardrailAPIInputs
|
||||
from litellm.types.llms.openai import (
|
||||
ChatCompletionToolCallChunk,
|
||||
ChatCompletionToolParam,
|
||||
)
|
||||
from litellm.types.responses.main import (
|
||||
GenericResponseOutputItem,
|
||||
OutputFunctionToolCall,
|
||||
OutputText,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.integrations.custom_guardrail import CustomGuardrail
|
||||
|
|
@ -63,17 +77,37 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
Handles both string input and list of message objects.
|
||||
"""
|
||||
input_data: Optional[Union[str, "ResponseInputParam"]] = data.get("input")
|
||||
tools_to_check: List[ChatCompletionToolParam] = []
|
||||
if input_data is None:
|
||||
return data
|
||||
|
||||
structured_messages = (
|
||||
LiteLLMCompletionResponsesConfig.transform_responses_api_input_to_messages(
|
||||
input=input_data,
|
||||
responses_api_request=data,
|
||||
)
|
||||
)
|
||||
|
||||
# Handle simple string input
|
||||
if isinstance(input_data, str):
|
||||
guardrailed_texts, _ = await guardrail_to_apply.apply_guardrail(
|
||||
texts=[input_data],
|
||||
inputs = GenericGuardrailAPIInputs(texts=[input_data])
|
||||
|
||||
# Extract and transform tools if present
|
||||
|
||||
if "tools" in data and data["tools"]:
|
||||
self._extract_and_transform_tools(data["tools"], tools_to_check)
|
||||
if tools_to_check:
|
||||
inputs["tools"] = tools_to_check
|
||||
if structured_messages:
|
||||
inputs["structured_messages"] = structured_messages # type: ignore
|
||||
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_data
|
||||
verbose_proxy_logger.debug("OpenAI Responses API: Processed string input")
|
||||
return data
|
||||
|
|
@ -88,7 +122,7 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
# Track (message_index, content_index) for each text
|
||||
# content_index is None for string content, int for list content
|
||||
|
||||
# Step 1: Extract all text content and images
|
||||
# Step 1: Extract all text content, images, and tools
|
||||
for msg_idx, message in enumerate(input_data):
|
||||
self._extract_input_text_and_images(
|
||||
message=message,
|
||||
|
|
@ -98,18 +132,28 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
task_mappings=task_mappings,
|
||||
)
|
||||
|
||||
# Extract and transform tools if present
|
||||
if "tools" in data and data["tools"]:
|
||||
self._extract_and_transform_tools(data["tools"], tools_to_check)
|
||||
|
||||
# Step 2: Apply guardrail to all texts in batch
|
||||
if texts_to_check:
|
||||
guardrailed_texts, guardrailed_images = (
|
||||
await guardrail_to_apply.apply_guardrail(
|
||||
texts=texts_to_check,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
images=images_to_check if images_to_check else None,
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if images_to_check:
|
||||
inputs["images"] = images_to_check
|
||||
if tools_to_check:
|
||||
inputs["tools"] = tools_to_check
|
||||
if structured_messages:
|
||||
inputs["structured_messages"] = structured_messages # type: ignore
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
|
||||
# Step 3: Map guardrail responses back to original input structure
|
||||
await self._apply_guardrail_responses_to_input(
|
||||
messages=input_data,
|
||||
|
|
@ -123,6 +167,29 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
|
||||
return data
|
||||
|
||||
def _extract_and_transform_tools(
|
||||
self,
|
||||
tools: List[Dict[str, Any]],
|
||||
tools_to_check: List[ChatCompletionToolParam],
|
||||
) -> None:
|
||||
"""
|
||||
Extract and transform tools from Responses API format to Chat Completion format.
|
||||
|
||||
Uses the LiteLLM transformation function to convert Responses API tools
|
||||
to Chat Completion tools that can be passed to guardrails.
|
||||
"""
|
||||
if tools is not None and isinstance(tools, list):
|
||||
# Transform Responses API tools to Chat Completion tools
|
||||
(
|
||||
transformed_tools,
|
||||
_,
|
||||
) = LiteLLMCompletionResponsesConfig.transform_responses_api_tools_to_chat_completion_tools(
|
||||
tools # type: ignore
|
||||
)
|
||||
tools_to_check.extend(
|
||||
cast(List[ChatCompletionToolParam], transformed_tools)
|
||||
)
|
||||
|
||||
def _extract_input_text_and_images(
|
||||
self,
|
||||
message: Any, # Can be Dict[str, Any] or ResponseInputParam
|
||||
|
|
@ -202,7 +269,7 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
user_api_key_dict: Optional[Any] = None,
|
||||
) -> Any:
|
||||
"""
|
||||
Process output response by applying guardrails to text content.
|
||||
Process output response by applying guardrails to text content and tool calls.
|
||||
|
||||
Args:
|
||||
response: LiteLLM ResponsesAPIResponse object
|
||||
|
|
@ -215,22 +282,19 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
|
||||
Response Format Support:
|
||||
- response.output is a list of output items
|
||||
- Each output item has a content list with OutputText objects
|
||||
- Each output item can be:
|
||||
* GenericResponseOutputItem with a content list of OutputText objects
|
||||
* OutputFunctionToolCall with tool call data
|
||||
- Each OutputText object has a text field
|
||||
"""
|
||||
# Step 0: Check if response has any text content to process
|
||||
if not self._has_text_content(response):
|
||||
verbose_proxy_logger.warning(
|
||||
"OpenAI Responses API: No text content in response, skipping guardrail"
|
||||
)
|
||||
return response
|
||||
|
||||
texts_to_check: List[str] = []
|
||||
images_to_check: List[str] = []
|
||||
tool_calls_to_check: List[ChatCompletionToolCallChunk] = []
|
||||
task_mappings: List[Tuple[int, int]] = []
|
||||
# Track (output_item_index, content_index) for each text
|
||||
|
||||
# Step 1: Extract all text content from response output
|
||||
# Step 1: Extract all text content and tool calls from response output
|
||||
for output_idx, output_item in enumerate(response.output):
|
||||
self._extract_output_text_and_images(
|
||||
output_item=output_item,
|
||||
|
|
@ -238,10 +302,11 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
texts_to_check=texts_to_check,
|
||||
images_to_check=images_to_check,
|
||||
task_mappings=task_mappings,
|
||||
tool_calls_to_check=tool_calls_to_check,
|
||||
)
|
||||
|
||||
# Step 2: Apply guardrail to all texts in batch
|
||||
if texts_to_check:
|
||||
if texts_to_check or tool_calls_to_check:
|
||||
# Create a request_data dict with response info and user API key metadata
|
||||
request_data: dict = {"response": response}
|
||||
|
||||
|
|
@ -252,16 +317,21 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
if user_metadata:
|
||||
request_data["litellm_metadata"] = user_metadata
|
||||
|
||||
guardrailed_texts, guardrailed_images = (
|
||||
await guardrail_to_apply.apply_guardrail(
|
||||
texts=texts_to_check,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
images=images_to_check if images_to_check else None,
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
inputs = GenericGuardrailAPIInputs(texts=texts_to_check)
|
||||
if images_to_check:
|
||||
inputs["images"] = images_to_check
|
||||
if tool_calls_to_check:
|
||||
inputs["tool_calls"] = tool_calls_to_check
|
||||
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs=inputs,
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
|
||||
# Step 3: Map guardrail responses back to original response structure
|
||||
await self._apply_guardrail_responses_to_output(
|
||||
response=response,
|
||||
|
|
@ -275,6 +345,31 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
|
||||
return response
|
||||
|
||||
async def process_output_streaming_response(
|
||||
self,
|
||||
responses_so_far: List[Any],
|
||||
guardrail_to_apply: "CustomGuardrail",
|
||||
litellm_logging_obj: Optional[Any] = None,
|
||||
user_api_key_dict: Optional[Any] = None,
|
||||
) -> List[Any]:
|
||||
"""
|
||||
Process output streaming response by applying guardrails to text content.
|
||||
"""
|
||||
string_so_far = self.get_streaming_string_so_far(responses_so_far)
|
||||
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": [string_so_far]},
|
||||
request_data={},
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
return responses_so_far
|
||||
|
||||
def get_streaming_string_so_far(self, responses_so_far: List[Any]) -> str:
|
||||
"""
|
||||
Get the string so far from the responses so far.
|
||||
"""
|
||||
return "".join([response.get("text", "") for response in responses_so_far])
|
||||
|
||||
def _has_text_content(self, response: "ResponsesAPIResponse") -> bool:
|
||||
"""
|
||||
Check if response has any text content to process.
|
||||
|
|
@ -285,6 +380,17 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
return False
|
||||
|
||||
for output_item in response.output:
|
||||
if isinstance(output_item, BaseModel):
|
||||
try:
|
||||
generic_response_output_item = (
|
||||
GenericResponseOutputItem.model_validate(
|
||||
output_item.model_dump()
|
||||
)
|
||||
)
|
||||
if generic_response_output_item.content:
|
||||
output_item = generic_response_output_item
|
||||
except Exception:
|
||||
continue
|
||||
if isinstance(output_item, (GenericResponseOutputItem, dict)):
|
||||
content = (
|
||||
output_item.content
|
||||
|
|
@ -296,9 +402,11 @@ 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
|
||||
|
||||
|
|
@ -309,15 +417,68 @@ class OpenAIResponsesHandler(BaseTranslation):
|
|||
texts_to_check: List[str],
|
||||
images_to_check: List[str],
|
||||
task_mappings: List[Tuple[int, int]],
|
||||
tool_calls_to_check: Optional[List[ChatCompletionToolCallChunk]] = None,
|
||||
) -> None:
|
||||
"""
|
||||
Extract text content and images from a response output item.
|
||||
Extract text content, images, and tool calls from a response output item.
|
||||
|
||||
Override this method to customize text/image extraction logic.
|
||||
Override this method to customize text/image/tool extraction logic.
|
||||
"""
|
||||
# Check if this is a tool call (OutputFunctionToolCall)
|
||||
if isinstance(output_item, OutputFunctionToolCall):
|
||||
if tool_calls_to_check is not None:
|
||||
tool_call_dict = LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
|
||||
tool_call_item=output_item,
|
||||
index=output_idx,
|
||||
)
|
||||
tool_calls_to_check.append(
|
||||
cast(ChatCompletionToolCallChunk, tool_call_dict)
|
||||
)
|
||||
return
|
||||
elif (
|
||||
isinstance(output_item, BaseModel)
|
||||
and hasattr(output_item, "type")
|
||||
and getattr(output_item, "type") == "function_call"
|
||||
):
|
||||
if tool_calls_to_check is not None:
|
||||
tool_call_dict = LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
|
||||
tool_call_item=output_item,
|
||||
index=output_idx,
|
||||
)
|
||||
tool_calls_to_check.append(
|
||||
cast(ChatCompletionToolCallChunk, tool_call_dict)
|
||||
)
|
||||
return
|
||||
elif (
|
||||
isinstance(output_item, dict) and output_item.get("type") == "function_call"
|
||||
):
|
||||
# Handle dict representation of tool call
|
||||
if tool_calls_to_check is not None:
|
||||
# Convert dict to OutputFunctionToolCall for processing
|
||||
try:
|
||||
tool_call_obj = OutputFunctionToolCall(**output_item)
|
||||
tool_call_dict = LiteLLMCompletionResponsesConfig.convert_response_function_tool_call_to_chat_completion_tool_call(
|
||||
tool_call_item=tool_call_obj,
|
||||
index=output_idx,
|
||||
)
|
||||
tool_calls_to_check.append(
|
||||
cast(ChatCompletionToolCallChunk, tool_call_dict)
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return
|
||||
|
||||
# Handle both GenericResponseOutputItem and dict
|
||||
if isinstance(output_item, GenericResponseOutputItem):
|
||||
content = output_item.content
|
||||
content: Optional[Union[List[OutputText], List[dict]]] = None
|
||||
if isinstance(output_item, BaseModel):
|
||||
try:
|
||||
generic_response_output_item = GenericResponseOutputItem.model_validate(
|
||||
output_item.model_dump()
|
||||
)
|
||||
if generic_response_output_item.content:
|
||||
content = generic_response_output_item.content
|
||||
except Exception:
|
||||
return
|
||||
elif isinstance(output_item, dict):
|
||||
content = output_item.get("content", [])
|
||||
else:
|
||||
|
|
|
|||
|
|
@ -50,12 +50,13 @@ class OpenAITextToSpeechHandler(BaseTranslation):
|
|||
return data
|
||||
|
||||
if isinstance(input_text, str):
|
||||
guardrailed_texts, _ = await guardrail_to_apply.apply_guardrail(
|
||||
texts=[input_text],
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": [input_text]},
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
data["input"] = guardrailed_texts[0] if guardrailed_texts else input_text
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
|
|||
|
|
@ -88,12 +88,13 @@ class OpenAIAudioTranscriptionHandler(BaseTranslation):
|
|||
if user_metadata:
|
||||
request_data["litellm_metadata"] = user_metadata
|
||||
|
||||
guardrailed_texts, _ = await guardrail_to_apply.apply_guardrail(
|
||||
texts=[original_text],
|
||||
guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": [original_text]},
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
)
|
||||
guardrailed_texts = guardrailed_inputs.get("texts", [])
|
||||
response.text = guardrailed_texts[0] if guardrailed_texts else original_text
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
|
|
|
|||
129
litellm/llms/openai_like/README.md
Normal file
129
litellm/llms/openai_like/README.md
Normal file
|
|
@ -0,0 +1,129 @@
|
|||
# JSON-Based OpenAI-Compatible Provider Configuration
|
||||
|
||||
This directory contains the new JSON-based configuration system for OpenAI-compatible providers.
|
||||
|
||||
## Overview
|
||||
|
||||
Instead of creating a full Python module for simple OpenAI-compatible providers, you can now define them in a single JSON file.
|
||||
|
||||
## Files
|
||||
|
||||
- `providers.json` - Configuration file for all JSON-based providers
|
||||
- `json_loader.py` - Loads and parses the JSON configuration
|
||||
- `dynamic_config.py` - Generates Python config classes from JSON
|
||||
- `chat/` - Existing OpenAI-like chat completion handlers
|
||||
|
||||
## Adding a New Provider
|
||||
|
||||
### For Simple OpenAI-Compatible Providers
|
||||
|
||||
Edit `providers.json` and add your provider:
|
||||
|
||||
```json
|
||||
{
|
||||
"your_provider": {
|
||||
"base_url": "https://api.yourprovider.com/v1",
|
||||
"api_key_env": "YOUR_PROVIDER_API_KEY"
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
That's it! The provider will be automatically loaded and available.
|
||||
|
||||
### Optional Configuration Fields
|
||||
|
||||
```json
|
||||
{
|
||||
"your_provider": {
|
||||
"base_url": "https://api.yourprovider.com/v1",
|
||||
"api_key_env": "YOUR_PROVIDER_API_KEY",
|
||||
|
||||
// Optional: Override base_url via environment variable
|
||||
"api_base_env": "YOUR_PROVIDER_API_BASE",
|
||||
|
||||
// Optional: Which base class to use (default: "openai_gpt")
|
||||
"base_class": "openai_gpt", // or "openai_like"
|
||||
|
||||
// Optional: Parameter name mappings
|
||||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
},
|
||||
|
||||
// Optional: Parameter constraints
|
||||
"constraints": {
|
||||
"temperature_max": 1.0,
|
||||
"temperature_min": 0.0,
|
||||
"temperature_min_with_n_gt_1": 0.3
|
||||
},
|
||||
|
||||
// Optional: Special handling flags
|
||||
"special_handling": {
|
||||
"convert_content_list_to_string": true
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Example: PublicAI
|
||||
|
||||
The first JSON-configured provider:
|
||||
|
||||
```json
|
||||
{
|
||||
"publicai": {
|
||||
"base_url": "https://api.publicai.co/v1",
|
||||
"api_key_env": "PUBLICAI_API_KEY",
|
||||
"api_base_env": "PUBLICAI_API_BASE",
|
||||
"base_class": "openai_gpt",
|
||||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
},
|
||||
"special_handling": {
|
||||
"convert_content_list_to_string": true
|
||||
}
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
## Usage
|
||||
|
||||
```python
|
||||
import litellm
|
||||
|
||||
response = litellm.completion(
|
||||
model="publicai/swiss-ai/apertus-8b-instruct",
|
||||
messages=[{"role": "user", "content": "Hello"}],
|
||||
)
|
||||
```
|
||||
|
||||
## Benefits
|
||||
|
||||
- **Simple**: 2-5 lines of JSON vs 100+ lines of Python
|
||||
- **Fast**: Add a provider in 5 minutes
|
||||
- **Safe**: No Python code to mess up
|
||||
- **Consistent**: All providers follow the same pattern
|
||||
- **Maintainable**: Centralized configuration
|
||||
|
||||
## When to Use Python Instead
|
||||
|
||||
Use a Python config class if you need:
|
||||
- Custom authentication (OAuth, rotating tokens, etc.)
|
||||
- Complex request/response transformations
|
||||
- Provider-specific streaming logic
|
||||
- Advanced tool calling transformations
|
||||
|
||||
## Implementation Details
|
||||
|
||||
### How It Works
|
||||
|
||||
1. `json_loader.py` loads `providers.json` on import
|
||||
2. `dynamic_config.py` generates config classes on-demand
|
||||
3. Provider resolution checks JSON registry first
|
||||
4. ProviderConfigManager returns JSON-based configs
|
||||
|
||||
### Integration Points
|
||||
|
||||
The JSON system is integrated at:
|
||||
- `litellm/litellm_core_utils/get_llm_provider_logic.py` - Provider resolution
|
||||
- `litellm/utils.py` - ProviderConfigManager
|
||||
- `litellm/constants.py` - openai_compatible_providers list
|
||||
145
litellm/llms/openai_like/dynamic_config.py
Normal file
145
litellm/llms/openai_like/dynamic_config.py
Normal file
|
|
@ -0,0 +1,145 @@
|
|||
"""
|
||||
Dynamic configuration class generator for JSON-based providers.
|
||||
"""
|
||||
|
||||
from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, overload
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
handle_messages_with_content_list_to_str_conversion,
|
||||
)
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from litellm.llms.openai_like.chat.transformation import OpenAILikeChatConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
from .json_loader import SimpleProviderConfig
|
||||
|
||||
|
||||
def create_config_class(provider: SimpleProviderConfig):
|
||||
"""Generate config class dynamically from JSON configuration"""
|
||||
|
||||
# Choose base class
|
||||
base_class = (
|
||||
OpenAIGPTConfig if provider.base_class == "openai_gpt" else OpenAILikeChatConfig
|
||||
)
|
||||
|
||||
class JSONProviderConfig(base_class):
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: Literal[True]
|
||||
) -> Coroutine[Any, Any, List[AllMessageValues]]:
|
||||
...
|
||||
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self,
|
||||
messages: List[AllMessageValues],
|
||||
model: str,
|
||||
is_async: Literal[False] = False,
|
||||
) -> List[AllMessageValues]:
|
||||
...
|
||||
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: bool = False
|
||||
) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]:
|
||||
"""Transform messages based on special_handling config"""
|
||||
|
||||
# Handle content list to string conversion if configured
|
||||
if provider.special_handling.get("convert_content_list_to_string"):
|
||||
messages = handle_messages_with_content_list_to_str_conversion(messages)
|
||||
|
||||
if is_async:
|
||||
return super()._transform_messages(
|
||||
messages=messages, model=model, is_async=True
|
||||
)
|
||||
else:
|
||||
return super()._transform_messages(
|
||||
messages=messages, model=model, is_async=False
|
||||
)
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self, api_base: Optional[str], api_key: Optional[str]
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""Get API base and key from JSON config"""
|
||||
|
||||
# Resolve base URL
|
||||
resolved_base = api_base
|
||||
if not resolved_base and provider.api_base_env:
|
||||
resolved_base = get_secret_str(provider.api_base_env)
|
||||
if not resolved_base:
|
||||
resolved_base = provider.base_url
|
||||
|
||||
# Resolve API key
|
||||
resolved_key = api_key or get_secret_str(provider.api_key_env)
|
||||
|
||||
return resolved_base, resolved_key
|
||||
|
||||
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:
|
||||
"""Build complete URL for the API endpoint"""
|
||||
if not api_base:
|
||||
api_base = provider.base_url
|
||||
|
||||
if not api_base.endswith("/chat/completions"):
|
||||
api_base = f"{api_base}/chat/completions"
|
||||
|
||||
return api_base
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
"""Get supported OpenAI params from base class"""
|
||||
return super().get_supported_openai_params(model=model)
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""Apply parameter mappings and constraints"""
|
||||
|
||||
supported_params = self.get_supported_openai_params(model)
|
||||
|
||||
# Apply supported params
|
||||
for param, value in non_default_params.items():
|
||||
# Check parameter mappings first
|
||||
if param in provider.param_mappings:
|
||||
optional_params[provider.param_mappings[param]] = value
|
||||
elif param in supported_params:
|
||||
optional_params[param] = value
|
||||
|
||||
# Apply temperature constraints if present
|
||||
if "temperature" in optional_params:
|
||||
temp = optional_params["temperature"]
|
||||
constraints = provider.constraints
|
||||
|
||||
# Clamp to max
|
||||
if "temperature_max" in constraints:
|
||||
temp = min(temp, constraints["temperature_max"])
|
||||
|
||||
# Clamp to min
|
||||
if "temperature_min" in constraints:
|
||||
temp = max(temp, constraints["temperature_min"])
|
||||
|
||||
# Special case: temperature_min_with_n_gt_1
|
||||
if "temperature_min_with_n_gt_1" in constraints:
|
||||
n = optional_params.get("n", 1)
|
||||
if n > 1 and temp < constraints["temperature_min_with_n_gt_1"]:
|
||||
temp = constraints["temperature_min_with_n_gt_1"]
|
||||
|
||||
optional_params["temperature"] = temp
|
||||
|
||||
return optional_params
|
||||
|
||||
@property
|
||||
def custom_llm_provider(self) -> Optional[str]:
|
||||
return provider.slug
|
||||
|
||||
return JSONProviderConfig
|
||||
74
litellm/llms/openai_like/json_loader.py
Normal file
74
litellm/llms/openai_like/json_loader.py
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
"""
|
||||
JSON-based provider configuration loader for OpenAI-compatible providers.
|
||||
"""
|
||||
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Dict, Optional
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
|
||||
class SimpleProviderConfig:
|
||||
"""Simple data class for JSON provider config"""
|
||||
|
||||
def __init__(self, slug: str, data: dict):
|
||||
self.slug = slug
|
||||
self.base_url = data["base_url"]
|
||||
self.api_key_env = data["api_key_env"]
|
||||
self.api_base_env = data.get("api_base_env")
|
||||
self.base_class = data.get("base_class", "openai_gpt")
|
||||
self.param_mappings = data.get("param_mappings", {})
|
||||
self.constraints = data.get("constraints", {})
|
||||
self.special_handling = data.get("special_handling", {})
|
||||
|
||||
|
||||
class JSONProviderRegistry:
|
||||
"""Load providers from JSON once on import"""
|
||||
|
||||
_providers: Dict[str, SimpleProviderConfig] = {}
|
||||
_loaded = False
|
||||
|
||||
@classmethod
|
||||
def load(cls):
|
||||
"""Load providers from JSON configuration file"""
|
||||
if cls._loaded:
|
||||
return
|
||||
|
||||
json_path = Path(__file__).parent / "providers.json"
|
||||
|
||||
if not json_path.exists():
|
||||
# No JSON file yet, that's okay
|
||||
cls._loaded = True
|
||||
return
|
||||
|
||||
try:
|
||||
with open(json_path) as f:
|
||||
data = json.load(f)
|
||||
|
||||
for slug, config in data.items():
|
||||
cls._providers[slug] = SimpleProviderConfig(slug, config)
|
||||
|
||||
cls._loaded = True
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Warning: Failed to load JSON provider configs: {e}")
|
||||
cls._loaded = True
|
||||
|
||||
@classmethod
|
||||
def get(cls, slug: str) -> Optional[SimpleProviderConfig]:
|
||||
"""Get a provider configuration by slug"""
|
||||
return cls._providers.get(slug)
|
||||
|
||||
@classmethod
|
||||
def exists(cls, slug: str) -> bool:
|
||||
"""Check if a provider is defined via JSON"""
|
||||
return slug in cls._providers
|
||||
|
||||
@classmethod
|
||||
def list_providers(cls) -> list:
|
||||
"""List all registered provider slugs"""
|
||||
return list(cls._providers.keys())
|
||||
|
||||
|
||||
# Load on import
|
||||
JSONProviderRegistry.load()
|
||||
14
litellm/llms/openai_like/providers.json
Normal file
14
litellm/llms/openai_like/providers.json
Normal file
|
|
@ -0,0 +1,14 @@
|
|||
{
|
||||
"publicai": {
|
||||
"base_url": "https://api.publicai.co/v1",
|
||||
"api_key_env": "PUBLICAI_API_KEY",
|
||||
"api_base_env": "PUBLICAI_API_BASE",
|
||||
"base_class": "openai_gpt",
|
||||
"param_mappings": {
|
||||
"max_completion_tokens": "max_tokens"
|
||||
},
|
||||
"special_handling": {
|
||||
"convert_content_list_to_string": true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -7,7 +7,9 @@ More information on our website: https://endpoints.ai.cloud.ovh.net
|
|||
from typing import Optional, Union, List
|
||||
|
||||
import httpx
|
||||
from litellm import ModelResponseStream, OpenAIGPTConfig, get_model_info, verbose_logger
|
||||
from litellm.utils import ModelResponseStream, get_model_info
|
||||
from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.llms.ovhcloud.utils import OVHCloudException
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.llms.base_llm.chat.transformation import BaseLLMException
|
||||
|
|
|
|||
|
|
@ -118,8 +118,8 @@ class PassThroughEndpointHandler(BaseTranslation):
|
|||
return data
|
||||
|
||||
# Apply guardrail (pass-through doesn't modify the text, just checks it)
|
||||
await guardrail_to_apply.apply_guardrail(
|
||||
texts=[text_to_check],
|
||||
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": [text_to_check]},
|
||||
request_data=data,
|
||||
input_type="request",
|
||||
logging_obj=litellm_logging_obj,
|
||||
|
|
@ -178,8 +178,8 @@ class PassThroughEndpointHandler(BaseTranslation):
|
|||
request_data["litellm_metadata"] = user_metadata
|
||||
|
||||
# Apply guardrail (pass-through doesn't modify the text, just checks it)
|
||||
await guardrail_to_apply.apply_guardrail(
|
||||
texts=[text_to_check],
|
||||
_guardrailed_inputs = await guardrail_to_apply.apply_guardrail(
|
||||
inputs={"texts": [text_to_check]},
|
||||
request_data=request_data,
|
||||
input_type="response",
|
||||
logging_obj=litellm_logging_obj,
|
||||
|
|
|
|||
|
|
@ -1,114 +0,0 @@
|
|||
"""
|
||||
Translates from OpenAI's `/v1/chat/completions` to PublicAI's `/v1/chat/completions`
|
||||
"""
|
||||
|
||||
from typing import Any, Coroutine, List, Literal, Optional, Tuple, Union, overload
|
||||
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
handle_messages_with_content_list_to_str_conversion,
|
||||
)
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
|
||||
from ...openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
||||
|
||||
class PublicAIChatConfig(OpenAIGPTConfig):
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: Literal[True]
|
||||
) -> Coroutine[Any, Any, List[AllMessageValues]]:
|
||||
...
|
||||
|
||||
@overload
|
||||
def _transform_messages(
|
||||
self,
|
||||
messages: List[AllMessageValues],
|
||||
model: str,
|
||||
is_async: Literal[False] = False,
|
||||
) -> List[AllMessageValues]:
|
||||
...
|
||||
|
||||
def _transform_messages(
|
||||
self, messages: List[AllMessageValues], model: str, is_async: bool = False
|
||||
) -> Union[List[AllMessageValues], Coroutine[Any, Any, List[AllMessageValues]]]:
|
||||
"""
|
||||
PublicAI does not support content in list format.
|
||||
"""
|
||||
messages = handle_messages_with_content_list_to_str_conversion(messages)
|
||||
if is_async:
|
||||
return super()._transform_messages(
|
||||
messages=messages, model=model, is_async=True
|
||||
)
|
||||
else:
|
||||
return super()._transform_messages(
|
||||
messages=messages, model=model, is_async=False
|
||||
)
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self, api_base: Optional[str], api_key: Optional[str]
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
api_base = (
|
||||
api_base
|
||||
or get_secret_str("PUBLICAI_API_BASE")
|
||||
or "https://platform.publicai.co/v1"
|
||||
) # type: ignore
|
||||
dynamic_api_key = api_key or get_secret_str("PUBLICAI_API_KEY")
|
||||
return api_base, dynamic_api_key
|
||||
|
||||
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:
|
||||
"""
|
||||
If api_base is not provided, use the default PublicAI /chat/completions endpoint.
|
||||
"""
|
||||
if not api_base:
|
||||
api_base = "https://platform.publicai.co/v1"
|
||||
|
||||
if not api_base.endswith("/chat/completions"):
|
||||
api_base = f"{api_base}/chat/completions"
|
||||
|
||||
return api_base
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> list:
|
||||
"""
|
||||
Get the supported OpenAI params for PublicAI models
|
||||
|
||||
PublicAI limitations:
|
||||
- functions parameter is not supported (use tools instead)
|
||||
"""
|
||||
excluded_params: List[str] = ["functions"]
|
||||
|
||||
base_openai_params = super().get_supported_openai_params(model=model)
|
||||
final_params: List[str] = []
|
||||
for param in base_openai_params:
|
||||
if param not in excluded_params:
|
||||
final_params.append(param)
|
||||
|
||||
return final_params
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OpenAI parameters to PublicAI parameters
|
||||
"""
|
||||
supported_openai_params = self.get_supported_openai_params(model)
|
||||
for param, value in non_default_params.items():
|
||||
if param == "max_completion_tokens":
|
||||
optional_params["max_tokens"] = value
|
||||
elif param in supported_openai_params:
|
||||
optional_params[param] = value
|
||||
|
||||
return optional_params
|
||||
|
||||
|
|
@ -8,7 +8,8 @@ Docs: https://docs.together.ai/reference/completions-1
|
|||
|
||||
from typing import Optional
|
||||
|
||||
from litellm import get_model_info, verbose_logger
|
||||
from litellm.utils import get_model_info
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
from ..openai.chat.gpt_transformation import OpenAIGPTConfig
|
||||
|
||||
|
|
|
|||
|
|
@ -61,6 +61,10 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
stream=None,
|
||||
auth_header=None,
|
||||
url=default_api_base,
|
||||
model=None,
|
||||
vertex_project=vertex_project or project_id,
|
||||
vertex_location=vertex_location or "us-central1",
|
||||
vertex_api_version="v1",
|
||||
)
|
||||
|
||||
headers = {
|
||||
|
|
@ -166,6 +170,10 @@ class VertexAIBatchPrediction(VertexLLM):
|
|||
stream=None,
|
||||
auth_header=None,
|
||||
url=default_api_base,
|
||||
model=None,
|
||||
vertex_project=vertex_project or project_id,
|
||||
vertex_location=vertex_location or "us-central1",
|
||||
vertex_api_version="v1",
|
||||
)
|
||||
|
||||
headers = {
|
||||
|
|
|
|||
|
|
@ -5,7 +5,8 @@ from typing import Any, Dict, List, Literal, Optional, Set, Tuple, Union, get_ty
|
|||
import httpx
|
||||
|
||||
import litellm
|
||||
from litellm import supports_response_schema, supports_system_messages, verbose_logger
|
||||
from litellm.utils import supports_response_schema, supports_system_messages
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.constants import DEFAULT_MAX_RECURSE_DEPTH
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_defs
|
||||
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
|
||||
|
|
@ -31,9 +32,12 @@ class VertexAIModelRoute(str, Enum):
|
|||
PARTNER_MODELS = "partner_models"
|
||||
GEMINI = "gemini"
|
||||
GEMMA = "gemma"
|
||||
BGE = "bge"
|
||||
MODEL_GARDEN = "model_garden"
|
||||
NON_GEMINI = "non_gemini"
|
||||
OPENAI_COMPATIBLE = "openai"
|
||||
|
||||
VERTEX_AI_MODEL_ROUTES = [f"{route.value}/" for route in VertexAIModelRoute]
|
||||
|
||||
def get_vertex_ai_model_route(
|
||||
model: str, litellm_params: Optional[dict] = None
|
||||
|
|
@ -60,6 +64,9 @@ def get_vertex_ai_model_route(
|
|||
|
||||
>>> get_vertex_ai_model_route("openai/gpt-oss-120b")
|
||||
VertexAIModelRoute.MODEL_GARDEN
|
||||
|
||||
>>> get_vertex_ai_model_route("1234567890", {"api_base": "http://10.96.32.8"})
|
||||
VertexAIModelRoute.GEMINI # Numeric endpoints with api_base use HTTP path
|
||||
"""
|
||||
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
|
||||
VertexAIPartnerModels,
|
||||
|
|
@ -69,11 +76,20 @@ def get_vertex_ai_model_route(
|
|||
if litellm_params and litellm_params.get("base_model") is not None:
|
||||
if "gemini" in litellm_params["base_model"]:
|
||||
return VertexAIModelRoute.GEMINI
|
||||
|
||||
|
||||
# Check if numeric endpoint ID with custom api_base (PSC endpoint)
|
||||
# Route to GEMINI (HTTP path) to support PSC endpoints properly
|
||||
if model.isdigit() and litellm_params and litellm_params.get("api_base"):
|
||||
return VertexAIModelRoute.GEMINI
|
||||
|
||||
# Check for partner models (llama, mistral, claude, etc.)
|
||||
if VertexAIPartnerModels.is_vertex_partner_model(model=model):
|
||||
return VertexAIModelRoute.PARTNER_MODELS
|
||||
|
||||
|
||||
# Check for BGE models
|
||||
if "bge/" in model or "bge" in model.lower():
|
||||
return VertexAIModelRoute.BGE
|
||||
|
||||
# Check for gemma models
|
||||
if "gemma/" in model:
|
||||
return VertexAIModelRoute.GEMMA
|
||||
|
|
@ -136,6 +152,69 @@ all_gemini_url_modes = Literal[
|
|||
]
|
||||
|
||||
|
||||
def get_vertex_base_model_name(model: str) -> str:
|
||||
"""
|
||||
Strip routing prefixes from model name for PSC/endpoint URL construction.
|
||||
|
||||
Patterns like "bge/", "gemma/", "openai/" are used for internal routing but
|
||||
should not appear in the actual endpoint URL. Routing prefixes are derived
|
||||
from VertexAIModelRoute enum values.
|
||||
|
||||
Args:
|
||||
model: The model name with potential prefix (e.g., "bge/123456", "gemma/gemma-3-12b-it")
|
||||
|
||||
Returns:
|
||||
str: The model name without routing prefix (e.g., "123456", "gemma-3-12b-it")
|
||||
|
||||
Examples:
|
||||
>>> get_vertex_base_model_name("bge/378943383978115072")
|
||||
"378943383978115072"
|
||||
|
||||
>>> get_vertex_base_model_name("gemma/gemma-3-12b-it")
|
||||
"gemma-3-12b-it"
|
||||
|
||||
>>> get_vertex_base_model_name("openai/gpt-oss-120b")
|
||||
"gpt-oss-120b"
|
||||
|
||||
>>> get_vertex_base_model_name("1234567890")
|
||||
"1234567890"
|
||||
"""
|
||||
# Derive routing prefixes from VertexAIModelRoute enum
|
||||
# Map specific routes to their prefixes (some routes like PARTNER_MODELS, GEMINI don't have prefixes)
|
||||
for route in VERTEX_AI_MODEL_ROUTES:
|
||||
if model.startswith(route):
|
||||
return model.replace(route, "", 1)
|
||||
|
||||
return model
|
||||
|
||||
|
||||
def _get_embedding_url(
|
||||
model: str,
|
||||
vertex_project: Optional[str],
|
||||
vertex_location: Optional[str],
|
||||
vertex_api_version: Literal["v1", "v1beta1"],
|
||||
) -> Tuple[str, str]:
|
||||
"""
|
||||
Get URL for embedding models.
|
||||
|
||||
Handles special patterns:
|
||||
- bge/endpoint_id -> strips to endpoint_id for endpoints/ routing
|
||||
- numeric model -> routes to endpoints/
|
||||
- regular model -> routes to publishers/google/models/
|
||||
"""
|
||||
endpoint = "predict"
|
||||
|
||||
# Strip routing prefixes (bge/, gemma/, etc.) for endpoint URL construction
|
||||
model = get_vertex_base_model_name(model=model)
|
||||
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
|
||||
if model.isdigit():
|
||||
# https://us-central1-aiplatform.googleapis.com/v1/projects/$PROJECT_ID/locations/us-central1/endpoints/$ENDPOINT_ID:predict
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
|
||||
|
||||
return url, endpoint
|
||||
|
||||
|
||||
def _get_vertex_url(
|
||||
mode: all_gemini_url_modes,
|
||||
model: str,
|
||||
|
|
@ -148,6 +227,7 @@ def _get_vertex_url(
|
|||
endpoint: Optional[str] = None
|
||||
|
||||
model = litellm.VertexGeminiConfig.get_model_for_vertex_ai_url(model=model)
|
||||
|
||||
if mode == "chat":
|
||||
### SET RUNTIME ENDPOINT ###
|
||||
endpoint = "generateContent"
|
||||
|
|
@ -172,11 +252,12 @@ def _get_vertex_url(
|
|||
if stream is True:
|
||||
url += "?alt=sse"
|
||||
elif mode == "embedding":
|
||||
endpoint = "predict"
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
|
||||
if model.isdigit():
|
||||
# https://us-central1-aiplatform.googleapis.com/v1/projects/$PROJECT_ID/locations/us-central1/endpoints/$ENDPOINT_ID:predict
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/{vertex_api_version}/projects/{vertex_project}/locations/{vertex_location}/endpoints/{model}:{endpoint}"
|
||||
return _get_embedding_url(
|
||||
model=model,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_api_version=vertex_api_version,
|
||||
)
|
||||
elif mode == "image_generation":
|
||||
endpoint = "predict"
|
||||
url = f"https://{vertex_location}-aiplatform.googleapis.com/v1/projects/{vertex_project}/locations/{vertex_location}/publishers/google/models/{model}:{endpoint}"
|
||||
|
|
@ -199,18 +280,24 @@ def _get_gemini_url(
|
|||
stream: Optional[bool],
|
||||
gemini_api_key: Optional[str],
|
||||
) -> Tuple[str, str]:
|
||||
from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import (
|
||||
VertexGeminiConfig,
|
||||
)
|
||||
|
||||
_gemini_model_name = "models/{}".format(model)
|
||||
api_version = "v1alpha" if VertexGeminiConfig._is_gemini_3_or_newer(model) else "v1beta"
|
||||
|
||||
if mode == "chat":
|
||||
endpoint = "generateContent"
|
||||
if stream is True:
|
||||
endpoint = "streamGenerateContent"
|
||||
url = "https://generativelanguage.googleapis.com/v1beta/{}:{}?key={}&alt=sse".format(
|
||||
_gemini_model_name, endpoint, gemini_api_key
|
||||
url = "https://generativelanguage.googleapis.com/{}/{}:{}?key={}&alt=sse".format(
|
||||
api_version, _gemini_model_name, endpoint, gemini_api_key
|
||||
)
|
||||
else:
|
||||
url = (
|
||||
"https://generativelanguage.googleapis.com/v1beta/{}:{}?key={}".format(
|
||||
_gemini_model_name, endpoint, gemini_api_key
|
||||
"https://generativelanguage.googleapis.com/{}/{}:{}?key={}".format(
|
||||
api_version, _gemini_model_name, endpoint, gemini_api_key
|
||||
)
|
||||
)
|
||||
elif mode == "embedding":
|
||||
|
|
@ -862,4 +949,4 @@ class VertexAITokenCounter(BaseTokenCounter):
|
|||
original_response=result,
|
||||
)
|
||||
|
||||
return None
|
||||
return None
|
||||
|
|
@ -85,6 +85,10 @@ class ContextCachingEndpoints(VertexBase):
|
|||
stream=None,
|
||||
auth_header=auth_header,
|
||||
url=url,
|
||||
model=None,
|
||||
vertex_project=vertex_project,
|
||||
vertex_location=vertex_location,
|
||||
vertex_api_version="v1beta1" if custom_llm_provider == "vertex_ai_beta" else "v1",
|
||||
)
|
||||
|
||||
def check_cache(
|
||||
|
|
|
|||
|
|
@ -5,7 +5,7 @@ Why separate file? Make it easy to see how transformation works
|
|||
"""
|
||||
|
||||
import os
|
||||
from typing import TYPE_CHECKING, List, Literal, Optional, Tuple, Union, cast
|
||||
from typing import TYPE_CHECKING, Dict, List, Literal, Optional, Tuple, Union, cast
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
|
|
@ -28,7 +28,6 @@ from litellm.types.files import (
|
|||
get_file_type_from_extension,
|
||||
is_gemini_1_5_accepted_file_type,
|
||||
)
|
||||
from litellm.types.utils import LlmProviders
|
||||
from litellm.types.llms.openai import (
|
||||
AllMessageValues,
|
||||
ChatCompletionAssistantMessage,
|
||||
|
|
@ -48,7 +47,7 @@ from litellm.types.llms.vertex_ai import (
|
|||
ToolConfig,
|
||||
Tools,
|
||||
)
|
||||
from litellm.types.utils import GenericImageParsingChunk
|
||||
from litellm.types.utils import GenericImageParsingChunk, LlmProviders
|
||||
|
||||
from ..common_utils import (
|
||||
_check_text_in_content,
|
||||
|
|
@ -64,24 +63,21 @@ else:
|
|||
LiteLLMLoggingObj = Any
|
||||
|
||||
|
||||
def _map_openai_detail_to_media_resolution(
|
||||
def _convert_detail_to_media_resolution_enum(
|
||||
detail: Optional[str],
|
||||
) -> Optional[Literal["low", "medium", "high"]]:
|
||||
"""
|
||||
Map OpenAI's "detail" parameter to Gemini's "media_resolution" parameter.
|
||||
"""
|
||||
) -> Optional[Dict[str, str]]:
|
||||
if detail == "low":
|
||||
return "low"
|
||||
return {"level": "MEDIA_RESOLUTION_LOW"}
|
||||
elif detail == "high":
|
||||
return "high"
|
||||
# "auto" or None means let the model decide, so we don't set media_resolution
|
||||
return {"level": "MEDIA_RESOLUTION_HIGH"}
|
||||
return None
|
||||
|
||||
|
||||
def _process_gemini_image(
|
||||
image_url: str,
|
||||
format: Optional[str] = None,
|
||||
media_resolution: Optional[Literal["low", "medium", "high"]] = None,
|
||||
media_resolution_enum: Optional[Dict[str, str]] = None,
|
||||
model: Optional[str] = None,
|
||||
) -> PartType:
|
||||
"""
|
||||
Given an image URL, return the appropriate PartType for Gemini
|
||||
|
|
@ -105,31 +101,43 @@ def _process_gemini_image(
|
|||
else:
|
||||
mime_type = format
|
||||
file_data = FileDataType(mime_type=mime_type, file_uri=image_url)
|
||||
|
||||
return PartType(file_data=file_data)
|
||||
part: PartType = {"file_data": file_data}
|
||||
|
||||
if media_resolution_enum is not None and model is not None:
|
||||
from .vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
if VertexGeminiConfig._is_gemini_3_or_newer(model):
|
||||
part_dict = dict(part)
|
||||
part_dict["media_resolution"] = media_resolution_enum
|
||||
return cast(PartType, part_dict)
|
||||
return part
|
||||
elif (
|
||||
"https://" in image_url
|
||||
and (image_type := format or _get_image_mime_type_from_url(image_url))
|
||||
is not None
|
||||
):
|
||||
file_data = FileDataType(file_uri=image_url, mime_type=image_type)
|
||||
return PartType(file_data=file_data)
|
||||
part: PartType = {"file_data": file_data}
|
||||
|
||||
if media_resolution_enum is not None and model is not None:
|
||||
from .vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
if VertexGeminiConfig._is_gemini_3_or_newer(model):
|
||||
part_dict = dict(part)
|
||||
part_dict["media_resolution"] = media_resolution_enum
|
||||
return cast(PartType, part_dict)
|
||||
return part
|
||||
elif "http://" in image_url or "https://" in image_url or "base64" in image_url:
|
||||
# https links for unsupported mime types and base64 images
|
||||
image = convert_to_anthropic_image_obj(image_url, format=format)
|
||||
_blob: BlobType = {"data": image["data"], "mime_type": image["media_type"]}
|
||||
if media_resolution is not None:
|
||||
_blob["media_resolution"] = media_resolution
|
||||
|
||||
# Convert snake_case keys to camelCase for JSON serialization
|
||||
# The TypedDict uses snake_case, but the API expects camelCase
|
||||
_blob_dict = dict(_blob)
|
||||
if "media_resolution" in _blob_dict:
|
||||
_blob_dict["mediaResolution"] = _blob_dict.pop("media_resolution")
|
||||
if "mime_type" in _blob_dict:
|
||||
_blob_dict["mimeType"] = _blob_dict.pop("mime_type")
|
||||
part: PartType = {"inline_data": cast(BlobType, _blob)}
|
||||
|
||||
return PartType(inline_data=cast(BlobType, _blob_dict))
|
||||
if media_resolution_enum is not None and model is not None:
|
||||
from .vertex_and_google_ai_studio_gemini import VertexGeminiConfig
|
||||
if VertexGeminiConfig._is_gemini_3_or_newer(model):
|
||||
part_dict = dict(part)
|
||||
part_dict["media_resolution"] = media_resolution_enum
|
||||
return cast(PartType, part_dict)
|
||||
return part
|
||||
raise Exception("Invalid image received - {}".format(image_url))
|
||||
except Exception as e:
|
||||
raise e
|
||||
|
|
@ -235,18 +243,19 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
|
|||
element = cast(ChatCompletionImageObject, element)
|
||||
img_element = element
|
||||
format: Optional[str] = None
|
||||
media_resolution: Optional[Literal["low", "medium", "high"]] = None
|
||||
media_resolution_enum: Optional[Dict[str, str]] = None
|
||||
if isinstance(img_element["image_url"], dict):
|
||||
image_url = img_element["image_url"]["url"]
|
||||
format = img_element["image_url"].get("format")
|
||||
detail = img_element["image_url"].get("detail")
|
||||
media_resolution = _map_openai_detail_to_media_resolution(detail)
|
||||
media_resolution_enum = _convert_detail_to_media_resolution_enum(detail)
|
||||
else:
|
||||
image_url = img_element["image_url"]
|
||||
_part = _process_gemini_image(
|
||||
image_url=image_url,
|
||||
format=format,
|
||||
media_resolution=media_resolution,
|
||||
media_resolution_enum=media_resolution_enum,
|
||||
model=model,
|
||||
)
|
||||
_parts.append(_part)
|
||||
elif element["type"] == "input_audio":
|
||||
|
|
@ -271,6 +280,7 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
|
|||
_part = _process_gemini_image(
|
||||
image_url=openai_image_str,
|
||||
format=audio_format_modified,
|
||||
model=model,
|
||||
)
|
||||
_parts.append(_part)
|
||||
elif element["type"] == "file":
|
||||
|
|
@ -287,6 +297,7 @@ def _gemini_convert_messages_with_history( # noqa: PLR0915
|
|||
_part = _process_gemini_image(
|
||||
image_url=passed_file,
|
||||
format=format,
|
||||
model=model,
|
||||
)
|
||||
_parts.append(_part)
|
||||
except Exception:
|
||||
|
|
|
|||
|
|
@ -2589,13 +2589,12 @@ class ModelResponseIterator:
|
|||
try:
|
||||
json_chunk = json.loads(chunk)
|
||||
|
||||
except json.JSONDecodeError as e:
|
||||
if (
|
||||
self.sent_first_chunk is False
|
||||
): # only check for accumulated json, on first chunk, else raise error. Prevent real errors from being masked.
|
||||
self.chunk_type = "accumulated_json"
|
||||
return self.handle_accumulated_json_chunk(chunk=chunk)
|
||||
raise e
|
||||
except json.JSONDecodeError:
|
||||
# Switch to accumulation mode for partial JSON chunks
|
||||
# This can happen at any point due to network fragmentation, not just first chunk
|
||||
# See: https://github.com/BerriAI/litellm/issues/16562
|
||||
self.chunk_type = "accumulated_json"
|
||||
return self.handle_accumulated_json_chunk(chunk=chunk)
|
||||
|
||||
if self.sent_first_chunk is False:
|
||||
self.sent_first_chunk = True
|
||||
|
|
|
|||
Some files were not shown because too many files have changed in this diff Show more
Loading…
Add table
Reference in a new issue