Merge branch 'litellm_oss_staging_02_04_2026' into litellm_fix_langfuse_otel_trace

This commit is contained in:
Krish Dholakia 2026-02-03 22:40:10 -08:00 • committed by GitHub
commit 607c9f02d7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
24 changed files with 1615 additions and 305 deletions

View file

@ -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/

View file

@ -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

View 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.
:::

View file

@ -469,6 +469,7 @@ const sidebars = {
label: "/a2a - A2A Agent Gateway",
items: [
"a2a",
"a2a_invoking_agents",
"a2a_cost_tracking",
"a2a_agent_permissions"
],

View file

@ -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)

View file

@ -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:

View file

@ -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",

View file

@ -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]

View file

@ -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,
)

View file

@ -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

View 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)

View 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

View file

@ -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}

View file

@ -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
View file

@ -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"

View file

@ -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 = "*"

View 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())

View file

@ -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"

View 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",
)

View 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)")

View file

@ -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();
});
});
});

View file

@ -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 />

View file

@ -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();
});
});
});

View file

@ -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>