mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge branch 'litellm_oss_staging_02_04_2026' into litellm_fix_langfuse_otel_trace
This commit is contained in:
commit
607c9f02d7
24 changed files with 1615 additions and 305 deletions
65
Makefile
65
Makefile
|
|
@ -1,7 +1,10 @@
|
|||
# LiteLLM Makefile
|
||||
# Simple Makefile for running tests and basic development tasks
|
||||
|
||||
.PHONY: help test test-unit test-integration test-unit-helm lint format install-dev install-proxy-dev install-test-deps install-helm-unittest check-circular-imports check-import-safety
|
||||
.PHONY: help test test-unit test-integration test-unit-helm \
|
||||
info lint lint-dev format \
|
||||
install-dev install-proxy-dev install-test-deps \
|
||||
install-helm-unittest check-circular-imports check-import-safety
|
||||
|
||||
# Default target
|
||||
help:
|
||||
|
|
@ -25,6 +28,13 @@ help:
|
|||
@echo " make test-integration - Run integration tests"
|
||||
@echo " make test-unit-helm - Run helm unit tests"
|
||||
|
||||
# Keep PIP simple for edge cases:
|
||||
PIP := $(shell command -v pip > /dev/null 2>&1 && echo "pip" || echo "python3 -m pip")
|
||||
|
||||
# Show info
|
||||
info:
|
||||
@echo "PIP: $(PIP)"
|
||||
|
||||
# Installation targets
|
||||
install-dev:
|
||||
poetry install --with dev
|
||||
|
|
@ -34,19 +44,19 @@ install-proxy-dev:
|
|||
|
||||
# CI-compatible installations (matches GitHub workflows exactly)
|
||||
install-dev-ci:
|
||||
pip install openai==2.8.0
|
||||
$(PIP) install openai==2.8.0
|
||||
poetry install --with dev
|
||||
pip install openai==2.8.0
|
||||
$(PIP) install openai==2.8.0
|
||||
|
||||
install-proxy-dev-ci:
|
||||
poetry install --with dev,proxy-dev --extras proxy
|
||||
pip install openai==2.8.0
|
||||
$(PIP) install openai==2.8.0
|
||||
|
||||
install-test-deps: install-proxy-dev
|
||||
poetry run pip install "pytest-retry==1.6.3"
|
||||
poetry run pip install pytest-xdist
|
||||
poetry run pip install openapi-core
|
||||
cd enterprise && poetry run pip install -e . && cd ..
|
||||
poetry run $(PIP) install "pytest-retry==1.6.3"
|
||||
poetry run $(PIP) install pytest-xdist
|
||||
poetry run $(PIP) install openapi-core
|
||||
cd enterprise && poetry run $(PIP) install -e . && cd ..
|
||||
|
||||
install-helm-unittest:
|
||||
helm plugin install https://github.com/helm-unittest/helm-unittest --version v0.4.4 || echo "ignore error if plugin exists"
|
||||
|
|
@ -62,8 +72,40 @@ format-check: install-dev
|
|||
lint-ruff: install-dev
|
||||
cd litellm && poetry run ruff check . && cd ..
|
||||
|
||||
# faster linter for developing ...
|
||||
# inspiration from:
|
||||
# https://github.com/astral-sh/ruff/discussions/10977
|
||||
# https://github.com/astral-sh/ruff/discussions/4049
|
||||
lint-format-changed: install-dev
|
||||
@git diff origin/main --unified=0 --no-color -- '*.py' | \
|
||||
perl -ne '\
|
||||
if (/^diff --git a\/(.*) b\//) { $$file = $$1; } \
|
||||
if (/^@@ .* \+(\d+)(?:,(\d+))? @@/) { \
|
||||
$$start = $$1; $$count = $$2 || 1; $$end = $$start + $$count - 1; \
|
||||
print "$$file:$$start:1-$$end:999\n"; \
|
||||
}' | \
|
||||
while read range; do \
|
||||
file="$${range%%:*}"; \
|
||||
lines="$${range#*:}"; \
|
||||
echo "Formatting $$file (lines $$lines)"; \
|
||||
poetry run ruff format --range "$$lines" "$$file"; \
|
||||
done
|
||||
|
||||
lint-ruff-dev: install-dev
|
||||
@tmpfile=$$(mktemp /tmp/ruff-dev.XXXXXX) && \
|
||||
cd litellm && \
|
||||
(poetry run ruff check . --output-format=pylint || true) > "$$tmpfile" && \
|
||||
poetry run diff-quality --violations=pylint "$$tmpfile" --compare-branch=origin/main && \
|
||||
cd .. ; \
|
||||
rm -f "$$tmpfile"
|
||||
|
||||
lint-ruff-FULL-dev: install-dev
|
||||
@files=$$(git diff --name-only origin/main -- '*.py'); \
|
||||
if [ -n "$$files" ]; then echo "$$files" | xargs poetry run ruff check; \
|
||||
else echo "No changed .py files to check."; fi
|
||||
|
||||
lint-mypy: install-dev
|
||||
poetry run pip install types-requests types-setuptools types-redis types-PyYAML
|
||||
poetry run $(PIP) install types-requests types-setuptools types-redis types-PyYAML
|
||||
cd litellm && poetry run mypy . --ignore-missing-imports && cd ..
|
||||
|
||||
lint-black: format-check
|
||||
|
|
@ -72,11 +114,14 @@ check-circular-imports: install-dev
|
|||
cd litellm && poetry run python ../tests/documentation_tests/test_circular_imports.py && cd ..
|
||||
|
||||
check-import-safety: install-dev
|
||||
poetry run python -c "from litellm import *" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)
|
||||
@poetry run python -c "from litellm import *; print('[from litellm import *] OK! no issues!');" || (echo '🚨 import failed, this means you introduced unprotected imports! 🚨'; exit 1)
|
||||
|
||||
# Combined linting (matches test-linting.yml workflow)
|
||||
lint: format-check lint-ruff lint-mypy check-circular-imports check-import-safety
|
||||
|
||||
# Faster linting for local development (only checks changed code)
|
||||
lint-dev: lint-format-changed lint-mypy check-circular-imports check-import-safety
|
||||
|
||||
# Testing targets
|
||||
test:
|
||||
poetry run pytest tests/
|
||||
|
|
|
|||
|
|
@ -68,116 +68,9 @@ Follow [this guide, to add your pydantic ai agent to LiteLLM Agent Gateway](./pr
|
|||
|
||||
## Invoking your Agents
|
||||
|
||||
Use the [A2A Python SDK](https://pypi.org/project/a2a-sdk) to invoke agents through LiteLLM.
|
||||
|
||||
This example shows how to:
|
||||
1. **List available agents** - Query `/v1/agents` to see which agents your key can access
|
||||
2. **Select an agent** - Pick an agent from the list
|
||||
3. **Invoke via A2A** - Use the A2A protocol to send messages to the agent
|
||||
|
||||
```python showLineNumbers title="invoke_a2a_agent.py"
|
||||
from uuid import uuid4
|
||||
import httpx
|
||||
import asyncio
|
||||
from a2a.client import A2ACardResolver, A2AClient
|
||||
from a2a.types import MessageSendParams, SendMessageRequest
|
||||
|
||||
# === CONFIGURE THESE ===
|
||||
LITELLM_BASE_URL = "http://localhost:4000" # Your LiteLLM proxy URL
|
||||
LITELLM_VIRTUAL_KEY = "sk-1234" # Your LiteLLM Virtual Key
|
||||
# =======================
|
||||
|
||||
async def main():
|
||||
headers = {"Authorization": f"Bearer {LITELLM_VIRTUAL_KEY}"}
|
||||
|
||||
async with httpx.AsyncClient(headers=headers) as client:
|
||||
# Step 1: List available agents
|
||||
response = await client.get(f"{LITELLM_BASE_URL}/v1/agents")
|
||||
agents = response.json()
|
||||
|
||||
print("Available agents:")
|
||||
for agent in agents:
|
||||
print(f" - {agent['agent_name']} (ID: {agent['agent_id']})")
|
||||
|
||||
if not agents:
|
||||
print("No agents available for this key")
|
||||
return
|
||||
|
||||
# Step 2: Select an agent and invoke it
|
||||
selected_agent = agents[0]
|
||||
agent_id = selected_agent["agent_id"]
|
||||
agent_name = selected_agent["agent_name"]
|
||||
print(f"\nInvoking: {agent_name}")
|
||||
|
||||
# Step 3: Use A2A protocol to invoke the agent
|
||||
base_url = f"{LITELLM_BASE_URL}/a2a/{agent_id}"
|
||||
resolver = A2ACardResolver(httpx_client=client, base_url=base_url)
|
||||
agent_card = await resolver.get_agent_card()
|
||||
a2a_client = A2AClient(httpx_client=client, agent_card=agent_card)
|
||||
|
||||
request = SendMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(
|
||||
message={
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hello, what can you do?"}],
|
||||
"messageId": uuid4().hex,
|
||||
}
|
||||
),
|
||||
)
|
||||
response = await a2a_client.send_message(request)
|
||||
print(f"Response: {response.model_dump(mode='json', exclude_none=True, indent=4)}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
```
|
||||
|
||||
### Streaming Responses
|
||||
|
||||
For streaming responses, use `send_message_streaming`:
|
||||
|
||||
```python showLineNumbers title="invoke_a2a_agent_streaming.py"
|
||||
from uuid import uuid4
|
||||
import httpx
|
||||
import asyncio
|
||||
from a2a.client import A2ACardResolver, A2AClient
|
||||
from a2a.types import MessageSendParams, SendStreamingMessageRequest
|
||||
|
||||
# === CONFIGURE THESE ===
|
||||
LITELLM_BASE_URL = "http://localhost:4000" # Your LiteLLM proxy URL
|
||||
LITELLM_VIRTUAL_KEY = "sk-1234" # Your LiteLLM Virtual Key
|
||||
LITELLM_AGENT_NAME = "ij-local" # Agent name registered in LiteLLM
|
||||
# =======================
|
||||
|
||||
async def main():
|
||||
base_url = f"{LITELLM_BASE_URL}/a2a/{LITELLM_AGENT_NAME}"
|
||||
headers = {"Authorization": f"Bearer {LITELLM_VIRTUAL_KEY}"}
|
||||
|
||||
async with httpx.AsyncClient(headers=headers) as httpx_client:
|
||||
# Resolve agent card and create client
|
||||
resolver = A2ACardResolver(httpx_client=httpx_client, base_url=base_url)
|
||||
agent_card = await resolver.get_agent_card()
|
||||
client = A2AClient(httpx_client=httpx_client, agent_card=agent_card)
|
||||
|
||||
# Send a streaming message
|
||||
request = SendStreamingMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(
|
||||
message={
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hello, what can you do?"}],
|
||||
"messageId": uuid4().hex,
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
# Stream the response
|
||||
async for chunk in client.send_message_streaming(request):
|
||||
print(chunk.model_dump(mode="json", exclude_none=True))
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
```
|
||||
See the [Invoking A2A Agents](./a2a_invoking_agents) guide to learn how to call your agents using:
|
||||
- **A2A SDK** - Native A2A protocol with full support for tasks and artifacts
|
||||
- **OpenAI SDK** - Familiar `/chat/completions` interface with `a2a/` model prefix
|
||||
|
||||
## Tracking Agent Logs
|
||||
|
||||
|
|
|
|||
280
docs/my-website/docs/a2a_invoking_agents.md
Normal file
280
docs/my-website/docs/a2a_invoking_agents.md
Normal file
|
|
@ -0,0 +1,280 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# Invoking A2A Agents
|
||||
|
||||
Learn how to invoke A2A agents through LiteLLM using different methods.
|
||||
|
||||
:::tip Deploy Your Own A2A Agent
|
||||
|
||||
Want to test with your own agent? Deploy this template A2A agent powered by Google Gemini:
|
||||
|
||||
[**shin-bot-litellm/a2a-gemini-agent**](https://github.com/shin-bot-litellm/a2a-gemini-agent) - Simple deployable A2A agent with streaming support
|
||||
|
||||
:::
|
||||
|
||||
## A2A SDK
|
||||
|
||||
Use the [A2A Python SDK](https://pypi.org/project/a2a-sdk) to invoke agents through LiteLLM using the A2A protocol.
|
||||
|
||||
### Non-Streaming
|
||||
|
||||
This example shows how to:
|
||||
1. **List available agents** - Query `/v1/agents` to see which agents your key can access
|
||||
2. **Select an agent** - Pick an agent from the list
|
||||
3. **Invoke via A2A** - Use the A2A protocol to send messages to the agent
|
||||
|
||||
```python showLineNumbers title="invoke_a2a_agent.py"
|
||||
from uuid import uuid4
|
||||
import httpx
|
||||
import asyncio
|
||||
from a2a.client import A2ACardResolver, A2AClient
|
||||
from a2a.types import MessageSendParams, SendMessageRequest
|
||||
|
||||
# === CONFIGURE THESE ===
|
||||
LITELLM_BASE_URL = "http://localhost:4000" # Your LiteLLM proxy URL
|
||||
LITELLM_VIRTUAL_KEY = "sk-1234" # Your LiteLLM Virtual Key
|
||||
# =======================
|
||||
|
||||
async def main():
|
||||
headers = {"Authorization": f"Bearer {LITELLM_VIRTUAL_KEY}"}
|
||||
|
||||
async with httpx.AsyncClient(headers=headers) as client:
|
||||
# Step 1: List available agents
|
||||
response = await client.get(f"{LITELLM_BASE_URL}/v1/agents")
|
||||
agents = response.json()
|
||||
|
||||
print("Available agents:")
|
||||
for agent in agents:
|
||||
print(f" - {agent['agent_name']} (ID: {agent['agent_id']})")
|
||||
|
||||
if not agents:
|
||||
print("No agents available for this key")
|
||||
return
|
||||
|
||||
# Step 2: Select an agent and invoke it
|
||||
selected_agent = agents[0]
|
||||
agent_id = selected_agent["agent_id"]
|
||||
agent_name = selected_agent["agent_name"]
|
||||
print(f"\nInvoking: {agent_name}")
|
||||
|
||||
# Step 3: Use A2A protocol to invoke the agent
|
||||
base_url = f"{LITELLM_BASE_URL}/a2a/{agent_id}"
|
||||
resolver = A2ACardResolver(httpx_client=client, base_url=base_url)
|
||||
agent_card = await resolver.get_agent_card()
|
||||
a2a_client = A2AClient(httpx_client=client, agent_card=agent_card)
|
||||
|
||||
request = SendMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(
|
||||
message={
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hello, what can you do?"}],
|
||||
"messageId": uuid4().hex,
|
||||
}
|
||||
),
|
||||
)
|
||||
response = await a2a_client.send_message(request)
|
||||
print(f"Response: {response.model_dump(mode='json', exclude_none=True, indent=4)}")
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
```
|
||||
|
||||
### Streaming
|
||||
|
||||
For streaming responses, use `send_message_streaming`:
|
||||
|
||||
```python showLineNumbers title="invoke_a2a_agent_streaming.py"
|
||||
from uuid import uuid4
|
||||
import httpx
|
||||
import asyncio
|
||||
from a2a.client import A2ACardResolver, A2AClient
|
||||
from a2a.types import MessageSendParams, SendStreamingMessageRequest
|
||||
|
||||
# === CONFIGURE THESE ===
|
||||
LITELLM_BASE_URL = "http://localhost:4000" # Your LiteLLM proxy URL
|
||||
LITELLM_VIRTUAL_KEY = "sk-1234" # Your LiteLLM Virtual Key
|
||||
LITELLM_AGENT_NAME = "ij-local" # Agent name registered in LiteLLM
|
||||
# =======================
|
||||
|
||||
async def main():
|
||||
base_url = f"{LITELLM_BASE_URL}/a2a/{LITELLM_AGENT_NAME}"
|
||||
headers = {"Authorization": f"Bearer {LITELLM_VIRTUAL_KEY}"}
|
||||
|
||||
async with httpx.AsyncClient(headers=headers) as httpx_client:
|
||||
# Resolve agent card and create client
|
||||
resolver = A2ACardResolver(httpx_client=httpx_client, base_url=base_url)
|
||||
agent_card = await resolver.get_agent_card()
|
||||
client = A2AClient(httpx_client=httpx_client, agent_card=agent_card)
|
||||
|
||||
# Send a streaming message
|
||||
request = SendStreamingMessageRequest(
|
||||
id=str(uuid4()),
|
||||
params=MessageSendParams(
|
||||
message={
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Tell me a long story"}],
|
||||
"messageId": uuid4().hex,
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
# Stream the response
|
||||
async for chunk in client.send_message_streaming(request):
|
||||
print(chunk.model_dump(mode="json", exclude_none=True))
|
||||
|
||||
if __name__ == "__main__":
|
||||
asyncio.run(main())
|
||||
```
|
||||
|
||||
## /chat/completions API (OpenAI SDK)
|
||||
|
||||
You can also invoke A2A agents using the familiar OpenAI SDK by using the `a2a/` model prefix.
|
||||
|
||||
### Non-Streaming
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python" default>
|
||||
|
||||
```python showLineNumbers title="openai_non_streaming.py"
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(
|
||||
api_key="sk-1234", # Your LiteLLM Virtual Key
|
||||
base_url="http://localhost:4000" # Your LiteLLM proxy URL
|
||||
)
|
||||
|
||||
response = client.chat.completions.create(
|
||||
model="a2a/my-agent", # Use a2a/ prefix with your agent name
|
||||
messages=[
|
||||
{"role": "user", "content": "Hello, what can you do?"}
|
||||
]
|
||||
)
|
||||
|
||||
print(response.choices[0].message.content)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="typescript" label="TypeScript">
|
||||
|
||||
```typescript showLineNumbers title="openai_non_streaming.ts"
|
||||
import OpenAI from 'openai';
|
||||
|
||||
const client = new OpenAI({
|
||||
apiKey: 'sk-1234', // Your LiteLLM Virtual Key
|
||||
baseURL: 'http://localhost:4000' // Your LiteLLM proxy URL
|
||||
});
|
||||
|
||||
const response = await client.chat.completions.create({
|
||||
model: 'a2a/my-agent', // Use a2a/ prefix with your agent name
|
||||
messages: [
|
||||
{ role: 'user', content: 'Hello, what can you do?' }
|
||||
]
|
||||
});
|
||||
|
||||
console.log(response.choices[0].message.content);
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="curl" label="cURL">
|
||||
|
||||
```bash showLineNumbers title="curl_non_streaming.sh"
|
||||
curl -X POST http://localhost:4000/v1/chat/completions \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "a2a/my-agent",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Hello, what can you do?"}
|
||||
]
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
### Streaming
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="python" label="Python" default>
|
||||
|
||||
```python showLineNumbers title="openai_streaming.py"
|
||||
import openai
|
||||
|
||||
client = openai.OpenAI(
|
||||
api_key="sk-1234", # Your LiteLLM Virtual Key
|
||||
base_url="http://localhost:4000" # Your LiteLLM proxy URL
|
||||
)
|
||||
|
||||
stream = client.chat.completions.create(
|
||||
model="a2a/my-agent", # Use a2a/ prefix with your agent name
|
||||
messages=[
|
||||
{"role": "user", "content": "Tell me a long story"}
|
||||
],
|
||||
stream=True
|
||||
)
|
||||
|
||||
for chunk in stream:
|
||||
if chunk.choices[0].delta.content:
|
||||
print(chunk.choices[0].delta.content, end="", flush=True)
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="typescript" label="TypeScript">
|
||||
|
||||
```typescript showLineNumbers title="openai_streaming.ts"
|
||||
import OpenAI from 'openai';
|
||||
|
||||
const client = new OpenAI({
|
||||
apiKey: 'sk-1234', // Your LiteLLM Virtual Key
|
||||
baseURL: 'http://localhost:4000' // Your LiteLLM proxy URL
|
||||
});
|
||||
|
||||
const stream = await client.chat.completions.create({
|
||||
model: 'a2a/my-agent', // Use a2a/ prefix with your agent name
|
||||
messages: [
|
||||
{ role: 'user', content: 'Tell me a long story' }
|
||||
],
|
||||
stream: true
|
||||
});
|
||||
|
||||
for await (const chunk of stream) {
|
||||
const content = chunk.choices[0]?.delta?.content;
|
||||
if (content) {
|
||||
process.stdout.write(content);
|
||||
}
|
||||
}
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="curl" label="cURL">
|
||||
|
||||
```bash showLineNumbers title="curl_streaming.sh"
|
||||
curl -X POST http://localhost:4000/v1/chat/completions \
|
||||
-H "Authorization: Bearer sk-1234" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"model": "a2a/my-agent",
|
||||
"messages": [
|
||||
{"role": "user", "content": "Tell me a long story"}
|
||||
],
|
||||
"stream": true
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
## Key Differences
|
||||
|
||||
| Method | Use Case | Advantages |
|
||||
|--------|----------|------------|
|
||||
| **A2A SDK** | Native A2A protocol integration | • Full A2A protocol support<br/>• Access to task states and artifacts<br/>• Context management |
|
||||
| **OpenAI SDK** | Familiar OpenAI-style interface | • Drop-in replacement for OpenAI calls<br/>• Easier migration from LLM to agent workflows<br/>• Works with existing OpenAI tooling |
|
||||
|
||||
:::tip Model Prefix
|
||||
|
||||
When using the OpenAI SDK, always prefix your agent name with `a2a/` (e.g., `a2a/my-agent`) to route requests to the A2A agent instead of an LLM provider.
|
||||
|
||||
:::
|
||||
|
|
@ -469,6 +469,7 @@ const sidebars = {
|
|||
label: "/a2a - A2A Agent Gateway",
|
||||
items: [
|
||||
"a2a",
|
||||
"a2a_invoking_agents",
|
||||
"a2a_cost_tracking",
|
||||
"a2a_agent_permissions"
|
||||
],
|
||||
|
|
|
|||
|
|
@ -1008,13 +1008,9 @@ class OpenTelemetry(CustomLogger):
|
|||
from opentelemetry._logs import SeverityNumber, get_logger, get_logger_provider
|
||||
|
||||
try:
|
||||
from opentelemetry.sdk._logs import (
|
||||
LogRecord as SdkLogRecord, # type: ignore[attr-defined] # OTEL < 1.39.0
|
||||
)
|
||||
from opentelemetry.sdk._logs import LogRecord as SdkLogRecord # type: ignore[attr-defined] # OTEL < 1.39.0
|
||||
except ImportError:
|
||||
from opentelemetry.sdk._logs._internal import (
|
||||
LogRecord as SdkLogRecord,
|
||||
) # OTEL >= 1.39.0
|
||||
from opentelemetry.sdk._logs._internal import LogRecord as SdkLogRecord # type: ignore[attr-defined, no-redef] # OTEL >= 1.39.0
|
||||
|
||||
otel_logger = get_logger(LITELLM_LOGGER_NAME)
|
||||
|
||||
|
|
|
|||
|
|
@ -94,7 +94,7 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
|
|||
if state == "completed":
|
||||
return "stop"
|
||||
elif state == "failed":
|
||||
return "error"
|
||||
return "stop" # Map failed state to 'stop' (valid finish_reason)
|
||||
|
||||
# Check for [DONE] marker
|
||||
if chunk.get("done") is True:
|
||||
|
|
|
|||
|
|
@ -2,10 +2,9 @@
|
|||
A2A Protocol Transformation for LiteLLM
|
||||
"""
|
||||
import uuid
|
||||
from typing import Any, Dict, Iterator, List, Optional, Union, cast
|
||||
from typing import Any, Dict, Iterator, List, Optional, Union
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.llms.base_llm.base_model_iterator import BaseModelResponseIterator
|
||||
from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException
|
||||
|
|
@ -27,6 +26,68 @@ class A2AConfig(BaseConfig):
|
|||
Handles transformation between OpenAI and A2A JSON-RPC 2.0 formats.
|
||||
"""
|
||||
|
||||
@staticmethod
|
||||
def resolve_agent_config_from_registry(
|
||||
model: str,
|
||||
api_base: Optional[str],
|
||||
api_key: Optional[str],
|
||||
headers: Optional[Dict[str, Any]],
|
||||
optional_params: Dict[str, Any],
|
||||
) -> tuple[Optional[str], Optional[str], Optional[Dict[str, Any]]]:
|
||||
"""
|
||||
Resolve agent configuration from registry if model format is "a2a/<agent-name>".
|
||||
|
||||
Extracts agent name from model string and looks up configuration in the
|
||||
agent registry (if available in proxy context).
|
||||
|
||||
Args:
|
||||
model: Model string (e.g., "a2a/my-agent")
|
||||
api_base: Explicit api_base (takes precedence over registry)
|
||||
api_key: Explicit api_key (takes precedence over registry)
|
||||
headers: Explicit headers (takes precedence over registry)
|
||||
optional_params: Dict to merge additional litellm_params into
|
||||
|
||||
Returns:
|
||||
Tuple of (api_base, api_key, headers) with registry values filled in
|
||||
"""
|
||||
# Extract agent name from model (e.g., "a2a/my-agent" -> "my-agent")
|
||||
agent_name = model.split("/", 1)[1] if "/" in model else None
|
||||
|
||||
# Only lookup if agent name exists and some config is missing
|
||||
if not agent_name or (api_base is not None and api_key is not None and headers is not None):
|
||||
return api_base, api_key, headers
|
||||
|
||||
# Try registry lookup (only available in proxy context)
|
||||
try:
|
||||
from litellm.proxy.agent_endpoints.agent_registry import (
|
||||
global_agent_registry,
|
||||
)
|
||||
|
||||
agent = global_agent_registry.get_agent_by_name(agent_name)
|
||||
if agent:
|
||||
# Get api_base from agent card URL
|
||||
if api_base is None and agent.agent_card_params:
|
||||
api_base = agent.agent_card_params.get("url")
|
||||
|
||||
# Get api_key, headers, and other params from litellm_params
|
||||
if agent.litellm_params:
|
||||
if api_key is None:
|
||||
api_key = agent.litellm_params.get("api_key")
|
||||
|
||||
if headers is None:
|
||||
agent_headers = agent.litellm_params.get("headers")
|
||||
if agent_headers:
|
||||
headers = agent_headers
|
||||
|
||||
# Merge other litellm_params (timeout, max_retries, etc.)
|
||||
for key, value in agent.litellm_params.items():
|
||||
if key not in ["api_key", "api_base", "headers", "model"] and key not in optional_params:
|
||||
optional_params[key] = value
|
||||
except ImportError:
|
||||
pass # Registry not available (not running in proxy context)
|
||||
|
||||
return api_base, api_key, headers
|
||||
|
||||
def get_supported_openai_params(self, model: str) -> List[str]:
|
||||
"""Return list of supported OpenAI parameters"""
|
||||
return [
|
||||
|
|
@ -46,9 +107,14 @@ class A2AConfig(BaseConfig):
|
|||
"""
|
||||
Map OpenAI parameters to A2A parameters.
|
||||
|
||||
For A2A protocol, we don't need to map most parameters since
|
||||
they're handled in the transform_request method.
|
||||
For A2A protocol, we need to map the stream parameter so
|
||||
transform_request can determine which JSON-RPC method to use.
|
||||
"""
|
||||
# Map stream parameter
|
||||
for param, value in non_default_params.items():
|
||||
if param == "stream" and value is True:
|
||||
optional_params["stream"] = value
|
||||
|
||||
return optional_params
|
||||
|
||||
def validate_environment(
|
||||
|
|
@ -160,8 +226,9 @@ class A2AConfig(BaseConfig):
|
|||
|
||||
# Build JSON-RPC 2.0 request
|
||||
# For A2A protocol, the method is "message/send" for non-streaming
|
||||
# and "message/stream" for streaming (handled by optional_params["stream"])
|
||||
method = "message/stream" if optional_params.get("stream") else "message/send"
|
||||
# and "message/stream" for streaming
|
||||
stream = optional_params.get("stream", False)
|
||||
method = "message/stream" if stream else "message/send"
|
||||
|
||||
request_data = {
|
||||
"jsonrpc": "2.0",
|
||||
|
|
|
|||
|
|
@ -113,6 +113,8 @@ def extract_text_from_a2a_response(
|
|||
# 1. Direct message: {"result": {"kind": "message", "parts": [...]}}
|
||||
# 2. Nested message: {"result": {"message": {"parts": [...]}}}
|
||||
# 3. Task with artifacts: {"result": {"kind": "task", "artifacts": [{"parts": [...]}]}}
|
||||
# 4. Task with status message: {"result": {"kind": "task", "status": {"message": {"parts": [...]}}}}
|
||||
# 5. Streaming artifact-update: {"result": {"kind": "artifact-update", "artifact": {"parts": [...]}}}
|
||||
|
||||
# Check if result itself has parts (direct message)
|
||||
if "parts" in result:
|
||||
|
|
@ -123,7 +125,23 @@ def extract_text_from_a2a_response(
|
|||
if message:
|
||||
return extract_text_from_a2a_message(message, depth=0, max_depth=max_depth)
|
||||
|
||||
# Handle task result with artifacts
|
||||
# Check for streaming artifact-update (singular artifact)
|
||||
artifact = result.get("artifact")
|
||||
if artifact and isinstance(artifact, dict):
|
||||
return extract_text_from_a2a_message(
|
||||
artifact, depth=0, max_depth=max_depth
|
||||
)
|
||||
|
||||
# Check for task status message (common in Gemini A2A agents)
|
||||
status = result.get("status", {})
|
||||
if isinstance(status, dict):
|
||||
status_message = status.get("message")
|
||||
if status_message:
|
||||
return extract_text_from_a2a_message(
|
||||
status_message, depth=0, max_depth=max_depth
|
||||
)
|
||||
|
||||
# Handle task result with artifacts (plural, array)
|
||||
artifacts = result.get("artifacts", [])
|
||||
if artifacts and len(artifacts) > 0:
|
||||
first_artifact = artifacts[0]
|
||||
|
|
|
|||
|
|
@ -22,7 +22,6 @@ from litellm.llms.custom_httpx.http_handler import (
|
|||
AsyncHTTPHandler,
|
||||
get_ssl_configuration,
|
||||
)
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
||||
|
||||
class OpenAIError(BaseLLMException):
|
||||
|
|
@ -205,67 +204,30 @@ class BaseOpenAILLM:
|
|||
if litellm.aclient_session is not None:
|
||||
return litellm.aclient_session
|
||||
|
||||
# Use the global cached client system to prevent memory leaks (issue #14540)
|
||||
# This routes through get_async_httpx_client() which provides TTL-based caching
|
||||
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
|
||||
# Get unified SSL configuration
|
||||
ssl_config = get_ssl_configuration()
|
||||
|
||||
try:
|
||||
# Get SSL config and include in params for proper cache key
|
||||
ssl_config = get_ssl_configuration()
|
||||
params = {"ssl_verify": ssl_config} if ssl_config is not None else {}
|
||||
params["disable_aiohttp_transport"] = litellm.disable_aiohttp_transport
|
||||
|
||||
# Get a cached AsyncHTTPHandler which manages the httpx.AsyncClient
|
||||
cached_handler = get_async_httpx_client(
|
||||
llm_provider=LlmProviders.OPENAI, # Cache key includes provider
|
||||
params=params, # Include SSL config in cache key
|
||||
return httpx.AsyncClient(
|
||||
verify=ssl_config,
|
||||
transport=AsyncHTTPHandler._create_async_transport(
|
||||
ssl_context=ssl_config
|
||||
if isinstance(ssl_config, ssl.SSLContext)
|
||||
else None,
|
||||
ssl_verify=ssl_config if isinstance(ssl_config, bool) else None,
|
||||
shared_session=shared_session,
|
||||
)
|
||||
# Return the underlying httpx client from the handler
|
||||
return cached_handler.client
|
||||
except (ImportError, AttributeError, KeyError) as e:
|
||||
# Fallback to creating a client directly if caching system unavailable
|
||||
# This preserves backwards compatibility
|
||||
verbose_logger.debug(
|
||||
f"Client caching unavailable ({type(e).__name__}), using direct client creation"
|
||||
)
|
||||
ssl_config = get_ssl_configuration()
|
||||
return httpx.AsyncClient(
|
||||
verify=ssl_config,
|
||||
transport=AsyncHTTPHandler._create_async_transport(
|
||||
ssl_context=ssl_config
|
||||
if isinstance(ssl_config, ssl.SSLContext)
|
||||
else None,
|
||||
ssl_verify=ssl_config if isinstance(ssl_config, bool) else None,
|
||||
shared_session=shared_session,
|
||||
),
|
||||
follow_redirects=True,
|
||||
)
|
||||
),
|
||||
follow_redirects=True,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _get_sync_http_client() -> Optional[httpx.Client]:
|
||||
if litellm.client_session is not None:
|
||||
return litellm.client_session
|
||||
|
||||
# Use the global cached client system to prevent memory leaks (issue #14540)
|
||||
from litellm.llms.custom_httpx.http_handler import _get_httpx_client
|
||||
# Get unified SSL configuration
|
||||
ssl_config = get_ssl_configuration()
|
||||
|
||||
try:
|
||||
# Get SSL config and include in params for proper cache key
|
||||
ssl_config = get_ssl_configuration()
|
||||
params = {"ssl_verify": ssl_config} if ssl_config is not None else None
|
||||
|
||||
# Get a cached HTTPHandler which manages the httpx.Client
|
||||
cached_handler = _get_httpx_client(params=params)
|
||||
# Return the underlying httpx client from the handler
|
||||
return cached_handler.client
|
||||
except (ImportError, AttributeError, KeyError) as e:
|
||||
# Fallback to creating a client directly if caching system unavailable
|
||||
verbose_logger.debug(
|
||||
f"Client caching unavailable ({type(e).__name__}), using direct client creation"
|
||||
)
|
||||
ssl_config = get_ssl_configuration()
|
||||
return httpx.Client(
|
||||
verify=ssl_config,
|
||||
follow_redirects=True,
|
||||
)
|
||||
return httpx.Client(
|
||||
verify=ssl_config,
|
||||
follow_redirects=True,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -2201,14 +2201,24 @@ def completion( # type: ignore # noqa: PLR0915
|
|||
)
|
||||
elif custom_llm_provider == "a2a":
|
||||
# A2A (Agent-to-Agent) Protocol
|
||||
api_base = (
|
||||
api_base
|
||||
or litellm.api_base
|
||||
or get_secret_str("A2A_API_BASE")
|
||||
# Resolve agent configuration from registry if model format is "a2a/<agent-name>"
|
||||
api_base, api_key, headers = litellm.A2AConfig.resolve_agent_config_from_registry(
|
||||
model=model,
|
||||
api_base=api_base,
|
||||
api_key=api_key,
|
||||
headers=headers,
|
||||
optional_params=optional_params,
|
||||
)
|
||||
|
||||
|
||||
# Fall back to environment variables and defaults
|
||||
api_base = api_base or litellm.api_base or get_secret_str("A2A_API_BASE")
|
||||
|
||||
if api_base is None:
|
||||
raise Exception("api_base is required for A2A provider")
|
||||
raise Exception(
|
||||
"api_base is required for A2A provider. "
|
||||
"Either provide api_base parameter, set A2A_API_BASE environment variable, "
|
||||
"or register the agent in the proxy with model='a2a/<agent-name>'."
|
||||
)
|
||||
|
||||
headers = headers or litellm.headers
|
||||
|
||||
|
|
|
|||
53
litellm/proxy/agent_endpoints/a2a_routing.py
Normal file
53
litellm/proxy/agent_endpoints/a2a_routing.py
Normal file
|
|
@ -0,0 +1,53 @@
|
|||
"""
|
||||
A2A Agent Routing
|
||||
|
||||
Handles routing for A2A agents (models with "a2a/<agent-name>" prefix).
|
||||
Looks up agents in the registry and injects their API base URL.
|
||||
"""
|
||||
|
||||
from typing import Any, Optional
|
||||
|
||||
import litellm
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
|
||||
async def route_a2a_agent_request(data: dict, route_type: str) -> Optional[Any]:
|
||||
"""
|
||||
Route A2A agent requests directly to litellm with injected API base.
|
||||
|
||||
Returns None if not an A2A request (allows normal routing to continue).
|
||||
"""
|
||||
# Import here to avoid circular imports
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.proxy.route_llm_request import (
|
||||
ROUTE_ENDPOINT_MAPPING,
|
||||
ProxyModelNotFoundError,
|
||||
)
|
||||
|
||||
model_name = data.get("model", "")
|
||||
|
||||
# Check if this is an A2A agent request
|
||||
if not isinstance(model_name, str) or not model_name.startswith("a2a/"):
|
||||
return None
|
||||
|
||||
# Extract agent name (e.g., "a2a/my-agent" -> "my-agent")
|
||||
agent_name = model_name[4:]
|
||||
|
||||
# Look up agent in registry
|
||||
agent = global_agent_registry.get_agent_by_name(agent_name)
|
||||
if agent is None:
|
||||
verbose_proxy_logger.error(f"[A2A] Agent '{agent_name}' not found in registry")
|
||||
route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type)
|
||||
raise ProxyModelNotFoundError(route=route_name, model_name=model_name)
|
||||
|
||||
# Get API base URL from agent config
|
||||
if not agent.agent_card_params or "url" not in agent.agent_card_params:
|
||||
verbose_proxy_logger.error(f"[A2A] Agent '{agent_name}' has no URL configured")
|
||||
route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type)
|
||||
raise ProxyModelNotFoundError(route=route_name, model_name=model_name)
|
||||
|
||||
# Inject API base and route to litellm
|
||||
data["api_base"] = agent.agent_card_params["url"]
|
||||
verbose_proxy_logger.debug(f"[A2A] Routing {model_name} to {data['api_base']}")
|
||||
|
||||
return getattr(litellm, f"{route_type}")(**data)
|
||||
96
litellm/proxy/agent_endpoints/model_list_helpers.py
Normal file
96
litellm/proxy/agent_endpoints/model_list_helpers.py
Normal file
|
|
@ -0,0 +1,96 @@
|
|||
"""
|
||||
Helper functions for appending A2A agents to model lists.
|
||||
|
||||
Used by proxy model endpoints to make agents appear in UI alongside models.
|
||||
"""
|
||||
from typing import List
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
|
||||
ModelGroupInfoProxy,
|
||||
)
|
||||
|
||||
|
||||
async def append_agents_to_model_group(
|
||||
model_groups: List[ModelGroupInfoProxy],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> List[ModelGroupInfoProxy]:
|
||||
"""
|
||||
Append A2A agents to model groups list for UI display.
|
||||
|
||||
Converts agents to model format with "a2a/<agent-name>" naming
|
||||
so they appear in playground and work with LiteLLM routing.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
|
||||
AgentRequestHandler,
|
||||
)
|
||||
|
||||
allowed_agent_ids = await AgentRequestHandler.get_allowed_agents(
|
||||
user_api_key_auth=user_api_key_dict
|
||||
)
|
||||
|
||||
for agent_id in allowed_agent_ids:
|
||||
agent = global_agent_registry.get_agent_by_id(agent_id)
|
||||
if agent is not None:
|
||||
model_groups.append(
|
||||
ModelGroupInfoProxy(
|
||||
model_group=f"a2a/{agent.agent_name}",
|
||||
mode="chat",
|
||||
providers=["a2a"],
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Error appending agents to model_group/info: {e}"
|
||||
)
|
||||
|
||||
return model_groups
|
||||
|
||||
|
||||
async def append_agents_to_model_info(
|
||||
models: List[dict],
|
||||
user_api_key_dict: UserAPIKeyAuth,
|
||||
) -> List[dict]:
|
||||
"""
|
||||
Append A2A agents to model info list for UI display.
|
||||
|
||||
Converts agents to model format with "a2a/<agent-name>" naming
|
||||
so they appear in models page and work with LiteLLM routing.
|
||||
"""
|
||||
try:
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.proxy.agent_endpoints.auth.agent_permission_handler import (
|
||||
AgentRequestHandler,
|
||||
)
|
||||
|
||||
allowed_agent_ids = await AgentRequestHandler.get_allowed_agents(
|
||||
user_api_key_auth=user_api_key_dict
|
||||
)
|
||||
|
||||
for agent_id in allowed_agent_ids:
|
||||
agent = global_agent_registry.get_agent_by_id(agent_id)
|
||||
if agent is not None:
|
||||
models.append({
|
||||
"model_name": f"a2a/{agent.agent_name}",
|
||||
"litellm_params": {
|
||||
"model": f"a2a/{agent.agent_name}",
|
||||
"custom_llm_provider": "a2a",
|
||||
},
|
||||
"model_info": {
|
||||
"id": agent.agent_id,
|
||||
"mode": "chat",
|
||||
"db_model": True,
|
||||
"created_by": agent.created_by,
|
||||
"created_at": agent.created_at,
|
||||
"updated_at": agent.updated_at,
|
||||
},
|
||||
})
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.debug(
|
||||
f"Error appending agents to v2/model/info: {e}"
|
||||
)
|
||||
|
||||
return models
|
||||
|
|
@ -239,6 +239,10 @@ from litellm.proxy._types import *
|
|||
from litellm.proxy.agent_endpoints.a2a_endpoints import router as a2a_router
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.proxy.agent_endpoints.endpoints import router as agent_endpoints_router
|
||||
from litellm.proxy.agent_endpoints.model_list_helpers import (
|
||||
append_agents_to_model_group,
|
||||
append_agents_to_model_info,
|
||||
)
|
||||
from litellm.proxy.analytics_endpoints.analytics_endpoints import (
|
||||
router as analytics_router,
|
||||
)
|
||||
|
|
@ -8616,6 +8620,15 @@ async def model_info_v2(
|
|||
)
|
||||
|
||||
verbose_proxy_logger.debug("all_models: %s", all_models)
|
||||
|
||||
# Append A2A agents to models list
|
||||
all_models = await append_agents_to_model_info(
|
||||
models=all_models,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Update total count to include agents
|
||||
search_total_count = len(all_models)
|
||||
|
||||
return _paginate_models_response(
|
||||
all_models=all_models,
|
||||
|
|
@ -9456,6 +9469,12 @@ async def model_group_info(
|
|||
model_groups: List[ModelGroupInfoProxy] = _get_model_group_info(
|
||||
llm_router=llm_router, all_models_str=all_models_str, model_group=model_group
|
||||
)
|
||||
|
||||
# Append A2A agents to model groups
|
||||
model_groups = await append_agents_to_model_group(
|
||||
model_groups=model_groups,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
return {"data": model_groups}
|
||||
|
||||
|
|
|
|||
|
|
@ -12,6 +12,11 @@ else:
|
|||
LitellmRouter = Any
|
||||
|
||||
|
||||
def _is_a2a_agent_model(model_name: Any) -> bool:
|
||||
"""Check if the model name is for an A2A agent (a2a/ prefix)."""
|
||||
return isinstance(model_name, str) and model_name.startswith("a2a/")
|
||||
|
||||
|
||||
ROUTE_ENDPOINT_MAPPING = {
|
||||
"acompletion": "/chat/completions",
|
||||
"atext_completion": "/completions",
|
||||
|
|
@ -322,6 +327,12 @@ async def route_request(
|
|||
except Exception:
|
||||
# If router fails (e.g., model not found in router), fall back to direct call
|
||||
return getattr(litellm, f"{route_type}")(**data)
|
||||
elif _is_a2a_agent_model(data.get("model", "")):
|
||||
from litellm.proxy.agent_endpoints.a2a_routing import (
|
||||
route_a2a_agent_request,
|
||||
)
|
||||
|
||||
return await route_a2a_agent_request(data, route_type)
|
||||
|
||||
elif user_model is not None:
|
||||
return getattr(litellm, f"{route_type}")(**data)
|
||||
|
|
|
|||
81
poetry.lock
generated
81
poetry.lock
generated
|
|
@ -723,6 +723,18 @@ markers = {main = "platform_python_implementation != \"PyPy\" or extra == \"prox
|
|||
[package.dependencies]
|
||||
pycparser = {version = "*", markers = "implementation_name != \"PyPy\""}
|
||||
|
||||
[[package]]
|
||||
name = "chardet"
|
||||
version = "5.2.0"
|
||||
description = "Universal encoding detector for Python 3"
|
||||
optional = false
|
||||
python-versions = ">=3.7"
|
||||
groups = ["dev"]
|
||||
files = [
|
||||
{file = "chardet-5.2.0-py3-none-any.whl", hash = "sha256:e1cf59446890a00105fe7b7912492ea04b6e6f06d4b742b2c788469e34c82970"},
|
||||
{file = "chardet-5.2.0.tar.gz", hash = "sha256:1b3b6ff479a8c414bc3fa2c0852995695c4a026dcd6d0633b2dd092ca39c1cf7"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
name = "charset-normalizer"
|
||||
version = "3.4.4"
|
||||
|
|
@ -1300,6 +1312,27 @@ dev = ["autoflake", "black", "build", "databricks-connect", "httpx", "ipython",
|
|||
notebook = ["ipython (>=8,<10)", "ipywidgets (>=8,<9)"]
|
||||
openai = ["httpx", "langchain-openai ; python_version > \"3.7\"", "openai"]
|
||||
|
||||
[[package]]
|
||||
name = "diff-cover"
|
||||
version = "9.7.2"
|
||||
description = "Run coverage and linting reports on diffs"
|
||||
optional = false
|
||||
python-versions = ">=3.9"
|
||||
groups = ["dev"]
|
||||
files = [
|
||||
{file = "diff_cover-9.7.2-py3-none-any.whl", hash = "sha256:cd6498620c747c2493a6c83c14362c32868bfd91cd8d0dd093f136070ec4ffc5"},
|
||||
{file = "diff_cover-9.7.2.tar.gz", hash = "sha256:872c820d2ecbf79c61d52c7dc70419015e0ab9289589566c791dd270fc0c6e3b"},
|
||||
]
|
||||
|
||||
[package.dependencies]
|
||||
chardet = ">=3.0.0"
|
||||
Jinja2 = ">=2.7.1"
|
||||
pluggy = ">=0.13.1,<2"
|
||||
Pygments = ">=2.19.1,<3.0.0"
|
||||
|
||||
[package.extras]
|
||||
toml = ["tomli (>=1.2.1)"]
|
||||
|
||||
[[package]]
|
||||
name = "diskcache"
|
||||
version = "5.6.3"
|
||||
|
|
@ -3085,7 +3118,7 @@ version = "3.1.6"
|
|||
description = "A very fast and expressive template engine."
|
||||
optional = false
|
||||
python-versions = ">=3.7"
|
||||
groups = ["main", "proxy-dev"]
|
||||
groups = ["main", "dev", "proxy-dev"]
|
||||
files = [
|
||||
{file = "jinja2-3.1.6-py3-none-any.whl", hash = "sha256:85ece4451f492d0c13c5dd7c13a64681a86afae63a5f347908daf103ce6d2f67"},
|
||||
{file = "jinja2-3.1.6.tar.gz", hash = "sha256:0137fb05990d35f1275a587e9aee6d56da821fc83491a0fb838183be43f66d6d"},
|
||||
|
|
@ -3515,7 +3548,7 @@ version = "3.0.3"
|
|||
description = "Safely add untrusted strings to HTML/XML markup."
|
||||
optional = false
|
||||
python-versions = ">=3.9"
|
||||
groups = ["main", "proxy-dev"]
|
||||
groups = ["main", "dev", "proxy-dev"]
|
||||
files = [
|
||||
{file = "markupsafe-3.0.3-cp310-cp310-macosx_10_9_x86_64.whl", hash = "sha256:2f981d352f04553a7171b8e44369f2af4055f888dfb147d55e42d29e29e74559"},
|
||||
{file = "markupsafe-3.0.3-cp310-cp310-macosx_11_0_arm64.whl", hash = "sha256:e1c1493fb6e50ab01d20a22826e57520f1284df32f2d8601fdd90b6304601419"},
|
||||
|
|
@ -5516,14 +5549,14 @@ files = [
|
|||
name = "pygments"
|
||||
version = "2.19.2"
|
||||
description = "Pygments is a syntax highlighting package written in Python."
|
||||
optional = true
|
||||
optional = false
|
||||
python-versions = ">=3.8"
|
||||
groups = ["main"]
|
||||
markers = "extra == \"utils\" or extra == \"proxy\""
|
||||
groups = ["main", "dev"]
|
||||
files = [
|
||||
{file = "pygments-2.19.2-py3-none-any.whl", hash = "sha256:86540386c03d588bb81d44bc3928634ff26449851e99741617ecb9037ee5ec0b"},
|
||||
{file = "pygments-2.19.2.tar.gz", hash = "sha256:636cb2477cec7f8952536970bc533bc43743542f70392ae026374600add5b887"},
|
||||
]
|
||||
markers = {main = "extra == \"utils\" or extra == \"proxy\""}
|
||||
|
||||
[package.extras]
|
||||
windows-terminal = ["colorama (>=0.4.6)"]
|
||||
|
|
@ -6601,29 +6634,29 @@ pyasn1 = ">=0.1.3"
|
|||
|
||||
[[package]]
|
||||
name = "ruff"
|
||||
version = "0.1.15"
|
||||
version = "0.2.2"
|
||||
description = "An extremely fast Python linter and code formatter, written in Rust."
|
||||
optional = false
|
||||
python-versions = ">=3.7"
|
||||
groups = ["dev"]
|
||||
files = [
|
||||
{file = "ruff-0.1.15-py3-none-macosx_10_12_x86_64.macosx_11_0_arm64.macosx_10_12_universal2.whl", hash = "sha256:5fe8d54df166ecc24106db7dd6a68d44852d14eb0729ea4672bb4d96c320b7df"},
|
||||
{file = "ruff-0.1.15-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:6f0bfbb53c4b4de117ac4d6ddfd33aa5fc31beeaa21d23c45c6dd249faf9126f"},
|
||||
{file = "ruff-0.1.15-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:e0d432aec35bfc0d800d4f70eba26e23a352386be3a6cf157083d18f6f5881c8"},
|
||||
{file = "ruff-0.1.15-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:9405fa9ac0e97f35aaddf185a1be194a589424b8713e3b97b762336ec79ff807"},
|
||||
{file = "ruff-0.1.15-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:c66ec24fe36841636e814b8f90f572a8c0cb0e54d8b5c2d0e300d28a0d7bffec"},
|
||||
{file = "ruff-0.1.15-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:6f8ad828f01e8dd32cc58bc28375150171d198491fc901f6f98d2a39ba8e3ff5"},
|
||||
{file = "ruff-0.1.15-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:86811954eec63e9ea162af0ffa9f8d09088bab51b7438e8b6488b9401863c25e"},
|
||||
{file = "ruff-0.1.15-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:fd4025ac5e87d9b80e1f300207eb2fd099ff8200fa2320d7dc066a3f4622dc6b"},
|
||||
{file = "ruff-0.1.15-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:b17b93c02cdb6aeb696effecea1095ac93f3884a49a554a9afa76bb125c114c1"},
|
||||
{file = "ruff-0.1.15-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:ddb87643be40f034e97e97f5bc2ef7ce39de20e34608f3f829db727a93fb82c5"},
|
||||
{file = "ruff-0.1.15-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:abf4822129ed3a5ce54383d5f0e964e7fef74a41e48eb1dfad404151efc130a2"},
|
||||
{file = "ruff-0.1.15-py3-none-musllinux_1_2_i686.whl", hash = "sha256:6c629cf64bacfd136c07c78ac10a54578ec9d1bd2a9d395efbee0935868bf852"},
|
||||
{file = "ruff-0.1.15-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:1bab866aafb53da39c2cadfb8e1c4550ac5340bb40300083eb8967ba25481447"},
|
||||
{file = "ruff-0.1.15-py3-none-win32.whl", hash = "sha256:2417e1cb6e2068389b07e6fa74c306b2810fe3ee3476d5b8a96616633f40d14f"},
|
||||
{file = "ruff-0.1.15-py3-none-win_amd64.whl", hash = "sha256:3837ac73d869efc4182d9036b1405ef4c73d9b1f88da2413875e34e0d6919587"},
|
||||
{file = "ruff-0.1.15-py3-none-win_arm64.whl", hash = "sha256:9a933dfb1c14ec7a33cceb1e49ec4a16b51ce3c20fd42663198746efc0427360"},
|
||||
{file = "ruff-0.1.15.tar.gz", hash = "sha256:f6dfa8c1b21c913c326919056c390966648b680966febcb796cc9d1aaab8564e"},
|
||||
{file = "ruff-0.2.2-py3-none-macosx_10_12_x86_64.macosx_11_0_arm64.macosx_10_12_universal2.whl", hash = "sha256:0a9efb032855ffb3c21f6405751d5e147b0c6b631e3ca3f6b20f917572b97eb6"},
|
||||
{file = "ruff-0.2.2-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:d450b7fbff85913f866a5384d8912710936e2b96da74541c82c1b458472ddb39"},
|
||||
{file = "ruff-0.2.2-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.whl", hash = "sha256:ecd46e3106850a5c26aee114e562c329f9a1fbe9e4821b008c4404f64ff9ce73"},
|
||||
{file = "ruff-0.2.2-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:5e22676a5b875bd72acd3d11d5fa9075d3a5f53b877fe7b4793e4673499318ba"},
|
||||
{file = "ruff-0.2.2-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:1695700d1e25a99d28f7a1636d85bafcc5030bba9d0578c0781ba1790dbcf51c"},
|
||||
{file = "ruff-0.2.2-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:b0c232af3d0bd8f521806223723456ffebf8e323bd1e4e82b0befb20ba18388e"},
|
||||
{file = "ruff-0.2.2-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:f63d96494eeec2fc70d909393bcd76c69f35334cdbd9e20d089fb3f0640216ca"},
|
||||
{file = "ruff-0.2.2-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:6a61ea0ff048e06de273b2e45bd72629f470f5da8f71daf09fe481278b175001"},
|
||||
{file = "ruff-0.2.2-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:5e1439c8f407e4f356470e54cdecdca1bd5439a0673792dbe34a2b0a551a2fe3"},
|
||||
{file = "ruff-0.2.2-py3-none-musllinux_1_2_aarch64.whl", hash = "sha256:940de32dc8853eba0f67f7198b3e79bc6ba95c2edbfdfac2144c8235114d6726"},
|
||||
{file = "ruff-0.2.2-py3-none-musllinux_1_2_armv7l.whl", hash = "sha256:0c126da55c38dd917621552ab430213bdb3273bb10ddb67bc4b761989210eb6e"},
|
||||
{file = "ruff-0.2.2-py3-none-musllinux_1_2_i686.whl", hash = "sha256:3b65494f7e4bed2e74110dac1f0d17dc8e1f42faaa784e7c58a98e335ec83d7e"},
|
||||
{file = "ruff-0.2.2-py3-none-musllinux_1_2_x86_64.whl", hash = "sha256:1ec49be4fe6ddac0503833f3ed8930528e26d1e60ad35c2446da372d16651ce9"},
|
||||
{file = "ruff-0.2.2-py3-none-win32.whl", hash = "sha256:d920499b576f6c68295bc04e7b17b6544d9d05f196bb3aac4358792ef6f34325"},
|
||||
{file = "ruff-0.2.2-py3-none-win_amd64.whl", hash = "sha256:cc9a91ae137d687f43a44c900e5d95e9617cb37d4c989e462980ba27039d239d"},
|
||||
{file = "ruff-0.2.2-py3-none-win_arm64.whl", hash = "sha256:c9d15fc41e6054bfc7200478720570078f0b41c9ae4f010bcc16bd6f4d1aacdd"},
|
||||
{file = "ruff-0.2.2.tar.gz", hash = "sha256:e62ed7f36b3068a30ba39193a14274cd706bc486fad521276458022f7bccb31d"},
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
@ -8490,4 +8523,4 @@ utils = ["numpydoc"]
|
|||
[metadata]
|
||||
lock-version = "2.1"
|
||||
python-versions = ">=3.9,<4.0"
|
||||
content-hash = "95fd27dc139d0e52e70093220c50582f16c78e5977ec77f4297f50a30df964c6"
|
||||
content-hash = "797603dcfef0a79781c7d3cba5dfe18f6aea4aa792220f47487ebc7bd04ae2e3"
|
||||
|
|
|
|||
|
|
@ -139,6 +139,7 @@ litellm = 'litellm:run_server'
|
|||
litellm-proxy = 'litellm.proxy.client.cli:cli'
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
diff-cover = "^9.0"
|
||||
flake8 = "^6.1.0"
|
||||
black = "^23.12.0"
|
||||
mypy = "^1.0"
|
||||
|
|
@ -149,7 +150,7 @@ pytest-retry = "^1.6.3"
|
|||
requests-mock = "^1.12.1"
|
||||
responses = "^0.25.7"
|
||||
respx = "^0.22.0"
|
||||
ruff = "^0.1.0"
|
||||
ruff = "^0.2.1"
|
||||
types-requests = "*"
|
||||
types-setuptools = "*"
|
||||
types-redis = "*"
|
||||
|
|
|
|||
46
tests/test_litellm/llms/test_lifecycle_fix.py
Normal file
46
tests/test_litellm/llms/test_lifecycle_fix.py
Normal file
|
|
@ -0,0 +1,46 @@
|
|||
"""
|
||||
Verifies that the httpx client used by AsyncOpenAI is NOT closed
|
||||
when AsyncHTTPHandler instances are garbage collected.
|
||||
"""
|
||||
import asyncio
|
||||
import gc
|
||||
import httpx
|
||||
from litellm.llms.openai.common_utils import BaseOpenAILLM
|
||||
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler
|
||||
|
||||
|
||||
async def test_httpx_client_not_closed_by_handler_gc():
|
||||
"""
|
||||
Before the fix: _get_async_http_client() returned handler.client,
|
||||
so when handler was GC'd its __del__ closed the client.
|
||||
After the fix: returns a standalone httpx.AsyncClient, no handler involved.
|
||||
"""
|
||||
# Get the client the same way AsyncOpenAI would
|
||||
client = BaseOpenAILLM._get_async_http_client()
|
||||
assert isinstance(client, httpx.AsyncClient)
|
||||
|
||||
# Simulate what the old code did: create an AsyncHTTPHandler and GC it
|
||||
handler = AsyncHTTPHandler()
|
||||
handler_client = handler.client
|
||||
del handler
|
||||
gc.collect()
|
||||
|
||||
# The client from _get_async_http_client should still be open
|
||||
# because it's NOT tied to any AsyncHTTPHandler
|
||||
assert not client.is_closed, "Client was closed prematurely!"
|
||||
|
||||
# Verify it can actually send (build a request without sending)
|
||||
try:
|
||||
req = client.build_request("GET", "https://example.com")
|
||||
print("PASS: Client is still usable after handler GC")
|
||||
except RuntimeError as e:
|
||||
if "closed" in str(e):
|
||||
print(f"FAIL: {e}")
|
||||
raise
|
||||
raise
|
||||
|
||||
await client.aclose()
|
||||
print("All checks passed!")
|
||||
|
||||
|
||||
asyncio.run(test_httpx_client_not_closed_by_handler_gc())
|
||||
|
|
@ -0,0 +1,110 @@
|
|||
"""
|
||||
Test appending A2A agents to model lists.
|
||||
|
||||
Maps to: litellm/proxy/agent_endpoints/model_list_helpers.py
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../../.."))
|
||||
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.agent_endpoints.model_list_helpers import (
|
||||
append_agents_to_model_group,
|
||||
append_agents_to_model_info,
|
||||
)
|
||||
from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth
|
||||
from litellm.types.agents import AgentResponse
|
||||
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
|
||||
ModelGroupInfoProxy,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_append_agents_to_model_group():
|
||||
"""Test agents are converted to model group format with a2a/ prefix"""
|
||||
|
||||
# Mock agent data
|
||||
mock_agent = AgentResponse(
|
||||
agent_id="test-agent-id",
|
||||
agent_name="my-agent",
|
||||
agent_card_params={"url": "http://example.com"},
|
||||
litellm_params=None,
|
||||
)
|
||||
|
||||
# Mock AgentRequestHandler at its source location
|
||||
mock_get_allowed_agents = AsyncMock(return_value=["test-agent-id"])
|
||||
|
||||
# Mock global_agent_registry
|
||||
mock_registry = Mock()
|
||||
mock_registry.get_agent_by_id = Mock(return_value=mock_agent)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.get_allowed_agents",
|
||||
mock_get_allowed_agents,
|
||||
):
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry",
|
||||
mock_registry,
|
||||
):
|
||||
model_groups = []
|
||||
user_api_key_dict = Mock(spec=UserAPIKeyAuth)
|
||||
|
||||
result = await append_agents_to_model_group(
|
||||
model_groups=model_groups,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify agent was converted with a2a/ prefix
|
||||
assert len(result) == 1
|
||||
assert result[0].model_group == "a2a/my-agent"
|
||||
assert result[0].mode == "chat"
|
||||
assert result[0].providers == ["a2a"]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_append_agents_to_model_info():
|
||||
"""Test agents are converted to model info format with a2a/ prefix"""
|
||||
|
||||
# Mock agent data
|
||||
mock_agent = AgentResponse(
|
||||
agent_id="agent-123",
|
||||
agent_name="test-agent",
|
||||
agent_card_params={"url": "http://example.com"},
|
||||
litellm_params=None,
|
||||
created_by="user-123",
|
||||
)
|
||||
|
||||
# Mock AgentRequestHandler at its source location
|
||||
mock_get_allowed_agents = AsyncMock(return_value=["agent-123"])
|
||||
|
||||
# Mock global_agent_registry
|
||||
mock_registry = Mock()
|
||||
mock_registry.get_agent_by_id = Mock(return_value=mock_agent)
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.get_allowed_agents",
|
||||
mock_get_allowed_agents,
|
||||
):
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.agent_registry.global_agent_registry",
|
||||
mock_registry,
|
||||
):
|
||||
models = []
|
||||
user_api_key_dict = Mock(spec=UserAPIKeyAuth)
|
||||
|
||||
result = await append_agents_to_model_info(
|
||||
models=models,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
|
||||
# Verify agent was converted with a2a/ prefix
|
||||
assert len(result) == 1
|
||||
assert result[0]["model_name"] == "a2a/test-agent"
|
||||
assert result[0]["litellm_params"]["model"] == "a2a/test-agent"
|
||||
assert result[0]["litellm_params"]["custom_llm_provider"] == "a2a"
|
||||
assert result[0]["model_info"]["id"] == "agent-123"
|
||||
assert result[0]["model_info"]["mode"] == "chat"
|
||||
105
tests/test_litellm/proxy/test_route_a2a_models.py
Normal file
105
tests/test_litellm/proxy/test_route_a2a_models.py
Normal file
|
|
@ -0,0 +1,105 @@
|
|||
"""
|
||||
Test A2A model routing in proxy.
|
||||
|
||||
Maps to: litellm/proxy/agent_endpoints/a2a_routing.py
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../../.."))
|
||||
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.agent_endpoints.a2a_routing import route_a2a_agent_request
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_a2a_model_bypasses_router():
|
||||
"""Test that a2a/ prefixed models bypass router and go directly to litellm with api_base"""
|
||||
|
||||
# Mock data for chat completion with a2a model
|
||||
data = {
|
||||
"model": "a2a/test-agent",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
}
|
||||
|
||||
# Mock router that doesn't have the a2a model
|
||||
mock_router = Mock()
|
||||
mock_router.model_names = ["gpt-4", "gpt-3.5-turbo"]
|
||||
mock_router.deployment_names = []
|
||||
mock_router.has_model_id = Mock(return_value=False)
|
||||
mock_router.model_group_alias = None
|
||||
mock_router.router_general_settings = Mock(pass_through_all_models=False)
|
||||
mock_router.default_deployment = None
|
||||
mock_router.pattern_router = Mock(patterns=[])
|
||||
mock_router.map_team_model = Mock(return_value=None)
|
||||
|
||||
# Mock agent in registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
mock_agent = AgentResponse(
|
||||
agent_id="test-agent-id",
|
||||
agent_name="test-agent",
|
||||
agent_card_params={"url": "http://agent.example.com"},
|
||||
litellm_params=None,
|
||||
)
|
||||
|
||||
mock_registry = Mock()
|
||||
mock_registry.get_agent_by_name = Mock(return_value=mock_agent)
|
||||
|
||||
# Mock litellm.acompletion to verify it's called
|
||||
mock_acompletion = AsyncMock(return_value={"id": "test-response"})
|
||||
|
||||
with patch("litellm.acompletion", mock_acompletion):
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.a2a_routing.global_agent_registry",
|
||||
mock_registry,
|
||||
):
|
||||
result = await route_request(
|
||||
data=data,
|
||||
llm_router=mock_router,
|
||||
user_model=None,
|
||||
route_type="acompletion",
|
||||
)
|
||||
|
||||
# Verify litellm.acompletion was called with api_base injected
|
||||
mock_acompletion.assert_called_once()
|
||||
call_kwargs = mock_acompletion.call_args.kwargs
|
||||
assert call_kwargs["model"] == "a2a/test-agent"
|
||||
assert call_kwargs["api_base"] == "http://agent.example.com"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_route_non_a2a_model_raises_error_if_not_in_router():
|
||||
"""Test that non-a2a models that aren't in router raise an error"""
|
||||
|
||||
# Mock data for chat completion with model not in router
|
||||
data = {
|
||||
"model": "unknown-model",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
}
|
||||
|
||||
# Mock router without the model
|
||||
mock_router = Mock()
|
||||
mock_router.model_names = ["gpt-4", "gpt-3.5-turbo"]
|
||||
mock_router.deployment_names = []
|
||||
mock_router.has_model_id = Mock(return_value=False)
|
||||
mock_router.model_group_alias = None
|
||||
mock_router.router_general_settings = Mock(pass_through_all_models=False)
|
||||
mock_router.default_deployment = None
|
||||
mock_router.pattern_router = Mock(patterns=[])
|
||||
mock_router.map_team_model = Mock(return_value=None)
|
||||
|
||||
# Should raise ProxyModelNotFoundError
|
||||
from litellm.proxy.route_llm_request import ProxyModelNotFoundError
|
||||
|
||||
with pytest.raises(ProxyModelNotFoundError):
|
||||
await route_request(
|
||||
data=data,
|
||||
llm_router=mock_router,
|
||||
user_model=None,
|
||||
route_type="acompletion",
|
||||
)
|
||||
73
tests/test_litellm/test_a2a_registry_lookup.py
Normal file
73
tests/test_litellm/test_a2a_registry_lookup.py
Normal file
|
|
@ -0,0 +1,73 @@
|
|||
"""
|
||||
Test A2A provider registry lookup functionality.
|
||||
|
||||
Maps to: litellm/llms/a2a/chat/transformation.py
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.llms.a2a.chat.transformation import A2AConfig
|
||||
|
||||
|
||||
def test_resolve_agent_config_from_registry_static_method():
|
||||
"""Test the static helper method for registry resolution"""
|
||||
|
||||
# Test 1: No agent name in model
|
||||
api_base, api_key, headers = A2AConfig.resolve_agent_config_from_registry(
|
||||
model="a2a",
|
||||
api_base="http://test.com",
|
||||
api_key=None,
|
||||
headers=None,
|
||||
optional_params={}
|
||||
)
|
||||
assert api_base == "http://test.com"
|
||||
|
||||
# Test 2: All params provided - should not lookup registry
|
||||
api_base, api_key, headers = A2AConfig.resolve_agent_config_from_registry(
|
||||
model="a2a/test-agent",
|
||||
api_base="http://explicit.com",
|
||||
api_key="explicit-key",
|
||||
headers={"X-Test": "value"},
|
||||
optional_params={}
|
||||
)
|
||||
assert api_base == "http://explicit.com"
|
||||
assert api_key == "explicit-key"
|
||||
|
||||
|
||||
def test_a2a_registry_integration():
|
||||
"""Test registry lookup in proxy context"""
|
||||
|
||||
try:
|
||||
from litellm.proxy.agent_endpoints.agent_registry import global_agent_registry
|
||||
from litellm.types.agents import AgentResponse
|
||||
|
||||
# Create test agent
|
||||
test_agent = AgentResponse(
|
||||
agent_id="test-id",
|
||||
agent_name="test-agent",
|
||||
agent_card_params={"url": "http://registry-url.example.com:9999"},
|
||||
litellm_params={"api_key": "registry-key"},
|
||||
)
|
||||
|
||||
# Register and test
|
||||
original_agents = global_agent_registry.agent_list.copy()
|
||||
global_agent_registry.register_agent(test_agent)
|
||||
|
||||
try:
|
||||
litellm.completion(
|
||||
model="a2a/test-agent",
|
||||
messages=[{"role": "user", "content": "Hello"}]
|
||||
)
|
||||
except Exception as e:
|
||||
# Should use registry URL (connection error expected)
|
||||
assert "registry-url.example.com" in str(e) or "APIConnectionError" in str(type(e).__name__)
|
||||
finally:
|
||||
global_agent_registry.agent_list = original_agents
|
||||
|
||||
except ImportError:
|
||||
pytest.skip("Registry not available (not in proxy context)")
|
||||
|
|
@ -1,24 +1,61 @@
|
|||
import { fireEvent, waitFor } from "@testing-library/react";
|
||||
import { describe, expect, it, vi } from "vitest";
|
||||
import { screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { renderWithProviders } from "../../../tests/test-utils";
|
||||
import { KeyResponse } from "../key_team_helpers/key_list";
|
||||
import { KeyEditView } from "./key_edit_view";
|
||||
|
||||
// Mock window.matchMedia
|
||||
Object.defineProperty(window, "matchMedia", {
|
||||
writable: true,
|
||||
value: vi.fn().mockImplementation((query) => ({
|
||||
matches: false,
|
||||
media: query,
|
||||
onchange: null,
|
||||
addListener: vi.fn(),
|
||||
removeListener: vi.fn(),
|
||||
addEventListener: vi.fn(),
|
||||
removeEventListener: vi.fn(),
|
||||
dispatchEvent: vi.fn(),
|
||||
})),
|
||||
vi.mock("../networking", async () => {
|
||||
const actual = await vi.importActual("../networking");
|
||||
return {
|
||||
...actual,
|
||||
getPromptsList: vi.fn().mockResolvedValue({
|
||||
prompts: [{ prompt_id: "prompt-1" }, { prompt_id: "prompt-2" }],
|
||||
}),
|
||||
modelAvailableCall: vi.fn().mockResolvedValue({
|
||||
data: [{ id: "gpt-4" }, { id: "gpt-3.5-turbo" }],
|
||||
}),
|
||||
tagListCall: vi.fn().mockResolvedValue({
|
||||
tag1: { name: "tag1", description: "Test tag 1" },
|
||||
tag2: { name: "tag2", description: "Test tag 2" },
|
||||
}),
|
||||
getGuardrailsList: vi.fn().mockResolvedValue({
|
||||
guardrails: [{ guardrail_name: "guardrail-1" }],
|
||||
}),
|
||||
getPoliciesList: vi.fn().mockResolvedValue({
|
||||
policies: [{ policy_name: "policy-1" }],
|
||||
}),
|
||||
getPassThroughEndpointsCall: vi.fn().mockResolvedValue({
|
||||
endpoints: [],
|
||||
}),
|
||||
vectorStoreListCall: vi.fn().mockResolvedValue({
|
||||
data: [],
|
||||
}),
|
||||
mcpToolsCall: vi.fn().mockResolvedValue({
|
||||
data: [],
|
||||
}),
|
||||
agentListCall: vi.fn().mockResolvedValue({
|
||||
data: [],
|
||||
}),
|
||||
fetchMCPServers: vi.fn().mockResolvedValue([]),
|
||||
fetchMCPAccessGroups: vi.fn().mockResolvedValue([]),
|
||||
listMCPTools: vi.fn().mockResolvedValue({
|
||||
tools: [],
|
||||
error: null,
|
||||
message: null,
|
||||
stack_trace: null,
|
||||
}),
|
||||
getAgentsList: vi.fn().mockResolvedValue({
|
||||
agents: [],
|
||||
}),
|
||||
getAgentAccessGroups: vi.fn().mockResolvedValue([]),
|
||||
};
|
||||
});
|
||||
|
||||
vi.mock("../organisms/create_key_button", () => ({
|
||||
fetchTeamModels: vi.fn().mockResolvedValue(["team-model-1", "team-model-2"]),
|
||||
}));
|
||||
|
||||
describe("KeyEditView", () => {
|
||||
const MOCK_KEY_DATA: KeyResponse = {
|
||||
token: "test-token-123",
|
||||
|
|
@ -93,8 +130,8 @@ describe("KeyEditView", () => {
|
|||
const { getByText } = renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => {}}
|
||||
onSubmit={async () => {}}
|
||||
onCancel={() => { }}
|
||||
onSubmit={async () => { }}
|
||||
accessToken={""}
|
||||
userID={""}
|
||||
userRole={""}
|
||||
|
|
@ -111,8 +148,8 @@ describe("KeyEditView", () => {
|
|||
const { getByText } = renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => {}}
|
||||
onSubmit={async () => {}}
|
||||
onCancel={() => { }}
|
||||
onSubmit={async () => { }}
|
||||
accessToken={""}
|
||||
userID={""}
|
||||
userRole={""}
|
||||
|
|
@ -129,8 +166,8 @@ describe("KeyEditView", () => {
|
|||
const { getByLabelText } = renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => {}}
|
||||
onSubmit={async () => {}}
|
||||
onCancel={() => { }}
|
||||
onSubmit={async () => { }}
|
||||
accessToken={""}
|
||||
userID={""}
|
||||
userRole={""}
|
||||
|
|
@ -144,13 +181,17 @@ describe("KeyEditView", () => {
|
|||
});
|
||||
});
|
||||
|
||||
beforeEach(() => {
|
||||
vi.clearAllMocks();
|
||||
});
|
||||
|
||||
it("should call onCancel when cancel button is clicked", async () => {
|
||||
const onCancelMock = vi.fn();
|
||||
const { getByText } = renderWithProviders(
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={onCancelMock}
|
||||
onSubmit={async () => {}}
|
||||
onSubmit={async () => { }}
|
||||
accessToken={""}
|
||||
userID={""}
|
||||
userRole={""}
|
||||
|
|
@ -159,12 +200,272 @@ describe("KeyEditView", () => {
|
|||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(getByText("Cancel")).toBeInTheDocument();
|
||||
expect(screen.getByRole("button", { name: /cancel/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const cancelButton = getByText("Cancel");
|
||||
fireEvent.click(cancelButton);
|
||||
const cancelButton = screen.getByRole("button", { name: /cancel/i });
|
||||
await userEvent.click(cancelButton);
|
||||
|
||||
expect(onCancelMock).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
it("should display key alias input field", async () => {
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => { }}
|
||||
onSubmit={async () => { }}
|
||||
accessToken={""}
|
||||
userID={""}
|
||||
userRole={""}
|
||||
premiumUser={false}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByLabelText("Key Alias")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should display models select field", async () => {
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => { }}
|
||||
onSubmit={async () => { }}
|
||||
accessToken={""}
|
||||
userID={""}
|
||||
userRole={""}
|
||||
premiumUser={false}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Models")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should display max budget input field", async () => {
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => { }}
|
||||
onSubmit={async () => { }}
|
||||
accessToken={""}
|
||||
userID={""}
|
||||
userRole={""}
|
||||
premiumUser={false}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByLabelText("Max Budget (USD)")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should display allowed routes input field", async () => {
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => { }}
|
||||
onSubmit={async () => { }}
|
||||
accessToken={""}
|
||||
userID={""}
|
||||
userRole={""}
|
||||
premiumUser={false}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByLabelText(/allowed routes/i)).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should call onSubmit with form values when form is submitted", async () => {
|
||||
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => { }}
|
||||
onSubmit={onSubmitMock}
|
||||
accessToken={"test-token"}
|
||||
userID={"test-user"}
|
||||
userRole={"admin"}
|
||||
premiumUser={false}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByRole("button", { name: /save changes/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const submitButton = screen.getByRole("button", { name: /save changes/i });
|
||||
await userEvent.click(submitButton);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(onSubmitMock).toHaveBeenCalled();
|
||||
});
|
||||
});
|
||||
|
||||
it("should disable models field when management routes are selected", async () => {
|
||||
const keyDataWithManagementRoutes = {
|
||||
...MOCK_KEY_DATA,
|
||||
allowed_routes: ["management_routes"],
|
||||
};
|
||||
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={keyDataWithManagementRoutes}
|
||||
onCancel={() => { }}
|
||||
onSubmit={async () => { }}
|
||||
accessToken={""}
|
||||
userID={""}
|
||||
userRole={""}
|
||||
premiumUser={false}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Models field is disabled for this key type")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should disable models field when info routes are selected", async () => {
|
||||
const keyDataWithInfoRoutes = {
|
||||
...MOCK_KEY_DATA,
|
||||
allowed_routes: ["info_routes"],
|
||||
};
|
||||
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={keyDataWithInfoRoutes}
|
||||
onCancel={() => { }}
|
||||
onSubmit={async () => { }}
|
||||
accessToken={""}
|
||||
userID={""}
|
||||
userRole={""}
|
||||
premiumUser={false}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Models field is disabled for this key type")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should disable guardrails selector when user is not premium", async () => {
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => { }}
|
||||
onSubmit={async () => { }}
|
||||
accessToken={"test-token"}
|
||||
userID={""}
|
||||
userRole={""}
|
||||
premiumUser={false}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Guardrails")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should parse comma-separated allowed routes on submit", async () => {
|
||||
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => { }}
|
||||
onSubmit={onSubmitMock}
|
||||
accessToken={"test-token"}
|
||||
userID={"test-user"}
|
||||
userRole={"admin"}
|
||||
premiumUser={false}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByLabelText(/allowed routes/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const allowedRoutesInput = screen.getByLabelText(/allowed routes/i);
|
||||
await userEvent.clear(allowedRoutesInput);
|
||||
await userEvent.type(allowedRoutesInput, "route1, route2, route3");
|
||||
|
||||
const submitButton = screen.getByRole("button", { name: /save changes/i });
|
||||
await userEvent.click(submitButton);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(onSubmitMock).toHaveBeenCalled();
|
||||
const callArgs = onSubmitMock.mock.calls[0][0];
|
||||
expect(Array.isArray(callArgs.allowed_routes)).toBe(true);
|
||||
expect(callArgs.allowed_routes).toEqual(["route1", "route2", "route3"]);
|
||||
});
|
||||
});
|
||||
|
||||
it("should handle empty allowed routes string on submit", async () => {
|
||||
const onSubmitMock = vi.fn().mockResolvedValue(undefined);
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => { }}
|
||||
onSubmit={onSubmitMock}
|
||||
accessToken={"test-token"}
|
||||
userID={"test-user"}
|
||||
userRole={"admin"}
|
||||
premiumUser={false}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByLabelText(/allowed routes/i)).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const allowedRoutesInput = screen.getByLabelText(/allowed routes/i);
|
||||
await userEvent.clear(allowedRoutesInput);
|
||||
|
||||
const submitButton = screen.getByRole("button", { name: /save changes/i });
|
||||
await userEvent.click(submitButton);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(onSubmitMock).toHaveBeenCalled();
|
||||
const callArgs = onSubmitMock.mock.calls[0][0];
|
||||
expect(callArgs.allowed_routes).toEqual([]);
|
||||
});
|
||||
});
|
||||
|
||||
|
||||
it("should disable cancel button during submission", async () => {
|
||||
const onSubmitMock = vi.fn(
|
||||
() =>
|
||||
new Promise<void>((resolve) => {
|
||||
setTimeout(resolve, 100);
|
||||
}),
|
||||
);
|
||||
|
||||
renderWithProviders(
|
||||
<KeyEditView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onCancel={() => { }}
|
||||
onSubmit={onSubmitMock}
|
||||
accessToken={"test-token"}
|
||||
userID={"test-user"}
|
||||
userRole={"admin"}
|
||||
premiumUser={false}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByRole("button", { name: /cancel/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const submitButton = screen.getByRole("button", { name: /save changes/i });
|
||||
await userEvent.click(submitButton);
|
||||
|
||||
await waitFor(() => {
|
||||
const cancelButton = screen.getByRole("button", { name: /cancel/i });
|
||||
expect(cancelButton).toBeDisabled();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -35,7 +35,6 @@ interface KeyEditViewProps {
|
|||
// Add this helper function
|
||||
const getAvailableModelsForKey = (keyData: KeyResponse, teams: any[] | null): string[] => {
|
||||
// If no teams data is available, return empty array
|
||||
console.log("getAvailableModelsForKey:", teams);
|
||||
if (!teams || !keyData.team_id) {
|
||||
return [];
|
||||
}
|
||||
|
|
@ -172,7 +171,9 @@ export function KeyEditView({
|
|||
: [],
|
||||
auto_rotate: keyData.auto_rotate || false,
|
||||
...(keyData.rotation_interval && { rotation_interval: keyData.rotation_interval }),
|
||||
allowed_routes: keyData.allowed_routes,
|
||||
allowed_routes: Array.isArray(keyData.allowed_routes) && keyData.allowed_routes.length > 0
|
||||
? keyData.allowed_routes.join(", ")
|
||||
: "",
|
||||
};
|
||||
|
||||
useEffect(() => {
|
||||
|
|
@ -197,7 +198,9 @@ export function KeyEditView({
|
|||
: [],
|
||||
auto_rotate: keyData.auto_rotate || false,
|
||||
...(keyData.rotation_interval && { rotation_interval: keyData.rotation_interval }),
|
||||
allowed_routes: keyData.allowed_routes,
|
||||
allowed_routes: Array.isArray(keyData.allowed_routes) && keyData.allowed_routes.length > 0
|
||||
? keyData.allowed_routes.join(", ")
|
||||
: "",
|
||||
});
|
||||
}, [keyData, form]);
|
||||
|
||||
|
|
@ -226,11 +229,24 @@ export function KeyEditView({
|
|||
fetchTags();
|
||||
}, [accessToken]);
|
||||
|
||||
console.log("premiumUser:", premiumUser);
|
||||
|
||||
const handleSubmit = async (values: any) => {
|
||||
try {
|
||||
setIsKeySaving(true);
|
||||
|
||||
// Parse allowed_routes from comma-separated string to array
|
||||
if (typeof values.allowed_routes === "string") {
|
||||
const trimmedInput = values.allowed_routes.trim();
|
||||
if (trimmedInput === "") {
|
||||
values.allowed_routes = [];
|
||||
} else {
|
||||
values.allowed_routes = trimmedInput
|
||||
.split(",")
|
||||
.map((route: string) => route.trim())
|
||||
.filter((route: string) => route.length > 0);
|
||||
}
|
||||
}
|
||||
// If it's already an array (shouldn't happen, but handle it), keep as is
|
||||
|
||||
await onSubmit(values);
|
||||
} finally {
|
||||
setIsKeySaving(false);
|
||||
|
|
@ -251,7 +267,11 @@ export function KeyEditView({
|
|||
}
|
||||
>
|
||||
{({ getFieldValue, setFieldValue }) => {
|
||||
const allowedRoutes = getFieldValue("allowed_routes") || [];
|
||||
const allowedRoutesValue = getFieldValue("allowed_routes") || "";
|
||||
// Convert string to array for checking
|
||||
const allowedRoutes = typeof allowedRoutesValue === "string" && allowedRoutesValue.trim() !== ""
|
||||
? allowedRoutesValue.split(",").map((r: string) => r.trim()).filter((r: string) => r.length > 0)
|
||||
: [];
|
||||
const isDisabled = allowedRoutes.includes("management_routes") || allowedRoutes.includes("info_routes");
|
||||
const models = getFieldValue("models") || [];
|
||||
|
||||
|
|
@ -290,7 +310,11 @@ export function KeyEditView({
|
|||
shouldUpdate={(prevValues, currentValues) => prevValues.allowed_routes !== currentValues.allowed_routes}
|
||||
>
|
||||
{({ getFieldValue, setFieldValue }) => {
|
||||
const allowedRoutes = getFieldValue("allowed_routes");
|
||||
const allowedRoutesValue = getFieldValue("allowed_routes") || "";
|
||||
// Convert string to array for getKeyTypeFromRoutes
|
||||
const allowedRoutes = typeof allowedRoutesValue === "string" && allowedRoutesValue.trim() !== ""
|
||||
? allowedRoutesValue.split(",").map((r: string) => r.trim()).filter((r: string) => r.length > 0)
|
||||
: [];
|
||||
const keyTypeValue = getKeyTypeFromRoutes(allowedRoutes);
|
||||
|
||||
return (
|
||||
|
|
@ -302,13 +326,13 @@ export function KeyEditView({
|
|||
onChange={(value) => {
|
||||
switch (value) {
|
||||
case "default":
|
||||
setFieldValue("allowed_routes", []);
|
||||
setFieldValue("allowed_routes", "");
|
||||
break;
|
||||
case "llm_api":
|
||||
setFieldValue("allowed_routes", ["llm_api_routes"]);
|
||||
setFieldValue("allowed_routes", "llm_api_routes");
|
||||
break;
|
||||
case "management":
|
||||
setFieldValue("allowed_routes", ["management_routes"]);
|
||||
setFieldValue("allowed_routes", "management_routes");
|
||||
setFieldValue("models", []);
|
||||
break;
|
||||
}
|
||||
|
|
@ -344,6 +368,22 @@ export function KeyEditView({
|
|||
</Form.Item>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item
|
||||
label={
|
||||
<span>
|
||||
Allowed Routes{" "}
|
||||
<Tooltip title="List of allowed routes for the key (comma-separated). Can be specific routes (e.g., '/chat/completions') or route patterns (e.g., 'llm_api_routes', 'management_routes', '/keys/*'). Leave empty to allow all routes.">
|
||||
<InfoCircleOutlined style={{ marginLeft: "4px" }} />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="allowed_routes"
|
||||
>
|
||||
<Input
|
||||
placeholder="Enter allowed routes (comma-separated). Special values: llm_api_routes, management_routes. Examples: llm_api_routes, /chat/completions, /keys/*. Leave empty to allow all routes"
|
||||
/>
|
||||
</Form.Item>
|
||||
|
||||
<Form.Item label="Max Budget (USD)" name="max_budget">
|
||||
<NumericalInput step={0.01} style={{ width: "100%" }} placeholder="Enter a numerical value" />
|
||||
</Form.Item>
|
||||
|
|
@ -473,7 +513,7 @@ export function KeyEditView({
|
|||
!premiumUser
|
||||
? "Premium feature - Upgrade to set allowed pass through routes by key"
|
||||
: Array.isArray(keyData.metadata?.allowed_passthrough_routes) &&
|
||||
keyData.metadata.allowed_passthrough_routes.length > 0
|
||||
keyData.metadata.allowed_passthrough_routes.length > 0
|
||||
? `Current: ${keyData.metadata.allowed_passthrough_routes.join(", ")}`
|
||||
: "Select or enter allowed pass through routes"
|
||||
}
|
||||
|
|
@ -590,11 +630,6 @@ export function KeyEditView({
|
|||
<Input />
|
||||
</Form.Item>
|
||||
|
||||
{/* Hidden form field for allowed_routes */}
|
||||
<Form.Item name="allowed_routes" hidden>
|
||||
<Input />
|
||||
</Form.Item>
|
||||
|
||||
{/* Hidden form field for disabled callbacks */}
|
||||
<Form.Item name="disabled_callbacks" hidden>
|
||||
<Input />
|
||||
|
|
|
|||
|
|
@ -1,6 +1,7 @@
|
|||
import useTeams from "@/app/(dashboard)/hooks/useTeams";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import useTeams from "@/app/(dashboard)/hooks/useTeams";
|
||||
import { render, screen, waitFor } from "@testing-library/react";
|
||||
import userEvent from "@testing-library/user-event";
|
||||
import { beforeEach, describe, expect, it, vi } from "vitest";
|
||||
import { KeyResponse, Team } from "../key_team_helpers/key_list";
|
||||
import KeyInfoView from "./key_info_view";
|
||||
|
|
@ -13,6 +14,21 @@ vi.mock("@/app/(dashboard)/hooks/useAuthorized", () => ({
|
|||
default: vi.fn(),
|
||||
}));
|
||||
|
||||
vi.mock("../networking", () => ({
|
||||
keyDeleteCall: vi.fn().mockResolvedValue({}),
|
||||
keyUpdateCall: vi.fn().mockResolvedValue({}),
|
||||
getPolicyInfoWithGuardrails: vi.fn().mockResolvedValue({
|
||||
resolved_guardrails: ["guardrail-1", "guardrail-2"],
|
||||
}),
|
||||
}));
|
||||
|
||||
vi.mock("@/utils/dataUtils", () => ({
|
||||
copyToClipboard: vi.fn().mockResolvedValue(true),
|
||||
formatNumberWithCommas: vi.fn((value: number, decimals?: number) => {
|
||||
return value.toFixed(decimals ?? 2);
|
||||
}),
|
||||
}));
|
||||
|
||||
describe("KeyInfoView", () => {
|
||||
beforeEach(() => {
|
||||
vi.mocked(useTeams).mockReturnValue({
|
||||
|
|
@ -105,34 +121,34 @@ describe("KeyInfoView", () => {
|
|||
it("should render tags", async () => {
|
||||
vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock);
|
||||
|
||||
const { getByText } = render(
|
||||
render(
|
||||
<KeyInfoView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onClose={() => {}}
|
||||
onClose={() => { }}
|
||||
keyId={"test-key-id"}
|
||||
onKeyDataUpdate={() => {}}
|
||||
onKeyDataUpdate={() => { }}
|
||||
teams={[]}
|
||||
/>,
|
||||
);
|
||||
await waitFor(() => {
|
||||
expect(getByText("test-tag")).toBeInTheDocument();
|
||||
expect(screen.getByText("test-tag")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should not render tags in metadata textarea", async () => {
|
||||
vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock);
|
||||
|
||||
const { container, getByText } = render(
|
||||
const { container } = render(
|
||||
<KeyInfoView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onClose={() => {}}
|
||||
onClose={() => { }}
|
||||
keyId={"test-key-id"}
|
||||
onKeyDataUpdate={() => {}}
|
||||
onKeyDataUpdate={() => { }}
|
||||
teams={[]}
|
||||
/>,
|
||||
);
|
||||
await waitFor(() => {
|
||||
expect(getByText("Metadata")).toBeInTheDocument();
|
||||
expect(screen.getByText("Metadata")).toBeInTheDocument();
|
||||
const metadataBlock = container.querySelector("pre");
|
||||
expect(metadataBlock).toBeInTheDocument();
|
||||
expect(metadataBlock?.textContent?.trim()).toBe("{}");
|
||||
|
|
@ -153,7 +169,7 @@ describe("KeyInfoView", () => {
|
|||
|
||||
const keyData = { ...MOCK_KEY_DATA, user_id: "other-user-id" };
|
||||
render(
|
||||
<KeyInfoView keyData={keyData} onClose={() => {}} keyId={"test-key-id"} onKeyDataUpdate={() => {}} teams={[]} />,
|
||||
<KeyInfoView keyData={keyData} onClose={() => { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
|
|
@ -182,6 +198,7 @@ describe("KeyInfoView", () => {
|
|||
role: "admin",
|
||||
},
|
||||
],
|
||||
spend: 0,
|
||||
};
|
||||
|
||||
vi.mocked(useTeams).mockReturnValue({
|
||||
|
|
@ -197,7 +214,7 @@ describe("KeyInfoView", () => {
|
|||
|
||||
const keyData = { ...MOCK_KEY_DATA, team_id: teamId, user_id: "other-user-id" };
|
||||
render(
|
||||
<KeyInfoView keyData={keyData} onClose={() => {}} keyId={"test-key-id"} onKeyDataUpdate={() => {}} teams={[]} />,
|
||||
<KeyInfoView keyData={keyData} onClose={() => { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
|
|
@ -221,7 +238,7 @@ describe("KeyInfoView", () => {
|
|||
const ownerUserId = "owner-user-id";
|
||||
const keyData = { ...MOCK_KEY_DATA, user_id: ownerUserId };
|
||||
render(
|
||||
<KeyInfoView keyData={keyData} onClose={() => {}} keyId={"test-key-id"} onKeyDataUpdate={() => {}} teams={[]} />,
|
||||
<KeyInfoView keyData={keyData} onClose={() => { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
|
|
@ -244,7 +261,7 @@ describe("KeyInfoView", () => {
|
|||
|
||||
const keyData = { ...MOCK_KEY_DATA, user_id: "owner-user-id" };
|
||||
render(
|
||||
<KeyInfoView keyData={keyData} onClose={() => {}} keyId={"test-key-id"} onKeyDataUpdate={() => {}} teams={[]} />,
|
||||
<KeyInfoView keyData={keyData} onClose={() => { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
|
|
@ -268,7 +285,7 @@ describe("KeyInfoView", () => {
|
|||
const ownerUserId = "internal-viewer-user-id";
|
||||
const keyData = { ...MOCK_KEY_DATA, user_id: ownerUserId };
|
||||
render(
|
||||
<KeyInfoView keyData={keyData} onClose={() => {}} keyId={"test-key-id"} onKeyDataUpdate={() => {}} teams={[]} />,
|
||||
<KeyInfoView keyData={keyData} onClose={() => { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
|
|
@ -296,6 +313,7 @@ describe("KeyInfoView", () => {
|
|||
role: "admin",
|
||||
},
|
||||
],
|
||||
spend: 0,
|
||||
};
|
||||
|
||||
vi.mocked(useTeams).mockReturnValue({
|
||||
|
|
@ -309,10 +327,9 @@ describe("KeyInfoView", () => {
|
|||
userRole: "user",
|
||||
});
|
||||
|
||||
// Key has a different team_id that doesn't match any team in teamsData
|
||||
const keyData = { ...MOCK_KEY_DATA, team_id: "non-matching-team-id", user_id: "other-user-id" };
|
||||
render(
|
||||
<KeyInfoView keyData={keyData} onClose={() => {}} keyId={"test-key-id"} onKeyDataUpdate={() => {}} teams={[]} />,
|
||||
<KeyInfoView keyData={keyData} onClose={() => { }} keyId={"test-key-id"} onKeyDataUpdate={() => { }} teams={[]} />,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
|
|
@ -320,4 +337,129 @@ describe("KeyInfoView", () => {
|
|||
expect(screen.queryByText("Delete Key")).not.toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should call onClose when back button is clicked", async () => {
|
||||
vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock);
|
||||
const onCloseMock = vi.fn();
|
||||
|
||||
render(
|
||||
<KeyInfoView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onClose={onCloseMock}
|
||||
keyId={"test-key-id"}
|
||||
onKeyDataUpdate={() => { }}
|
||||
teams={[]}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByRole("button", { name: /back to keys/i })).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const backButton = screen.getByRole("button", { name: /back to keys/i });
|
||||
await userEvent.click(backButton);
|
||||
|
||||
expect(onCloseMock).toHaveBeenCalledTimes(1);
|
||||
});
|
||||
|
||||
|
||||
it("should show edit button in settings tab when user has write access", async () => {
|
||||
vi.mocked(useAuthorized).mockReturnValue({
|
||||
...baseUseAuthorizedMock,
|
||||
userRole: "Admin",
|
||||
});
|
||||
|
||||
render(
|
||||
<KeyInfoView
|
||||
keyData={MOCK_KEY_DATA}
|
||||
onClose={() => { }}
|
||||
keyId={"test-key-id"}
|
||||
onKeyDataUpdate={() => { }}
|
||||
teams={[]}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
const settingsTab = screen.getByRole("tab", { name: /settings/i });
|
||||
expect(settingsTab).toBeInTheDocument();
|
||||
});
|
||||
|
||||
const settingsTab = screen.getByRole("tab", { name: /settings/i });
|
||||
await userEvent.click(settingsTab);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
|
||||
it("should display guardrails when present", async () => {
|
||||
vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock);
|
||||
|
||||
const keyDataWithGuardrails = {
|
||||
...MOCK_KEY_DATA,
|
||||
metadata: {
|
||||
...MOCK_KEY_DATA.metadata,
|
||||
guardrails: ["guardrail-1", "guardrail-2"],
|
||||
},
|
||||
};
|
||||
|
||||
render(
|
||||
<KeyInfoView
|
||||
keyData={keyDataWithGuardrails}
|
||||
onClose={() => { }}
|
||||
keyId={"test-key-id"}
|
||||
onKeyDataUpdate={() => { }}
|
||||
teams={[]}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Guardrails")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should display policies when present", async () => {
|
||||
vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock);
|
||||
|
||||
const keyDataWithPolicies = {
|
||||
...MOCK_KEY_DATA,
|
||||
metadata: {
|
||||
...MOCK_KEY_DATA.metadata,
|
||||
policies: ["policy-1"],
|
||||
},
|
||||
};
|
||||
|
||||
render(
|
||||
<KeyInfoView
|
||||
keyData={keyDataWithPolicies}
|
||||
onClose={() => { }}
|
||||
keyId={"test-key-id"}
|
||||
onKeyDataUpdate={() => { }}
|
||||
teams={[]}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Policies")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
|
||||
it("should display no key found message when keyData is undefined", async () => {
|
||||
vi.mocked(useAuthorized).mockReturnValue(baseUseAuthorizedMock);
|
||||
|
||||
render(
|
||||
<KeyInfoView
|
||||
keyData={undefined}
|
||||
onClose={() => { }}
|
||||
keyId={"test-key-id"}
|
||||
onKeyDataUpdate={() => { }}
|
||||
teams={[]}
|
||||
/>,
|
||||
);
|
||||
|
||||
await waitFor(() => {
|
||||
expect(screen.getByText("Key not found")).toBeInTheDocument();
|
||||
});
|
||||
});
|
||||
});
|
||||
|
|
|
|||
|
|
@ -1,9 +1,10 @@
|
|||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
import useTeams from "@/app/(dashboard)/hooks/useTeams";
|
||||
import { formatNumberWithCommas, copyToClipboard as utilCopyToClipboard } from "@/utils/dataUtils";
|
||||
import { mapEmptyStringToNull } from "@/utils/keyUpdateUtils";
|
||||
import { ArrowLeftIcon, RefreshIcon, TrashIcon } from "@heroicons/react/outline";
|
||||
import { Badge, Button, Card, Grid, Tab, TabGroup, TabList, TabPanel, TabPanels, Text, Title } from "@tremor/react";
|
||||
import { Button as AntdButton, Form, Tooltip } from "antd";
|
||||
import { Button as AntdButton, Form, Tag, Tooltip } from "antd";
|
||||
import { CheckIcon, CopyIcon } from "lucide-react";
|
||||
import { useEffect, useState } from "react";
|
||||
import { isProxyAdminRole, isUserTeamAdminForSingleTeam, rolesWithWriteAccess } from "../../utils/roles";
|
||||
|
|
@ -14,12 +15,11 @@ import { extractLoggingSettings, formatMetadataForDisplay, stripTagsFromMetadata
|
|||
import { KeyResponse } from "../key_team_helpers/key_list";
|
||||
import LoggingSettingsView from "../logging_settings_view";
|
||||
import NotificationManager from "../molecules/notifications_manager";
|
||||
import { keyDeleteCall, keyUpdateCall, getPolicyInfoWithGuardrails } from "../networking";
|
||||
import { getPolicyInfoWithGuardrails, keyDeleteCall, keyUpdateCall } from "../networking";
|
||||
import ObjectPermissionsView from "../object_permissions_view";
|
||||
import { RegenerateKeyModal } from "../organisms/regenerate_key_modal";
|
||||
import { parseErrorMessage } from "../shared/errorUtils";
|
||||
import { KeyEditView } from "./key_edit_view";
|
||||
import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized";
|
||||
|
||||
interface KeyInfoViewProps {
|
||||
keyId: string;
|
||||
|
|
@ -206,8 +206,8 @@ export default function KeyInfoView({
|
|||
...(formValues.logging_settings ? { logging: formValues.logging_settings } : {}),
|
||||
...(formValues.disabled_callbacks?.length > 0
|
||||
? {
|
||||
litellm_disabled_callbacks: mapDisplayToInternalNames(formValues.disabled_callbacks),
|
||||
}
|
||||
litellm_disabled_callbacks: mapDisplayToInternalNames(formValues.disabled_callbacks),
|
||||
}
|
||||
: {}),
|
||||
};
|
||||
} catch (error) {
|
||||
|
|
@ -225,8 +225,8 @@ export default function KeyInfoView({
|
|||
...(formValues.logging_settings ? { logging: formValues.logging_settings } : {}),
|
||||
...(formValues.disabled_callbacks?.length > 0
|
||||
? {
|
||||
litellm_disabled_callbacks: mapDisplayToInternalNames(formValues.disabled_callbacks),
|
||||
}
|
||||
litellm_disabled_callbacks: mapDisplayToInternalNames(formValues.disabled_callbacks),
|
||||
}
|
||||
: {}),
|
||||
};
|
||||
}
|
||||
|
|
@ -334,7 +334,6 @@ export default function KeyInfoView({
|
|||
});
|
||||
return `${dateStr} at ${timeStr}`;
|
||||
};
|
||||
console.log("userRole", userRole);
|
||||
|
||||
const canModifyKey =
|
||||
isProxyAdminRole(userRole || "") ||
|
||||
|
|
@ -364,11 +363,10 @@ export default function KeyInfoView({
|
|||
size="small"
|
||||
icon={copiedStates["key-id"] ? <CheckIcon size={12} /> : <CopyIcon size={12} />}
|
||||
onClick={() => copyToClipboard(currentKeyData.token_id || currentKeyData.token, "key-id")}
|
||||
className={`ml-2 transition-all duration-200${
|
||||
copiedStates["key-id"]
|
||||
? "text-green-600 bg-green-50 border-green-200"
|
||||
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
|
||||
}`}
|
||||
className={`ml-2 transition-all duration-200${copiedStates["key-id"]
|
||||
? "text-green-600 bg-green-50 border-green-200"
|
||||
: "text-gray-500 hover:text-gray-700 hover:bg-gray-100"
|
||||
}`}
|
||||
/>
|
||||
</div>
|
||||
|
||||
|
|
@ -691,10 +689,10 @@ export default function KeyInfoView({
|
|||
<div className="flex flex-wrap gap-2 mt-1">
|
||||
{Array.isArray(currentKeyData.metadata?.tags) && currentKeyData.metadata.tags.length > 0
|
||||
? currentKeyData.metadata.tags.map((tag, index) => (
|
||||
<span key={index} className="px-2 mr-2 py-1 bg-blue-100 rounded text-xs">
|
||||
{tag}
|
||||
</span>
|
||||
))
|
||||
<span key={index} className="px-2 mr-2 py-1 bg-blue-100 rounded text-xs">
|
||||
{tag}
|
||||
</span>
|
||||
))
|
||||
: "No tags specified"}
|
||||
</div>
|
||||
</div>
|
||||
|
|
@ -704,24 +702,39 @@ export default function KeyInfoView({
|
|||
<Text>
|
||||
{Array.isArray(currentKeyData.metadata?.prompts) && currentKeyData.metadata.prompts.length > 0
|
||||
? currentKeyData.metadata.prompts.map((prompt, index) => (
|
||||
<span key={index} className="px-2 mr-2 py-1 bg-blue-100 rounded text-xs">
|
||||
{prompt}
|
||||
</span>
|
||||
))
|
||||
<span key={index} className="px-2 mr-2 py-1 bg-blue-100 rounded text-xs">
|
||||
{prompt}
|
||||
</span>
|
||||
))
|
||||
: "No prompts specified"}
|
||||
</Text>
|
||||
</div>
|
||||
|
||||
<div>
|
||||
<Text className="font-medium">Allowed Routes</Text>
|
||||
<div className="flex flex-wrap gap-2 mt-1">
|
||||
{Array.isArray(currentKeyData.allowed_routes) && currentKeyData.allowed_routes.length > 0 ? (
|
||||
currentKeyData.allowed_routes.map((route, index) => (
|
||||
<span key={index} className="px-2 py-1 bg-blue-100 rounded text-xs">
|
||||
{route}
|
||||
</span>
|
||||
))
|
||||
) : (
|
||||
<Tag color="green">All routes allowed</Tag>
|
||||
)}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div>
|
||||
<Text className="font-medium">Allowed Pass Through Routes</Text>
|
||||
<Text>
|
||||
{Array.isArray(currentKeyData.metadata?.allowed_passthrough_routes) &&
|
||||
currentKeyData.metadata.allowed_passthrough_routes.length > 0
|
||||
currentKeyData.metadata.allowed_passthrough_routes.length > 0
|
||||
? currentKeyData.metadata.allowed_passthrough_routes.map((route, index) => (
|
||||
<span key={index} className="px-2 mr-2 py-1 bg-blue-100 rounded text-xs">
|
||||
{route}
|
||||
</span>
|
||||
))
|
||||
<span key={index} className="px-2 mr-2 py-1 bg-blue-100 rounded text-xs">
|
||||
{route}
|
||||
</span>
|
||||
))
|
||||
: "No pass through routes specified"}
|
||||
</Text>
|
||||
</div>
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue