mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Merge remote-tracking branch 'origin' into litellm_daily_agent_table
This commit is contained in:
commit
91056c1d7e
63 changed files with 5239 additions and 1789 deletions
231
docs/my-website/docs/anthropic_count_tokens.md
Normal file
231
docs/my-website/docs/anthropic_count_tokens.md
Normal file
|
|
@ -0,0 +1,231 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# /v1/messages/count_tokens
|
||||
|
||||
## Overview
|
||||
|
||||
Anthropic-compatible token counting endpoint. Count tokens for messages before sending them to the model.
|
||||
|
||||
| Feature | Supported | Notes |
|
||||
|---------|-----------|-------|
|
||||
| Cost Tracking | ❌ | Token counting only, no cost incurred |
|
||||
| Logging | ✅ | Works across all integrations |
|
||||
| End-user Tracking | ✅ | |
|
||||
| Supported Providers | Anthropic, Vertex AI (Claude), Bedrock (Claude), Gemini, Vertex AI | Auto-routes to provider-specific token counting APIs |
|
||||
|
||||
## Quick Start
|
||||
|
||||
### 1. Start LiteLLM Proxy
|
||||
|
||||
```bash
|
||||
litellm --config /path/to/config.yaml
|
||||
|
||||
# RUNNING on http://0.0.0.0:4000
|
||||
```
|
||||
|
||||
### 2. Count Tokens
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="curl" label="curl">
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/messages/count_tokens" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, how are you?"}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="python" label="Python (httpx)">
|
||||
|
||||
```python
|
||||
import httpx
|
||||
|
||||
response = httpx.post(
|
||||
"http://localhost:4000/v1/messages/count_tokens",
|
||||
headers={
|
||||
"Content-Type": "application/json",
|
||||
"Authorization": "Bearer sk-1234"
|
||||
},
|
||||
json={
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, how are you?"}
|
||||
]
|
||||
}
|
||||
)
|
||||
|
||||
print(response.json())
|
||||
# {"input_tokens": 14}
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
**Expected Response:**
|
||||
|
||||
```json
|
||||
{
|
||||
"input_tokens": 14
|
||||
}
|
||||
```
|
||||
|
||||
## LiteLLM Proxy Configuration
|
||||
|
||||
Add models to your `config.yaml`:
|
||||
|
||||
```yaml
|
||||
model_list:
|
||||
- model_name: claude-3-5-sonnet
|
||||
litellm_params:
|
||||
model: anthropic/claude-3-5-sonnet-20241022
|
||||
api_key: os.environ/ANTHROPIC_API_KEY
|
||||
|
||||
- model_name: claude-vertex
|
||||
litellm_params:
|
||||
model: vertex_ai/claude-3-5-sonnet-v2@20241022
|
||||
vertex_project: my-project
|
||||
vertex_location: us-east5
|
||||
|
||||
- model_name: claude-bedrock
|
||||
litellm_params:
|
||||
model: bedrock/anthropic.claude-3-5-sonnet-20241022-v2:0
|
||||
aws_region_name: us-west-2
|
||||
```
|
||||
|
||||
## Request Parameters
|
||||
|
||||
| Parameter | Type | Required | Description |
|
||||
|-----------|------|----------|-------------|
|
||||
| `model` | string | ✅ | The model to use for token counting |
|
||||
| `messages` | array | ✅ | Array of messages in Anthropic format |
|
||||
|
||||
### Messages Format
|
||||
|
||||
```json
|
||||
{
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello!"},
|
||||
{"role": "assistant", "content": "Hi there!"},
|
||||
{"role": "user", "content": "How are you?"}
|
||||
]
|
||||
}
|
||||
```
|
||||
|
||||
## Response Format
|
||||
|
||||
```json
|
||||
{
|
||||
"input_tokens": <number>
|
||||
}
|
||||
```
|
||||
|
||||
| Field | Type | Description |
|
||||
|-------|------|-------------|
|
||||
| `input_tokens` | integer | Number of tokens in the input messages |
|
||||
|
||||
## Supported Providers
|
||||
|
||||
The `/v1/messages/count_tokens` endpoint automatically routes to the appropriate provider-specific token counting API:
|
||||
|
||||
| Provider | Token Counting Method |
|
||||
|----------|----------------------|
|
||||
| Anthropic | [Anthropic Token Counting API](https://docs.anthropic.com/en/docs/build-with-claude/token-counting) |
|
||||
| Vertex AI (Claude) | Vertex AI Partner Models Token Counter |
|
||||
| Bedrock (Claude) | AWS Bedrock CountTokens API |
|
||||
| Gemini | Google AI Studio countTokens API |
|
||||
| Vertex AI (Gemini) | Vertex AI countTokens API |
|
||||
|
||||
## Examples
|
||||
|
||||
### Count Tokens with System Message
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/messages/count_tokens" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "You are a helpful assistant. Please help me write a haiku about programming."}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
### Count Tokens for Multi-turn Conversation
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/messages/count_tokens" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "What is the capital of France?"},
|
||||
{"role": "assistant", "content": "The capital of France is Paris."},
|
||||
{"role": "user", "content": "What is its population?"}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
### Using with Vertex AI Claude
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/messages/count_tokens" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "claude-vertex",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, world!"}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
### Using with Bedrock Claude
|
||||
|
||||
```bash
|
||||
curl -X POST "http://localhost:4000/v1/messages/count_tokens" \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-d '{
|
||||
"model": "claude-bedrock",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, world!"}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
## Comparison with Anthropic Passthrough
|
||||
|
||||
LiteLLM provides two ways to count tokens:
|
||||
|
||||
| Endpoint | Description | Use Case |
|
||||
|----------|-------------|----------|
|
||||
| `/v1/messages/count_tokens` | LiteLLM's Anthropic-compatible endpoint | Works with all supported providers (Anthropic, Vertex AI, Bedrock, etc.) |
|
||||
| `/anthropic/v1/messages/count_tokens` | [Pass-through to Anthropic API](./pass_through/anthropic_completion.md#example-2-token-counting-api) | Direct Anthropic API access with native headers |
|
||||
|
||||
### Pass-through Example
|
||||
|
||||
For direct Anthropic API access with full native headers:
|
||||
|
||||
```bash
|
||||
curl --request POST \
|
||||
--url http://0.0.0.0:4000/anthropic/v1/messages/count_tokens \
|
||||
--header "x-api-key: $LITELLM_API_KEY" \
|
||||
--header "anthropic-version: 2023-06-01" \
|
||||
--header "anthropic-beta: token-counting-2024-11-01" \
|
||||
--header "content-type: application/json" \
|
||||
--data '{
|
||||
"model": "claude-3-5-sonnet-20241022",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, world"}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
|
@ -13,36 +13,36 @@ https://github.com/BerriAI/litellm
|
|||
- Track spend & set budgets per project [LiteLLM Proxy Server](https://docs.litellm.ai/docs/simple_proxy)
|
||||
|
||||
## How to use LiteLLM
|
||||
You can use litellm through either:
|
||||
1. [LiteLLM Proxy Server](#litellm-proxy-server-llm-gateway) - Server (LLM Gateway) to call 100+ LLMs, load balance, cost tracking across projects
|
||||
2. [LiteLLM python SDK](#basic-usage) - Python Client to call 100+ LLMs, load balance, cost tracking
|
||||
|
||||
### **When to use LiteLLM Proxy Server (LLM Gateway)**
|
||||
You can use LiteLLM through either the Proxy Server or Python SDK. Both gives you a unified interface to access multiple LLMs (100+ LLMs). Choose the option that best fits your needs:
|
||||
|
||||
:::tip
|
||||
<table style={{width: '100%', tableLayout: 'fixed'}}>
|
||||
<thead>
|
||||
<tr>
|
||||
<th style={{width: '14%'}}></th>
|
||||
<th style={{width: '43%'}}><strong><a href="#litellm-proxy-server-llm-gateway">LiteLLM Proxy Server</a></strong></th>
|
||||
<th style={{width: '43%'}}><strong><a href="#basic-usage">LiteLLM Python SDK</a></strong></th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody>
|
||||
<tr>
|
||||
<td style={{width: '14%'}}><strong>Use Case</strong></td>
|
||||
<td style={{width: '43%'}}>Central service (LLM Gateway) to access multiple LLMs</td>
|
||||
<td style={{width: '43%'}}>Use LiteLLM directly in your Python code</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{width: '14%'}}><strong>Who Uses It?</strong></td>
|
||||
<td style={{width: '43%'}}>Gen AI Enablement / ML Platform Teams</td>
|
||||
<td style={{width: '43%'}}>Developers building LLM projects</td>
|
||||
</tr>
|
||||
<tr>
|
||||
<td style={{width: '14%'}}><strong>Key Features</strong></td>
|
||||
<td style={{width: '43%'}}>• Centralized API gateway with authentication & authorization<br />• Multi-tenant cost tracking and spend management per project/user<br />• Per-project customization (logging, guardrails, caching)<br />• Virtual keys for secure access control<br />• Admin dashboard UI for monitoring and management</td>
|
||||
<td style={{width: '43%'}}>• Direct Python library integration in your codebase<br />• Router with retry/fallback logic across multiple deployments (e.g. Azure/OpenAI) - <a href="https://docs.litellm.ai/docs/routing">Router</a><br />• Application-level load balancing and cost tracking<br />• Exception handling with OpenAI-compatible errors<br />• Observability callbacks (Lunary, MLflow, Langfuse, etc.)</td>
|
||||
</tr>
|
||||
</tbody>
|
||||
</table>
|
||||
|
||||
Use LiteLLM Proxy Server if you want a **central service (LLM Gateway) to access multiple LLMs**
|
||||
|
||||
Typically used by Gen AI Enablement / ML PLatform Teams
|
||||
|
||||
:::
|
||||
|
||||
- LiteLLM Proxy gives you a unified interface to access multiple LLMs (100+ LLMs)
|
||||
- Track LLM Usage and setup guardrails
|
||||
- Customize Logging, Guardrails, Caching per project
|
||||
|
||||
### **When to use LiteLLM Python SDK**
|
||||
|
||||
:::tip
|
||||
|
||||
Use LiteLLM Python SDK if you want to use LiteLLM in your **python code**
|
||||
|
||||
Typically used by developers building llm projects
|
||||
|
||||
:::
|
||||
|
||||
- LiteLLM SDK gives you a unified interface to access multiple LLMs (100+ LLMs)
|
||||
- Retry/fallback logic across multiple deployments (e.g. Azure/OpenAI) - [Router](https://docs.litellm.ai/docs/routing)
|
||||
|
||||
## **LiteLLM Python SDK**
|
||||
|
||||
|
|
|
|||
240
docs/my-website/docs/providers/langgraph.md
Normal file
240
docs/my-website/docs/providers/langgraph.md
Normal file
|
|
@ -0,0 +1,240 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# LangGraph
|
||||
|
||||
Call LangGraph agents through LiteLLM using the OpenAI chat completions format.
|
||||
|
||||
| Property | Details |
|
||||
|----------|---------|
|
||||
| Description | LangGraph is a framework for building stateful, multi-actor applications with LLMs. LiteLLM supports calling LangGraph agents via their streaming and non-streaming endpoints. |
|
||||
| Provider Route on LiteLLM | `langgraph/{agent_id}` |
|
||||
| Provider Doc | [LangGraph Platform ↗](https://langchain-ai.github.io/langgraph/cloud/quick_start/) |
|
||||
|
||||
**Prerequisites:** You need a running LangGraph server. See [Setting Up a Local LangGraph Server](#setting-up-a-local-langgraph-server) below.
|
||||
|
||||
## Quick Start
|
||||
|
||||
### Model Format
|
||||
|
||||
```shell showLineNumbers title="Model Format"
|
||||
langgraph/{agent_id}
|
||||
```
|
||||
|
||||
**Example:**
|
||||
- `langgraph/agent` - calls the default agent
|
||||
|
||||
### LiteLLM Python SDK
|
||||
|
||||
```python showLineNumbers title="Basic LangGraph Completion"
|
||||
import litellm
|
||||
|
||||
response = litellm.completion(
|
||||
model="langgraph/agent",
|
||||
messages=[
|
||||
{"role": "user", "content": "What is 25 * 4?"}
|
||||
],
|
||||
api_base="http://localhost:2024",
|
||||
)
|
||||
|
||||
print(response.choices[0].message.content)
|
||||
```
|
||||
|
||||
```python showLineNumbers title="Streaming LangGraph Response"
|
||||
import litellm
|
||||
|
||||
response = litellm.completion(
|
||||
model="langgraph/agent",
|
||||
messages=[
|
||||
{"role": "user", "content": "What is the weather in Tokyo?"}
|
||||
],
|
||||
api_base="http://localhost:2024",
|
||||
stream=True,
|
||||
)
|
||||
|
||||
for chunk in response:
|
||||
if chunk.choices[0].delta.content:
|
||||
print(chunk.choices[0].delta.content, end="")
|
||||
```
|
||||
|
||||
### LiteLLM Proxy
|
||||
|
||||
#### 1. Configure your model in config.yaml
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="config-yaml" label="config.yaml">
|
||||
|
||||
```yaml showLineNumbers title="LiteLLM Proxy Configuration"
|
||||
model_list:
|
||||
- model_name: langgraph-agent
|
||||
litellm_params:
|
||||
model: langgraph/agent
|
||||
api_base: http://localhost:2024
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
#### 2. Start the LiteLLM Proxy
|
||||
|
||||
```bash showLineNumbers title="Start LiteLLM Proxy"
|
||||
litellm --config config.yaml
|
||||
```
|
||||
|
||||
#### 3. Make requests to your LangGraph agent
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="curl" label="Curl">
|
||||
|
||||
```bash showLineNumbers title="Basic Request"
|
||||
curl http://localhost:4000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer $LITELLM_API_KEY" \
|
||||
-d '{
|
||||
"model": "langgraph-agent",
|
||||
"messages": [
|
||||
{"role": "user", "content": "What is 25 * 4?"}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
```bash showLineNumbers title="Streaming Request"
|
||||
curl http://localhost:4000/v1/chat/completions \
|
||||
-H "Content-Type: application/json" \
|
||||
-H "Authorization: Bearer $LITELLM_API_KEY" \
|
||||
-d '{
|
||||
"model": "langgraph-agent",
|
||||
"messages": [
|
||||
{"role": "user", "content": "What is the weather in Tokyo?"}
|
||||
],
|
||||
"stream": true
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
|
||||
<TabItem value="openai-sdk" label="OpenAI Python SDK">
|
||||
|
||||
```python showLineNumbers title="Using OpenAI SDK with LiteLLM Proxy"
|
||||
from openai import OpenAI
|
||||
|
||||
client = OpenAI(
|
||||
base_url="http://localhost:4000",
|
||||
api_key="your-litellm-api-key"
|
||||
)
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model="langgraph-agent",
|
||||
messages=[
|
||||
{"role": "user", "content": "What is 25 * 4?"}
|
||||
]
|
||||
)
|
||||
|
||||
print(response.choices[0].message.content)
|
||||
```
|
||||
|
||||
```python showLineNumbers title="Streaming with OpenAI SDK"
|
||||
from openai import OpenAI
|
||||
|
||||
client = OpenAI(
|
||||
base_url="http://localhost:4000",
|
||||
api_key="your-litellm-api-key"
|
||||
)
|
||||
|
||||
stream = client.chat.completions.create(
|
||||
model="langgraph-agent",
|
||||
messages=[
|
||||
{"role": "user", "content": "What is the weather in Tokyo?"}
|
||||
],
|
||||
stream=True
|
||||
)
|
||||
|
||||
for chunk in stream:
|
||||
if chunk.choices[0].delta.content is not None:
|
||||
print(chunk.choices[0].delta.content, end="")
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Environment Variables
|
||||
|
||||
| Variable | Description |
|
||||
|----------|-------------|
|
||||
| `LANGGRAPH_API_BASE` | Base URL of your LangGraph server (default: `http://localhost:2024`) |
|
||||
| `LANGGRAPH_API_KEY` | Optional API key for authentication |
|
||||
|
||||
## Supported Parameters
|
||||
|
||||
| Parameter | Type | Description |
|
||||
|-----------|------|-------------|
|
||||
| `model` | string | The agent ID in format `langgraph/{agent_id}` |
|
||||
| `messages` | array | Chat messages in OpenAI format |
|
||||
| `stream` | boolean | Enable streaming responses |
|
||||
| `api_base` | string | LangGraph server URL |
|
||||
| `api_key` | string | Optional API key |
|
||||
|
||||
|
||||
## Setting Up a Local LangGraph Server
|
||||
|
||||
Before using LiteLLM with LangGraph, you need a running LangGraph server.
|
||||
|
||||
### Prerequisites
|
||||
|
||||
- Python 3.11+
|
||||
- An LLM API key (OpenAI or Google Gemini)
|
||||
|
||||
### 1. Install the LangGraph CLI
|
||||
|
||||
```bash
|
||||
pip install "langgraph-cli[inmem]"
|
||||
```
|
||||
|
||||
### 2. Create a new LangGraph project
|
||||
|
||||
```bash
|
||||
langgraph new my-agent --template new-langgraph-project-python
|
||||
cd my-agent
|
||||
```
|
||||
|
||||
### 3. Install dependencies
|
||||
|
||||
```bash
|
||||
pip install -e .
|
||||
```
|
||||
|
||||
### 4. Set your API key
|
||||
|
||||
```bash
|
||||
echo "OPENAI_API_KEY=your_key_here" > .env
|
||||
```
|
||||
|
||||
### 5. Start the server
|
||||
|
||||
```bash
|
||||
langgraph dev
|
||||
```
|
||||
|
||||
The server will start at `http://localhost:2024`.
|
||||
|
||||
### Verify the server is running
|
||||
|
||||
```bash
|
||||
curl -s --request POST \
|
||||
--url "http://localhost:2024/runs/wait" \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{
|
||||
"assistant_id": "agent",
|
||||
"input": {
|
||||
"messages": [{"role": "human", "content": "Hello!"}]
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
|
||||
|
||||
## Further Reading
|
||||
|
||||
- [LangGraph Platform Documentation](https://langchain-ai.github.io/langgraph/cloud/quick_start/)
|
||||
- [LangGraph GitHub](https://github.com/langchain-ai/langgraph)
|
||||
|
||||
|
|
@ -130,6 +130,17 @@ GENERIC_INCLUDE_CLIENT_ID = "false" # some providers enforce that the client_id
|
|||
GENERIC_SCOPE = "openid profile email" # default scope openid is sometimes not enough to retrieve basic user info like first_name and last_name located in profile scope
|
||||
```
|
||||
|
||||
**Assigning User Roles via SSO**
|
||||
|
||||
Use `GENERIC_USER_ROLE_ATTRIBUTE` to specify which attribute in the SSO token contains the user's role. The role value must be one of the following supported LiteLLM roles:
|
||||
|
||||
- `proxy_admin` - Admin over the platform
|
||||
- `proxy_admin_viewer` - Can login, view all keys, view all spend (read-only)
|
||||
- `internal_user` - Can login, view/create/delete their own keys, view their spend
|
||||
- `internal_user_view_only` - Can login, view their own keys, view their own spend
|
||||
|
||||
Nested attribute paths are supported (e.g., `claims.role` or `attributes.litellm_role`).
|
||||
|
||||
- Set Redirect URI, if your provider requires it
|
||||
- Set a redirect url = `<your proxy base url>/sso/callback`
|
||||
```shell
|
||||
|
|
|
|||
|
|
@ -641,6 +641,7 @@ router_settings:
|
|||
| LANGFUSE_PUBLIC_KEY | Public key for Langfuse authentication
|
||||
| LANGFUSE_RELEASE | Release version of Langfuse integration
|
||||
| LANGFUSE_SECRET_KEY | Secret key for Langfuse authentication
|
||||
| LANGFUSE_PROPAGATE_TRACE_ID | Flag to enable propagating trace ID to Langfuse. Default is False
|
||||
| LANGSMITH_API_KEY | API key for Langsmith platform
|
||||
| LANGSMITH_BASE_URL | Base URL for Langsmith service
|
||||
| LANGSMITH_BATCH_SIZE | Batch size for operations in Langsmith
|
||||
|
|
@ -855,6 +856,8 @@ router_settings:
|
|||
| 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
|
||||
| SPEND_LOG_QUEUE_POLL_INTERVAL | Polling interval in seconds for spend log queue. Default is 2.0
|
||||
| SPEND_LOG_QUEUE_SIZE_THRESHOLD | Threshold for spend log queue size before processing. Default is 100
|
||||
| COROUTINE_CHECKER_MAX_SIZE_IN_MEMORY | Maximum size for CoroutineChecker in-memory cache. Default is 1000
|
||||
| DEFAULT_SHARED_HEALTH_CHECK_TTL | Time-to-live in seconds for cached health check results in shared health check mode. Default is 300 (5 minutes)
|
||||
| DEFAULT_SHARED_HEALTH_CHECK_LOCK_TTL | Time-to-live in seconds for health check lock in shared health check mode. Default is 60 (1 minute)
|
||||
|
|
|
|||
|
|
@ -81,6 +81,13 @@ CMD ["--port", "4000", "--config", "./proxy_server_config.yaml", "--num_workers"
|
|||
export MAX_REQUESTS_BEFORE_RESTART=10000
|
||||
```
|
||||
|
||||
> **Tip:** When using `--max_requests_before_restart`, the `--run_gunicorn` flag is more stable and mature as it uses Gunicorn's battle-tested worker recycling mechanism instead of Uvicorn's implementation.
|
||||
|
||||
```shell
|
||||
# Use Gunicorn for more stable worker recycling
|
||||
CMD ["--port", "4000", "--config", "./proxy_server_config.yaml", "--num_workers", "$(nproc)", "--run_gunicorn", "--max_requests_before_restart", "10000"]
|
||||
```
|
||||
|
||||
|
||||
## 4. Use Redis 'port','host', 'password'. NOT 'redis_url'
|
||||
|
||||
|
|
|
|||
|
|
@ -493,6 +493,7 @@ const sidebars = {
|
|||
]
|
||||
},
|
||||
"anthropic_unified",
|
||||
"anthropic_count_tokens",
|
||||
"moderation",
|
||||
"ocr",
|
||||
{
|
||||
|
|
@ -705,6 +706,7 @@ const sidebars = {
|
|||
"providers/infinity",
|
||||
"providers/jina_ai",
|
||||
"providers/lambda_ai",
|
||||
"providers/langgraph",
|
||||
"providers/lemonade",
|
||||
"providers/llamafile",
|
||||
"providers/lm_studio",
|
||||
|
|
|
|||
|
|
@ -315,6 +315,7 @@ model LiteLLM_SpendLogs {
|
|||
session_id String?
|
||||
status String?
|
||||
mcp_namespaced_tool_name String?
|
||||
agent_id String?
|
||||
proxy_server_request Json? @default("{}")
|
||||
@@index([startTime])
|
||||
@@index([end_user])
|
||||
|
|
|
|||
|
|
@ -1,5 +1,8 @@
|
|||
"""
|
||||
Cost calculator for A2A (Agent-to-Agent) calls.
|
||||
|
||||
Supports dynamic cost parameters that allow platform owners
|
||||
to define custom costs per agent query or per token.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Optional
|
||||
|
|
@ -20,17 +23,81 @@ class A2ACostCalculator:
|
|||
"""
|
||||
Calculate the cost of an A2A send_message call.
|
||||
|
||||
Default is 0.0. In the future, users can configure cost per agent call.
|
||||
Supports multiple cost parameters for platform owners:
|
||||
- cost_per_query: Fixed cost per query
|
||||
- input_cost_per_token + output_cost_per_token: Token-based pricing
|
||||
|
||||
Priority order:
|
||||
1. response_cost - if set directly (backward compatibility)
|
||||
2. cost_per_query - fixed cost per query
|
||||
3. input_cost_per_token + output_cost_per_token - token-based cost
|
||||
4. Default to 0.0
|
||||
|
||||
Args:
|
||||
litellm_logging_obj: The LiteLLM logging object containing call details
|
||||
|
||||
Returns:
|
||||
float: The cost of the A2A 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
|
||||
)
|
||||
model_call_details = litellm_logging_obj.model_call_details
|
||||
|
||||
# Check if user set a custom response cost (backward compatibility)
|
||||
response_cost = model_call_details.get("response_cost", None)
|
||||
if response_cost is not None:
|
||||
return response_cost
|
||||
return float(response_cost)
|
||||
|
||||
# Get litellm_params for cost parameters
|
||||
litellm_params = model_call_details.get("litellm_params", {}) or {}
|
||||
|
||||
# Check for cost_per_query (fixed cost per query)
|
||||
if litellm_params.get("cost_per_query") is not None:
|
||||
return float(litellm_params["cost_per_query"])
|
||||
|
||||
# Check for token-based pricing
|
||||
input_cost_per_token = litellm_params.get("input_cost_per_token")
|
||||
output_cost_per_token = litellm_params.get("output_cost_per_token")
|
||||
|
||||
if input_cost_per_token is not None or output_cost_per_token is not None:
|
||||
return A2ACostCalculator._calculate_token_based_cost(
|
||||
model_call_details=model_call_details,
|
||||
input_cost_per_token=input_cost_per_token,
|
||||
output_cost_per_token=output_cost_per_token,
|
||||
)
|
||||
|
||||
# Default to 0.0 for A2A calls
|
||||
return 0.0
|
||||
|
||||
@staticmethod
|
||||
def _calculate_token_based_cost(
|
||||
model_call_details: dict,
|
||||
input_cost_per_token: Optional[float],
|
||||
output_cost_per_token: Optional[float],
|
||||
) -> float:
|
||||
"""
|
||||
Calculate cost based on token usage and per-token pricing.
|
||||
|
||||
Args:
|
||||
model_call_details: The model call details containing usage
|
||||
input_cost_per_token: Cost per input token (can be None, defaults to 0)
|
||||
output_cost_per_token: Cost per output token (can be None, defaults to 0)
|
||||
|
||||
Returns:
|
||||
float: The calculated cost
|
||||
"""
|
||||
# Get usage from model_call_details
|
||||
usage = model_call_details.get("usage")
|
||||
if usage is None:
|
||||
return 0.0
|
||||
|
||||
# Get token counts
|
||||
prompt_tokens = getattr(usage, "prompt_tokens", 0) or 0
|
||||
completion_tokens = getattr(usage, "completion_tokens", 0) or 0
|
||||
|
||||
# Calculate costs
|
||||
input_cost = prompt_tokens * (float(input_cost_per_token) if input_cost_per_token else 0.0)
|
||||
output_cost = completion_tokens * (float(output_cost_per_token) if output_cost_per_token else 0.0)
|
||||
|
||||
return input_cost + output_cost
|
||||
|
|
|
|||
74
litellm/a2a_protocol/litellm_completion_bridge/README.md
Normal file
74
litellm/a2a_protocol/litellm_completion_bridge/README.md
Normal file
|
|
@ -0,0 +1,74 @@
|
|||
# A2A to LiteLLM Completion Bridge
|
||||
|
||||
Routes A2A protocol requests through `litellm.acompletion`, enabling any LiteLLM-supported provider to be invoked via A2A.
|
||||
|
||||
## Flow
|
||||
|
||||
```
|
||||
A2A Request → Transform → litellm.acompletion → Transform → A2A Response
|
||||
```
|
||||
|
||||
## SDK Usage
|
||||
|
||||
Use the existing `asend_message` and `asend_message_streaming` functions with `litellm_params`:
|
||||
|
||||
```python
|
||||
from litellm.a2a_protocol import asend_message, asend_message_streaming
|
||||
from a2a.types import SendMessageRequest, SendStreamingMessageRequest, MessageSendParams
|
||||
from uuid import uuid4
|
||||
|
||||
# Non-streaming
|
||||
request = SendMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(
|
||||
message={"role": "user", "parts": [{"kind": "text", "text": "Hello!"}], "messageId": uuid4().hex}
|
||||
)
|
||||
)
|
||||
response = await asend_message(
|
||||
request=request,
|
||||
api_base="http://localhost:2024",
|
||||
litellm_params={"custom_llm_provider": "langgraph", "model": "agent"},
|
||||
)
|
||||
|
||||
# Streaming
|
||||
stream_request = SendStreamingMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(
|
||||
message={"role": "user", "parts": [{"kind": "text", "text": "Hello!"}], "messageId": uuid4().hex}
|
||||
)
|
||||
)
|
||||
async for chunk in asend_message_streaming(
|
||||
request=stream_request,
|
||||
api_base="http://localhost:2024",
|
||||
litellm_params={"custom_llm_provider": "langgraph", "model": "agent"},
|
||||
):
|
||||
print(chunk)
|
||||
```
|
||||
|
||||
## Proxy Usage
|
||||
|
||||
Configure an agent with `custom_llm_provider` in `litellm_params`:
|
||||
|
||||
```yaml
|
||||
agents:
|
||||
- agent_name: my-langgraph-agent
|
||||
agent_card_params:
|
||||
name: "LangGraph Agent"
|
||||
url: "http://localhost:2024" # Used as api_base
|
||||
litellm_params:
|
||||
custom_llm_provider: langgraph
|
||||
model: agent
|
||||
```
|
||||
|
||||
When an A2A request hits `/a2a/{agent_id}/message/send`, the bridge:
|
||||
|
||||
1. Detects `custom_llm_provider` in agent's `litellm_params`
|
||||
2. Transforms A2A message → OpenAI messages
|
||||
3. Calls `litellm.acompletion(model="langgraph/agent", api_base="http://localhost:2024")`
|
||||
4. Transforms response → A2A format
|
||||
|
||||
## Classes
|
||||
|
||||
- `A2ACompletionBridgeTransformation` - Static methods for message format conversion
|
||||
- `A2ACompletionBridgeHandler` - Static methods for handling requests (streaming/non-streaming)
|
||||
|
||||
23
litellm/a2a_protocol/litellm_completion_bridge/__init__.py
Normal file
23
litellm/a2a_protocol/litellm_completion_bridge/__init__.py
Normal file
|
|
@ -0,0 +1,23 @@
|
|||
"""
|
||||
A2A to LiteLLM Completion Bridge.
|
||||
|
||||
This module provides transformation between A2A protocol messages and
|
||||
LiteLLM completion API, enabling any LiteLLM-supported provider to be
|
||||
invoked via the A2A protocol.
|
||||
"""
|
||||
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
A2ACompletionBridgeHandler,
|
||||
handle_a2a_completion,
|
||||
handle_a2a_completion_streaming,
|
||||
)
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
||||
A2ACompletionBridgeTransformation,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"A2ACompletionBridgeTransformation",
|
||||
"A2ACompletionBridgeHandler",
|
||||
"handle_a2a_completion",
|
||||
"handle_a2a_completion_streaming",
|
||||
]
|
||||
229
litellm/a2a_protocol/litellm_completion_bridge/handler.py
Normal file
229
litellm/a2a_protocol/litellm_completion_bridge/handler.py
Normal file
|
|
@ -0,0 +1,229 @@
|
|||
"""
|
||||
Handler for A2A to LiteLLM completion bridge.
|
||||
|
||||
Routes A2A requests through litellm.acompletion based on custom_llm_provider.
|
||||
|
||||
A2A Streaming Events (in order):
|
||||
1. Task event (kind: "task") - Initial task creation with status "submitted"
|
||||
2. Status update (kind: "status-update") - Status change to "working"
|
||||
3. Artifact update (kind: "artifact-update") - Content/artifact delivery
|
||||
4. Status update (kind: "status-update") - Final status "completed" with final=true
|
||||
"""
|
||||
|
||||
from typing import Any, AsyncIterator, Dict, Optional
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
||||
A2ACompletionBridgeTransformation,
|
||||
A2AStreamingContext,
|
||||
)
|
||||
|
||||
|
||||
class A2ACompletionBridgeHandler:
|
||||
"""
|
||||
Static methods for handling A2A requests via LiteLLM completion.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
async def handle_non_streaming(
|
||||
request_id: str,
|
||||
params: Dict[str, Any],
|
||||
litellm_params: Dict[str, Any],
|
||||
api_base: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Handle non-streaming A2A request via litellm.acompletion.
|
||||
|
||||
Args:
|
||||
request_id: A2A JSON-RPC request ID
|
||||
params: A2A MessageSendParams containing the message
|
||||
litellm_params: Agent's litellm_params (custom_llm_provider, model, etc.)
|
||||
api_base: API base URL from agent_card_params
|
||||
|
||||
Returns:
|
||||
A2A SendMessageResponse dict
|
||||
"""
|
||||
# Extract message from params
|
||||
message = params.get("message", {})
|
||||
|
||||
# Transform A2A message to OpenAI format
|
||||
openai_messages = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(
|
||||
message
|
||||
)
|
||||
|
||||
# Get completion params
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
model = litellm_params.get("model", "agent")
|
||||
|
||||
# Build full model string if provider specified
|
||||
# Skip prepending if model already starts with the provider prefix
|
||||
if custom_llm_provider and not model.startswith(f"{custom_llm_provider}/"):
|
||||
full_model = f"{custom_llm_provider}/{model}"
|
||||
else:
|
||||
full_model = model
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A completion bridge: model={full_model}, api_base={api_base}"
|
||||
)
|
||||
|
||||
# Call litellm.acompletion
|
||||
response = await litellm.acompletion(
|
||||
model=full_model,
|
||||
messages=openai_messages,
|
||||
api_base=api_base,
|
||||
stream=False,
|
||||
)
|
||||
|
||||
# Transform response to A2A format
|
||||
a2a_response = A2ACompletionBridgeTransformation.openai_response_to_a2a_response(
|
||||
response=response,
|
||||
request_id=request_id,
|
||||
)
|
||||
|
||||
verbose_logger.info(f"A2A completion bridge completed: request_id={request_id}")
|
||||
|
||||
return a2a_response
|
||||
|
||||
@staticmethod
|
||||
async def handle_streaming(
|
||||
request_id: str,
|
||||
params: Dict[str, Any],
|
||||
litellm_params: Dict[str, Any],
|
||||
api_base: Optional[str] = None,
|
||||
) -> AsyncIterator[Dict[str, Any]]:
|
||||
"""
|
||||
Handle streaming A2A request via litellm.acompletion with stream=True.
|
||||
|
||||
Emits proper A2A streaming events:
|
||||
1. Task event (kind: "task") - Initial task with status "submitted"
|
||||
2. Status update (kind: "status-update") - Status "working"
|
||||
3. Artifact update (kind: "artifact-update") - Content delivery
|
||||
4. Status update (kind: "status-update") - Final "completed" status
|
||||
|
||||
Args:
|
||||
request_id: A2A JSON-RPC request ID
|
||||
params: A2A MessageSendParams containing the message
|
||||
litellm_params: Agent's litellm_params (custom_llm_provider, model, etc.)
|
||||
api_base: API base URL from agent_card_params
|
||||
|
||||
Yields:
|
||||
A2A streaming response events
|
||||
"""
|
||||
# Extract message from params
|
||||
message = params.get("message", {})
|
||||
|
||||
# Create streaming context
|
||||
ctx = A2AStreamingContext(
|
||||
request_id=request_id,
|
||||
input_message=message,
|
||||
)
|
||||
|
||||
# Transform A2A message to OpenAI format
|
||||
openai_messages = A2ACompletionBridgeTransformation.a2a_message_to_openai_messages(
|
||||
message
|
||||
)
|
||||
|
||||
# Get completion params
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
model = litellm_params.get("model", "agent")
|
||||
|
||||
# Build full model string if provider specified
|
||||
# Skip prepending if model already starts with the provider prefix
|
||||
if custom_llm_provider and not model.startswith(f"{custom_llm_provider}/"):
|
||||
full_model = f"{custom_llm_provider}/{model}"
|
||||
else:
|
||||
full_model = model
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A completion bridge streaming: model={full_model}, api_base={api_base}"
|
||||
)
|
||||
|
||||
# 1. Emit initial task event (kind: "task", status: "submitted")
|
||||
task_event = A2ACompletionBridgeTransformation.create_task_event(ctx)
|
||||
yield task_event
|
||||
|
||||
# 2. Emit status update (kind: "status-update", status: "working")
|
||||
working_event = A2ACompletionBridgeTransformation.create_status_update_event(
|
||||
ctx=ctx,
|
||||
state="working",
|
||||
final=False,
|
||||
message_text="Processing request...",
|
||||
)
|
||||
yield working_event
|
||||
|
||||
# Call litellm.acompletion with streaming
|
||||
response = await litellm.acompletion(
|
||||
model=full_model,
|
||||
messages=openai_messages,
|
||||
api_base=api_base,
|
||||
stream=True,
|
||||
)
|
||||
|
||||
# 3. Accumulate content and emit artifact update
|
||||
accumulated_text = ""
|
||||
chunk_count = 0
|
||||
async for chunk in response: # type: ignore[union-attr]
|
||||
chunk_count += 1
|
||||
|
||||
# Extract delta content
|
||||
content = ""
|
||||
if chunk is not None and hasattr(chunk, "choices") and chunk.choices:
|
||||
choice = chunk.choices[0]
|
||||
if hasattr(choice, "delta") and choice.delta:
|
||||
content = choice.delta.content or ""
|
||||
|
||||
if content:
|
||||
accumulated_text += content
|
||||
|
||||
# Emit artifact update with accumulated content
|
||||
if accumulated_text:
|
||||
artifact_event = A2ACompletionBridgeTransformation.create_artifact_update_event(
|
||||
ctx=ctx,
|
||||
text=accumulated_text,
|
||||
)
|
||||
yield artifact_event
|
||||
|
||||
# 4. Emit final status update (kind: "status-update", status: "completed", final: true)
|
||||
completed_event = A2ACompletionBridgeTransformation.create_status_update_event(
|
||||
ctx=ctx,
|
||||
state="completed",
|
||||
final=True,
|
||||
)
|
||||
yield completed_event
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A completion bridge streaming completed: request_id={request_id}, chunks={chunk_count}"
|
||||
)
|
||||
|
||||
|
||||
# Convenience functions that delegate to the class methods
|
||||
async def handle_a2a_completion(
|
||||
request_id: str,
|
||||
params: Dict[str, Any],
|
||||
litellm_params: Dict[str, Any],
|
||||
api_base: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""Convenience function for non-streaming A2A completion."""
|
||||
return await A2ACompletionBridgeHandler.handle_non_streaming(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
|
||||
async def handle_a2a_completion_streaming(
|
||||
request_id: str,
|
||||
params: Dict[str, Any],
|
||||
litellm_params: Dict[str, Any],
|
||||
api_base: Optional[str] = None,
|
||||
) -> AsyncIterator[Dict[str, Any]]:
|
||||
"""Convenience function for streaming A2A completion."""
|
||||
async for chunk in A2ACompletionBridgeHandler.handle_streaming(
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
api_base=api_base,
|
||||
):
|
||||
yield chunk
|
||||
286
litellm/a2a_protocol/litellm_completion_bridge/transformation.py
Normal file
286
litellm/a2a_protocol/litellm_completion_bridge/transformation.py
Normal file
|
|
@ -0,0 +1,286 @@
|
|||
"""
|
||||
Transformation utilities for A2A <-> OpenAI message format conversion.
|
||||
|
||||
A2A Message Format:
|
||||
{
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hello!"}],
|
||||
"messageId": "abc123"
|
||||
}
|
||||
|
||||
OpenAI Message Format:
|
||||
{"role": "user", "content": "Hello!"}
|
||||
|
||||
A2A Streaming Events:
|
||||
- Task event (kind: "task") - Initial task creation with status "submitted"
|
||||
- Status update (kind: "status-update") - Status changes (working, completed)
|
||||
- Artifact update (kind: "artifact-update") - Content/artifact delivery
|
||||
"""
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Optional
|
||||
from uuid import uuid4
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
|
||||
class A2AStreamingContext:
|
||||
"""
|
||||
Context holder for A2A streaming state.
|
||||
Tracks task_id, context_id, and message accumulation.
|
||||
"""
|
||||
|
||||
def __init__(self, request_id: str, input_message: Dict[str, Any]):
|
||||
self.request_id = request_id
|
||||
self.task_id = str(uuid4())
|
||||
self.context_id = str(uuid4())
|
||||
self.input_message = input_message
|
||||
self.accumulated_text = ""
|
||||
self.has_emitted_task = False
|
||||
self.has_emitted_working = False
|
||||
|
||||
|
||||
class A2ACompletionBridgeTransformation:
|
||||
"""
|
||||
Static methods for transforming between A2A and OpenAI message formats.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def a2a_message_to_openai_messages(
|
||||
a2a_message: Dict[str, Any],
|
||||
) -> List[Dict[str, str]]:
|
||||
"""
|
||||
Transform an A2A message to OpenAI message format.
|
||||
|
||||
Args:
|
||||
a2a_message: A2A message with role, parts, and messageId
|
||||
|
||||
Returns:
|
||||
List of OpenAI-format messages
|
||||
"""
|
||||
role = a2a_message.get("role", "user")
|
||||
parts = a2a_message.get("parts", [])
|
||||
|
||||
# Map A2A roles to OpenAI roles
|
||||
openai_role = role
|
||||
if role == "user":
|
||||
openai_role = "user"
|
||||
elif role == "assistant":
|
||||
openai_role = "assistant"
|
||||
elif role == "system":
|
||||
openai_role = "system"
|
||||
|
||||
# Extract text content from parts
|
||||
content_parts = []
|
||||
for part in parts:
|
||||
kind = part.get("kind", "")
|
||||
if kind == "text":
|
||||
text = part.get("text", "")
|
||||
content_parts.append(text)
|
||||
|
||||
content = "\n".join(content_parts) if content_parts else ""
|
||||
|
||||
verbose_logger.debug(
|
||||
f"A2A -> OpenAI transform: role={role} -> {openai_role}, content_length={len(content)}"
|
||||
)
|
||||
|
||||
return [{"role": openai_role, "content": content}]
|
||||
|
||||
@staticmethod
|
||||
def openai_response_to_a2a_response(
|
||||
response: Any,
|
||||
request_id: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Transform a LiteLLM ModelResponse to A2A SendMessageResponse format.
|
||||
|
||||
Args:
|
||||
response: LiteLLM ModelResponse object
|
||||
request_id: Original A2A request ID
|
||||
|
||||
Returns:
|
||||
A2A SendMessageResponse dict
|
||||
"""
|
||||
# Extract content from response
|
||||
content = ""
|
||||
if hasattr(response, "choices") and response.choices:
|
||||
choice = response.choices[0]
|
||||
if hasattr(choice, "message") and choice.message:
|
||||
content = choice.message.content or ""
|
||||
|
||||
# Build A2A message
|
||||
a2a_message = {
|
||||
"role": "agent",
|
||||
"parts": [{"kind": "text", "text": content}],
|
||||
"messageId": uuid4().hex,
|
||||
}
|
||||
|
||||
# Build A2A response
|
||||
a2a_response = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"result": {
|
||||
"message": a2a_message,
|
||||
},
|
||||
}
|
||||
|
||||
verbose_logger.debug(
|
||||
f"OpenAI -> A2A transform: content_length={len(content)}"
|
||||
)
|
||||
|
||||
return a2a_response
|
||||
|
||||
@staticmethod
|
||||
def _get_timestamp() -> str:
|
||||
"""Get current timestamp in ISO format with timezone."""
|
||||
return datetime.now(timezone.utc).isoformat()
|
||||
|
||||
@staticmethod
|
||||
def create_task_event(
|
||||
ctx: A2AStreamingContext,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create the initial task event with status 'submitted'.
|
||||
|
||||
This is the first event emitted in an A2A streaming response.
|
||||
"""
|
||||
return {
|
||||
"id": ctx.request_id,
|
||||
"jsonrpc": "2.0",
|
||||
"result": {
|
||||
"contextId": ctx.context_id,
|
||||
"history": [
|
||||
{
|
||||
"contextId": ctx.context_id,
|
||||
"kind": "message",
|
||||
"messageId": ctx.input_message.get("messageId", uuid4().hex),
|
||||
"parts": ctx.input_message.get("parts", []),
|
||||
"role": ctx.input_message.get("role", "user"),
|
||||
"taskId": ctx.task_id,
|
||||
}
|
||||
],
|
||||
"id": ctx.task_id,
|
||||
"kind": "task",
|
||||
"status": {
|
||||
"state": "submitted",
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def create_status_update_event(
|
||||
ctx: A2AStreamingContext,
|
||||
state: str,
|
||||
final: bool = False,
|
||||
message_text: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create a status update event.
|
||||
|
||||
Args:
|
||||
ctx: Streaming context
|
||||
state: Status state ('working', 'completed')
|
||||
final: Whether this is the final event
|
||||
message_text: Optional message text for 'working' status
|
||||
"""
|
||||
status: Dict[str, Any] = {
|
||||
"state": state,
|
||||
"timestamp": A2ACompletionBridgeTransformation._get_timestamp(),
|
||||
}
|
||||
|
||||
# Add message for 'working' status
|
||||
if state == "working" and message_text:
|
||||
status["message"] = {
|
||||
"contextId": ctx.context_id,
|
||||
"kind": "message",
|
||||
"messageId": str(uuid4()),
|
||||
"parts": [{"kind": "text", "text": message_text}],
|
||||
"role": "agent",
|
||||
"taskId": ctx.task_id,
|
||||
}
|
||||
|
||||
return {
|
||||
"id": ctx.request_id,
|
||||
"jsonrpc": "2.0",
|
||||
"result": {
|
||||
"contextId": ctx.context_id,
|
||||
"final": final,
|
||||
"kind": "status-update",
|
||||
"status": status,
|
||||
"taskId": ctx.task_id,
|
||||
},
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def create_artifact_update_event(
|
||||
ctx: A2AStreamingContext,
|
||||
text: str,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Create an artifact update event with content.
|
||||
|
||||
Args:
|
||||
ctx: Streaming context
|
||||
text: The text content for the artifact
|
||||
"""
|
||||
return {
|
||||
"id": ctx.request_id,
|
||||
"jsonrpc": "2.0",
|
||||
"result": {
|
||||
"artifact": {
|
||||
"artifactId": str(uuid4()),
|
||||
"name": "response",
|
||||
"parts": [{"kind": "text", "text": text}],
|
||||
},
|
||||
"contextId": ctx.context_id,
|
||||
"kind": "artifact-update",
|
||||
"taskId": ctx.task_id,
|
||||
},
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def openai_chunk_to_a2a_chunk(
|
||||
chunk: Any,
|
||||
request_id: Optional[str] = None,
|
||||
is_final: bool = False,
|
||||
) -> Optional[Dict[str, Any]]:
|
||||
"""
|
||||
Transform a LiteLLM streaming chunk to A2A streaming format.
|
||||
|
||||
NOTE: This method is deprecated for streaming. Use the event-based
|
||||
methods (create_task_event, create_status_update_event,
|
||||
create_artifact_update_event) instead for proper A2A streaming.
|
||||
|
||||
Args:
|
||||
chunk: LiteLLM ModelResponse chunk
|
||||
request_id: Original A2A request ID
|
||||
is_final: Whether this is the final chunk
|
||||
|
||||
Returns:
|
||||
A2A streaming chunk dict or None if no content
|
||||
"""
|
||||
# Extract delta content
|
||||
content = ""
|
||||
if chunk is not None and hasattr(chunk, "choices") and chunk.choices:
|
||||
choice = chunk.choices[0]
|
||||
if hasattr(choice, "delta") and choice.delta:
|
||||
content = choice.delta.content or ""
|
||||
|
||||
if not content and not is_final:
|
||||
return None
|
||||
|
||||
# Build A2A streaming chunk (legacy format)
|
||||
a2a_chunk = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": request_id,
|
||||
"result": {
|
||||
"message": {
|
||||
"role": "agent",
|
||||
"parts": [{"kind": "text", "text": content}],
|
||||
"messageId": uuid4().hex,
|
||||
},
|
||||
"final": is_final,
|
||||
},
|
||||
}
|
||||
|
||||
return a2a_chunk
|
||||
|
|
@ -5,9 +5,14 @@ Provides standalone functions with @client decorator for LiteLLM logging integra
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import datetime
|
||||
from typing import TYPE_CHECKING, Any, AsyncIterator, Coroutine, Dict, Optional, Union
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.a2a_protocol.streaming_iterator import A2AStreamingIterator
|
||||
from litellm.a2a_protocol.utils import A2ARequestUtils
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
get_async_httpx_client,
|
||||
httpxSpecialProvider,
|
||||
|
|
@ -38,6 +43,49 @@ except ImportError:
|
|||
pass
|
||||
|
||||
|
||||
def _set_usage_on_logging_obj(
|
||||
kwargs: Dict[str, Any],
|
||||
prompt_tokens: int,
|
||||
completion_tokens: int,
|
||||
) -> None:
|
||||
"""
|
||||
Set usage on litellm_logging_obj for standard logging payload.
|
||||
|
||||
Args:
|
||||
kwargs: The kwargs dict containing litellm_logging_obj
|
||||
prompt_tokens: Number of input tokens
|
||||
completion_tokens: Number of output tokens
|
||||
"""
|
||||
litellm_logging_obj = kwargs.get("litellm_logging_obj")
|
||||
if litellm_logging_obj is not None:
|
||||
usage = litellm.Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=prompt_tokens + completion_tokens,
|
||||
)
|
||||
litellm_logging_obj.model_call_details["usage"] = usage
|
||||
|
||||
|
||||
def _set_agent_id_on_logging_obj(
|
||||
kwargs: Dict[str, Any],
|
||||
agent_id: Optional[str],
|
||||
) -> None:
|
||||
"""
|
||||
Set agent_id on litellm_logging_obj for SpendLogs tracking.
|
||||
|
||||
Args:
|
||||
kwargs: The kwargs dict containing litellm_logging_obj
|
||||
agent_id: The A2A agent ID
|
||||
"""
|
||||
if agent_id is None:
|
||||
return
|
||||
|
||||
litellm_logging_obj = kwargs.get("litellm_logging_obj")
|
||||
if litellm_logging_obj is not None:
|
||||
# Set agent_id directly on model_call_details (same pattern as custom_llm_provider)
|
||||
litellm_logging_obj.model_call_details["agent_id"] = agent_id
|
||||
|
||||
|
||||
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.
|
||||
|
|
@ -72,46 +120,105 @@ def _get_a2a_model_info(a2a_client: Any, kwargs: Dict[str, Any]) -> str:
|
|||
|
||||
@client
|
||||
async def asend_message(
|
||||
a2a_client: "A2AClientType",
|
||||
request: "SendMessageRequest",
|
||||
a2a_client: Optional["A2AClientType"] = None,
|
||||
request: Optional["SendMessageRequest"] = None,
|
||||
api_base: Optional[str] = None,
|
||||
litellm_params: Optional[Dict[str, Any]] = None,
|
||||
agent_id: Optional[str] = None,
|
||||
**kwargs: Any,
|
||||
) -> LiteLLMSendMessageResponse:
|
||||
"""
|
||||
Async: Send a message to an A2A agent.
|
||||
|
||||
Uses the @client decorator for LiteLLM logging and tracking.
|
||||
If litellm_params contains custom_llm_provider, routes through the completion bridge.
|
||||
|
||||
Args:
|
||||
a2a_client: An initialized a2a.client.A2AClient instance
|
||||
request: SendMessageRequest from a2a.types
|
||||
a2a_client: An initialized a2a.client.A2AClient instance (optional if using completion bridge)
|
||||
request: SendMessageRequest from a2a.types (optional if using completion bridge with api_base)
|
||||
api_base: API base URL (required for completion bridge, optional for standard A2A)
|
||||
litellm_params: Optional dict with custom_llm_provider, model, etc. for completion bridge
|
||||
agent_id: Optional agent ID for tracking in SpendLogs
|
||||
**kwargs: Additional arguments passed to the client decorator
|
||||
|
||||
Returns:
|
||||
LiteLLMSendMessageResponse (wraps a2a SendMessageResponse with _hidden_params)
|
||||
|
||||
Example:
|
||||
Example (standard A2A):
|
||||
```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,
|
||||
}
|
||||
message={"role": "user", "parts": [{"kind": "text", "text": "Hello!"}], "messageId": uuid4().hex}
|
||||
)
|
||||
)
|
||||
response = await asend_message(a2a_client=a2a_client, request=request)
|
||||
```
|
||||
|
||||
Example (completion bridge with LangGraph):
|
||||
```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(
|
||||
request=request,
|
||||
api_base="http://localhost:2024",
|
||||
litellm_params={"custom_llm_provider": "langgraph", "model": "agent"},
|
||||
)
|
||||
```
|
||||
"""
|
||||
litellm_params = litellm_params or {}
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
|
||||
# Route through completion bridge if custom_llm_provider is set
|
||||
if custom_llm_provider:
|
||||
if request is None:
|
||||
raise ValueError("request is required for completion bridge")
|
||||
# api_base is optional for providers that derive endpoint from model (e.g., bedrock/agentcore)
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A using completion bridge: provider={custom_llm_provider}, api_base={api_base}"
|
||||
)
|
||||
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
A2ACompletionBridgeHandler,
|
||||
)
|
||||
|
||||
# Extract params from request
|
||||
params = request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params)
|
||||
|
||||
response_dict = await A2ACompletionBridgeHandler.handle_non_streaming(
|
||||
request_id=str(request.id),
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
# Convert to LiteLLMSendMessageResponse
|
||||
return LiteLLMSendMessageResponse.from_dict(response_dict)
|
||||
|
||||
# Standard A2A client flow
|
||||
if request is None:
|
||||
raise ValueError("request is required")
|
||||
|
||||
# Create A2A client if not provided but api_base is available
|
||||
if a2a_client is None:
|
||||
if api_base is None:
|
||||
raise ValueError("Either a2a_client or api_base is required for standard A2A flow")
|
||||
a2a_client = await create_a2a_client(base_url=api_base)
|
||||
|
||||
agent_name = _get_a2a_model_info(a2a_client, kwargs)
|
||||
|
||||
verbose_logger.info(f"A2A send_message request_id={request.id}, agent={agent_name}")
|
||||
|
|
@ -123,6 +230,23 @@ async def asend_message(
|
|||
# Wrap in LiteLLM response type for _hidden_params support
|
||||
response = LiteLLMSendMessageResponse.from_a2a_response(a2a_response)
|
||||
|
||||
# Calculate token usage from request and response
|
||||
response_dict = a2a_response.model_dump(mode="json", exclude_none=True)
|
||||
prompt_tokens, completion_tokens, _ = A2ARequestUtils.calculate_usage_from_request_response(
|
||||
request=request,
|
||||
response_dict=response_dict,
|
||||
)
|
||||
|
||||
# Set usage on logging obj for standard logging payload
|
||||
_set_usage_on_logging_obj(
|
||||
kwargs=kwargs,
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
)
|
||||
|
||||
# Set agent_id on logging obj for SpendLogs tracking
|
||||
_set_agent_id_on_logging_obj(kwargs=kwargs, agent_id=agent_id)
|
||||
|
||||
return response
|
||||
|
||||
|
||||
|
|
@ -157,31 +281,141 @@ def send_message(
|
|||
|
||||
|
||||
async def asend_message_streaming(
|
||||
a2a_client: "A2AClientType",
|
||||
request: "SendStreamingMessageRequest",
|
||||
) -> AsyncIterator["SendStreamingMessageResponse"]:
|
||||
a2a_client: Optional["A2AClientType"] = None,
|
||||
request: Optional["SendStreamingMessageRequest"] = None,
|
||||
api_base: Optional[str] = None,
|
||||
litellm_params: Optional[Dict[str, Any]] = None,
|
||||
agent_id: Optional[str] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
proxy_server_request: Optional[Dict[str, Any]] = None,
|
||||
) -> AsyncIterator[Any]:
|
||||
"""
|
||||
Async: Send a streaming message to an A2A agent.
|
||||
|
||||
If litellm_params contains custom_llm_provider, routes through the completion bridge.
|
||||
|
||||
Args:
|
||||
a2a_client: An initialized a2a.client.A2AClient instance
|
||||
a2a_client: An initialized a2a.client.A2AClient instance (optional if using completion bridge)
|
||||
request: SendStreamingMessageRequest from a2a.types
|
||||
api_base: API base URL (required for completion bridge)
|
||||
litellm_params: Optional dict with custom_llm_provider, model, etc. for completion bridge
|
||||
agent_id: Optional agent ID for tracking in SpendLogs
|
||||
metadata: Optional metadata dict (contains user_api_key, user_id, team_id, etc.)
|
||||
proxy_server_request: Optional proxy server request data
|
||||
|
||||
Yields:
|
||||
SendStreamingMessageResponse chunks from the agent
|
||||
|
||||
Example (completion bridge with LangGraph):
|
||||
```python
|
||||
from litellm.a2a_protocol import asend_message_streaming
|
||||
from a2a.types import SendStreamingMessageRequest, MessageSendParams
|
||||
from uuid import uuid4
|
||||
|
||||
request = SendStreamingMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(
|
||||
message={"role": "user", "parts": [{"kind": "text", "text": "Hello!"}], "messageId": uuid4().hex}
|
||||
)
|
||||
)
|
||||
async for chunk in asend_message_streaming(
|
||||
request=request,
|
||||
api_base="http://localhost:2024",
|
||||
litellm_params={"custom_llm_provider": "langgraph", "model": "agent"},
|
||||
):
|
||||
print(chunk)
|
||||
```
|
||||
"""
|
||||
litellm_params = litellm_params or {}
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
|
||||
# Route through completion bridge if custom_llm_provider is set
|
||||
if custom_llm_provider:
|
||||
if request is None:
|
||||
raise ValueError("request is required for completion bridge")
|
||||
# api_base is optional for providers that derive endpoint from model (e.g., bedrock/agentcore)
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A streaming using completion bridge: provider={custom_llm_provider}"
|
||||
)
|
||||
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
A2ACompletionBridgeHandler,
|
||||
)
|
||||
|
||||
# Extract params from request
|
||||
params = request.params.model_dump(mode="json") if hasattr(request.params, "model_dump") else dict(request.params)
|
||||
|
||||
async for chunk in A2ACompletionBridgeHandler.handle_streaming(
|
||||
request_id=str(request.id),
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
api_base=api_base,
|
||||
):
|
||||
yield chunk
|
||||
return
|
||||
|
||||
# Standard A2A client flow
|
||||
if request is None:
|
||||
raise ValueError("request is required")
|
||||
|
||||
# Create A2A client if not provided but api_base is available
|
||||
if a2a_client is None:
|
||||
if api_base is None:
|
||||
raise ValueError("Either a2a_client or api_base is required for standard A2A flow")
|
||||
a2a_client = await create_a2a_client(base_url=api_base)
|
||||
|
||||
verbose_logger.info(f"A2A send_message_streaming request_id={request.id}")
|
||||
|
||||
# Track for logging
|
||||
import datetime
|
||||
|
||||
start_time = datetime.datetime.now()
|
||||
stream = a2a_client.send_message_streaming(request)
|
||||
|
||||
chunk_count = 0
|
||||
async for chunk in stream:
|
||||
chunk_count += 1
|
||||
yield chunk
|
||||
# Build logging object for streaming completion callbacks
|
||||
agent_card = getattr(a2a_client, "_litellm_agent_card", None) or getattr(a2a_client, "agent_card", None)
|
||||
agent_name = getattr(agent_card, "name", "unknown") if agent_card else "unknown"
|
||||
model = f"a2a_agent/{agent_name}"
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A send_message_streaming completed, request_id={request.id}, chunks={chunk_count}"
|
||||
logging_obj = Logging(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "streaming-request"}],
|
||||
stream=False, # complete response logging after stream ends
|
||||
call_type="asend_message_streaming",
|
||||
start_time=start_time,
|
||||
litellm_call_id=str(request.id),
|
||||
function_id=str(request.id),
|
||||
)
|
||||
logging_obj.model = model
|
||||
logging_obj.custom_llm_provider = "a2a_agent"
|
||||
logging_obj.model_call_details["model"] = model
|
||||
logging_obj.model_call_details["custom_llm_provider"] = "a2a_agent"
|
||||
if agent_id:
|
||||
logging_obj.model_call_details["agent_id"] = agent_id
|
||||
|
||||
# Propagate litellm_params for spend logging (includes cost_per_query, etc.)
|
||||
_litellm_params = litellm_params.copy() if litellm_params else {}
|
||||
# Merge metadata into litellm_params.metadata (required for proxy cost tracking)
|
||||
if metadata:
|
||||
_litellm_params["metadata"] = metadata
|
||||
if proxy_server_request:
|
||||
_litellm_params["proxy_server_request"] = proxy_server_request
|
||||
|
||||
logging_obj.litellm_params = _litellm_params
|
||||
logging_obj.optional_params = _litellm_params # used by cost calc
|
||||
logging_obj.model_call_details["litellm_params"] = _litellm_params
|
||||
logging_obj.model_call_details["metadata"] = metadata or {}
|
||||
|
||||
iterator = A2AStreamingIterator(
|
||||
stream=stream,
|
||||
request=request,
|
||||
logging_obj=logging_obj,
|
||||
agent_name=agent_name,
|
||||
)
|
||||
|
||||
async for chunk in iterator:
|
||||
yield chunk
|
||||
|
||||
|
||||
async def create_a2a_client(
|
||||
|
|
@ -296,3 +530,5 @@ async def aget_agent_card(
|
|||
f"Fetched agent card: {agent_card.name if hasattr(agent_card, 'name') else 'unknown'}"
|
||||
)
|
||||
return agent_card
|
||||
|
||||
|
||||
|
|
|
|||
173
litellm/a2a_protocol/streaming_iterator.py
Normal file
173
litellm/a2a_protocol/streaming_iterator.py
Normal file
|
|
@ -0,0 +1,173 @@
|
|||
"""
|
||||
A2A Streaming Iterator with token tracking and logging support.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
from typing import TYPE_CHECKING, Any, AsyncIterator, Dict, List, Optional
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.a2a_protocol.cost_calculator import A2ACostCalculator
|
||||
from litellm.a2a_protocol.utils import A2ARequestUtils
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.thread_pool_executor import executor
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from a2a.types import SendStreamingMessageRequest, SendStreamingMessageResponse
|
||||
|
||||
|
||||
class A2AStreamingIterator:
|
||||
"""
|
||||
Async iterator for A2A streaming responses with token tracking.
|
||||
|
||||
Collects chunks, extracts text, and logs usage on completion.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
stream: AsyncIterator["SendStreamingMessageResponse"],
|
||||
request: "SendStreamingMessageRequest",
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
agent_name: str = "unknown",
|
||||
):
|
||||
self.stream = stream
|
||||
self.request = request
|
||||
self.logging_obj = logging_obj
|
||||
self.agent_name = agent_name
|
||||
self.start_time = datetime.now()
|
||||
|
||||
# Collect chunks for token counting
|
||||
self.chunks: List[Any] = []
|
||||
self.collected_text_parts: List[str] = []
|
||||
self.final_chunk: Optional[Any] = None
|
||||
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self) -> "SendStreamingMessageResponse":
|
||||
try:
|
||||
chunk = await self.stream.__anext__()
|
||||
|
||||
# Store chunk
|
||||
self.chunks.append(chunk)
|
||||
|
||||
# Extract text from chunk for token counting
|
||||
self._collect_text_from_chunk(chunk)
|
||||
|
||||
# Check if this is the final chunk (completed status)
|
||||
if self._is_completed_chunk(chunk):
|
||||
self.final_chunk = chunk
|
||||
|
||||
return chunk
|
||||
|
||||
except StopAsyncIteration:
|
||||
# Stream ended - handle logging
|
||||
if self.final_chunk is None and self.chunks:
|
||||
self.final_chunk = self.chunks[-1]
|
||||
await self._handle_stream_complete()
|
||||
raise
|
||||
|
||||
def _collect_text_from_chunk(self, chunk: Any) -> None:
|
||||
"""Extract text from a streaming chunk and add to collected parts."""
|
||||
try:
|
||||
chunk_dict = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
|
||||
text = A2ARequestUtils.extract_text_from_response(chunk_dict)
|
||||
if text:
|
||||
self.collected_text_parts.append(text)
|
||||
except Exception:
|
||||
verbose_logger.debug("Failed to extract text from A2A streaming chunk")
|
||||
|
||||
def _is_completed_chunk(self, chunk: Any) -> bool:
|
||||
"""Check if chunk indicates stream completion."""
|
||||
try:
|
||||
chunk_dict = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {}
|
||||
result = chunk_dict.get("result", {})
|
||||
if isinstance(result, dict):
|
||||
status = result.get("status", {})
|
||||
if isinstance(status, dict):
|
||||
return status.get("state") == "completed"
|
||||
except Exception:
|
||||
pass
|
||||
return False
|
||||
|
||||
async def _handle_stream_complete(self) -> None:
|
||||
"""Handle logging and token counting when stream completes."""
|
||||
try:
|
||||
end_time = datetime.now()
|
||||
|
||||
# Calculate tokens from collected text
|
||||
input_message = A2ARequestUtils.get_input_message_from_request(self.request)
|
||||
input_text = A2ARequestUtils.extract_text_from_message(input_message)
|
||||
prompt_tokens = A2ARequestUtils.count_tokens(input_text)
|
||||
|
||||
# Use the last (most complete) text from chunks
|
||||
output_text = self.collected_text_parts[-1] if self.collected_text_parts else ""
|
||||
completion_tokens = A2ARequestUtils.count_tokens(output_text)
|
||||
|
||||
total_tokens = prompt_tokens + completion_tokens
|
||||
|
||||
# Create usage object
|
||||
usage = litellm.Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=total_tokens,
|
||||
)
|
||||
|
||||
# Set usage on logging obj
|
||||
self.logging_obj.model_call_details["usage"] = usage
|
||||
# Mark stream flag for downstream callbacks
|
||||
self.logging_obj.model_call_details["stream"] = False
|
||||
|
||||
# Calculate cost using A2ACostCalculator
|
||||
response_cost = A2ACostCalculator.calculate_a2a_cost(self.logging_obj)
|
||||
self.logging_obj.model_call_details["response_cost"] = response_cost
|
||||
|
||||
# Build result for logging
|
||||
result = self._build_logging_result(usage)
|
||||
|
||||
# Call success handlers - they will build standard_logging_object
|
||||
asyncio.create_task(
|
||||
self.logging_obj.async_success_handler(
|
||||
result=result,
|
||||
start_time=self.start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=None,
|
||||
)
|
||||
)
|
||||
|
||||
executor.submit(
|
||||
self.logging_obj.success_handler,
|
||||
result=result,
|
||||
cache_hit=None,
|
||||
start_time=self.start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
verbose_logger.info(
|
||||
f"A2A streaming completed: prompt_tokens={prompt_tokens}, "
|
||||
f"completion_tokens={completion_tokens}, total_tokens={total_tokens}, "
|
||||
f"response_cost={response_cost}"
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.debug(f"Error in A2A streaming completion handler: {e}")
|
||||
|
||||
def _build_logging_result(self, usage: litellm.Usage) -> Dict[str, Any]:
|
||||
"""Build a result dict for logging."""
|
||||
result: Dict[str, Any] = {
|
||||
"id": getattr(self.request, "id", "unknown"),
|
||||
"jsonrpc": "2.0",
|
||||
"usage": usage.model_dump() if hasattr(usage, "model_dump") else dict(usage),
|
||||
}
|
||||
|
||||
# Add final chunk result if available
|
||||
if self.final_chunk:
|
||||
try:
|
||||
chunk_dict = self.final_chunk.model_dump(mode="json", exclude_none=True)
|
||||
result["result"] = chunk_dict.get("result", {})
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return result
|
||||
|
||||
138
litellm/a2a_protocol/utils.py
Normal file
138
litellm/a2a_protocol/utils.py
Normal file
|
|
@ -0,0 +1,138 @@
|
|||
"""
|
||||
Utility functions for A2A protocol.
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Tuple, Union
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from a2a.types import SendMessageRequest, SendStreamingMessageRequest
|
||||
|
||||
|
||||
class A2ARequestUtils:
|
||||
"""Utility class for A2A request/response processing."""
|
||||
|
||||
@staticmethod
|
||||
def extract_text_from_message(message: Any) -> str:
|
||||
"""
|
||||
Extract text content from A2A message parts.
|
||||
|
||||
Args:
|
||||
message: A2A message dict or object with 'parts' containing text parts
|
||||
|
||||
Returns:
|
||||
Concatenated text from all text parts
|
||||
"""
|
||||
if message is None:
|
||||
return ""
|
||||
|
||||
# Handle both dict and object access
|
||||
if isinstance(message, dict):
|
||||
parts = message.get("parts", [])
|
||||
else:
|
||||
parts = getattr(message, "parts", []) or []
|
||||
|
||||
text_parts: List[str] = []
|
||||
for part in parts:
|
||||
if isinstance(part, dict):
|
||||
if part.get("kind") == "text":
|
||||
text_parts.append(part.get("text", ""))
|
||||
else:
|
||||
if getattr(part, "kind", None) == "text":
|
||||
text_parts.append(getattr(part, "text", ""))
|
||||
|
||||
return " ".join(text_parts)
|
||||
|
||||
@staticmethod
|
||||
def extract_text_from_response(response_dict: Dict[str, Any]) -> str:
|
||||
"""
|
||||
Extract text content from A2A response result.
|
||||
|
||||
Args:
|
||||
response_dict: A2A response dict with 'result' containing message
|
||||
|
||||
Returns:
|
||||
Text from response message parts
|
||||
"""
|
||||
result = response_dict.get("result", {})
|
||||
if not isinstance(result, dict):
|
||||
return ""
|
||||
|
||||
message = result.get("message", {})
|
||||
return A2ARequestUtils.extract_text_from_message(message)
|
||||
|
||||
@staticmethod
|
||||
def get_input_message_from_request(
|
||||
request: "Union[SendMessageRequest, SendStreamingMessageRequest]",
|
||||
) -> Any:
|
||||
"""
|
||||
Extract the input message from an A2A request.
|
||||
|
||||
Args:
|
||||
request: The A2A SendMessageRequest or SendStreamingMessageRequest
|
||||
|
||||
Returns:
|
||||
The message object/dict or None
|
||||
"""
|
||||
params = getattr(request, "params", None)
|
||||
if params is None:
|
||||
return None
|
||||
return getattr(params, "message", None)
|
||||
|
||||
@staticmethod
|
||||
def count_tokens(text: str) -> int:
|
||||
"""
|
||||
Count tokens in text using litellm.token_counter.
|
||||
|
||||
Args:
|
||||
text: Text to count tokens for
|
||||
|
||||
Returns:
|
||||
Token count, or 0 if counting fails
|
||||
"""
|
||||
if not text:
|
||||
return 0
|
||||
try:
|
||||
return litellm.token_counter(text=text)
|
||||
except Exception:
|
||||
verbose_logger.debug("Failed to count tokens")
|
||||
return 0
|
||||
|
||||
@staticmethod
|
||||
def calculate_usage_from_request_response(
|
||||
request: "Union[SendMessageRequest, SendStreamingMessageRequest]",
|
||||
response_dict: Dict[str, Any],
|
||||
) -> Tuple[int, int, int]:
|
||||
"""
|
||||
Calculate token usage from A2A request and response.
|
||||
|
||||
Args:
|
||||
request: The A2A SendMessageRequest or SendStreamingMessageRequest
|
||||
response_dict: The A2A response as a dict
|
||||
|
||||
Returns:
|
||||
Tuple of (prompt_tokens, completion_tokens, total_tokens)
|
||||
"""
|
||||
# Count input tokens
|
||||
input_message = A2ARequestUtils.get_input_message_from_request(request)
|
||||
input_text = A2ARequestUtils.extract_text_from_message(input_message)
|
||||
prompt_tokens = A2ARequestUtils.count_tokens(input_text)
|
||||
|
||||
# Count output tokens
|
||||
output_text = A2ARequestUtils.extract_text_from_response(response_dict)
|
||||
completion_tokens = A2ARequestUtils.count_tokens(output_text)
|
||||
|
||||
total_tokens = prompt_tokens + completion_tokens
|
||||
|
||||
return prompt_tokens, completion_tokens, total_tokens
|
||||
|
||||
|
||||
# Backwards compatibility aliases
|
||||
def extract_text_from_a2a_message(message: Any) -> str:
|
||||
return A2ARequestUtils.extract_text_from_message(message)
|
||||
|
||||
|
||||
def extract_text_from_a2a_response(response_dict: Dict[str, Any]) -> str:
|
||||
return A2ARequestUtils.extract_text_from_response(response_dict)
|
||||
|
|
@ -42,6 +42,7 @@ def get_litellm_params(
|
|||
input_cost_per_token=None,
|
||||
output_cost_per_token=None,
|
||||
output_cost_per_second=None,
|
||||
cost_per_query=None,
|
||||
cooldown_time=None,
|
||||
text_completion=None,
|
||||
azure_ad_token_provider=None,
|
||||
|
|
@ -87,6 +88,7 @@ def get_litellm_params(
|
|||
"input_cost_per_second": input_cost_per_second,
|
||||
"output_cost_per_token": output_cost_per_token,
|
||||
"output_cost_per_second": output_cost_per_second,
|
||||
"cost_per_query": cost_per_query,
|
||||
"cooldown_time": cooldown_time,
|
||||
"text_completion": text_completion,
|
||||
"azure_ad_token_provider": azure_ad_token_provider,
|
||||
|
|
|
|||
|
|
@ -872,6 +872,14 @@ def _get_openai_compatible_provider_info( # noqa: PLR0915
|
|||
full_model, api_base, api_key, "ragflow"
|
||||
)
|
||||
model = full_model
|
||||
elif custom_llm_provider == "langgraph":
|
||||
# LangGraph is a custom provider, just need to set api_base
|
||||
api_base = (
|
||||
api_base
|
||||
or get_secret_str("LANGGRAPH_API_BASE")
|
||||
or "http://localhost:2024"
|
||||
)
|
||||
dynamic_api_key = api_key or get_secret_str("LANGGRAPH_API_KEY")
|
||||
|
||||
if api_base is not None and not isinstance(api_base, str):
|
||||
raise Exception("api base needs to be a string. api_base={}".format(api_base))
|
||||
|
|
|
|||
|
|
@ -1651,6 +1651,11 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
result = self._handle_non_streaming_google_genai_generate_content_response_logging(
|
||||
result=result
|
||||
)
|
||||
elif (
|
||||
self.call_type == CallTypes.asend_message.value
|
||||
or self.call_type == CallTypes.send_message.value
|
||||
):
|
||||
result = self._handle_a2a_response_logging(result=result)
|
||||
|
||||
logging_result = self.normalize_logging_result(result=result)
|
||||
|
||||
|
|
@ -3243,6 +3248,29 @@ class Logging(LiteLLMLoggingBaseClass):
|
|||
)
|
||||
return result
|
||||
|
||||
def _handle_a2a_response_logging(self, result: Any) -> Any:
|
||||
"""
|
||||
Handles logging for A2A (Agent-to-Agent) responses.
|
||||
|
||||
Adds usage from model_call_details to the result if available.
|
||||
Uses Pydantic's model_copy to avoid modifying the original response.
|
||||
|
||||
Args:
|
||||
result: The LiteLLMSendMessageResponse from the A2A call
|
||||
|
||||
Returns:
|
||||
The response object with usage added if available
|
||||
"""
|
||||
# Get usage from model_call_details (set by asend_message)
|
||||
usage = self.model_call_details.get("usage")
|
||||
if usage is None:
|
||||
return result
|
||||
|
||||
# Deep copy result and add usage
|
||||
result_copy = result.model_copy(deep=True)
|
||||
result_copy.usage = usage.model_dump() if hasattr(usage, "model_dump") else dict(usage)
|
||||
return result_copy
|
||||
|
||||
|
||||
def _get_masked_values(
|
||||
sensitive_object: dict,
|
||||
|
|
@ -3806,7 +3834,10 @@ def _init_custom_logger_compatible_class( # noqa: PLR0915
|
|||
from litellm.integrations.opentelemetry import (
|
||||
OpenTelemetryConfig,
|
||||
)
|
||||
from litellm.integrations.weave.weave_otel import WeaveOtelLogger, get_weave_otel_config
|
||||
from litellm.integrations.weave.weave_otel import (
|
||||
WeaveOtelLogger,
|
||||
get_weave_otel_config,
|
||||
)
|
||||
|
||||
weave_otel_config = get_weave_otel_config()
|
||||
|
||||
|
|
|
|||
4
litellm/llms/langgraph/__init__.py
Normal file
4
litellm/llms/langgraph/__init__.py
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
from litellm.llms.langgraph.chat.transformation import LangGraphConfig
|
||||
|
||||
__all__ = ["LangGraphConfig"]
|
||||
|
||||
4
litellm/llms/langgraph/chat/__init__.py
Normal file
4
litellm/llms/langgraph/chat/__init__.py
Normal file
|
|
@ -0,0 +1,4 @@
|
|||
from litellm.llms.langgraph.chat.transformation import LangGraphConfig
|
||||
|
||||
__all__ = ["LangGraphConfig"]
|
||||
|
||||
235
litellm/llms/langgraph/chat/sse_iterator.py
Normal file
235
litellm/llms/langgraph/chat/sse_iterator.py
Normal file
|
|
@ -0,0 +1,235 @@
|
|||
"""
|
||||
SSE Stream Iterator for LangGraph.
|
||||
|
||||
Handles Server-Sent Events (SSE) streaming responses from LangGraph.
|
||||
"""
|
||||
|
||||
import json
|
||||
import uuid
|
||||
from typing import TYPE_CHECKING, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.types.utils import Delta, ModelResponse, StreamingChoices
|
||||
|
||||
if TYPE_CHECKING:
|
||||
pass
|
||||
|
||||
|
||||
class LangGraphSSEStreamIterator:
|
||||
"""
|
||||
Iterator for LangGraph SSE streaming responses.
|
||||
Supports both sync and async iteration.
|
||||
|
||||
LangGraph stream format with stream_mode="messages-tuple":
|
||||
Each SSE event is a tuple: (event_type, data)
|
||||
Common event types: "messages", "metadata"
|
||||
"""
|
||||
|
||||
def __init__(self, response: httpx.Response, model: str):
|
||||
self.response = response
|
||||
self.model = model
|
||||
self.finished = False
|
||||
self.line_iterator = None
|
||||
self.async_line_iterator = None
|
||||
|
||||
def __iter__(self):
|
||||
"""Initialize sync iteration."""
|
||||
self.line_iterator = self.response.iter_lines()
|
||||
return self
|
||||
|
||||
def __aiter__(self):
|
||||
"""Initialize async iteration."""
|
||||
self.async_line_iterator = self.response.aiter_lines()
|
||||
return self
|
||||
|
||||
def _parse_sse_line(self, line: str) -> Optional[ModelResponse]:
|
||||
"""
|
||||
Parse a single SSE line and return a ModelResponse chunk if applicable.
|
||||
|
||||
LangGraph SSE format can vary:
|
||||
- data: [...] (tuple format)
|
||||
- event: ...\ndata: ...
|
||||
"""
|
||||
line = line.strip()
|
||||
if not line:
|
||||
return None
|
||||
|
||||
# Handle SSE data lines
|
||||
if line.startswith("data:"):
|
||||
json_str = line[5:].strip()
|
||||
if not json_str:
|
||||
return None
|
||||
|
||||
try:
|
||||
data = json.loads(json_str)
|
||||
return self._process_data(data)
|
||||
except json.JSONDecodeError:
|
||||
verbose_logger.debug(f"Skipping non-JSON SSE line: {line[:100]}")
|
||||
return None
|
||||
|
||||
return None
|
||||
|
||||
def _process_data(self, data) -> Optional[ModelResponse]:
|
||||
"""
|
||||
Process parsed data from SSE stream.
|
||||
|
||||
LangGraph uses tuple format: [event_type, payload]
|
||||
"""
|
||||
# Handle tuple format: ["messages", ...]
|
||||
if isinstance(data, list) and len(data) >= 2:
|
||||
event_type = data[0]
|
||||
payload = data[1]
|
||||
|
||||
if event_type == "messages":
|
||||
return self._process_messages_event(payload)
|
||||
elif event_type == "metadata":
|
||||
# Metadata event, might contain usage info
|
||||
return self._process_metadata_event(payload)
|
||||
|
||||
# Handle dict format (alternative response format)
|
||||
elif isinstance(data, dict):
|
||||
if "content" in data:
|
||||
return self._create_content_chunk(data.get("content", ""))
|
||||
elif "messages" in data:
|
||||
messages = data.get("messages", [])
|
||||
if messages:
|
||||
last_msg = messages[-1]
|
||||
if isinstance(last_msg, dict) and last_msg.get("type") == "ai":
|
||||
return self._create_content_chunk(last_msg.get("content", ""))
|
||||
|
||||
return None
|
||||
|
||||
def _process_messages_event(self, payload) -> Optional[ModelResponse]:
|
||||
"""
|
||||
Process a messages event from the stream.
|
||||
|
||||
payload format: [[message_object, metadata], ...]
|
||||
"""
|
||||
if isinstance(payload, list):
|
||||
for item in payload:
|
||||
if isinstance(item, list) and len(item) >= 1:
|
||||
msg = item[0]
|
||||
if isinstance(msg, dict):
|
||||
msg_type = msg.get("type", "")
|
||||
content = msg.get("content", "")
|
||||
|
||||
# Only return AI messages with content
|
||||
if msg_type == "ai" and content:
|
||||
return self._create_content_chunk(content)
|
||||
elif msg_type == "AIMessageChunk" and content:
|
||||
return self._create_content_chunk(content)
|
||||
elif isinstance(item, dict):
|
||||
msg_type = item.get("type", "")
|
||||
content = item.get("content", "")
|
||||
if msg_type in ("ai", "AIMessageChunk") and content:
|
||||
return self._create_content_chunk(content)
|
||||
|
||||
return None
|
||||
|
||||
def _process_metadata_event(self, payload) -> Optional[ModelResponse]:
|
||||
"""
|
||||
Process a metadata event, which may signal the end of the stream.
|
||||
"""
|
||||
if isinstance(payload, dict):
|
||||
# Check if this is a final event
|
||||
if "run_id" in payload:
|
||||
self.finished = True
|
||||
return self._create_final_chunk()
|
||||
return None
|
||||
|
||||
def _create_content_chunk(self, text: str) -> ModelResponse:
|
||||
"""Create a ModelResponse chunk with content."""
|
||||
chunk = ModelResponse(
|
||||
id=f"chatcmpl-{uuid.uuid4()}",
|
||||
created=0,
|
||||
model=self.model,
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
chunk.choices = [
|
||||
StreamingChoices(
|
||||
finish_reason=None,
|
||||
index=0,
|
||||
delta=Delta(content=text, role="assistant"),
|
||||
)
|
||||
]
|
||||
|
||||
return chunk
|
||||
|
||||
def _create_final_chunk(self) -> ModelResponse:
|
||||
"""Create a final ModelResponse chunk with finish_reason."""
|
||||
chunk = ModelResponse(
|
||||
id=f"chatcmpl-{uuid.uuid4()}",
|
||||
created=0,
|
||||
model=self.model,
|
||||
object="chat.completion.chunk",
|
||||
)
|
||||
|
||||
chunk.choices = [
|
||||
StreamingChoices(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
delta=Delta(),
|
||||
)
|
||||
]
|
||||
|
||||
return chunk
|
||||
|
||||
def __next__(self) -> ModelResponse:
|
||||
"""Sync iteration - parse SSE events and yield ModelResponse chunks."""
|
||||
try:
|
||||
if self.line_iterator is None:
|
||||
raise StopIteration
|
||||
|
||||
for line in self.line_iterator:
|
||||
result = self._parse_sse_line(line)
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
# Stream ended naturally - send final chunk if not already finished
|
||||
if not self.finished:
|
||||
self.finished = True
|
||||
return self._create_final_chunk()
|
||||
|
||||
raise StopIteration
|
||||
|
||||
except StopIteration:
|
||||
raise
|
||||
except httpx.StreamConsumed:
|
||||
raise StopIteration
|
||||
except httpx.StreamClosed:
|
||||
raise StopIteration
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Error in LangGraph SSE stream: {str(e)}")
|
||||
raise StopIteration
|
||||
|
||||
async def __anext__(self) -> ModelResponse:
|
||||
"""Async iteration - parse SSE events and yield ModelResponse chunks."""
|
||||
try:
|
||||
if self.async_line_iterator is None:
|
||||
raise StopAsyncIteration
|
||||
|
||||
async for line in self.async_line_iterator:
|
||||
result = self._parse_sse_line(line)
|
||||
if result is not None:
|
||||
return result
|
||||
|
||||
# Stream ended naturally - send final chunk if not already finished
|
||||
if not self.finished:
|
||||
self.finished = True
|
||||
return self._create_final_chunk()
|
||||
|
||||
raise StopAsyncIteration
|
||||
|
||||
except StopAsyncIteration:
|
||||
raise
|
||||
except httpx.StreamConsumed:
|
||||
raise StopAsyncIteration
|
||||
except httpx.StreamClosed:
|
||||
raise StopAsyncIteration
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Error in LangGraph SSE stream: {str(e)}")
|
||||
raise StopAsyncIteration
|
||||
|
||||
509
litellm/llms/langgraph/chat/transformation.py
Normal file
509
litellm/llms/langgraph/chat/transformation.py
Normal file
|
|
@ -0,0 +1,509 @@
|
|||
"""
|
||||
Transformation for LangGraph API.
|
||||
|
||||
LangGraph provides streaming (/runs/stream) and non-streaming (/runs/wait) endpoints
|
||||
for running agents.
|
||||
|
||||
Streaming endpoint: POST /runs/stream
|
||||
Non-streaming endpoint: POST /runs/wait
|
||||
"""
|
||||
|
||||
import json
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
import httpx
|
||||
|
||||
from litellm._logging import verbose_logger
|
||||
from litellm.litellm_core_utils.prompt_templates.common_utils import (
|
||||
convert_content_list_to_str,
|
||||
)
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.llms.langgraph.chat.sse_iterator import LangGraphSSEStreamIterator
|
||||
from litellm.types.llms.openai import AllMessageValues
|
||||
from litellm.types.utils import Choices, Message, ModelResponse, Usage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
LiteLLMLoggingObj = _LiteLLMLoggingObj
|
||||
else:
|
||||
LiteLLMLoggingObj = Any
|
||||
HTTPHandler = Any
|
||||
AsyncHTTPHandler = Any
|
||||
CustomStreamWrapper = Any
|
||||
|
||||
|
||||
class LangGraphError(BaseLLMException):
|
||||
"""Exception class for LangGraph API errors."""
|
||||
|
||||
pass
|
||||
|
||||
|
||||
class LangGraphConfig(BaseConfig):
|
||||
"""
|
||||
Configuration for LangGraph API.
|
||||
|
||||
LangGraph is a framework for building stateful, multi-actor applications with LLMs.
|
||||
It provides a streaming and non-streaming API for running agents.
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
super().__init__(**kwargs)
|
||||
|
||||
def _get_openai_compatible_provider_info(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
) -> Tuple[Optional[str], Optional[str]]:
|
||||
"""
|
||||
Get LangGraph API base and key from params or environment.
|
||||
|
||||
Returns:
|
||||
Tuple of (api_base, api_key)
|
||||
"""
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
||||
api_base = (
|
||||
api_base
|
||||
or get_secret_str("LANGGRAPH_API_BASE")
|
||||
or "http://localhost:2024"
|
||||
)
|
||||
|
||||
api_key = api_key or get_secret_str("LANGGRAPH_API_KEY")
|
||||
|
||||
return api_base, api_key
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
"""
|
||||
LangGraph supports minimal OpenAI params since it's an agent runtime.
|
||||
"""
|
||||
return ["stream"]
|
||||
|
||||
def map_openai_params(
|
||||
self,
|
||||
non_default_params: dict,
|
||||
optional_params: dict,
|
||||
model: str,
|
||||
drop_params: bool,
|
||||
) -> dict:
|
||||
"""
|
||||
Map OpenAI params to LangGraph params.
|
||||
"""
|
||||
return optional_params
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
model: str,
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
stream: Optional[bool] = None,
|
||||
) -> str:
|
||||
"""
|
||||
Get the complete URL for the LangGraph request.
|
||||
|
||||
Streaming: /runs/stream
|
||||
Non-streaming: /runs/wait
|
||||
"""
|
||||
if api_base is None:
|
||||
raise ValueError(
|
||||
"api_base is required for LangGraph. Set it via LANGGRAPH_API_BASE env var or api_base parameter."
|
||||
)
|
||||
|
||||
# Remove trailing slash if present
|
||||
api_base = api_base.rstrip("/")
|
||||
|
||||
# Choose endpoint based on streaming mode
|
||||
if stream:
|
||||
return f"{api_base}/runs/stream"
|
||||
else:
|
||||
return f"{api_base}/runs/wait"
|
||||
|
||||
def _get_assistant_id(self, model: str, optional_params: dict) -> str:
|
||||
"""
|
||||
Get the assistant ID from model or optional_params.
|
||||
|
||||
model format: "langgraph/assistant_id" or just "assistant_id"
|
||||
"""
|
||||
assistant_id = optional_params.get("assistant_id")
|
||||
if assistant_id:
|
||||
return assistant_id
|
||||
|
||||
# Extract from model name
|
||||
if "/" in model:
|
||||
parts = model.split("/", 1)
|
||||
if len(parts) == 2:
|
||||
return parts[1]
|
||||
return model
|
||||
|
||||
def _convert_messages_to_langgraph_format(
|
||||
self, messages: List[AllMessageValues]
|
||||
) -> List[Dict[str, str]]:
|
||||
"""
|
||||
Convert OpenAI-format messages to LangGraph format.
|
||||
|
||||
OpenAI format: {"role": "user", "content": "..."}
|
||||
LangGraph format: {"role": "human", "content": "..."}
|
||||
"""
|
||||
langgraph_messages = []
|
||||
for msg in messages:
|
||||
role = msg.get("role", "user")
|
||||
content = msg.get("content", "")
|
||||
|
||||
# Convert OpenAI roles to LangGraph roles
|
||||
if role == "user":
|
||||
langgraph_role = "human"
|
||||
elif role == "assistant":
|
||||
langgraph_role = "assistant"
|
||||
elif role == "system":
|
||||
langgraph_role = "system"
|
||||
else:
|
||||
langgraph_role = "human"
|
||||
|
||||
# Handle content that might be a list
|
||||
if isinstance(content, list):
|
||||
content = convert_content_list_to_str(msg)
|
||||
|
||||
langgraph_messages.append({"role": langgraph_role, "content": content})
|
||||
|
||||
return langgraph_messages
|
||||
|
||||
def transform_request(
|
||||
self,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
headers: dict,
|
||||
) -> dict:
|
||||
"""
|
||||
Transform the request to LangGraph format.
|
||||
|
||||
LangGraph request format:
|
||||
{
|
||||
"assistant_id": "agent",
|
||||
"input": {
|
||||
"messages": [{"role": "human", "content": "..."}]
|
||||
},
|
||||
"stream_mode": "messages-tuple" # for streaming
|
||||
}
|
||||
"""
|
||||
assistant_id = self._get_assistant_id(model, optional_params)
|
||||
langgraph_messages = self._convert_messages_to_langgraph_format(messages)
|
||||
|
||||
payload: Dict[str, Any] = {
|
||||
"assistant_id": assistant_id,
|
||||
"input": {"messages": langgraph_messages},
|
||||
}
|
||||
|
||||
# Add stream_mode for streaming requests
|
||||
stream = litellm_params.get("stream", False)
|
||||
if stream:
|
||||
stream_mode = optional_params.get("stream_mode", "messages-tuple")
|
||||
payload["stream_mode"] = stream_mode
|
||||
|
||||
# Add optional config if provided
|
||||
if "config" in optional_params:
|
||||
payload["config"] = optional_params["config"]
|
||||
|
||||
# Add optional metadata if provided
|
||||
if "metadata" in optional_params:
|
||||
payload["metadata"] = optional_params["metadata"]
|
||||
|
||||
# Add thread_id if provided (for stateful conversations)
|
||||
if "thread_id" in optional_params:
|
||||
payload["thread_id"] = optional_params["thread_id"]
|
||||
|
||||
verbose_logger.debug(f"LangGraph request payload: {payload}")
|
||||
return payload
|
||||
|
||||
def _extract_content_from_response(self, response_json: dict) -> str:
|
||||
"""
|
||||
Extract content from LangGraph non-streaming response.
|
||||
|
||||
Response format varies, but commonly:
|
||||
{
|
||||
"messages": [...], # or could be in different structure
|
||||
"values": {...}
|
||||
}
|
||||
"""
|
||||
# Try to get the last AI message from the response
|
||||
messages = response_json.get("messages", [])
|
||||
if isinstance(messages, list) and messages:
|
||||
# Find the last AI/assistant message
|
||||
for msg in reversed(messages):
|
||||
if isinstance(msg, dict):
|
||||
msg_type = msg.get("type", "")
|
||||
role = msg.get("role", "")
|
||||
if msg_type == "ai" or role == "assistant":
|
||||
return msg.get("content", "")
|
||||
|
||||
# Check values for output
|
||||
values = response_json.get("values", {})
|
||||
if isinstance(values, dict):
|
||||
output_messages = values.get("messages", [])
|
||||
if isinstance(output_messages, list) and output_messages:
|
||||
for msg in reversed(output_messages):
|
||||
if isinstance(msg, dict):
|
||||
msg_type = msg.get("type", "")
|
||||
if msg_type == "ai":
|
||||
return msg.get("content", "")
|
||||
|
||||
# Fallback: try to serialize the whole response
|
||||
verbose_logger.warning(
|
||||
"Could not extract content from LangGraph response, returning raw"
|
||||
)
|
||||
return json.dumps(response_json)
|
||||
|
||||
def get_streaming_response(
|
||||
self,
|
||||
model: str,
|
||||
raw_response: httpx.Response,
|
||||
) -> LangGraphSSEStreamIterator:
|
||||
"""
|
||||
Return a streaming iterator for SSE responses.
|
||||
"""
|
||||
return LangGraphSSEStreamIterator(response=raw_response, model=model)
|
||||
|
||||
def get_sync_custom_stream_wrapper(
|
||||
self,
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_base: str,
|
||||
headers: dict,
|
||||
data: dict,
|
||||
messages: list,
|
||||
client: Optional[Union[HTTPHandler, "AsyncHTTPHandler"]] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
signed_json_body: Optional[bytes] = None,
|
||||
) -> CustomStreamWrapper:
|
||||
"""
|
||||
Get a CustomStreamWrapper for synchronous streaming.
|
||||
"""
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
HTTPHandler,
|
||||
_get_httpx_client,
|
||||
)
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
if client is None or not isinstance(client, HTTPHandler):
|
||||
client = _get_httpx_client(params={})
|
||||
|
||||
verbose_logger.debug(f"Making sync streaming request to: {api_base}")
|
||||
|
||||
# Make streaming request
|
||||
response = client.post(
|
||||
api_base,
|
||||
headers=headers,
|
||||
data=json.dumps(data),
|
||||
stream=True,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise LangGraphError(
|
||||
status_code=response.status_code, message=str(response.read())
|
||||
)
|
||||
|
||||
# Create iterator for SSE stream
|
||||
completion_stream = self.get_streaming_response(
|
||||
model=model, raw_response=response
|
||||
)
|
||||
|
||||
streaming_response = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
# LOGGING
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
original_response="first stream response received",
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
|
||||
return streaming_response
|
||||
|
||||
async def get_async_custom_stream_wrapper(
|
||||
self,
|
||||
model: str,
|
||||
custom_llm_provider: str,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
api_base: str,
|
||||
headers: dict,
|
||||
data: dict,
|
||||
messages: list,
|
||||
client: Optional["AsyncHTTPHandler"] = None,
|
||||
json_mode: Optional[bool] = None,
|
||||
signed_json_body: Optional[bytes] = None,
|
||||
) -> CustomStreamWrapper:
|
||||
"""
|
||||
Get a CustomStreamWrapper for asynchronous streaming.
|
||||
"""
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.utils import CustomStreamWrapper
|
||||
|
||||
if client is None or not isinstance(client, AsyncHTTPHandler):
|
||||
client = get_async_httpx_client(
|
||||
llm_provider=cast(Any, "langgraph"), params={}
|
||||
)
|
||||
|
||||
verbose_logger.debug(f"Making async streaming request to: {api_base}")
|
||||
|
||||
# Make async streaming request
|
||||
response = await client.post(
|
||||
api_base,
|
||||
headers=headers,
|
||||
data=json.dumps(data),
|
||||
stream=True,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
if response.status_code != 200:
|
||||
raise LangGraphError(
|
||||
status_code=response.status_code, message=str(await response.aread())
|
||||
)
|
||||
|
||||
# Create iterator for SSE stream
|
||||
completion_stream = self.get_streaming_response(
|
||||
model=model, raw_response=response
|
||||
)
|
||||
|
||||
streaming_response = CustomStreamWrapper(
|
||||
completion_stream=completion_stream,
|
||||
model=model,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
logging_obj=logging_obj,
|
||||
)
|
||||
|
||||
# LOGGING
|
||||
logging_obj.post_call(
|
||||
input=messages,
|
||||
api_key="",
|
||||
original_response="first stream response received",
|
||||
additional_args={"complete_input_dict": data},
|
||||
)
|
||||
|
||||
return streaming_response
|
||||
|
||||
@property
|
||||
def has_custom_stream_wrapper(self) -> bool:
|
||||
"""Indicates that this config has custom streaming support."""
|
||||
return True
|
||||
|
||||
@property
|
||||
def supports_stream_param_in_request_body(self) -> bool:
|
||||
"""
|
||||
LangGraph does not use a stream param in request body.
|
||||
Streaming is determined by the endpoint URL.
|
||||
"""
|
||||
return False
|
||||
|
||||
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 the LangGraph response to LiteLLM ModelResponse format.
|
||||
"""
|
||||
try:
|
||||
response_json = raw_response.json()
|
||||
verbose_logger.debug(f"LangGraph response: {response_json}")
|
||||
|
||||
content = self._extract_content_from_response(response_json)
|
||||
|
||||
# Create the message
|
||||
message = Message(content=content, role="assistant")
|
||||
|
||||
# Create choices
|
||||
choice = Choices(finish_reason="stop", index=0, message=message)
|
||||
|
||||
# Update model response
|
||||
model_response.choices = [choice]
|
||||
model_response.model = model
|
||||
|
||||
# LangGraph doesn't provide token usage, so we estimate it
|
||||
try:
|
||||
from litellm.utils import token_counter
|
||||
|
||||
prompt_tokens = token_counter(model="gpt-3.5-turbo", messages=messages)
|
||||
completion_tokens = token_counter(
|
||||
model="gpt-3.5-turbo", text=content, count_response_tokens=True
|
||||
)
|
||||
total_tokens = prompt_tokens + completion_tokens
|
||||
|
||||
usage = Usage(
|
||||
prompt_tokens=prompt_tokens,
|
||||
completion_tokens=completion_tokens,
|
||||
total_tokens=total_tokens,
|
||||
)
|
||||
setattr(model_response, "usage", usage)
|
||||
except Exception as e:
|
||||
verbose_logger.warning(f"Failed to calculate token usage: {str(e)}")
|
||||
|
||||
return model_response
|
||||
|
||||
except Exception as e:
|
||||
verbose_logger.error(f"Error processing LangGraph response: {str(e)}")
|
||||
raise LangGraphError(
|
||||
message=f"Error processing response: {str(e)}",
|
||||
status_code=raw_response.status_code,
|
||||
)
|
||||
|
||||
def validate_environment(
|
||||
self,
|
||||
headers: dict,
|
||||
model: str,
|
||||
messages: List[AllMessageValues],
|
||||
optional_params: dict,
|
||||
litellm_params: dict,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""
|
||||
Validate and set up environment for LangGraph requests.
|
||||
"""
|
||||
headers["Content-Type"] = "application/json"
|
||||
|
||||
# Add API key if provided
|
||||
if api_key:
|
||||
headers["Authorization"] = f"Bearer {api_key}"
|
||||
|
||||
return headers
|
||||
|
||||
def get_error_class(
|
||||
self, error_message: str, status_code: int, headers: Union[dict, httpx.Headers]
|
||||
) -> BaseLLMException:
|
||||
return LangGraphError(status_code=status_code, message=error_message)
|
||||
|
||||
def should_fake_stream(
|
||||
self,
|
||||
model: Optional[str],
|
||||
stream: Optional[bool],
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
LangGraph has native streaming support, so we don't need to fake stream.
|
||||
"""
|
||||
return False
|
||||
|
||||
|
|
@ -1,18 +1,21 @@
|
|||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union
|
||||
from io import BufferedReader
|
||||
from typing import cast
|
||||
from typing import TYPE_CHECKING, Any, Dict, List, Optional, Tuple, Union, cast
|
||||
|
||||
import httpx
|
||||
from httpx._types import RequestFiles
|
||||
|
||||
import litellm
|
||||
from litellm.llms.base_llm.videos.transformation import BaseVideoConfig
|
||||
from litellm.types.videos.main import VideoCreateOptionalRequestParams
|
||||
from litellm.llms.openai.image_edit.transformation import ImageEditRequestUtils
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.llms.openai import CreateVideoRequest
|
||||
from litellm.types.router import GenericLiteLLMParams
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
from litellm.types.videos.main import VideoObject
|
||||
from litellm.types.videos.utils import encode_video_id_with_provider, extract_original_video_id
|
||||
import litellm
|
||||
from litellm.llms.openai.image_edit.transformation import ImageEditRequestUtils
|
||||
from litellm.types.videos.main import VideoCreateOptionalRequestParams, VideoObject
|
||||
from litellm.types.videos.utils import (
|
||||
encode_video_id_with_provider,
|
||||
extract_original_video_id,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
|
||||
|
||||
|
|
|
|||
|
|
@ -2,21 +2,22 @@ from __future__ import annotations
|
|||
|
||||
import json
|
||||
import time
|
||||
from typing import AsyncIterator, Iterator, Optional
|
||||
|
||||
import httpx
|
||||
|
||||
from typing import Iterator, Optional, AsyncIterator
|
||||
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
from litellm.types.llms.openai import OpenAIChatCompletionChunk
|
||||
|
||||
from ...custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
|
||||
|
||||
# -------------------------------
|
||||
# Errors
|
||||
# -------------------------------
|
||||
class GenAIHubOrchestrationError(Exception):
|
||||
class GenAIHubOrchestrationError(BaseLLMException):
|
||||
def __init__(self, status_code: int, message: str):
|
||||
super().__init__(message)
|
||||
super().__init__(status_code=status_code, message=message)
|
||||
self.status_code = status_code
|
||||
self.message = message
|
||||
|
||||
|
|
|
|||
|
|
@ -2,11 +2,11 @@
|
|||
Translates from OpenAI's `/v1/embeddings` to IBM's `/text/embeddings` route.
|
||||
"""
|
||||
|
||||
from typing import Optional, List, Dict, Literal
|
||||
from pydantic import BaseModel, Field
|
||||
from functools import cached_property
|
||||
from typing import Dict, List, Literal, Optional, Union
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from litellm.llms.base_llm.embedding.transformation import (
|
||||
BaseEmbeddingConfig,
|
||||
|
|
@ -55,7 +55,7 @@ class EmbeddingsModules(BaseModel):
|
|||
|
||||
|
||||
class EmbeddingInput(BaseModel):
|
||||
text: str | List[str]
|
||||
text: Union[str, List[str]]
|
||||
type: Literal["text", "document", "query"] = "text"
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -176,7 +176,6 @@ from .llms.databricks.embed.handler import DatabricksEmbeddingHandler
|
|||
from .llms.deprecated_providers import aleph_alpha, palm
|
||||
from .llms.gemini.common_utils import get_api_key_from_env
|
||||
from .llms.groq.chat.handler import GroqChatCompletion
|
||||
from .llms.sap.chat.handler import GenAIHubOrchestration
|
||||
from .llms.heroku.chat.transformation import HerokuChatConfig
|
||||
from .llms.huggingface.embedding.handler import HuggingFaceEmbedding
|
||||
from .llms.lemonade.chat.transformation import LemonadeChatConfig
|
||||
|
|
@ -196,6 +195,7 @@ from .llms.predibase.chat.handler import PredibaseChatCompletion
|
|||
from .llms.replicate.chat.handler import completion as replicate_chat_completion
|
||||
from .llms.sagemaker.chat.handler import SagemakerChatHandler
|
||||
from .llms.sagemaker.completion.handler import SagemakerLLM
|
||||
from .llms.sap.chat.handler import GenAIHubOrchestration
|
||||
from .llms.vertex_ai import vertex_ai_non_gemini
|
||||
from .llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import VertexLLM
|
||||
from .llms.vertex_ai.gemini_embeddings.batch_embed_content_handler import (
|
||||
|
|
@ -3963,6 +3963,39 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
logging_obj=logging,
|
||||
)
|
||||
|
||||
elif custom_llm_provider == "langgraph":
|
||||
# LangGraph - Agent Runtime Provider
|
||||
from litellm.llms.langgraph.chat.transformation import LangGraphConfig
|
||||
|
||||
(
|
||||
api_base,
|
||||
api_key,
|
||||
) = LangGraphConfig()._get_openai_compatible_provider_info(
|
||||
api_base=api_base or litellm.api_base,
|
||||
api_key=api_key or litellm.api_key,
|
||||
)
|
||||
|
||||
headers = headers or litellm.headers
|
||||
|
||||
response = base_llm_http_handler.completion(
|
||||
model=model,
|
||||
stream=stream,
|
||||
messages=messages,
|
||||
acompletion=acompletion,
|
||||
api_base=api_base,
|
||||
model_response=model_response,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
shared_session=shared_session,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
timeout=timeout,
|
||||
headers=headers,
|
||||
encoding=encoding,
|
||||
api_key=api_key,
|
||||
logging_obj=logging,
|
||||
client=client,
|
||||
)
|
||||
|
||||
else:
|
||||
raise LiteLLMUnknownProvider(
|
||||
model=model, custom_llm_provider=custom_llm_provider
|
||||
|
|
|
|||
|
|
@ -2677,6 +2677,7 @@ class SpendLogsPayload(TypedDict):
|
|||
model_id: Optional[str]
|
||||
model_group: Optional[str]
|
||||
mcp_namespaced_tool_name: Optional[str]
|
||||
agent_id: Optional[str]
|
||||
api_base: str
|
||||
user: str
|
||||
metadata: str # json str
|
||||
|
|
|
|||
|
|
@ -46,22 +46,38 @@ def _get_agent(agent_id: str):
|
|||
|
||||
|
||||
async def _handle_stream_message(
|
||||
a2a_client: Any,
|
||||
api_base: Optional[str],
|
||||
request_id: str,
|
||||
params: dict,
|
||||
litellm_params: Optional[dict] = None,
|
||||
agent_id: Optional[str] = None,
|
||||
metadata: Optional[dict] = None,
|
||||
proxy_server_request: Optional[dict] = None,
|
||||
) -> StreamingResponse:
|
||||
"""Handle message/stream method."""
|
||||
"""Handle message/stream method via SDK functions."""
|
||||
from a2a.types import MessageSendParams, SendStreamingMessageRequest
|
||||
|
||||
a2a_request = SendStreamingMessageRequest(
|
||||
id=request_id,
|
||||
params=MessageSendParams(**params),
|
||||
)
|
||||
from litellm.a2a_protocol import asend_message_streaming
|
||||
|
||||
async def stream_response():
|
||||
try:
|
||||
async for chunk in a2a_client.send_message_streaming(a2a_request):
|
||||
yield json.dumps(chunk.model_dump(mode="json", exclude_none=True)) + "\n"
|
||||
a2a_request = SendStreamingMessageRequest(
|
||||
id=request_id,
|
||||
params=MessageSendParams(**params),
|
||||
)
|
||||
async for chunk in asend_message_streaming(
|
||||
request=a2a_request,
|
||||
api_base=api_base,
|
||||
litellm_params=litellm_params,
|
||||
agent_id=agent_id,
|
||||
metadata=metadata,
|
||||
proxy_server_request=proxy_server_request,
|
||||
):
|
||||
# Chunk may be dict or object depending on bridge vs standard path
|
||||
if hasattr(chunk, "model_dump"):
|
||||
yield json.dumps(chunk.model_dump(mode="json", exclude_none=True)) + "\n"
|
||||
else:
|
||||
yield json.dumps(chunk) + "\n"
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception(f"Error streaming A2A response: {e}")
|
||||
yield json.dumps({
|
||||
|
|
@ -153,7 +169,9 @@ async def invoke_agent_a2a(
|
|||
- message/send: Send a message and get a response
|
||||
- message/stream: Send a message and stream the response
|
||||
"""
|
||||
from litellm.a2a_protocol import asend_message, create_a2a_client
|
||||
from a2a.types import MessageSendParams, SendMessageRequest
|
||||
|
||||
from litellm.a2a_protocol import asend_message
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
|
||||
AgentRequestHandler,
|
||||
)
|
||||
|
|
@ -195,10 +213,17 @@ async def invoke_agent_a2a(
|
|||
# Get backend URL and agent name
|
||||
agent_url = agent.agent_card_params.get("url")
|
||||
agent_name = agent.agent_card_params.get("name", agent_id)
|
||||
if not agent_url:
|
||||
|
||||
# Get litellm_params (may include custom_llm_provider for completion bridge)
|
||||
litellm_params = agent.litellm_params or {}
|
||||
custom_llm_provider = litellm_params.get("custom_llm_provider")
|
||||
|
||||
# URL is required unless using completion bridge with a provider that derives endpoint from model
|
||||
# (e.g., bedrock/agentcore derives endpoint from ARN in model string)
|
||||
if not agent_url and not custom_llm_provider:
|
||||
return _jsonrpc_error(request_id, -32000, f"Agent '{agent_id}' has no URL configured", 500)
|
||||
|
||||
verbose_proxy_logger.info(f"Proxying A2A request to agent '{agent_id}' at {agent_url}")
|
||||
verbose_proxy_logger.info(f"Proxying A2A request to agent '{agent_id}' at {agent_url or 'completion-bridge'}")
|
||||
|
||||
# Set up data dict for litellm processing
|
||||
body.update({
|
||||
|
|
@ -216,28 +241,32 @@ async def invoke_agent_a2a(
|
|||
version=version,
|
||||
)
|
||||
|
||||
# Create A2A client
|
||||
a2a_client = await create_a2a_client(base_url=agent_url)
|
||||
|
||||
# Route through SDK functions
|
||||
if method == "message/send":
|
||||
from a2a.types import MessageSendParams, SendMessageRequest
|
||||
|
||||
a2a_request = SendMessageRequest(
|
||||
id=request_id,
|
||||
params=MessageSendParams(**params),
|
||||
)
|
||||
|
||||
# Pass litellm data through kwargs for proper logging
|
||||
response = await asend_message(
|
||||
a2a_client=a2a_client,
|
||||
request=a2a_request,
|
||||
api_base=agent_url,
|
||||
litellm_params=litellm_params,
|
||||
agent_id=agent.agent_id,
|
||||
metadata=data.get("metadata", {}),
|
||||
proxy_server_request=data.get("proxy_server_request"),
|
||||
)
|
||||
return JSONResponse(content=response.model_dump(mode="json", exclude_none=True))
|
||||
|
||||
elif method == "message/stream":
|
||||
return await _handle_stream_message(a2a_client, request_id, params)
|
||||
return await _handle_stream_message(
|
||||
api_base=agent_url,
|
||||
request_id=request_id,
|
||||
params=params,
|
||||
litellm_params=litellm_params,
|
||||
agent_id=agent.agent_id,
|
||||
metadata=data.get("metadata", {}),
|
||||
proxy_server_request=data.get("proxy_server_request"),
|
||||
)
|
||||
else:
|
||||
return _jsonrpc_error(request_id, -32601, f"Method '{method}' not found")
|
||||
|
||||
|
|
|
|||
|
|
@ -6,6 +6,7 @@ Module responsible for
|
|||
"""
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import json
|
||||
import os
|
||||
import random
|
||||
|
|
@ -164,14 +165,14 @@ class DBSpendUpdateWriter:
|
|||
asyncio.create_task(
|
||||
self._update_tag_db(
|
||||
response_cost=response_cost,
|
||||
request_tags=payload.get("request_tags"),
|
||||
request_tags=copy.deepcopy(payload.get("request_tags")),
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
)
|
||||
|
||||
if disable_spend_logs is False:
|
||||
await self._insert_spend_log_to_db(
|
||||
payload=payload,
|
||||
payload=copy.deepcopy(payload),
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
else:
|
||||
|
|
@ -181,14 +182,14 @@ class DBSpendUpdateWriter:
|
|||
|
||||
asyncio.create_task(
|
||||
self.add_spend_log_transaction_to_daily_user_transaction(
|
||||
payload=payload,
|
||||
payload=copy.deepcopy(payload),
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
)
|
||||
|
||||
asyncio.create_task(
|
||||
self.add_spend_log_transaction_to_daily_end_user_transaction(
|
||||
payload=payload,
|
||||
payload=copy.deepcopy(payload),
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
)
|
||||
|
|
@ -202,20 +203,20 @@ class DBSpendUpdateWriter:
|
|||
|
||||
asyncio.create_task(
|
||||
self.add_spend_log_transaction_to_daily_team_transaction(
|
||||
payload=payload,
|
||||
payload=copy.deepcopy(payload),
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
)
|
||||
asyncio.create_task(
|
||||
self.add_spend_log_transaction_to_daily_org_transaction(
|
||||
payload=payload,
|
||||
payload=copy.deepcopy(payload),
|
||||
org_id=org_id,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
)
|
||||
asyncio.create_task(
|
||||
self.add_spend_log_transaction_to_daily_tag_transaction(
|
||||
payload=payload,
|
||||
payload=copy.deepcopy(payload),
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -29,6 +29,7 @@ from litellm.integrations.custom_logger import CustomLogger
|
|||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.auth.auth_utils import get_model_rate_limit_from_metadata
|
||||
from litellm.types.llms.openai import BaseLiteLLMOpenAIResponseObject
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from opentelemetry.trace import Span as _Span
|
||||
|
|
@ -1232,6 +1233,28 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
|
||||
return pipeline_operations
|
||||
|
||||
def _get_total_tokens_from_usage(self, usage: Any | None, rate_limit_type: Literal["output", "input", "total"]) -> int:
|
||||
# Get total tokens from response
|
||||
total_tokens = 0
|
||||
# spot fix for /responses api
|
||||
if usage:
|
||||
if isinstance(usage, Usage):
|
||||
if rate_limit_type == "output":
|
||||
total_tokens = usage.completion_tokens
|
||||
elif rate_limit_type == "input":
|
||||
total_tokens = usage.prompt_tokens
|
||||
elif rate_limit_type == "total":
|
||||
total_tokens = usage.total_tokens
|
||||
elif isinstance(usage, dict):
|
||||
# Responses API usage comes as a dict in ResponsesAPIResponse
|
||||
if rate_limit_type == "output":
|
||||
total_tokens = usage.get("completion_tokens", 0)
|
||||
elif rate_limit_type == "input":
|
||||
total_tokens = usage.get("prompt_tokens", 0)
|
||||
elif rate_limit_type == "total":
|
||||
total_tokens = usage.get("total_tokens", 0)
|
||||
return total_tokens
|
||||
|
||||
async def _execute_token_increment_script(
|
||||
self,
|
||||
pipeline_operations: List["RedisPipelineIncrementOperation"],
|
||||
|
|
@ -1313,11 +1336,10 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
|
||||
def get_rate_limit_type(self) -> Literal["output", "input", "total"]:
|
||||
from litellm.proxy.proxy_server import general_settings
|
||||
|
||||
specified_rate_limit_type = general_settings.get(
|
||||
"token_rate_limit_type", "output"
|
||||
"token_rate_limit_type", "total"
|
||||
)
|
||||
if not specified_rate_limit_type or specified_rate_limit_type not in [
|
||||
if specified_rate_limit_type not in [
|
||||
"output",
|
||||
"input",
|
||||
"total",
|
||||
|
|
@ -1336,7 +1358,6 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
get_model_group_from_litellm_kwargs,
|
||||
)
|
||||
from litellm.types.caching import RedisPipelineIncrementOperation
|
||||
from litellm.types.utils import ModelResponse, Usage
|
||||
|
||||
rate_limit_type = self.get_rate_limit_type()
|
||||
|
||||
|
|
@ -1372,13 +1393,7 @@ class _PROXY_MaxParallelRequestsHandler_v3(CustomLogger):
|
|||
response_obj, BaseLiteLLMOpenAIResponseObject
|
||||
):
|
||||
_usage = getattr(response_obj, "usage", None)
|
||||
if _usage and isinstance(_usage, Usage):
|
||||
if rate_limit_type == "output":
|
||||
total_tokens = _usage.completion_tokens
|
||||
elif rate_limit_type == "input":
|
||||
total_tokens = _usage.prompt_tokens
|
||||
elif rate_limit_type == "total":
|
||||
total_tokens = _usage.total_tokens
|
||||
total_tokens = self._get_total_tokens_from_usage(usage=_usage, rate_limit_type=rate_limit_type)
|
||||
|
||||
# Create pipeline operations for TPM increments
|
||||
pipeline_operations: List[RedisPipelineIncrementOperation] = []
|
||||
|
|
|
|||
|
|
@ -821,6 +821,40 @@ async def add_new_model(
|
|||
model_params: Deployment,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Add a new model to the proxy.
|
||||
|
||||
Parameters:
|
||||
- model_name: str - The name users will use to call this model (required)
|
||||
- litellm_params: dict - LiteLLM-specific parameters (required)
|
||||
- model: str - The actual model identifier, e.g., "azure/my-deployment-name" (required - this is the only required field in litellm_params)
|
||||
- api_key: str - API key for the provider (optional)
|
||||
- api_base: str - API base URL (optional)
|
||||
- Other optional params: api_version, timeout, max_retries, etc.
|
||||
- model_info: dict - Additional model metadata returned in /v1/model/info (optional)
|
||||
|
||||
Example curl:
|
||||
|
||||
```bash
|
||||
curl -L -X POST 'http://0.0.0.0:4000/model/new' \
|
||||
-H 'Authorization: Bearer LITELLM_VIRTUAL_KEY' \
|
||||
-H 'Content-Type: application/json' \
|
||||
-d '{
|
||||
"model_name": "my-azure-model",
|
||||
"litellm_params": {
|
||||
"model": "azure/my-deployment-name",
|
||||
"api_key": "my-azure-api-key",
|
||||
"api_base": "https://my-endpoint.openai.azure.com"
|
||||
},
|
||||
"model_info": {
|
||||
"my_custom_key": "my_custom_value"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
Returns:
|
||||
- The created model entry with model_id
|
||||
"""
|
||||
from litellm.proxy.proxy_server import (
|
||||
general_settings,
|
||||
premium_user,
|
||||
|
|
|
|||
|
|
@ -265,6 +265,10 @@ def generic_response_convertor(
|
|||
"GENERIC_USER_PROVIDER_ATTRIBUTE", "provider"
|
||||
)
|
||||
|
||||
generic_user_role_attribute_name = os.getenv(
|
||||
"GENERIC_USER_ROLE_ATTRIBUTE", "role"
|
||||
)
|
||||
|
||||
verbose_proxy_logger.debug(
|
||||
f" generic_user_id_attribute_name: {generic_user_id_attribute_name}\n generic_user_email_attribute_name: {generic_user_email_attribute_name}"
|
||||
)
|
||||
|
|
@ -277,6 +281,17 @@ def generic_response_convertor(
|
|||
team_ids = jwt_handler.get_team_ids_from_jwt(cast(dict, response))
|
||||
all_teams.extend(team_ids)
|
||||
|
||||
# Extract user role from SSO response
|
||||
user_role_from_sso = get_nested_value(response, generic_user_role_attribute_name)
|
||||
user_role: Optional[LitellmUserRoles] = None
|
||||
if user_role_from_sso is not None:
|
||||
role = get_litellm_user_role(user_role_from_sso)
|
||||
if role is not None:
|
||||
user_role = role
|
||||
verbose_proxy_logger.debug(
|
||||
f"Found valid LitellmUserRoles '{role.value}' from SSO attribute '{generic_user_role_attribute_name}'"
|
||||
)
|
||||
|
||||
return CustomOpenID(
|
||||
id=get_nested_value(response, generic_user_id_attribute_name),
|
||||
display_name=get_nested_value(
|
||||
|
|
@ -287,7 +302,7 @@ def generic_response_convertor(
|
|||
last_name=get_nested_value(response, generic_user_last_name_attribute_name),
|
||||
provider=get_nested_value(response, generic_provider_attribute_name),
|
||||
team_ids=all_teams,
|
||||
user_role=None,
|
||||
user_role=user_role,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -4,6 +4,10 @@ model_list:
|
|||
model: openai/gpt-4o-mini
|
||||
tpm: 1000
|
||||
|
||||
# LangGraph models
|
||||
- model_name: langgraph/*
|
||||
litellm_params:
|
||||
model: langgraph/*
|
||||
|
||||
litellm_settings:
|
||||
callbacks: ["dynamic_rate_limiter_v3"]
|
||||
|
|
|
|||
|
|
@ -170,8 +170,8 @@ from litellm.constants import (
|
|||
PROXY_BUDGET_RESCHEDULER_MIN_TIME,
|
||||
)
|
||||
from litellm.exceptions import RejectedRequestError
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.integrations.custom_guardrail import ModifyResponseException
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting
|
||||
from litellm.litellm_core_utils.core_helpers import (
|
||||
_get_parent_otel_span_from_kwargs,
|
||||
|
|
@ -1085,11 +1085,12 @@ def mount_swagger_ui():
|
|||
mount_swagger_ui()
|
||||
|
||||
docs_url = _get_docs_url()
|
||||
root_redirect_url = os.getenv("ROOT_REDIRECT_URL")
|
||||
if docs_url != "/" and root_redirect_url:
|
||||
root_redirect_url: Optional[str] = os.getenv("ROOT_REDIRECT_URL")
|
||||
if docs_url != "/" and root_redirect_url is not None:
|
||||
|
||||
@app.get("/", include_in_schema=False)
|
||||
async def root_redirect():
|
||||
return RedirectResponse(url=root_redirect_url)
|
||||
return RedirectResponse(url=root_redirect_url) # type: ignore[arg-type]
|
||||
|
||||
from typing import Dict
|
||||
|
||||
|
|
@ -4439,7 +4440,7 @@ class ProxyStartupEvent:
|
|||
### MONITOR SPEND LOGS QUEUE (queue-size-based job) ###
|
||||
if general_settings.get("disable_spend_logs", False) is False:
|
||||
from litellm.proxy.utils import _monitor_spend_logs_queue
|
||||
|
||||
|
||||
# Start background task to monitor spend logs queue size
|
||||
asyncio.create_task(
|
||||
_monitor_spend_logs_queue(
|
||||
|
|
@ -4485,63 +4486,12 @@ class ProxyStartupEvent:
|
|||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
)
|
||||
await proxy_config.get_credentials(prisma_client=prisma_client)
|
||||
if (
|
||||
proxy_logging_obj is not None
|
||||
and proxy_logging_obj.slack_alerting_instance.alerting is not None
|
||||
and prisma_client is not None
|
||||
):
|
||||
print("Alerting: Initializing Weekly/Monthly Spend Reports") # noqa
|
||||
### Schedule weekly/monthly spend reports ###
|
||||
### Schedule spend reports ###
|
||||
spend_report_frequency: str = (
|
||||
general_settings.get("spend_report_frequency", "7d") or "7d"
|
||||
)
|
||||
|
||||
# Parse the frequency
|
||||
days = int(spend_report_frequency[:-1])
|
||||
if spend_report_frequency[-1].lower() != "d":
|
||||
raise ValueError(
|
||||
"spend_report_frequency must be specified in days, e.g., '1d', '7d'"
|
||||
)
|
||||
|
||||
scheduler.add_job(
|
||||
proxy_logging_obj.slack_alerting_instance.send_weekly_spend_report,
|
||||
"interval",
|
||||
days=days,
|
||||
# REMOVED jitter parameter - major cause of memory leak
|
||||
# Use random start time instead for distribution
|
||||
next_run_time=datetime.now()
|
||||
+ timedelta(
|
||||
seconds=10 + random.randint(0, 300)
|
||||
), # Random 0-5 min offset
|
||||
args=[spend_report_frequency],
|
||||
id="weekly_spend_report_job",
|
||||
replace_existing=True,
|
||||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
)
|
||||
|
||||
scheduler.add_job(
|
||||
proxy_logging_obj.slack_alerting_instance.send_monthly_spend_report,
|
||||
"cron",
|
||||
day=1,
|
||||
id="monthly_spend_report_job",
|
||||
replace_existing=True,
|
||||
)
|
||||
|
||||
# Beta Feature - only used when prometheus api is in .env
|
||||
if os.getenv("PROMETHEUS_URL"):
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
scheduler.add_job(
|
||||
proxy_logging_obj.slack_alerting_instance.send_fallback_stats_from_prometheus,
|
||||
"cron",
|
||||
hour=PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS,
|
||||
minute=0,
|
||||
timezone=ZoneInfo("America/Los_Angeles"), # Pacific Time
|
||||
id="prometheus_fallback_stats_job",
|
||||
replace_existing=True,
|
||||
)
|
||||
await proxy_logging_obj.slack_alerting_instance.send_fallback_stats_from_prometheus()
|
||||
await cls._initialize_slack_alerting_jobs(
|
||||
scheduler=scheduler,
|
||||
general_settings=general_settings,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
prisma_client=prisma_client,
|
||||
)
|
||||
|
||||
await cls._initialize_spend_tracking_background_jobs(scheduler=scheduler)
|
||||
|
||||
|
|
@ -4681,6 +4631,65 @@ class ProxyStartupEvent:
|
|||
"Key rotation disabled (set LITELLM_KEY_ROTATION_ENABLED=true to enable)"
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def _initialize_slack_alerting_jobs(
|
||||
cls,
|
||||
scheduler: AsyncIOScheduler,
|
||||
general_settings: dict,
|
||||
proxy_logging_obj: ProxyLogging,
|
||||
prisma_client: PrismaClient,
|
||||
):
|
||||
"""Initialize Slack alerting background jobs for spend reports."""
|
||||
if (
|
||||
proxy_logging_obj is not None
|
||||
and proxy_logging_obj.slack_alerting_instance.alerting is not None
|
||||
and prisma_client is not None
|
||||
):
|
||||
print("Alerting: Initializing Weekly/Monthly Spend Reports") # noqa
|
||||
spend_report_frequency: str = (
|
||||
general_settings.get("spend_report_frequency", "7d") or "7d"
|
||||
)
|
||||
|
||||
days = int(spend_report_frequency[:-1])
|
||||
if spend_report_frequency[-1].lower() != "d":
|
||||
raise ValueError(
|
||||
"spend_report_frequency must be specified in days, e.g., '1d', '7d'"
|
||||
)
|
||||
|
||||
scheduler.add_job(
|
||||
proxy_logging_obj.slack_alerting_instance.send_weekly_spend_report,
|
||||
"interval",
|
||||
days=days,
|
||||
next_run_time=datetime.now()
|
||||
+ timedelta(seconds=10 + random.randint(0, 300)),
|
||||
args=[spend_report_frequency],
|
||||
id="weekly_spend_report_job",
|
||||
replace_existing=True,
|
||||
misfire_grace_time=APSCHEDULER_MISFIRE_GRACE_TIME,
|
||||
)
|
||||
|
||||
scheduler.add_job(
|
||||
proxy_logging_obj.slack_alerting_instance.send_monthly_spend_report,
|
||||
"cron",
|
||||
day=1,
|
||||
id="monthly_spend_report_job",
|
||||
replace_existing=True,
|
||||
)
|
||||
|
||||
if os.getenv("PROMETHEUS_URL"):
|
||||
from zoneinfo import ZoneInfo
|
||||
|
||||
scheduler.add_job(
|
||||
proxy_logging_obj.slack_alerting_instance.send_fallback_stats_from_prometheus,
|
||||
"cron",
|
||||
hour=PROMETHEUS_FALLBACK_STATS_SEND_TIME_HOURS,
|
||||
minute=0,
|
||||
timezone=ZoneInfo("America/Los_Angeles"),
|
||||
id="prometheus_fallback_stats_job",
|
||||
replace_existing=True,
|
||||
)
|
||||
await proxy_logging_obj.slack_alerting_instance.send_fallback_stats_from_prometheus()
|
||||
|
||||
@classmethod
|
||||
async def _setup_prisma_client(
|
||||
cls,
|
||||
|
|
@ -5123,14 +5132,14 @@ async def completion( # noqa: PLR0915
|
|||
|
||||
if _data.get("stream", None) is not None and _data["stream"] is True:
|
||||
_text_response = litellm.ModelResponse()
|
||||
_text_response.choices[0].text = e.message
|
||||
_text_response.model = e.model # type: ignore
|
||||
_text_response.choices[0].text = e.message # type: ignore[attr-defined]
|
||||
_text_response.model = e.model # type: ignore[assignment]
|
||||
_usage = litellm.Usage(
|
||||
prompt_tokens=0,
|
||||
completion_tokens=0,
|
||||
total_tokens=0,
|
||||
)
|
||||
_text_response.usage = _usage # type: ignore
|
||||
_text_response.usage = _usage # type: ignore[assignment]
|
||||
_iterator = litellm.utils.ModelResponseIterator(
|
||||
model_response=_text_response, convert_to_delta=True
|
||||
)
|
||||
|
|
|
|||
76
litellm/proxy/public_endpoints/agent_create_fields.json
Normal file
76
litellm/proxy/public_endpoints/agent_create_fields.json
Normal file
|
|
@ -0,0 +1,76 @@
|
|||
[
|
||||
{
|
||||
"agent_type": "a2a",
|
||||
"agent_type_display_name": "A2A Standard",
|
||||
"description": "Standard A2A protocol",
|
||||
"logo_url": "/assets/logos/a2a_agent.png",
|
||||
"credential_fields": [],
|
||||
"litellm_params_template": {}
|
||||
},
|
||||
{
|
||||
"agent_type": "langgraph",
|
||||
"agent_type_display_name": "LangGraph",
|
||||
"description": "Connect to LangGraph agents via the LangGraph Platform API",
|
||||
"logo_url": "/assets/logos/langgraph.png",
|
||||
"model_template": "langgraph/{assistant_id}",
|
||||
"credential_fields": [
|
||||
{
|
||||
"key": "assistant_id",
|
||||
"label": "Assistant ID",
|
||||
"placeholder": "agent",
|
||||
"tooltip": "The assistant/agent ID from your LangGraph deployment",
|
||||
"required": true,
|
||||
"field_type": "text",
|
||||
"default_value": "agent",
|
||||
"include_in_litellm_params": false
|
||||
},
|
||||
{
|
||||
"key": "api_base",
|
||||
"label": "LangGraph API Base",
|
||||
"placeholder": "http://localhost:2024",
|
||||
"tooltip": "The base URL for your LangGraph server (e.g., http://localhost:2024 or your deployed LangGraph Cloud URL)",
|
||||
"required": true,
|
||||
"field_type": "text",
|
||||
"default_value": "http://localhost:2024",
|
||||
"include_in_litellm_params": true
|
||||
},
|
||||
{
|
||||
"key": "api_key",
|
||||
"label": "LangGraph API Key",
|
||||
"placeholder": null,
|
||||
"tooltip": "API key for authenticating with your LangGraph server (optional for local development)",
|
||||
"required": false,
|
||||
"field_type": "password",
|
||||
"default_value": null,
|
||||
"include_in_litellm_params": true
|
||||
}
|
||||
],
|
||||
"litellm_params_template": {
|
||||
"custom_llm_provider": "langgraph"
|
||||
}
|
||||
},
|
||||
{
|
||||
"agent_type": "bedrock_agentcore",
|
||||
"agent_type_display_name": "Bedrock AgentCore",
|
||||
"description": "Connect to Amazon Bedrock AgentCore hosted agent runtimes",
|
||||
"logo_url": "/assets/logos/bedrock.svg",
|
||||
"inherit_credentials_from_provider": "Bedrock",
|
||||
"model_template": "bedrock/agentcore/{agent_runtime_arn}",
|
||||
"credential_fields": [
|
||||
{
|
||||
"key": "agent_runtime_arn",
|
||||
"label": "Agent Runtime ARN",
|
||||
"placeholder": "arn:aws:bedrock-agentcore:us-west-2:123456789012:runtime/my-agent-runtime",
|
||||
"tooltip": "The ARN of your Bedrock AgentCore runtime. Find this in your AWS Bedrock console under AgentCore.",
|
||||
"required": true,
|
||||
"field_type": "text",
|
||||
"default_value": null,
|
||||
"include_in_litellm_params": false
|
||||
}
|
||||
],
|
||||
"litellm_params_template": {
|
||||
"custom_llm_provider": "bedrock"
|
||||
}
|
||||
}
|
||||
]
|
||||
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
from typing import List
|
||||
import os
|
||||
import json
|
||||
import os
|
||||
from typing import List
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
|
|
@ -12,6 +12,7 @@ from litellm.types.proxy.management_endpoints.model_management_endpoints import
|
|||
ModelGroupInfoProxy,
|
||||
)
|
||||
from litellm.types.proxy.public_endpoints.public_endpoints import (
|
||||
AgentCreateInfo,
|
||||
ProviderCreateInfo,
|
||||
PublicModelHubInfo,
|
||||
)
|
||||
|
|
@ -167,3 +168,52 @@ async def get_litellm_model_cost_map():
|
|||
status_code=500,
|
||||
detail=f"Internal Server Error ({str(e)})",
|
||||
)
|
||||
|
||||
|
||||
@router.get(
|
||||
"/public/agents/fields",
|
||||
tags=["public", "[beta] Agents"],
|
||||
response_model=List[AgentCreateInfo],
|
||||
)
|
||||
async def get_agent_fields() -> List[AgentCreateInfo]:
|
||||
"""
|
||||
Return agent type metadata required by the dashboard create-agent flow.
|
||||
|
||||
If an agent has `inherit_credentials_from_provider`, the provider's credential
|
||||
fields are automatically appended to the agent's credential_fields.
|
||||
"""
|
||||
base_path = os.path.join(
|
||||
os.path.dirname(os.path.dirname(os.path.dirname(__file__))),
|
||||
"proxy",
|
||||
"public_endpoints",
|
||||
)
|
||||
|
||||
agent_create_fields_path = os.path.join(base_path, "agent_create_fields.json")
|
||||
provider_create_fields_path = os.path.join(base_path, "provider_create_fields.json")
|
||||
|
||||
with open(agent_create_fields_path, "r") as f:
|
||||
agent_create_fields = json.load(f)
|
||||
|
||||
with open(provider_create_fields_path, "r") as f:
|
||||
provider_create_fields = json.load(f)
|
||||
|
||||
# Build a lookup map for providers by name
|
||||
provider_map = {p["provider"]: p for p in provider_create_fields}
|
||||
|
||||
# Merge inherited credential fields
|
||||
for agent in agent_create_fields:
|
||||
inherit_from = agent.get("inherit_credentials_from_provider")
|
||||
if inherit_from and inherit_from in provider_map:
|
||||
provider = provider_map[inherit_from]
|
||||
# Copy provider fields and mark them for inclusion in litellm_params
|
||||
inherited_fields = []
|
||||
for field in provider.get("credential_fields", []):
|
||||
field_copy = field.copy()
|
||||
field_copy["include_in_litellm_params"] = True
|
||||
inherited_fields.append(field_copy)
|
||||
# Append provider credential fields after agent's own fields
|
||||
agent["credential_fields"] = agent.get("credential_fields", []) + inherited_fields
|
||||
# Remove the inherit field from response (not needed by frontend)
|
||||
agent.pop("inherit_credentials_from_provider", None)
|
||||
|
||||
return agent_create_fields
|
||||
|
|
|
|||
|
|
@ -315,6 +315,7 @@ model LiteLLM_SpendLogs {
|
|||
session_id String?
|
||||
status String?
|
||||
mcp_namespaced_tool_name String?
|
||||
agent_id String?
|
||||
proxy_server_request Json? @default("{}")
|
||||
@@index([startTime])
|
||||
@@index([end_user])
|
||||
|
|
|
|||
|
|
@ -225,13 +225,16 @@ def get_logging_payload( # noqa: PLR0915
|
|||
response_obj_dict = {}
|
||||
|
||||
# Handle OCR responses which use usage_info instead of usage
|
||||
usage: dict = {}
|
||||
if call_type in ["ocr", "aocr"]:
|
||||
usage = _extract_usage_for_ocr_call(response_obj, response_obj_dict)
|
||||
else:
|
||||
# Use response_obj_dict instead of response_obj to avoid calling .get() on Pydantic models
|
||||
usage = response_obj_dict.get("usage", None) or {}
|
||||
if isinstance(usage, litellm.Usage):
|
||||
usage = dict(usage)
|
||||
_usage = response_obj_dict.get("usage", None) or {}
|
||||
if isinstance(_usage, litellm.Usage):
|
||||
usage = dict(_usage)
|
||||
elif isinstance(_usage, dict):
|
||||
usage = _usage
|
||||
|
||||
id = get_spend_logs_id(call_type or "acompletion", response_obj_dict, kwargs)
|
||||
standard_logging_payload = cast(
|
||||
|
|
@ -369,6 +372,9 @@ def get_logging_payload( # noqa: PLR0915
|
|||
"namespaced_tool_name", None
|
||||
)
|
||||
|
||||
# Extract agent_id for A2A requests (set directly on model_call_details)
|
||||
agent_id: Optional[str] = kwargs.get("agent_id")
|
||||
|
||||
try:
|
||||
payload: SpendLogsPayload = SpendLogsPayload(
|
||||
request_id=str(id),
|
||||
|
|
@ -396,6 +402,7 @@ def get_logging_payload( # noqa: PLR0915
|
|||
model_group=_model_group,
|
||||
model_id=_model_id,
|
||||
mcp_namespaced_tool_name=mcp_namespaced_tool_name,
|
||||
agent_id=agent_id,
|
||||
requester_ip_address=clean_metadata.get("requester_ip_address", None),
|
||||
custom_llm_provider=kwargs.get("custom_llm_provider", ""),
|
||||
messages=_get_messages_for_spend_logs_payload(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,5 @@
|
|||
import asyncio
|
||||
import copy
|
||||
import gc
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
|
|
@ -3379,7 +3378,6 @@ class ProxyUpdateSpend:
|
|||
headers={"Content-Type": "application/json"},
|
||||
)
|
||||
del json_data
|
||||
gc.collect()
|
||||
if response.status_code == 200:
|
||||
prisma_client.spend_log_transactions = (
|
||||
prisma_client.spend_log_transactions[
|
||||
|
|
@ -3401,9 +3399,6 @@ class ProxyUpdateSpend:
|
|||
)
|
||||
# Explicitly clear batch memory
|
||||
del batch, batch_with_dates
|
||||
# Only run gc every 5 batches to reduce overhead
|
||||
if j % (BATCH_SIZE * 5) == 0:
|
||||
gc.collect()
|
||||
|
||||
prisma_client.spend_log_transactions = (
|
||||
prisma_client.spend_log_transactions[len(logs_to_process) :]
|
||||
|
|
@ -3429,7 +3424,6 @@ class ProxyUpdateSpend:
|
|||
finally:
|
||||
# Clean up logs_to_process after all processing is complete
|
||||
del logs_to_process
|
||||
gc.collect()
|
||||
|
||||
@staticmethod
|
||||
def disable_spend_updates() -> bool:
|
||||
|
|
|
|||
|
|
@ -221,6 +221,9 @@ class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase):
|
|||
result: Optional[Dict[str, Any]] = None
|
||||
error: Optional[Dict[str, Any]] = None
|
||||
|
||||
# LiteLLM usage tracking
|
||||
usage: Optional[Dict[str, Any]] = None
|
||||
|
||||
model_config = {"extra": "allow"}
|
||||
|
||||
# LiteLLM private attributes for logging/cost tracking
|
||||
|
|
@ -243,3 +246,16 @@ class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase):
|
|||
response_dict = response.model_dump(mode="json", exclude_none=True)
|
||||
|
||||
return cls(**response_dict)
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, response_dict: Dict[str, Any]) -> "LiteLLMSendMessageResponse":
|
||||
"""
|
||||
Create a LiteLLMSendMessageResponse from a dict.
|
||||
|
||||
Args:
|
||||
response_dict: Dict with A2A response structure
|
||||
|
||||
Returns:
|
||||
LiteLLMSendMessageResponse with _hidden_params support
|
||||
"""
|
||||
return cls(**response_dict)
|
||||
|
|
|
|||
68
litellm/types/llms/langgraph.py
Normal file
68
litellm/types/llms/langgraph.py
Normal file
|
|
@ -0,0 +1,68 @@
|
|||
"""
|
||||
Type definitions for LangGraph API.
|
||||
|
||||
LangGraph provides a streaming and non-streaming API for running agents.
|
||||
"""
|
||||
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
from typing_extensions import Literal, TypedDict
|
||||
|
||||
|
||||
# Request Types
|
||||
class LangGraphMessage(TypedDict, total=False):
|
||||
"""Message format for LangGraph input."""
|
||||
|
||||
role: Literal["human", "assistant", "system"]
|
||||
content: str
|
||||
|
||||
|
||||
class LangGraphInput(TypedDict, total=False):
|
||||
"""Input structure for LangGraph request."""
|
||||
|
||||
messages: List[LangGraphMessage]
|
||||
|
||||
|
||||
class LangGraphRequest(TypedDict, total=False):
|
||||
"""Request structure for LangGraph API."""
|
||||
|
||||
assistant_id: str
|
||||
input: LangGraphInput
|
||||
stream_mode: Optional[str]
|
||||
config: Optional[Dict[str, Any]]
|
||||
metadata: Optional[Dict[str, Any]]
|
||||
|
||||
|
||||
# Response Types - Streaming
|
||||
class LangGraphStreamEvent(TypedDict, total=False):
|
||||
"""Single event in a LangGraph stream response."""
|
||||
|
||||
event: str
|
||||
data: Any
|
||||
|
||||
|
||||
# Response Types - Non-streaming
|
||||
class LangGraphResponseMessage(TypedDict, total=False):
|
||||
"""Message in LangGraph response."""
|
||||
|
||||
type: str
|
||||
content: str
|
||||
id: Optional[str]
|
||||
name: Optional[str]
|
||||
|
||||
|
||||
class LangGraphResponse(TypedDict, total=False):
|
||||
"""Non-streaming response structure from LangGraph."""
|
||||
|
||||
messages: List[LangGraphResponseMessage]
|
||||
values: Dict[str, Any]
|
||||
|
||||
|
||||
# Parsed response for internal use
|
||||
class LangGraphParsedResponse(TypedDict):
|
||||
"""Parsed response from LangGraph."""
|
||||
|
||||
content: str
|
||||
role: str
|
||||
usage: Optional[Dict[str, int]]
|
||||
|
||||
|
|
@ -27,3 +27,25 @@ class ProviderCreateInfo(BaseModel):
|
|||
litellm_provider: str
|
||||
credential_fields: List[ProviderCredentialField]
|
||||
default_model_placeholder: Optional[str] = None
|
||||
|
||||
|
||||
class AgentCredentialField(BaseModel):
|
||||
key: str
|
||||
label: str
|
||||
placeholder: Optional[str] = None
|
||||
tooltip: Optional[str] = None
|
||||
required: bool = False
|
||||
field_type: Literal["text", "password", "select", "upload", "textarea"] = "text"
|
||||
options: Optional[List[str]] = None
|
||||
default_value: Optional[str] = None
|
||||
include_in_litellm_params: Optional[bool] = None
|
||||
|
||||
|
||||
class AgentCreateInfo(BaseModel):
|
||||
agent_type: str
|
||||
agent_type_display_name: str
|
||||
description: Optional[str] = None
|
||||
logo_url: Optional[str] = None
|
||||
credential_fields: List[AgentCredentialField]
|
||||
litellm_params_template: Optional[Dict[str, str]] = None
|
||||
model_template: Optional[str] = None
|
||||
|
|
|
|||
|
|
@ -3008,6 +3008,7 @@ class LlmProviders(str, Enum):
|
|||
LEMONADE = "lemonade"
|
||||
AMAZON_NOVA = "amazon_nova"
|
||||
A2A_AGENT = "a2a_agent"
|
||||
LANGGRAPH = "langgraph"
|
||||
|
||||
|
||||
# Create a set of all provider values for quick lookup
|
||||
|
|
|
|||
|
|
@ -97,6 +97,10 @@ from litellm.litellm_core_utils.core_helpers import (
|
|||
)
|
||||
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
|
||||
from litellm.litellm_core_utils.default_encoding import encoding
|
||||
from litellm.litellm_core_utils.dot_notation_indexing import (
|
||||
delete_nested_value,
|
||||
is_nested_path,
|
||||
)
|
||||
from litellm.litellm_core_utils.exception_mapping_utils import (
|
||||
_get_response_headers,
|
||||
exception_type,
|
||||
|
|
@ -138,10 +142,6 @@ from litellm.litellm_core_utils.redact_messages import (
|
|||
LiteLLMLoggingObject,
|
||||
redact_message_input_output_from_logging,
|
||||
)
|
||||
from litellm.litellm_core_utils.dot_notation_indexing import (
|
||||
delete_nested_value,
|
||||
is_nested_path,
|
||||
)
|
||||
from litellm.litellm_core_utils.rules import Rules
|
||||
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
|
||||
from litellm.litellm_core_utils.token_counter import get_modified_max_tokens
|
||||
|
|
@ -7288,6 +7288,10 @@ class ProviderConfigManager:
|
|||
return litellm.OVHCloudChatConfig()
|
||||
elif litellm.LlmProviders.AMAZON_NOVA == provider:
|
||||
return litellm.AmazonNovaChatConfig()
|
||||
elif litellm.LlmProviders.LANGGRAPH == provider:
|
||||
from litellm.llms.langgraph.chat.transformation import LangGraphConfig
|
||||
|
||||
return LangGraphConfig()
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -271,7 +271,6 @@ def video_generation( # noqa: PLR0915
|
|||
@client
|
||||
def video_content(
|
||||
video_id: str,
|
||||
api_base: Optional[str] = None,
|
||||
timeout: Optional[float] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
|
|
@ -384,8 +383,6 @@ def video_content(
|
|||
@client
|
||||
async def avideo_content(
|
||||
video_id: str,
|
||||
api_key: Optional[str] = None,
|
||||
api_base: Optional[str] = None,
|
||||
timeout: Optional[float] = None,
|
||||
custom_llm_provider: Optional[str] = None,
|
||||
# Use the following arguments if you need to pass additional parameters to the API that aren't available via kwargs.
|
||||
|
|
@ -400,8 +397,6 @@ async def avideo_content(
|
|||
|
||||
Parameters:
|
||||
- `video_id` (str): The identifier of the video whose content to download
|
||||
- `api_key` (Optional[str]): The API key to use for authentication
|
||||
- `api_base` (Optional[str]): The base URL for the API
|
||||
- `timeout` (Optional[float]): The timeout for the request in seconds
|
||||
- `custom_llm_provider` (Optional[str]): The LLM provider to use
|
||||
- `extra_headers` (Optional[Dict[str, Any]]): Additional headers
|
||||
|
|
@ -425,8 +420,6 @@ async def avideo_content(
|
|||
func = partial(
|
||||
video_content,
|
||||
video_id=video_id,
|
||||
api_key=api_key,
|
||||
api_base=api_base,
|
||||
timeout=timeout,
|
||||
custom_llm_provider=custom_llm_provider,
|
||||
extra_headers=extra_headers,
|
||||
|
|
|
|||
|
|
@ -18810,6 +18810,20 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"mistral/codestral-2508": {
|
||||
"input_cost_per_token": 3e-07,
|
||||
"litellm_provider": "mistral",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 9e-07,
|
||||
"source": "https://mistral.ai/news/codestral-25-08",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"mistral/codestral-latest": {
|
||||
"input_cost_per_token": 1e-06,
|
||||
"litellm_provider": "mistral",
|
||||
|
|
@ -18876,6 +18890,34 @@
|
|||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"mistral/labs-devstral-small-2512": {
|
||||
"input_cost_per_token": 1e-07,
|
||||
"litellm_provider": "mistral",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 3e-07,
|
||||
"source": "https://docs.mistral.ai/models/devstral-small-2-25-12",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"mistral/devstral-2512": {
|
||||
"input_cost_per_token": 4e-07,
|
||||
"litellm_provider": "mistral",
|
||||
"max_input_tokens": 256000,
|
||||
"max_output_tokens": 256000,
|
||||
"max_tokens": 256000,
|
||||
"mode": "chat",
|
||||
"output_cost_per_token": 2e-06,
|
||||
"source": "https://mistral.ai/news/devstral-2-vibe-cli",
|
||||
"supports_assistant_prefill": true,
|
||||
"supports_function_calling": true,
|
||||
"supports_response_schema": true,
|
||||
"supports_tool_choice": true
|
||||
},
|
||||
"mistral/magistral-medium-2506": {
|
||||
"input_cost_per_token": 2e-06,
|
||||
"litellm_provider": "mistral",
|
||||
|
|
|
|||
2100
poetry.lock
generated
2100
poetry.lock
generated
File diff suppressed because it is too large
Load diff
|
|
@ -1910,6 +1910,23 @@
|
|||
"rerank": false,
|
||||
"a2a": true
|
||||
}
|
||||
},
|
||||
"langgraph": {
|
||||
"display_name": "LangGraph (`langgraph`)",
|
||||
"url": "https://docs.litellm.ai/docs/providers/langgraph",
|
||||
"endpoints": {
|
||||
"chat_completions": true,
|
||||
"messages": true,
|
||||
"responses": true,
|
||||
"embeddings": false,
|
||||
"image_generations": false,
|
||||
"audio_transcriptions": false,
|
||||
"audio_speech": false,
|
||||
"moderations": false,
|
||||
"batches": false,
|
||||
"rerank": false,
|
||||
"a2a": true
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -315,6 +315,7 @@ model LiteLLM_SpendLogs {
|
|||
session_id String?
|
||||
status String?
|
||||
mcp_namespaced_tool_name String?
|
||||
agent_id String?
|
||||
proxy_server_request Json? @default("{}")
|
||||
@@index([startTime])
|
||||
@@index([end_user])
|
||||
|
|
|
|||
203
tests/agent_tests/test_a2a_completion_bridge.py
Normal file
203
tests/agent_tests/test_a2a_completion_bridge.py
Normal file
|
|
@ -0,0 +1,203 @@
|
|||
"""
|
||||
Test for A2A to LiteLLM Completion Bridge.
|
||||
|
||||
Tests the SDK-level functions that route A2A requests through litellm.acompletion.
|
||||
|
||||
Run with:
|
||||
pytest tests/agent_tests/test_a2a_completion_bridge.py -v -s
|
||||
|
||||
Prerequisites:
|
||||
- LangGraph server running on localhost:2024
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import litellm
|
||||
from a2a.types import MessageSendParams, SendMessageRequest, SendStreamingMessageRequest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a2a_completion_bridge_non_streaming():
|
||||
"""
|
||||
Test non-streaming A2A request via the completion bridge with LangGraph provider.
|
||||
"""
|
||||
from litellm.a2a_protocol import asend_message
|
||||
|
||||
litellm._turn_on_debug()
|
||||
|
||||
send_message_payload = {
|
||||
"message": {
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "What is 2 + 2?"}],
|
||||
"messageId": uuid4().hex,
|
||||
}
|
||||
}
|
||||
|
||||
request = SendMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(**send_message_payload), # type: ignore
|
||||
)
|
||||
|
||||
response = await asend_message(
|
||||
request=request,
|
||||
api_base="http://localhost:2024",
|
||||
litellm_params={"custom_llm_provider": "langgraph", "model": "agent"},
|
||||
)
|
||||
|
||||
# Validate response is LiteLLMSendMessageResponse
|
||||
assert response.jsonrpc == "2.0"
|
||||
assert response.id is not None
|
||||
assert response.result is not None
|
||||
assert "message" in response.result
|
||||
|
||||
message = response.result["message"]
|
||||
assert "role" in message
|
||||
assert message["role"] == "agent"
|
||||
assert "parts" in message
|
||||
assert len(message["parts"]) > 0
|
||||
assert message["parts"][0]["kind"] == "text"
|
||||
assert len(message["parts"][0]["text"]) > 0
|
||||
|
||||
print(f"Response: {response.model_dump(mode='json', exclude_none=True)}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a2a_completion_bridge_streaming():
|
||||
"""
|
||||
Test streaming A2A request via the completion bridge with LangGraph provider.
|
||||
|
||||
Validates proper A2A streaming format with events:
|
||||
1. Task event (kind: "task") - Initial task with status "submitted"
|
||||
2. Status update (kind: "status-update") - Status "working"
|
||||
3. Artifact update (kind: "artifact-update") - Content delivery
|
||||
4. Status update (kind: "status-update") - Final "completed" status
|
||||
"""
|
||||
from litellm.a2a_protocol import asend_message_streaming
|
||||
|
||||
litellm._turn_on_debug()
|
||||
|
||||
send_message_payload = {
|
||||
"message": {
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Count from 1 to 5."}],
|
||||
"messageId": uuid4().hex,
|
||||
}
|
||||
}
|
||||
|
||||
request = SendStreamingMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(**send_message_payload), # type: ignore
|
||||
)
|
||||
|
||||
chunks = []
|
||||
async for chunk in asend_message_streaming(
|
||||
request=request,
|
||||
api_base="http://localhost:2024",
|
||||
litellm_params={"custom_llm_provider": "langgraph", "model": "agent"},
|
||||
):
|
||||
chunks.append(chunk)
|
||||
print(f"Chunk: {chunk}")
|
||||
|
||||
# Validate we received proper A2A streaming events
|
||||
assert len(chunks) >= 4, f"Expected at least 4 chunks (task, working, artifact, completed), got {len(chunks)}"
|
||||
|
||||
# Validate chunk structure follows A2A spec
|
||||
for chunk in chunks:
|
||||
assert "jsonrpc" in chunk
|
||||
assert chunk["jsonrpc"] == "2.0"
|
||||
assert "id" in chunk
|
||||
assert "result" in chunk
|
||||
|
||||
# Validate first chunk is task event
|
||||
task_chunk = chunks[0]
|
||||
assert task_chunk["result"]["kind"] == "task", "First chunk should be task event"
|
||||
assert task_chunk["result"]["status"]["state"] == "submitted"
|
||||
assert "contextId" in task_chunk["result"]
|
||||
assert "id" in task_chunk["result"] # task id
|
||||
assert "history" in task_chunk["result"]
|
||||
|
||||
# Validate second chunk is working status update
|
||||
working_chunk = chunks[1]
|
||||
assert working_chunk["result"]["kind"] == "status-update", "Second chunk should be status-update"
|
||||
assert working_chunk["result"]["status"]["state"] == "working"
|
||||
assert "taskId" in working_chunk["result"]
|
||||
assert "contextId" in working_chunk["result"]
|
||||
assert working_chunk["result"]["final"] is False
|
||||
|
||||
# Validate artifact update chunk
|
||||
artifact_chunk = chunks[2]
|
||||
assert artifact_chunk["result"]["kind"] == "artifact-update", "Third chunk should be artifact-update"
|
||||
assert "artifact" in artifact_chunk["result"]
|
||||
assert "artifactId" in artifact_chunk["result"]["artifact"]
|
||||
assert "parts" in artifact_chunk["result"]["artifact"]
|
||||
assert len(artifact_chunk["result"]["artifact"]["parts"]) > 0
|
||||
assert artifact_chunk["result"]["artifact"]["parts"][0]["kind"] == "text"
|
||||
|
||||
# Validate final chunk is completed status update
|
||||
final_chunk = chunks[-1]
|
||||
assert final_chunk["result"]["kind"] == "status-update", "Last chunk should be status-update"
|
||||
assert final_chunk["result"]["status"]["state"] == "completed"
|
||||
assert final_chunk["result"]["final"] is True
|
||||
|
||||
print(f"Received {len(chunks)} chunks with proper A2A streaming format")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_a2a_completion_bridge_bedrock_agentcore():
|
||||
"""
|
||||
Test A2A request via the completion bridge with Bedrock AgentCore provider.
|
||||
|
||||
Uses the AgentCore runtime ARN to call a hosted agent.
|
||||
"""
|
||||
from litellm.a2a_protocol import asend_message_streaming
|
||||
|
||||
litellm._turn_on_debug()
|
||||
|
||||
# Bedrock AgentCore ARN (streaming-capable runtime)
|
||||
agentcore_arn = "arn:aws:bedrock-agentcore:us-west-2:888602223428:runtime/hosted_agent_r9jvp-3ySZuRHjLC"
|
||||
|
||||
send_message_payload = {
|
||||
"message": {
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Explain machine learning in simple terms"}],
|
||||
"messageId": uuid4().hex,
|
||||
}
|
||||
}
|
||||
|
||||
request = SendStreamingMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(**send_message_payload), # type: ignore
|
||||
)
|
||||
|
||||
chunks = []
|
||||
async for chunk in asend_message_streaming(
|
||||
request=request,
|
||||
api_base=None, # Not needed for Bedrock AgentCore
|
||||
litellm_params={
|
||||
"custom_llm_provider": "bedrock",
|
||||
"model": f"bedrock/agentcore/{agentcore_arn}",
|
||||
},
|
||||
):
|
||||
chunks.append(chunk)
|
||||
print(f"Chunk: {chunk}")
|
||||
|
||||
# Validate we received proper A2A streaming events
|
||||
assert len(chunks) >= 4, f"Expected at least 4 chunks, got {len(chunks)}"
|
||||
|
||||
# Validate first chunk is task event
|
||||
assert chunks[0]["result"]["kind"] == "task"
|
||||
assert chunks[0]["result"]["status"]["state"] == "submitted"
|
||||
|
||||
# Validate final chunk is completed status
|
||||
assert chunks[-1]["result"]["kind"] == "status-update"
|
||||
assert chunks[-1]["result"]["status"]["state"] == "completed"
|
||||
assert chunks[-1]["result"]["final"] is True
|
||||
|
||||
print(f"Received {len(chunks)} chunks from Bedrock AgentCore")
|
||||
|
||||
173
tests/llm_translation/test_langgraph.py
Normal file
173
tests/llm_translation/test_langgraph.py
Normal file
|
|
@ -0,0 +1,173 @@
|
|||
"""
|
||||
Tests for LangGraph provider integration.
|
||||
|
||||
These tests require a LangGraph server running locally on port 2024.
|
||||
To start a LangGraph server, follow the LangGraph documentation.
|
||||
|
||||
Example test server curl commands:
|
||||
Streaming:
|
||||
curl -s --request POST \
|
||||
--url "http://localhost:2024/runs/stream" \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{"assistant_id": "agent", "input": {"messages": [{"role": "human", "content": "What is 25 * 4?"}]}, "stream_mode": "messages-tuple"}'
|
||||
|
||||
Non-streaming:
|
||||
curl -s --request POST \
|
||||
--url "http://localhost:2024/runs/wait" \
|
||||
--header 'Content-Type: application/json' \
|
||||
--data '{"assistant_id": "agent", "input": {"messages": [{"role": "human", "content": "What is 25 * 4?"}]}}'
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_langgraph_acompletion_non_streaming():
|
||||
"""
|
||||
Test non-streaming acompletion call to LangGraph server.
|
||||
Uses the /runs/wait endpoint for synchronous response.
|
||||
"""
|
||||
api_base = os.environ.get("LANGGRAPH_API_BASE", "http://localhost:2024")
|
||||
|
||||
try:
|
||||
response = await litellm.acompletion(
|
||||
model="langgraph/agent",
|
||||
messages=[{"role": "user", "content": "What is 25 * 4?"}],
|
||||
api_base=api_base,
|
||||
stream=False,
|
||||
)
|
||||
|
||||
assert response is not None
|
||||
assert response.choices is not None
|
||||
assert len(response.choices) > 0
|
||||
assert response.choices[0].message is not None
|
||||
assert response.choices[0].message.content is not None
|
||||
assert len(response.choices[0].message.content) > 0
|
||||
|
||||
except Exception as e:
|
||||
pytest.skip(f"LangGraph server not available: {e}")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_langgraph_acompletion_streaming():
|
||||
"""
|
||||
Test streaming acompletion call to LangGraph server.
|
||||
Uses the /runs/stream endpoint with stream_mode="messages-tuple".
|
||||
"""
|
||||
api_base = os.environ.get("LANGGRAPH_API_BASE", "http://localhost:2024")
|
||||
|
||||
try:
|
||||
response = await litellm.acompletion(
|
||||
model="langgraph/agent",
|
||||
messages=[{"role": "user", "content": "What is the weather in Tokyo?"}],
|
||||
api_base=api_base,
|
||||
stream=True,
|
||||
)
|
||||
|
||||
full_content = ""
|
||||
chunk_count = 0
|
||||
|
||||
async for chunk in response:
|
||||
chunk_count += 1
|
||||
if (
|
||||
chunk.choices
|
||||
and chunk.choices[0].delta
|
||||
and chunk.choices[0].delta.content
|
||||
):
|
||||
full_content += chunk.choices[0].delta.content
|
||||
|
||||
assert chunk_count > 0, "Should receive at least one chunk"
|
||||
|
||||
except Exception as e:
|
||||
pytest.skip(f"LangGraph server not available: {e}")
|
||||
|
||||
|
||||
def test_langgraph_config_get_complete_url():
|
||||
"""
|
||||
Test that LangGraphConfig correctly generates URLs for streaming and non-streaming.
|
||||
"""
|
||||
from litellm.llms.langgraph.chat.transformation import LangGraphConfig
|
||||
|
||||
config = LangGraphConfig()
|
||||
|
||||
non_streaming_url = config.get_complete_url(
|
||||
api_base="http://localhost:2024",
|
||||
api_key=None,
|
||||
model="agent",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
stream=False,
|
||||
)
|
||||
assert non_streaming_url == "http://localhost:2024/runs/wait"
|
||||
|
||||
streaming_url = config.get_complete_url(
|
||||
api_base="http://localhost:2024",
|
||||
api_key=None,
|
||||
model="agent",
|
||||
optional_params={},
|
||||
litellm_params={},
|
||||
stream=True,
|
||||
)
|
||||
assert streaming_url == "http://localhost:2024/runs/stream"
|
||||
|
||||
|
||||
def test_langgraph_config_transform_request():
|
||||
"""
|
||||
Test that LangGraphConfig correctly transforms requests.
|
||||
"""
|
||||
from litellm.llms.langgraph.chat.transformation import LangGraphConfig
|
||||
|
||||
config = LangGraphConfig()
|
||||
|
||||
messages = [
|
||||
{"role": "system", "content": "You are a helpful assistant."},
|
||||
{"role": "user", "content": "What is 2 + 2?"},
|
||||
]
|
||||
|
||||
request = config.transform_request(
|
||||
model="langgraph/agent",
|
||||
messages=messages,
|
||||
optional_params={},
|
||||
litellm_params={"stream": False},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert request["assistant_id"] == "agent"
|
||||
assert "input" in request
|
||||
assert "messages" in request["input"]
|
||||
assert len(request["input"]["messages"]) == 2
|
||||
assert request["input"]["messages"][0]["role"] == "system"
|
||||
assert request["input"]["messages"][1]["role"] == "human"
|
||||
|
||||
streaming_request = config.transform_request(
|
||||
model="langgraph/agent",
|
||||
messages=messages,
|
||||
optional_params={},
|
||||
litellm_params={"stream": True},
|
||||
headers={},
|
||||
)
|
||||
|
||||
assert streaming_request["stream_mode"] == "messages-tuple"
|
||||
|
||||
|
||||
def test_langgraph_provider_detection():
|
||||
"""
|
||||
Test that the langgraph provider is correctly detected from model name.
|
||||
"""
|
||||
from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider
|
||||
|
||||
model, provider, api_key, api_base = get_llm_provider(
|
||||
model="langgraph/agent",
|
||||
api_base="http://localhost:2024",
|
||||
)
|
||||
|
||||
assert provider == "langgraph"
|
||||
assert model == "agent"
|
||||
|
||||
|
|
@ -0,0 +1,159 @@
|
|||
"""
|
||||
Test A2A completion bridge streaming transformation to proper A2A format.
|
||||
|
||||
Tests that the completion bridge emits proper A2A streaming events:
|
||||
1. Task event (kind: "task") - Initial task with status "submitted"
|
||||
2. Status update (kind: "status-update") - Status "working"
|
||||
3. Artifact update (kind: "artifact-update") - Content delivery
|
||||
4. Status update (kind: "status-update") - Final "completed" status
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
class TestA2AStreamingTransformation:
|
||||
"""Test the A2A streaming transformation creates proper events."""
|
||||
|
||||
def test_create_task_event(self):
|
||||
"""Test that create_task_event produces proper A2A task event structure."""
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
||||
A2ACompletionBridgeTransformation,
|
||||
A2AStreamingContext,
|
||||
)
|
||||
|
||||
input_message = {
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hello"}],
|
||||
"messageId": "msg-123",
|
||||
}
|
||||
ctx = A2AStreamingContext(request_id="req-456", input_message=input_message)
|
||||
|
||||
event = A2ACompletionBridgeTransformation.create_task_event(ctx)
|
||||
|
||||
# Validate structure
|
||||
assert event["jsonrpc"] == "2.0"
|
||||
assert event["id"] == "req-456"
|
||||
assert event["result"]["kind"] == "task"
|
||||
assert event["result"]["status"]["state"] == "submitted"
|
||||
assert "contextId" in event["result"]
|
||||
assert "id" in event["result"] # task id
|
||||
assert "history" in event["result"]
|
||||
assert len(event["result"]["history"]) == 1
|
||||
assert event["result"]["history"][0]["role"] == "user"
|
||||
|
||||
def test_create_status_update_working(self):
|
||||
"""Test that create_status_update_event produces proper working status."""
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
||||
A2ACompletionBridgeTransformation,
|
||||
A2AStreamingContext,
|
||||
)
|
||||
|
||||
ctx = A2AStreamingContext(
|
||||
request_id="req-456",
|
||||
input_message={"role": "user", "parts": []},
|
||||
)
|
||||
|
||||
event = A2ACompletionBridgeTransformation.create_status_update_event(
|
||||
ctx=ctx,
|
||||
state="working",
|
||||
final=False,
|
||||
message_text="Processing...",
|
||||
)
|
||||
|
||||
assert event["result"]["kind"] == "status-update"
|
||||
assert event["result"]["status"]["state"] == "working"
|
||||
assert event["result"]["final"] is False
|
||||
assert "taskId" in event["result"]
|
||||
assert "contextId" in event["result"]
|
||||
assert "timestamp" in event["result"]["status"]
|
||||
|
||||
def test_create_artifact_update(self):
|
||||
"""Test that create_artifact_update_event produces proper artifact event."""
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.transformation import (
|
||||
A2ACompletionBridgeTransformation,
|
||||
A2AStreamingContext,
|
||||
)
|
||||
|
||||
ctx = A2AStreamingContext(
|
||||
request_id="req-456",
|
||||
input_message={"role": "user", "parts": []},
|
||||
)
|
||||
|
||||
event = A2ACompletionBridgeTransformation.create_artifact_update_event(
|
||||
ctx=ctx,
|
||||
text="Hello, I am an AI assistant.",
|
||||
)
|
||||
|
||||
assert event["result"]["kind"] == "artifact-update"
|
||||
assert "artifact" in event["result"]
|
||||
assert "artifactId" in event["result"]["artifact"]
|
||||
assert event["result"]["artifact"]["name"] == "response"
|
||||
assert event["result"]["artifact"]["parts"][0]["kind"] == "text"
|
||||
assert event["result"]["artifact"]["parts"][0]["text"] == "Hello, I am an AI assistant."
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_handle_streaming_emits_proper_events():
|
||||
"""Test that handle_streaming emits events in correct order with proper structure."""
|
||||
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
|
||||
A2ACompletionBridgeHandler,
|
||||
)
|
||||
|
||||
# Mock litellm.acompletion to return a streaming response
|
||||
mock_chunk1 = MagicMock()
|
||||
mock_chunk1.choices = [MagicMock()]
|
||||
mock_chunk1.choices[0].delta = MagicMock()
|
||||
mock_chunk1.choices[0].delta.content = "Hello"
|
||||
|
||||
mock_chunk2 = MagicMock()
|
||||
mock_chunk2.choices = [MagicMock()]
|
||||
mock_chunk2.choices[0].delta = MagicMock()
|
||||
mock_chunk2.choices[0].delta.content = " world"
|
||||
|
||||
async def mock_streaming_response():
|
||||
yield mock_chunk1
|
||||
yield mock_chunk2
|
||||
|
||||
with patch("litellm.acompletion", new_callable=AsyncMock) as mock_acompletion:
|
||||
mock_acompletion.return_value = mock_streaming_response()
|
||||
|
||||
params = {
|
||||
"message": {
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hi"}],
|
||||
"messageId": "msg-123",
|
||||
}
|
||||
}
|
||||
|
||||
events = []
|
||||
async for event in A2ACompletionBridgeHandler.handle_streaming(
|
||||
request_id="req-456",
|
||||
params=params,
|
||||
litellm_params={"custom_llm_provider": "langgraph", "model": "agent"},
|
||||
api_base="http://localhost:2024",
|
||||
):
|
||||
events.append(event)
|
||||
|
||||
# Should have 4 events: task, working, artifact, completed
|
||||
assert len(events) == 4
|
||||
|
||||
# Event 1: task submitted
|
||||
assert events[0]["result"]["kind"] == "task"
|
||||
assert events[0]["result"]["status"]["state"] == "submitted"
|
||||
|
||||
# Event 2: status working
|
||||
assert events[1]["result"]["kind"] == "status-update"
|
||||
assert events[1]["result"]["status"]["state"] == "working"
|
||||
assert events[1]["result"]["final"] is False
|
||||
|
||||
# Event 3: artifact update with accumulated content
|
||||
assert events[2]["result"]["kind"] == "artifact-update"
|
||||
assert events[2]["result"]["artifact"]["parts"][0]["text"] == "Hello world"
|
||||
|
||||
# Event 4: status completed
|
||||
assert events[3]["result"]["kind"] == "status-update"
|
||||
assert events[3]["result"]["status"]["state"] == "completed"
|
||||
assert events[3]["result"]["final"] is True
|
||||
|
||||
350
tests/test_litellm/a2a_protocol/test_cost_calculator.py
Normal file
350
tests/test_litellm/a2a_protocol/test_cost_calculator.py
Normal file
|
|
@ -0,0 +1,350 @@
|
|||
"""
|
||||
Test A2A cost calculator with cost_per_query parameter.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from typing import Optional
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.integrations.custom_logger import CustomLogger
|
||||
|
||||
|
||||
class CostLogger(CustomLogger):
|
||||
"""Custom logger to capture response_cost."""
|
||||
|
||||
def __init__(self):
|
||||
self.response_cost: Optional[float] = None
|
||||
super().__init__()
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
slp = kwargs.get("standard_logging_object")
|
||||
if slp:
|
||||
self.response_cost = slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_asend_message_uses_cost_per_query():
|
||||
"""
|
||||
Test that asend_message uses cost_per_query param for response_cost.
|
||||
"""
|
||||
from litellm.a2a_protocol import asend_message
|
||||
|
||||
# Setup logger
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
cost_logger = CostLogger()
|
||||
litellm.callbacks = [cost_logger]
|
||||
|
||||
# Mock A2A client
|
||||
mock_client = MagicMock()
|
||||
mock_client._litellm_agent_card = MagicMock()
|
||||
mock_client._litellm_agent_card.name = "test-agent"
|
||||
|
||||
# Mock response with required fields
|
||||
mock_response = MagicMock()
|
||||
mock_response.model_dump = MagicMock(return_value={
|
||||
"id": "test-123",
|
||||
"jsonrpc": "2.0",
|
||||
"result": {"status": "completed"},
|
||||
})
|
||||
mock_client.send_message = AsyncMock(return_value=mock_response)
|
||||
|
||||
# Mock request
|
||||
mock_request = MagicMock()
|
||||
mock_request.id = "test-123"
|
||||
|
||||
# Call asend_message with cost_per_query
|
||||
await asend_message(
|
||||
a2a_client=mock_client,
|
||||
request=mock_request,
|
||||
cost_per_query=0.05,
|
||||
)
|
||||
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
assert cost_logger.response_cost == 0.05
|
||||
|
||||
|
||||
class TokenAndCostLogger(CustomLogger):
|
||||
"""Custom logger to capture both token counts and cost."""
|
||||
|
||||
def __init__(self):
|
||||
self.response_cost: Optional[float] = None
|
||||
self.prompt_tokens: Optional[int] = None
|
||||
self.completion_tokens: Optional[int] = None
|
||||
super().__init__()
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
slp = kwargs.get("standard_logging_object")
|
||||
if slp:
|
||||
self.response_cost = slp.get("response_cost") if isinstance(slp, dict) else getattr(slp, "response_cost", None)
|
||||
self.prompt_tokens = slp.get("prompt_tokens") if isinstance(slp, dict) else getattr(slp, "prompt_tokens", None)
|
||||
self.completion_tokens = slp.get("completion_tokens") if isinstance(slp, dict) else getattr(slp, "completion_tokens", None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_asend_message_uses_input_output_cost_per_token():
|
||||
"""
|
||||
Test that asend_message calculates cost using input_cost_per_token and output_cost_per_token.
|
||||
Validates exact cost calculation: cost = (prompt_tokens * input_cost) + (completion_tokens * output_cost)
|
||||
"""
|
||||
from litellm.a2a_protocol import asend_message
|
||||
|
||||
# Setup logger
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
token_cost_logger = TokenAndCostLogger()
|
||||
litellm.callbacks = [token_cost_logger]
|
||||
|
||||
# Mock A2A client
|
||||
mock_client = MagicMock()
|
||||
mock_client._litellm_agent_card = MagicMock()
|
||||
mock_client._litellm_agent_card.name = "test-agent"
|
||||
|
||||
# Realistic A2A response with message parts
|
||||
mock_response = MagicMock()
|
||||
mock_response.model_dump = MagicMock(return_value={
|
||||
"id": "test-123",
|
||||
"jsonrpc": "2.0",
|
||||
"result": {
|
||||
"status": {"state": "completed"},
|
||||
"message": {
|
||||
"role": "assistant",
|
||||
"parts": [{"kind": "text", "text": "Hello! I am your assistant. How can I help you today?"}],
|
||||
"messageId": "msg-456",
|
||||
}
|
||||
},
|
||||
})
|
||||
mock_client.send_message = AsyncMock(return_value=mock_response)
|
||||
|
||||
# Mock request with message parts
|
||||
mock_request = MagicMock()
|
||||
mock_request.id = "test-123"
|
||||
mock_request.params = MagicMock()
|
||||
mock_request.params.message = {
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hello, what can you do?"}],
|
||||
"messageId": "msg-123",
|
||||
}
|
||||
|
||||
# Define specific cost per token values
|
||||
input_cost_per_token = 0.00001 # $0.01 per 1000 tokens
|
||||
output_cost_per_token = 0.00002 # $0.02 per 1000 tokens
|
||||
|
||||
await asend_message(
|
||||
a2a_client=mock_client,
|
||||
request=mock_request,
|
||||
input_cost_per_token=input_cost_per_token,
|
||||
output_cost_per_token=output_cost_per_token,
|
||||
)
|
||||
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
# Get actual token counts from logger
|
||||
prompt_tokens = token_cost_logger.prompt_tokens
|
||||
completion_tokens = token_cost_logger.completion_tokens
|
||||
response_cost = token_cost_logger.response_cost
|
||||
|
||||
print(f"\n=== Token-Based Cost Results ===")
|
||||
print(f"prompt_tokens: {prompt_tokens}")
|
||||
print(f"completion_tokens: {completion_tokens}")
|
||||
print(f"input_cost_per_token: {input_cost_per_token}")
|
||||
print(f"output_cost_per_token: {output_cost_per_token}")
|
||||
print(f"response_cost: {response_cost}")
|
||||
|
||||
# Verify tokens were captured
|
||||
assert prompt_tokens is not None, "prompt_tokens should be captured"
|
||||
assert completion_tokens is not None, "completion_tokens should be captured"
|
||||
assert response_cost is not None, "response_cost should be captured"
|
||||
|
||||
# Calculate expected cost
|
||||
expected_cost = (prompt_tokens * input_cost_per_token) + (completion_tokens * output_cost_per_token)
|
||||
print(f"expected_cost: {expected_cost}")
|
||||
|
||||
# Verify exact cost calculation
|
||||
assert response_cost == expected_cost, f"response_cost {response_cost} should equal expected {expected_cost}"
|
||||
|
||||
|
||||
class AgentIdLogger(CustomLogger):
|
||||
"""Custom logger to capture agent_id from kwargs."""
|
||||
|
||||
def __init__(self):
|
||||
self.agent_id: Optional[str] = None
|
||||
self.kwargs: Optional[dict] = None
|
||||
super().__init__()
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
self.kwargs = kwargs
|
||||
self.agent_id = kwargs.get("agent_id")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_asend_message_passes_agent_id_to_callback():
|
||||
"""
|
||||
Test that asend_message passes agent_id to callbacks via kwargs.
|
||||
"""
|
||||
from litellm.a2a_protocol import asend_message
|
||||
|
||||
# Setup logger
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
agent_id_logger = AgentIdLogger()
|
||||
litellm.callbacks = [agent_id_logger]
|
||||
|
||||
# Mock A2A client
|
||||
mock_client = MagicMock()
|
||||
mock_client._litellm_agent_card = MagicMock()
|
||||
mock_client._litellm_agent_card.name = "test-agent"
|
||||
|
||||
# Mock response
|
||||
mock_response = MagicMock()
|
||||
mock_response.model_dump = MagicMock(return_value={
|
||||
"id": "test-123",
|
||||
"jsonrpc": "2.0",
|
||||
"result": {"status": "completed"},
|
||||
})
|
||||
mock_client.send_message = AsyncMock(return_value=mock_response)
|
||||
|
||||
# Mock request
|
||||
mock_request = MagicMock()
|
||||
mock_request.id = "test-123"
|
||||
|
||||
test_agent_id = "agent-uuid-12345"
|
||||
|
||||
# Call asend_message with agent_id
|
||||
await asend_message(
|
||||
a2a_client=mock_client,
|
||||
request=mock_request,
|
||||
agent_id=test_agent_id,
|
||||
)
|
||||
|
||||
await asyncio.sleep(0.1)
|
||||
|
||||
# Verify agent_id was passed to callback
|
||||
assert agent_id_logger.agent_id == test_agent_id, f"Expected agent_id '{test_agent_id}', got '{agent_id_logger.agent_id}'"
|
||||
|
||||
|
||||
class MetadataLogger(CustomLogger):
|
||||
"""Custom logger to capture metadata from kwargs for proxy spend tracking."""
|
||||
|
||||
def __init__(self):
|
||||
self.metadata: Optional[dict] = None
|
||||
self.litellm_params: Optional[dict] = None
|
||||
self.user_api_key: Optional[str] = None
|
||||
self.user_id: Optional[str] = None
|
||||
self.team_id: Optional[str] = None
|
||||
super().__init__()
|
||||
|
||||
async def async_log_success_event(self, kwargs, response_obj, start_time, end_time):
|
||||
self.litellm_params = kwargs.get("litellm_params", {})
|
||||
self.metadata = self.litellm_params.get("metadata", {})
|
||||
self.user_api_key = self.metadata.get("user_api_key")
|
||||
self.user_id = self.metadata.get("user_api_key_user_id")
|
||||
self.team_id = self.metadata.get("user_api_key_team_id")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_asend_message_streaming_propagates_metadata():
|
||||
"""
|
||||
Test that asend_message_streaming propagates metadata to logging object.
|
||||
This ensures user_api_key, user_id, team_id are available for SpendLogs.
|
||||
"""
|
||||
from litellm.a2a_protocol import asend_message_streaming
|
||||
|
||||
# Setup logger
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
metadata_logger = MetadataLogger()
|
||||
litellm.logging_callback_manager.add_litellm_async_success_callback(metadata_logger)
|
||||
|
||||
# Mock A2A client
|
||||
mock_client = MagicMock()
|
||||
mock_client._litellm_agent_card = MagicMock()
|
||||
mock_client._litellm_agent_card.name = "test-agent"
|
||||
|
||||
# Mock streaming response
|
||||
async def mock_stream():
|
||||
yield MagicMock(model_dump=lambda mode, exclude_none: {"chunk": 1})
|
||||
yield MagicMock(model_dump=lambda mode, exclude_none: {"chunk": 2})
|
||||
|
||||
mock_client.send_message_streaming = MagicMock(return_value=mock_stream())
|
||||
|
||||
# Mock request
|
||||
mock_request = MagicMock()
|
||||
mock_request.id = "test-stream-metadata"
|
||||
mock_request.params = MagicMock()
|
||||
mock_request.params.message = {"role": "user", "parts": [{"kind": "text", "text": "Hello"}]}
|
||||
|
||||
# Metadata from proxy (contains user_api_key, user_id, team_id for SpendLogs)
|
||||
test_metadata = {
|
||||
"user_api_key": "sk-test-key-hash-12345",
|
||||
"user_api_key_user_id": "user-uuid-123",
|
||||
"user_api_key_team_id": "team-uuid-456",
|
||||
}
|
||||
|
||||
# Consume streaming response with metadata
|
||||
chunks = []
|
||||
async for chunk in asend_message_streaming(
|
||||
a2a_client=mock_client,
|
||||
request=mock_request,
|
||||
metadata=test_metadata,
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
await asyncio.sleep(0.2)
|
||||
|
||||
# Verify metadata was propagated to callback
|
||||
assert metadata_logger.user_api_key == "sk-test-key-hash-12345"
|
||||
assert metadata_logger.user_id == "user-uuid-123"
|
||||
assert metadata_logger.team_id == "team-uuid-456"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_asend_message_streaming_triggers_callbacks():
|
||||
"""
|
||||
Test that asend_message_streaming triggers callbacks after stream completes.
|
||||
"""
|
||||
from litellm.a2a_protocol import asend_message_streaming
|
||||
|
||||
# Setup logger - must use logging_callback_manager to properly register
|
||||
litellm.logging_callback_manager._reset_all_callbacks()
|
||||
callback_logger = AgentIdLogger()
|
||||
litellm.logging_callback_manager.add_litellm_async_success_callback(callback_logger)
|
||||
litellm.logging_callback_manager.add_litellm_success_callback(callback_logger)
|
||||
|
||||
# Mock A2A client
|
||||
mock_client = MagicMock()
|
||||
mock_client._litellm_agent_card = MagicMock()
|
||||
mock_client._litellm_agent_card.name = "test-agent"
|
||||
|
||||
# Mock streaming response
|
||||
async def mock_stream():
|
||||
yield MagicMock(model_dump=lambda mode, exclude_none: {"chunk": 1})
|
||||
yield MagicMock(model_dump=lambda mode, exclude_none: {"chunk": 2})
|
||||
|
||||
mock_client.send_message_streaming = MagicMock(return_value=mock_stream())
|
||||
|
||||
# Mock request
|
||||
mock_request = MagicMock()
|
||||
mock_request.id = "test-stream-123"
|
||||
mock_request.params = MagicMock()
|
||||
mock_request.params.message = {"role": "user", "parts": [{"kind": "text", "text": "Hello"}]}
|
||||
|
||||
test_agent_id = "test-agent-id-streaming"
|
||||
|
||||
# Consume streaming response
|
||||
chunks = []
|
||||
async for chunk in asend_message_streaming(
|
||||
a2a_client=mock_client,
|
||||
request=mock_request,
|
||||
agent_id=test_agent_id,
|
||||
):
|
||||
chunks.append(chunk)
|
||||
|
||||
await asyncio.sleep(0.2)
|
||||
|
||||
# Verify chunks were received
|
||||
assert len(chunks) == 2
|
||||
|
||||
# Verify callbacks WERE triggered after stream completed
|
||||
assert callback_logger.kwargs is not None, "Streaming should trigger callbacks after completion"
|
||||
assert callback_logger.agent_id == test_agent_id, f"Expected agent_id '{test_agent_id}', got '{callback_logger.agent_id}'"
|
||||
|
|
@ -1403,13 +1403,13 @@ async def test_async_log_success_event_increments_by_actual_tokens():
|
|||
end_time=None,
|
||||
)
|
||||
|
||||
# Verify increments happened with actual token count (50 completion tokens)
|
||||
# Verify increments happened with actual token count (60 total tokens)
|
||||
assert len(increment_calls) == 2, f"Expected 2 increment calls, got {len(increment_calls)}"
|
||||
|
||||
# Both should increment by 50 (completion_tokens, since rate_limit_type defaults to 'output')
|
||||
# Both should increment by 50 (total_tokens, since rate_limit_type defaults to 'total')
|
||||
for call in increment_calls:
|
||||
assert call["increment_value"] == 50, (
|
||||
f"Expected increment of 50 tokens, got {call['increment_value']} for key {call['key']}"
|
||||
assert call["increment_value"] == 60, (
|
||||
f"Expected increment of 60 tokens, got {call['increment_value']} for key {call['key']}"
|
||||
)
|
||||
|
||||
# Verify correct keys were used
|
||||
|
|
|
|||
|
|
@ -1583,6 +1583,231 @@ async def test_missing_descriptor_fallback():
|
|||
assert "Current limit: 2" in exc_info.value.detail
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_rate_limit_type_default_is_total(monkeypatch):
|
||||
"""
|
||||
Test that get_rate_limit_type returns 'total' as the default when no setting is specified.
|
||||
|
||||
This verifies the change from 'output' to 'total' as the default value.
|
||||
"""
|
||||
local_cache = DualCache()
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
|
||||
# Mock general_settings to return empty dict (no token_rate_limit_type set)
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
original_settings = getattr(proxy_server, 'general_settings', {})
|
||||
monkeypatch.setattr(proxy_server, 'general_settings', {})
|
||||
|
||||
try:
|
||||
result = parallel_request_handler.get_rate_limit_type()
|
||||
assert result == "total", f"Default rate limit type should be 'total', got '{result}'"
|
||||
finally:
|
||||
monkeypatch.setattr(proxy_server, 'general_settings', original_settings)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_rate_limit_type_invalid_falls_back_to_total(monkeypatch):
|
||||
"""
|
||||
Test that get_rate_limit_type falls back to 'total' when an invalid value is specified.
|
||||
"""
|
||||
local_cache = DualCache()
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
|
||||
# Mock general_settings to return an invalid token_rate_limit_type
|
||||
import litellm.proxy.proxy_server as proxy_server
|
||||
original_settings = getattr(proxy_server, 'general_settings', {})
|
||||
monkeypatch.setattr(proxy_server, 'general_settings', {'token_rate_limit_type': 'invalid_type'})
|
||||
|
||||
try:
|
||||
result = parallel_request_handler.get_rate_limit_type()
|
||||
assert result == "total", f"Invalid rate limit type should fall back to 'total', got '{result}'"
|
||||
finally:
|
||||
monkeypatch.setattr(proxy_server, 'general_settings', original_settings)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"token_rate_limit_type,expected_field",
|
||||
[
|
||||
("input", "prompt_tokens"),
|
||||
("output", "completion_tokens"),
|
||||
("total", "total_tokens"),
|
||||
],
|
||||
)
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_log_success_event_with_dict_usage(monkeypatch, token_rate_limit_type, expected_field):
|
||||
"""
|
||||
Test that async_log_success_event correctly handles usage as a dict (Responses API format).
|
||||
|
||||
The Responses API returns usage as a dict in ResponsesAPIResponse instead of a Usage object.
|
||||
This test verifies that token counting works correctly with dict-based usage.
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
_api_key = "sk-12345"
|
||||
_api_key = hash_token(_api_key)
|
||||
local_cache = DualCache()
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
|
||||
# Mock the get_rate_limit_type method
|
||||
def mock_get_rate_limit_type():
|
||||
return token_rate_limit_type
|
||||
|
||||
monkeypatch.setattr(
|
||||
parallel_request_handler, "get_rate_limit_type", mock_get_rate_limit_type
|
||||
)
|
||||
|
||||
# Create a mock response object with usage as a dict (Responses API format)
|
||||
mock_response = MagicMock()
|
||||
mock_response.usage = {
|
||||
"prompt_tokens": 25,
|
||||
"completion_tokens": 35,
|
||||
"total_tokens": 60
|
||||
}
|
||||
# Make isinstance check for BaseLiteLLMOpenAIResponseObject return True
|
||||
from litellm.types.utils import BaseLiteLLMOpenAIResponseObject
|
||||
mock_response.__class__ = type('MockResponse', (BaseLiteLLMOpenAIResponseObject,), {})
|
||||
|
||||
# Create mock kwargs for the success event
|
||||
mock_kwargs = {
|
||||
"standard_logging_object": {
|
||||
"metadata": {
|
||||
"user_api_key_hash": _api_key,
|
||||
"user_api_key_user_id": None,
|
||||
"user_api_key_team_id": None,
|
||||
"user_api_key_end_user_id": None,
|
||||
}
|
||||
},
|
||||
"model": "gpt-3.5-turbo",
|
||||
}
|
||||
|
||||
# Mock the pipeline increment method to capture the operations
|
||||
captured_operations = []
|
||||
|
||||
async def mock_increment_pipeline(increment_list, **kwargs):
|
||||
captured_operations.extend(increment_list)
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(
|
||||
parallel_request_handler.internal_usage_cache.dual_cache,
|
||||
"async_increment_cache_pipeline",
|
||||
mock_increment_pipeline,
|
||||
)
|
||||
|
||||
# Call the success event handler
|
||||
await parallel_request_handler.async_log_success_event(
|
||||
kwargs=mock_kwargs,
|
||||
response_obj=mock_response,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
# Find the TPM increment operation
|
||||
tpm_operation = None
|
||||
for op in captured_operations:
|
||||
if op["key"].endswith(":tokens"):
|
||||
tpm_operation = op
|
||||
break
|
||||
|
||||
assert tpm_operation is not None, "Should have a TPM increment operation"
|
||||
|
||||
# Check that the correct token count was used based on the rate limit type
|
||||
expected_tokens = {
|
||||
"input": 25, # prompt_tokens
|
||||
"output": 35, # completion_tokens
|
||||
"total": 60, # total_tokens
|
||||
}
|
||||
|
||||
assert (
|
||||
tpm_operation["increment_value"] == expected_tokens[token_rate_limit_type]
|
||||
), f"Expected {expected_tokens[token_rate_limit_type]} tokens for type '{token_rate_limit_type}', got {tpm_operation['increment_value']}"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_log_success_event_with_dict_usage_missing_fields(monkeypatch):
|
||||
"""
|
||||
Test that async_log_success_event handles dict usage with missing fields gracefully.
|
||||
|
||||
When usage dict is missing expected fields, it should default to 0.
|
||||
"""
|
||||
from unittest.mock import MagicMock
|
||||
|
||||
_api_key = "sk-12345"
|
||||
_api_key = hash_token(_api_key)
|
||||
local_cache = DualCache()
|
||||
parallel_request_handler = _PROXY_MaxParallelRequestsHandler(
|
||||
internal_usage_cache=InternalUsageCache(local_cache)
|
||||
)
|
||||
|
||||
# Mock the get_rate_limit_type method
|
||||
def mock_get_rate_limit_type():
|
||||
return "output"
|
||||
|
||||
monkeypatch.setattr(
|
||||
parallel_request_handler, "get_rate_limit_type", mock_get_rate_limit_type
|
||||
)
|
||||
|
||||
# Create a mock response object with usage as a dict missing some fields
|
||||
mock_response = MagicMock()
|
||||
mock_response.usage = {
|
||||
"prompt_tokens": 25,
|
||||
# completion_tokens is missing
|
||||
# total_tokens is missing
|
||||
}
|
||||
from litellm.types.utils import BaseLiteLLMOpenAIResponseObject
|
||||
mock_response.__class__ = type('MockResponse', (BaseLiteLLMOpenAIResponseObject,), {})
|
||||
|
||||
# Create mock kwargs for the success event
|
||||
mock_kwargs = {
|
||||
"standard_logging_object": {
|
||||
"metadata": {
|
||||
"user_api_key_hash": _api_key,
|
||||
"user_api_key_user_id": None,
|
||||
"user_api_key_team_id": None,
|
||||
"user_api_key_end_user_id": None,
|
||||
}
|
||||
},
|
||||
"model": "gpt-3.5-turbo",
|
||||
}
|
||||
|
||||
# Mock the pipeline increment method to capture the operations
|
||||
captured_operations = []
|
||||
|
||||
async def mock_increment_pipeline(increment_list, **kwargs):
|
||||
captured_operations.extend(increment_list)
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(
|
||||
parallel_request_handler.internal_usage_cache.dual_cache,
|
||||
"async_increment_cache_pipeline",
|
||||
mock_increment_pipeline,
|
||||
)
|
||||
|
||||
# Call the success event handler - should not raise exception
|
||||
await parallel_request_handler.async_log_success_event(
|
||||
kwargs=mock_kwargs,
|
||||
response_obj=mock_response,
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
)
|
||||
|
||||
# Find the TPM increment operation
|
||||
tpm_operation = None
|
||||
for op in captured_operations:
|
||||
if op["key"].endswith(":tokens"):
|
||||
tpm_operation = op
|
||||
break
|
||||
|
||||
assert tpm_operation is not None, "Should have a TPM increment operation"
|
||||
# Should default to 0 when field is missing
|
||||
assert tpm_operation["increment_value"] == 0, "Should default to 0 when completion_tokens is missing"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_execute_token_increment_script_cluster_compatibility():
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -2438,6 +2438,70 @@ class TestGenericResponseConvertorNestedAttributes:
|
|||
assert result.display_name == "user-sub-123" # Top-level attribute works
|
||||
|
||||
|
||||
class TestGenericResponseConvertorUserRole:
|
||||
"""Test generic_response_convertor user role extraction from SSO token"""
|
||||
|
||||
def test_generic_response_convertor_extracts_valid_user_role(self):
|
||||
"""
|
||||
Test that generic_response_convertor extracts a valid LiteLLM user role
|
||||
from the SSO token using the GENERIC_USER_ROLE_ATTRIBUTE env var.
|
||||
"""
|
||||
from litellm.proxy._types import LitellmUserRoles
|
||||
from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor
|
||||
|
||||
mock_jwt_handler = MagicMock(spec=JWTHandler)
|
||||
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
|
||||
|
||||
sso_response = {
|
||||
"preferred_username": "testuser",
|
||||
"email": "test@example.com",
|
||||
"sub": "Test User",
|
||||
"role": "proxy_admin",
|
||||
}
|
||||
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{"GENERIC_USER_ROLE_ATTRIBUTE": "role"},
|
||||
):
|
||||
result = generic_response_convertor(
|
||||
response=sso_response,
|
||||
jwt_handler=mock_jwt_handler,
|
||||
sso_jwt_handler=None,
|
||||
)
|
||||
|
||||
assert isinstance(result, CustomOpenID)
|
||||
assert result.user_role == LitellmUserRoles.PROXY_ADMIN
|
||||
|
||||
def test_generic_response_convertor_ignores_invalid_user_role(self):
|
||||
"""
|
||||
Test that generic_response_convertor ignores invalid role values
|
||||
and sets user_role to None.
|
||||
"""
|
||||
from litellm.proxy.management_endpoints.ui_sso import generic_response_convertor
|
||||
|
||||
mock_jwt_handler = MagicMock(spec=JWTHandler)
|
||||
mock_jwt_handler.get_team_ids_from_jwt.return_value = []
|
||||
|
||||
sso_response = {
|
||||
"preferred_username": "testuser",
|
||||
"email": "test@example.com",
|
||||
"role": "invalid_role_value",
|
||||
}
|
||||
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{"GENERIC_USER_ROLE_ATTRIBUTE": "role"},
|
||||
):
|
||||
result = generic_response_convertor(
|
||||
response=sso_response,
|
||||
jwt_handler=mock_jwt_handler,
|
||||
sso_jwt_handler=None,
|
||||
)
|
||||
|
||||
assert isinstance(result, CustomOpenID)
|
||||
assert result.user_role is None
|
||||
|
||||
|
||||
class TestGetGenericSSORedirectParams:
|
||||
"""Test _get_generic_sso_redirect_params state parameter priority handling"""
|
||||
|
||||
|
|
|
|||
|
|
@ -594,3 +594,41 @@ async def test_api_key_preserved_through_failure_hook_to_database():
|
|||
print("- Both SpendLogs AND DailyUserSpend will have correct api_key")
|
||||
print("="*80 + "\n")
|
||||
|
||||
|
||||
@patch("litellm.proxy.proxy_server.master_key", None)
|
||||
@patch("litellm.proxy.proxy_server.general_settings", {})
|
||||
def test_get_logging_payload_includes_agent_id_from_kwargs():
|
||||
"""
|
||||
Test that get_logging_payload extracts agent_id from kwargs and includes it in the payload.
|
||||
"""
|
||||
test_agent_id = "agent-uuid-12345"
|
||||
|
||||
kwargs = {
|
||||
"model": "a2a_agent/test-agent",
|
||||
"custom_llm_provider": "a2a_agent",
|
||||
"agent_id": test_agent_id,
|
||||
"litellm_params": {
|
||||
"metadata": {
|
||||
"user_api_key": "sk-test-key",
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
response_obj = {
|
||||
"id": "test-response-123",
|
||||
"jsonrpc": "2.0",
|
||||
"result": {"status": "completed"},
|
||||
}
|
||||
|
||||
start_time = datetime.datetime.now(timezone.utc)
|
||||
end_time = datetime.datetime.now(timezone.utc)
|
||||
|
||||
payload = get_logging_payload(
|
||||
kwargs=kwargs,
|
||||
response_obj=response_obj,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
)
|
||||
|
||||
assert payload["agent_id"] == test_agent_id, f"Expected agent_id '{test_agent_id}', got '{payload.get('agent_id')}'"
|
||||
|
||||
|
|
|
|||
|
|
@ -832,6 +832,37 @@ def test_video_content_handler_uses_get_for_openai():
|
|||
assert called_url == "https://api.openai.com/v1/videos/video_abc/content"
|
||||
|
||||
|
||||
def test_video_content_respects_api_base_and_api_key_from_kwargs():
|
||||
"""Test that video_content respects api_base and api_key from kwargs (simulating database entry)."""
|
||||
from litellm.videos.main import video_content
|
||||
|
||||
# Mock the handler to capture litellm_params
|
||||
captured_litellm_params = None
|
||||
|
||||
def capture_litellm_params(*args, **kwargs):
|
||||
nonlocal captured_litellm_params
|
||||
captured_litellm_params = kwargs.get("litellm_params")
|
||||
return b"mp4-bytes"
|
||||
|
||||
with patch('litellm.videos.main.base_llm_http_handler') as mock_handler:
|
||||
mock_handler.video_content_handler = capture_litellm_params
|
||||
|
||||
# Call video_content with api_base and api_key in kwargs (simulating database entry)
|
||||
# This simulates how the router passes model config from database via **kwargs
|
||||
result = video_content(
|
||||
video_id="video_test_123",
|
||||
custom_llm_provider="azure",
|
||||
api_base="https://test-resource.openai.azure.com/", # Passed via kwargs by router
|
||||
api_key="test-api-key-from-db", # Passed via kwargs by router
|
||||
)
|
||||
|
||||
# Verify that api_base and api_key from kwargs were included in litellm_params
|
||||
assert captured_litellm_params is not None
|
||||
assert captured_litellm_params.get("api_base") == "https://test-resource.openai.azure.com/"
|
||||
assert captured_litellm_params.get("api_key") == "test-api-key-from-db"
|
||||
assert result == b"mp4-bytes"
|
||||
|
||||
|
||||
def test_openai_video_config_has_async_transform():
|
||||
"""Ensure OpenAIVideoConfig exposes async_transform_video_content_response at runtime."""
|
||||
cfg = OpenAIVideoConfig()
|
||||
|
|
|
|||
BIN
ui/litellm-dashboard/public/assets/logos/langgraph.png
Normal file
BIN
ui/litellm-dashboard/public/assets/logos/langgraph.png
Normal file
Binary file not shown.
|
After Width: | Height: | Size: 5.4 KiB |
|
|
@ -1,7 +1,9 @@
|
|||
import React, { useState } from "react";
|
||||
import { Modal, Form, Button as AntButton, message } from "antd";
|
||||
import { createAgentCall } from "../networking";
|
||||
import React, { useState, useEffect } from "react";
|
||||
import { Modal, Form, message, Select } from "antd";
|
||||
import { Button } from "@tremor/react";
|
||||
import { createAgentCall, getAgentCreateMetadata, AgentCreateInfo } from "../networking";
|
||||
import AgentFormFields from "./agent_form_fields";
|
||||
import DynamicAgentFormFields, { buildDynamicAgentData } from "./dynamic_agent_form_fields";
|
||||
import { getDefaultFormValues, buildAgentDataFromForm } from "./agent_config";
|
||||
|
||||
interface AddAgentFormProps {
|
||||
|
|
@ -19,6 +21,29 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
|
|||
}) => {
|
||||
const [form] = Form.useForm();
|
||||
const [isSubmitting, setIsSubmitting] = useState(false);
|
||||
const [agentType, setAgentType] = useState<string>("a2a");
|
||||
const [agentTypeMetadata, setAgentTypeMetadata] = useState<AgentCreateInfo[]>([]);
|
||||
const [loadingMetadata, setLoadingMetadata] = useState(false);
|
||||
|
||||
// Fetch agent type metadata on mount
|
||||
useEffect(() => {
|
||||
const fetchMetadata = async () => {
|
||||
setLoadingMetadata(true);
|
||||
try {
|
||||
const metadata = await getAgentCreateMetadata();
|
||||
setAgentTypeMetadata(metadata);
|
||||
} catch (error) {
|
||||
console.error("Error fetching agent metadata:", error);
|
||||
} finally {
|
||||
setLoadingMetadata(false);
|
||||
}
|
||||
};
|
||||
fetchMetadata();
|
||||
}, []);
|
||||
|
||||
const selectedAgentTypeInfo = agentTypeMetadata.find(
|
||||
(info) => info.agent_type === agentType
|
||||
);
|
||||
|
||||
const handleSubmit = async (values: any) => {
|
||||
if (!accessToken) {
|
||||
|
|
@ -28,10 +53,18 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
|
|||
|
||||
setIsSubmitting(true);
|
||||
try {
|
||||
const agentData = buildAgentDataFromForm(values);
|
||||
let agentData: any;
|
||||
|
||||
if (agentType === "a2a") {
|
||||
agentData = buildAgentDataFromForm(values);
|
||||
} else if (selectedAgentTypeInfo) {
|
||||
agentData = buildDynamicAgentData(values, selectedAgentTypeInfo);
|
||||
}
|
||||
|
||||
await createAgentCall(accessToken, agentData);
|
||||
message.success("Agent created successfully");
|
||||
form.resetFields();
|
||||
setAgentType("a2a");
|
||||
onSuccess();
|
||||
onClose();
|
||||
} catch (error) {
|
||||
|
|
@ -44,42 +77,114 @@ const AddAgentForm: React.FC<AddAgentFormProps> = ({
|
|||
|
||||
const handleCancel = () => {
|
||||
form.resetFields();
|
||||
setAgentType("a2a");
|
||||
onClose();
|
||||
};
|
||||
|
||||
const handleAgentTypeChange = (value: string) => {
|
||||
setAgentType(value);
|
||||
form.resetFields();
|
||||
};
|
||||
|
||||
// Get the logo for the selected agent type for the header
|
||||
const selectedLogo = selectedAgentTypeInfo?.logo_url || agentTypeMetadata.find(a => a.agent_type === "a2a")?.logo_url;
|
||||
|
||||
return (
|
||||
<Modal
|
||||
title="Add New Agent"
|
||||
title={
|
||||
<div className="flex items-center space-x-3 pb-4 border-b border-gray-100">
|
||||
{selectedLogo && (
|
||||
<img
|
||||
src={selectedLogo}
|
||||
alt="Agent"
|
||||
className="w-6 h-6 object-contain"
|
||||
/>
|
||||
)}
|
||||
<h2 className="text-xl font-semibold text-gray-900">Add New Agent</h2>
|
||||
</div>
|
||||
}
|
||||
open={visible}
|
||||
onCancel={handleCancel}
|
||||
footer={null}
|
||||
width={800}
|
||||
width={900}
|
||||
className="top-8"
|
||||
styles={{
|
||||
body: { padding: "24px" },
|
||||
header: { padding: "24px 24px 0 24px", border: "none" },
|
||||
}}
|
||||
>
|
||||
<Form
|
||||
form={form}
|
||||
layout="vertical"
|
||||
onFinish={handleSubmit}
|
||||
initialValues={getDefaultFormValues()}
|
||||
>
|
||||
<AgentFormFields showAgentName={true} />
|
||||
|
||||
<Form.Item>
|
||||
<div style={{ display: "flex", justifyContent: "flex-end", gap: "8px" }}>
|
||||
<AntButton onClick={handleCancel}>
|
||||
Cancel
|
||||
</AntButton>
|
||||
<AntButton
|
||||
htmlType="submit"
|
||||
loading={isSubmitting}
|
||||
<div className="mt-4">
|
||||
<Form
|
||||
form={form}
|
||||
layout="vertical"
|
||||
onFinish={handleSubmit}
|
||||
initialValues={agentType === "a2a" ? getDefaultFormValues() : {}}
|
||||
className="space-y-4"
|
||||
>
|
||||
{/* Agent Type Selection */}
|
||||
<Form.Item
|
||||
label={<span className="text-sm font-medium text-gray-700">Agent Type</span>}
|
||||
required
|
||||
tooltip="Select the type of agent you want to create"
|
||||
>
|
||||
<Select
|
||||
value={agentType}
|
||||
onChange={handleAgentTypeChange}
|
||||
size="large"
|
||||
style={{ width: "100%" }}
|
||||
optionLabelProp="label"
|
||||
>
|
||||
Create Agent
|
||||
</AntButton>
|
||||
{agentTypeMetadata.map((info) => (
|
||||
<Select.Option
|
||||
key={info.agent_type}
|
||||
value={info.agent_type}
|
||||
label={
|
||||
<div className="flex items-center gap-2">
|
||||
<img src={info.logo_url || ""} alt="" className="w-4 h-4 object-contain" />
|
||||
<span>{info.agent_type_display_name}</span>
|
||||
</div>
|
||||
}
|
||||
>
|
||||
<div className="flex items-center gap-3 py-1">
|
||||
<img
|
||||
src={info.logo_url || ""}
|
||||
alt={info.agent_type_display_name}
|
||||
className="w-5 h-5 object-contain"
|
||||
/>
|
||||
<div>
|
||||
<div className="font-medium">{info.agent_type_display_name}</div>
|
||||
{info.description && (
|
||||
<div className="text-xs text-gray-500">{info.description}</div>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
</Select.Option>
|
||||
))}
|
||||
</Select>
|
||||
</Form.Item>
|
||||
|
||||
{/* Conditional Form Fields */}
|
||||
<div className="mt-6">
|
||||
{agentType === "a2a" ? (
|
||||
<AgentFormFields showAgentName={true} />
|
||||
) : selectedAgentTypeInfo ? (
|
||||
<DynamicAgentFormFields agentTypeInfo={selectedAgentTypeInfo} />
|
||||
) : null}
|
||||
</div>
|
||||
</Form.Item>
|
||||
</Form>
|
||||
|
||||
{/* Footer Buttons */}
|
||||
<div className="flex items-center justify-end space-x-3 pt-6 border-t border-gray-100 mt-6">
|
||||
<Button variant="secondary" onClick={handleCancel}>
|
||||
Cancel
|
||||
</Button>
|
||||
<Button variant="primary" loading={isSubmitting}>
|
||||
{isSubmitting ? "Creating..." : "Create Agent"}
|
||||
</Button>
|
||||
</div>
|
||||
</Form>
|
||||
</div>
|
||||
</Modal>
|
||||
);
|
||||
};
|
||||
|
||||
export default AddAgentForm;
|
||||
|
||||
|
|
|
|||
|
|
@ -0,0 +1,125 @@
|
|||
import React from "react";
|
||||
import { Form, Input, Select } from "antd";
|
||||
import { AgentCreateInfo, AgentCredentialFieldMetadata } from "../networking";
|
||||
|
||||
interface DynamicAgentFormFieldsProps {
|
||||
agentTypeInfo: AgentCreateInfo;
|
||||
}
|
||||
|
||||
/**
|
||||
* Form fields for dynamic agent types (e.g., LangGraph).
|
||||
* Renders common fields (agent name, display name, description) plus
|
||||
* credential fields defined by the agent type metadata.
|
||||
*/
|
||||
const DynamicAgentFormFields: React.FC<DynamicAgentFormFieldsProps> = ({
|
||||
agentTypeInfo,
|
||||
}) => {
|
||||
return (
|
||||
<>
|
||||
<Form.Item
|
||||
label="Agent Name"
|
||||
name="agent_name"
|
||||
rules={[{ required: true, message: "Please enter a unique agent name" }]}
|
||||
tooltip="Unique identifier for the agent"
|
||||
>
|
||||
<Input placeholder="e.g., my-langgraph-agent" />
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label="Description"
|
||||
name="description"
|
||||
tooltip="Brief description of what this agent does"
|
||||
>
|
||||
<Input.TextArea rows={2} placeholder="Describe what this agent does..." />
|
||||
</Form.Item>
|
||||
|
||||
{agentTypeInfo.credential_fields.map((field: AgentCredentialFieldMetadata) => (
|
||||
<Form.Item
|
||||
key={field.key}
|
||||
label={field.label}
|
||||
name={field.key}
|
||||
rules={field.required ? [{ required: true, message: `Please enter ${field.label}` }] : undefined}
|
||||
tooltip={field.tooltip}
|
||||
initialValue={field.default_value}
|
||||
>
|
||||
{field.field_type === "password" ? (
|
||||
<Input.Password placeholder={field.placeholder || ""} />
|
||||
) : field.field_type === "textarea" ? (
|
||||
<Input.TextArea rows={3} placeholder={field.placeholder || ""} />
|
||||
) : field.field_type === "select" && field.options ? (
|
||||
<Select placeholder={field.placeholder || ""}>
|
||||
{field.options.map((opt) => (
|
||||
<Select.Option key={opt} value={opt}>
|
||||
{opt}
|
||||
</Select.Option>
|
||||
))}
|
||||
</Select>
|
||||
) : (
|
||||
<Input placeholder={field.placeholder || ""} />
|
||||
)}
|
||||
</Form.Item>
|
||||
))}
|
||||
</>
|
||||
);
|
||||
};
|
||||
|
||||
/**
|
||||
* Builds agent data from form values for dynamic agent types.
|
||||
* Uses configuration from agentTypeInfo to determine which fields to include.
|
||||
*/
|
||||
export const buildDynamicAgentData = (
|
||||
values: any,
|
||||
agentTypeInfo: AgentCreateInfo
|
||||
) => {
|
||||
// Build litellm_params from template
|
||||
const litellmParams: Record<string, any> = {
|
||||
...(agentTypeInfo.litellm_params_template || {}),
|
||||
};
|
||||
|
||||
// Add credential fields marked with include_in_litellm_params
|
||||
for (const field of agentTypeInfo.credential_fields) {
|
||||
const value = values[field.key];
|
||||
if (value && field.include_in_litellm_params !== false) {
|
||||
litellmParams[field.key] = value;
|
||||
}
|
||||
}
|
||||
|
||||
// Apply model_template if defined (e.g., "bedrock/agentcore/{agent_runtime_arn}")
|
||||
if (agentTypeInfo.model_template) {
|
||||
let model = agentTypeInfo.model_template;
|
||||
// Replace {field_key} placeholders with actual values
|
||||
for (const field of agentTypeInfo.credential_fields) {
|
||||
const placeholder = `{${field.key}}`;
|
||||
if (model.includes(placeholder) && values[field.key]) {
|
||||
model = model.replace(placeholder, values[field.key]);
|
||||
}
|
||||
}
|
||||
litellmParams.model = model;
|
||||
}
|
||||
|
||||
return {
|
||||
agent_name: values.agent_name,
|
||||
agent_card_params: {
|
||||
protocolVersion: "1.0",
|
||||
name: values.display_name || values.agent_name,
|
||||
description: values.description || `${agentTypeInfo.agent_type_display_name} agent`,
|
||||
url: values.api_base || "",
|
||||
version: "1.0.0",
|
||||
defaultInputModes: ["text"],
|
||||
defaultOutputModes: ["text"],
|
||||
capabilities: {
|
||||
streaming: true,
|
||||
},
|
||||
skills: [{
|
||||
id: "chat",
|
||||
name: "Chat",
|
||||
description: "General chat capability",
|
||||
tags: ["chat", "conversation"],
|
||||
}],
|
||||
},
|
||||
litellm_params: litellmParams,
|
||||
};
|
||||
};
|
||||
|
||||
export default DynamicAgentFormFields;
|
||||
|
||||
|
|
@ -195,6 +195,28 @@ export interface ProviderCreateInfo {
|
|||
credential_fields: ProviderCredentialFieldMetadata[];
|
||||
}
|
||||
|
||||
export interface AgentCredentialFieldMetadata {
|
||||
key: string;
|
||||
label: string;
|
||||
placeholder?: string | null;
|
||||
tooltip?: string | null;
|
||||
required?: boolean;
|
||||
field_type?: "text" | "password" | "select" | "upload" | "textarea";
|
||||
options?: string[] | null;
|
||||
default_value?: string | null;
|
||||
include_in_litellm_params?: boolean;
|
||||
}
|
||||
|
||||
export interface AgentCreateInfo {
|
||||
agent_type: string;
|
||||
agent_type_display_name: string;
|
||||
description?: string | null;
|
||||
logo_url?: string | null;
|
||||
credential_fields: AgentCredentialFieldMetadata[];
|
||||
litellm_params_template?: Record<string, string> | null;
|
||||
model_template?: string | null;
|
||||
}
|
||||
|
||||
export interface PublicModelHubInfo {
|
||||
docs_title: string;
|
||||
custom_docs_description: string | null;
|
||||
|
|
@ -255,6 +277,26 @@ export const getProviderCreateMetadata = async (): Promise<ProviderCreateInfo[]>
|
|||
return jsonData;
|
||||
};
|
||||
|
||||
export const getAgentCreateMetadata = async (): Promise<AgentCreateInfo[]> => {
|
||||
/**
|
||||
* Fetch agent type metadata from the proxy's public endpoint.
|
||||
* This is used by the UI to dynamically render agent-specific credential fields.
|
||||
*/
|
||||
const url = proxyBaseUrl ? `${proxyBaseUrl}/public/agents/fields` : `/public/agents/fields`;
|
||||
const response = await fetch(url, {
|
||||
method: "GET",
|
||||
});
|
||||
|
||||
if (!response.ok) {
|
||||
const errorText = await response.text();
|
||||
console.error("Failed to fetch agent create metadata:", response.status, errorText);
|
||||
throw new Error("Failed to load agent configuration");
|
||||
}
|
||||
|
||||
const jsonData: AgentCreateInfo[] = await response.json();
|
||||
return jsonData;
|
||||
};
|
||||
|
||||
// Global variable for the header name
|
||||
let globalLitellmHeaderName: string = "Authorization";
|
||||
const MCP_AUTH_HEADER: string = "x-mcp-auth";
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue