mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
Merge pull request #22888 from BerriAI/litellm_a2a-custom-headers
[Feat] Add a2a custom headers
This commit is contained in:
commit
8b0375f99c
14 changed files with 1183 additions and 13 deletions
252
docs/my-website/docs/a2a_agent_headers.md
Normal file
252
docs/my-website/docs/a2a_agent_headers.md
Normal file
|
|
@ -0,0 +1,252 @@
|
|||
import Tabs from '@theme/Tabs';
|
||||
import TabItem from '@theme/TabItem';
|
||||
|
||||
# A2A Agent Authentication Headers
|
||||
|
||||
Forward authentication credentials (Bearer tokens, API keys, etc.) from clients to backend A2A agents.
|
||||
|
||||
## Overview
|
||||
|
||||
When LiteLLM proxies a request to a backend A2A agent, the agent may require its own authentication headers. There are three ways to supply them:
|
||||
|
||||
| Method | Who configures | How it works |
|
||||
|---|---|---|
|
||||
| **Static headers** | Admin (UI / API) | Always sent, regardless of client request |
|
||||
| **Forward client headers** | Admin (UI / API) | Header names to extract from client request and forward |
|
||||
| **Convention-based** | Client (no admin config) | Client sends `x-a2a-{agent_name}-{header}` — automatically routed |
|
||||
|
||||
All three methods can be combined. **Static headers always win** on key conflicts.
|
||||
|
||||
---
|
||||
|
||||
## Method 1 — Static Headers
|
||||
|
||||
Admin-configured headers that are always sent to the backend agent. Use this for server-to-server tokens or internal credentials that clients should never see or override.
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="ui" label="UI">
|
||||
|
||||
1. Go to **Agents** in the LiteLLM dashboard.
|
||||
2. Create or edit an agent.
|
||||
3. Open the **Authentication Headers** panel.
|
||||
4. Under **Static Headers**, click **Add Static Header** and fill in the header name and value.
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="api" label="REST API">
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:4000/v1/agents \
|
||||
-H "Authorization: Bearer sk-admin" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"agent_name": "my-agent",
|
||||
"agent_card_params": { ... },
|
||||
"static_headers": {
|
||||
"Authorization": "Bearer internal-server-token",
|
||||
"X-Internal-Service": "litellm-proxy"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
To update an existing agent:
|
||||
|
||||
```bash
|
||||
curl -X PATCH http://localhost:4000/v1/agents/{agent_id} \
|
||||
-H "Authorization: Bearer sk-admin" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"static_headers": {
|
||||
"Authorization": "Bearer new-token"
|
||||
}
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
**Client call — no special headers needed:**
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:4000/a2a/my-agent \
|
||||
-H "Authorization: Bearer sk-client-key" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"jsonrpc": "2.0", "id": "1", "method": "message/send",
|
||||
"params": { "message": { "role": "user", "parts": [{"kind": "text", "text": "Hello"}], "messageId": "msg-1" } }
|
||||
}'
|
||||
```
|
||||
|
||||
The backend agent receives `Authorization: Bearer internal-server-token` without the client ever knowing the value.
|
||||
|
||||
---
|
||||
|
||||
## Method 2 — Forward Client Headers
|
||||
|
||||
Admin specifies a list of header **names**. When the client sends a request that includes those headers, LiteLLM extracts their values and forwards them to the backend agent. The client controls the values; the admin controls which headers are eligible to be forwarded.
|
||||
|
||||
<Tabs>
|
||||
<TabItem value="ui" label="UI">
|
||||
|
||||
1. Go to **Agents** in the LiteLLM dashboard.
|
||||
2. Create or edit an agent.
|
||||
3. Open the **Authentication Headers** panel.
|
||||
4. Under **Forward Client Headers**, type header names and press **Enter** (e.g. `x-api-key`, `Authorization`).
|
||||
|
||||
</TabItem>
|
||||
<TabItem value="api" label="REST API">
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:4000/v1/agents \
|
||||
-H "Authorization: Bearer sk-admin" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"agent_name": "my-agent",
|
||||
"agent_card_params": { ... },
|
||||
"extra_headers": ["x-api-key", "x-user-token"]
|
||||
}'
|
||||
```
|
||||
|
||||
</TabItem>
|
||||
</Tabs>
|
||||
|
||||
**Client call — include the forwarded headers:**
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:4000/a2a/my-agent \
|
||||
-H "Authorization: Bearer sk-client-key" \
|
||||
-H "x-api-key: user-secret-value" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{ ... }'
|
||||
```
|
||||
|
||||
The backend agent receives `x-api-key: user-secret-value`.
|
||||
|
||||
:::note
|
||||
Header name matching is **case-insensitive**. If the client sends `X-API-Key` and `extra_headers` lists `x-api-key`, they match.
|
||||
:::
|
||||
|
||||
---
|
||||
|
||||
## Method 3 — Convention-Based Forwarding
|
||||
|
||||
Clients can forward headers to a specific agent without any admin pre-configuration by using the naming convention:
|
||||
|
||||
```
|
||||
x-a2a-{agent_name_or_id}-{header_name}: value
|
||||
```
|
||||
|
||||
LiteLLM parses these headers automatically and routes them to the matching agent only.
|
||||
|
||||
**Examples:**
|
||||
|
||||
| Client header sent | Agent name/ID | Forwarded as |
|
||||
|---|---|---|
|
||||
| `x-a2a-my-agent-authorization: Bearer tok` | `my-agent` | `authorization: Bearer tok` |
|
||||
| `x-a2a-my-agent-x-api-key: secret` | `my-agent` | `x-api-key: secret` |
|
||||
| `x-a2a-abc123-authorization: Bearer tok` | agent ID `abc123` | `authorization: Bearer tok` |
|
||||
|
||||
```bash
|
||||
curl -X POST http://localhost:4000/a2a/my-agent \
|
||||
-H "Authorization: Bearer sk-client-key" \
|
||||
-H "x-a2a-my-agent-authorization: Bearer agent-specific-token" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{ ... }'
|
||||
```
|
||||
|
||||
The `x-a2a-other-agent-authorization` header sent in the same request is **not** forwarded to `my-agent` — it is silently ignored.
|
||||
|
||||
:::tip Matches both agent name and agent ID
|
||||
Both the human-readable name (e.g. `my-agent`) and the UUID (e.g. `abc123-...`) are valid. Use whichever is convenient for the client.
|
||||
:::
|
||||
|
||||
---
|
||||
|
||||
## Merge Precedence
|
||||
|
||||
When multiple methods supply the same header name, **static headers win**:
|
||||
|
||||
```
|
||||
dynamic (forwarded/convention) → merged ← static (overlays, wins)
|
||||
```
|
||||
|
||||
Example:
|
||||
|
||||
| Source | `Authorization` value |
|
||||
|---|---|
|
||||
| Client sends (via `extra_headers` or convention) | `Bearer client-token` |
|
||||
| Admin-configured `static_headers` | `Bearer server-token` |
|
||||
| **What the backend agent receives** | **`Bearer server-token`** |
|
||||
|
||||
This ensures admin-controlled credentials cannot be overridden by client requests.
|
||||
|
||||
---
|
||||
|
||||
## Combining All Three Methods
|
||||
|
||||
```bash
|
||||
# Register agent with static + forwarded headers
|
||||
curl -X POST http://localhost:4000/v1/agents \
|
||||
-H "Authorization: Bearer sk-admin" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{
|
||||
"agent_name": "my-agent",
|
||||
"agent_card_params": { ... },
|
||||
"static_headers": {
|
||||
"X-Internal-Token": "secret123"
|
||||
},
|
||||
"extra_headers": ["x-user-id"]
|
||||
}'
|
||||
|
||||
# Client call using all three mechanisms
|
||||
curl -X POST http://localhost:4000/a2a/my-agent \
|
||||
-H "Authorization: Bearer sk-client-key" \
|
||||
-H "x-user-id: user-42" \
|
||||
-H "x-a2a-my-agent-x-request-id: req-abc" \
|
||||
-H "Content-Type: application/json" \
|
||||
-d '{ ... }'
|
||||
```
|
||||
|
||||
The backend agent receives:
|
||||
|
||||
```
|
||||
X-Internal-Token: secret123 ← static header (always)
|
||||
x-user-id: user-42 ← forwarded (in extra_headers)
|
||||
x-request-id: req-abc ← convention-based (x-a2a-my-agent-*)
|
||||
X-LiteLLM-Trace-Id: <uuid> ← LiteLLM internal
|
||||
X-LiteLLM-Agent-Id: <agent-id> ← LiteLLM internal
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## Header Isolation
|
||||
|
||||
Each agent invocation uses an isolated HTTP connection. Headers configured for agent A are **never** sent to agent B, even if both agents are running and receiving requests simultaneously.
|
||||
|
||||
---
|
||||
|
||||
## API Reference
|
||||
|
||||
### `POST /v1/agents` / `PATCH /v1/agents/{agent_id}`
|
||||
|
||||
| Field | Type | Description |
|
||||
|---|---|---|
|
||||
| `static_headers` | `object` | `{"Header-Name": "value"}` — always forwarded |
|
||||
| `extra_headers` | `string[]` | Header names to extract from client request and forward |
|
||||
|
||||
### Agent Response
|
||||
|
||||
Both fields are returned in `GET /v1/agents` and `GET /v1/agents/{agent_id}`:
|
||||
|
||||
```json
|
||||
{
|
||||
"agent_id": "...",
|
||||
"agent_name": "my-agent",
|
||||
"static_headers": { "X-Internal-Token": "secret123" },
|
||||
"extra_headers": ["x-user-id"],
|
||||
...
|
||||
}
|
||||
```
|
||||
|
||||
:::caution
|
||||
`static_headers` values are stored in the database and returned by the API. Treat them as you would any credential — do not store sensitive long-lived tokens here if your API is publicly accessible. Consider using short-lived tokens or environment-injected secrets instead.
|
||||
:::
|
||||
|
|
@ -539,6 +539,7 @@ const sidebars = {
|
|||
items: [
|
||||
"a2a",
|
||||
"a2a_invoking_agents",
|
||||
"a2a_agent_headers",
|
||||
"a2a_cost_tracking",
|
||||
"a2a_agent_permissions"
|
||||
],
|
||||
|
|
|
|||
|
|
@ -0,0 +1,5 @@
|
|||
-- Add static_headers and extra_headers to LiteLLM_AgentsTable
|
||||
|
||||
ALTER TABLE "LiteLLM_AgentsTable"
|
||||
ADD COLUMN IF NOT EXISTS "static_headers" JSONB DEFAULT '{}',
|
||||
ADD COLUMN IF NOT EXISTS "extra_headers" TEXT[] DEFAULT ARRAY[]::TEXT[];
|
||||
|
|
@ -9,6 +9,7 @@ import datetime
|
|||
import uuid
|
||||
from typing import TYPE_CHECKING, Any, AsyncIterator, Coroutine, Dict, Optional, Union
|
||||
|
||||
import httpx
|
||||
import litellm
|
||||
from litellm._logging import verbose_logger, verbose_proxy_logger
|
||||
from litellm.a2a_protocol.streaming_iterator import A2AStreamingIterator
|
||||
|
|
@ -212,6 +213,7 @@ async def asend_message(
|
|||
api_base: Optional[str] = None,
|
||||
litellm_params: Optional[Dict[str, Any]] = None,
|
||||
agent_id: Optional[str] = None,
|
||||
agent_extra_headers: Optional[Dict[str, str]] = None,
|
||||
**kwargs: Any,
|
||||
) -> LiteLLMSendMessageResponse:
|
||||
"""
|
||||
|
|
@ -293,9 +295,12 @@ async def asend_message(
|
|||
"Either a2a_client or api_base is required for standard A2A flow"
|
||||
)
|
||||
trace_id = trace_id or str(uuid.uuid4())
|
||||
extra_headers = {"X-LiteLLM-Trace-Id": trace_id}
|
||||
extra_headers: Dict[str, str] = {"X-LiteLLM-Trace-Id": trace_id}
|
||||
if agent_id:
|
||||
extra_headers["X-LiteLLM-Agent-Id"] = agent_id
|
||||
# Overlay agent-level headers (agent headers take precedence over LiteLLM internal ones)
|
||||
if agent_extra_headers:
|
||||
extra_headers.update(agent_extra_headers)
|
||||
a2a_client = await create_a2a_client(
|
||||
base_url=api_base, extra_headers=extra_headers
|
||||
)
|
||||
|
|
@ -442,6 +447,7 @@ async def asend_message_streaming(
|
|||
agent_id: Optional[str] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
proxy_server_request: Optional[Dict[str, Any]] = None,
|
||||
agent_extra_headers: Optional[Dict[str, str]] = None,
|
||||
) -> AsyncIterator[Any]:
|
||||
"""
|
||||
Async: Send a streaming message to an A2A agent.
|
||||
|
|
@ -523,7 +529,17 @@ async def asend_message_streaming(
|
|||
raise ValueError(
|
||||
"Either a2a_client or api_base is required for standard A2A flow"
|
||||
)
|
||||
a2a_client = await create_a2a_client(base_url=api_base)
|
||||
# Mirror the non-streaming path: always include trace and agent-id headers
|
||||
streaming_extra_headers: Dict[str, str] = {
|
||||
"X-LiteLLM-Trace-Id": str(request.id),
|
||||
}
|
||||
if agent_id:
|
||||
streaming_extra_headers["X-LiteLLM-Agent-Id"] = agent_id
|
||||
if agent_extra_headers:
|
||||
streaming_extra_headers.update(agent_extra_headers)
|
||||
a2a_client = await create_a2a_client(
|
||||
base_url=api_base, extra_headers=streaming_extra_headers
|
||||
)
|
||||
|
||||
# Type assertion: a2a_client is guaranteed to be non-None here
|
||||
assert a2a_client is not None
|
||||
|
|
@ -637,17 +653,17 @@ async def create_a2a_client(
|
|||
|
||||
verbose_logger.info(f"Creating A2A client for {base_url}")
|
||||
|
||||
# Use LiteLLM's cached httpx client
|
||||
http_handler = get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.A2A,
|
||||
params={"timeout": timeout},
|
||||
# Always create a fresh httpx client per A2A call so that per-agent auth
|
||||
# headers (extra_headers) are never shared across agents or requests.
|
||||
# Mutating a cached shared client would cause headers from one agent to
|
||||
# bleed into requests made to a different agent.
|
||||
httpx_client = httpx.AsyncClient(
|
||||
timeout=httpx.Timeout(timeout),
|
||||
headers=extra_headers or {},
|
||||
)
|
||||
httpx_client = http_handler.client
|
||||
|
||||
if extra_headers:
|
||||
httpx_client.headers.update(extra_headers)
|
||||
verbose_proxy_logger.debug(
|
||||
f"A2A client created with extra_headers={extra_headers}"
|
||||
f"A2A client created with extra_headers={list(extra_headers.keys())}"
|
||||
)
|
||||
|
||||
# Resolve agent card
|
||||
|
|
|
|||
|
|
@ -6,13 +6,14 @@ The A2A SDK can point to LiteLLM's URL and invoke agents registered with LiteLLM
|
|||
"""
|
||||
|
||||
import json
|
||||
from typing import Any, Optional
|
||||
from typing import Any, Dict, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request, Response
|
||||
from fastapi.responses import JSONResponse, StreamingResponse
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.utils import all_litellm_params
|
||||
|
||||
|
|
@ -55,6 +56,7 @@ async def _handle_stream_message(
|
|||
metadata: Optional[dict] = None,
|
||||
proxy_server_request: Optional[dict] = None,
|
||||
*,
|
||||
agent_extra_headers: Optional[Dict[str, str]] = None,
|
||||
user_api_key_dict: Optional[UserAPIKeyAuth] = None,
|
||||
request_data: Optional[dict] = None,
|
||||
proxy_logging_obj: Optional[Any] = None,
|
||||
|
|
@ -105,6 +107,7 @@ async def _handle_stream_message(
|
|||
agent_id=agent_id,
|
||||
metadata=metadata,
|
||||
proxy_server_request=proxy_server_request,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
|
||||
if (
|
||||
|
|
@ -385,6 +388,36 @@ async def invoke_agent_a2a(
|
|||
version=version,
|
||||
)
|
||||
|
||||
# Build merged headers for the backend agent
|
||||
static_headers: Dict[str, str] = dict(agent.static_headers or {})
|
||||
|
||||
raw_headers = dict(request.headers)
|
||||
normalized = {k.lower(): v for k, v in raw_headers.items()}
|
||||
|
||||
dynamic_headers: Dict[str, str] = {}
|
||||
|
||||
# 1. Admin-configured extra_headers: forward named headers from client request
|
||||
if agent.extra_headers:
|
||||
for header_name in agent.extra_headers:
|
||||
val = normalized.get(header_name.lower())
|
||||
if val is not None:
|
||||
dynamic_headers[header_name] = val
|
||||
|
||||
# 2. Convention-based forwarding: x-a2a-{agent_id_or_name}-{header_name}
|
||||
# Matches both agent_id (UUID) and agent_name (alias), case-insensitive.
|
||||
for alias in (agent.agent_id.lower(), agent.agent_name.lower()):
|
||||
prefix = f"x-a2a-{alias}-"
|
||||
for key, val in normalized.items():
|
||||
if key.startswith(prefix):
|
||||
header_name = key[len(prefix) :]
|
||||
if header_name:
|
||||
dynamic_headers[header_name] = val
|
||||
|
||||
agent_extra_headers = merge_agent_headers(
|
||||
dynamic_headers=dynamic_headers or None,
|
||||
static_headers=static_headers or None,
|
||||
)
|
||||
|
||||
# Route through SDK functions
|
||||
if method == "message/send":
|
||||
from a2a.types import MessageSendParams, SendMessageRequest
|
||||
|
|
@ -401,6 +434,7 @@ async def invoke_agent_a2a(
|
|||
metadata=data.get("metadata", {}),
|
||||
proxy_server_request=data.get("proxy_server_request"),
|
||||
litellm_logging_obj=logging_obj,
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
)
|
||||
|
||||
response = await proxy_logging_obj.post_call_success_hook(
|
||||
|
|
@ -425,6 +459,7 @@ async def invoke_agent_a2a(
|
|||
agent_id=agent.agent_id,
|
||||
metadata=data.get("metadata", {}),
|
||||
proxy_server_request=data.get("proxy_server_request"),
|
||||
agent_extra_headers=agent_extra_headers,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
request_data=data,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
|
|
|
|||
|
|
@ -128,6 +128,14 @@ class AgentRegistry:
|
|||
agent_copy, None, prisma_client
|
||||
)
|
||||
|
||||
# Serialize static_headers
|
||||
static_headers_obj = agent.get("static_headers")
|
||||
static_headers_val: Optional[str] = (
|
||||
safe_dumps(dict(static_headers_obj)) if static_headers_obj else None
|
||||
)
|
||||
|
||||
extra_headers_val: Optional[List[str]] = agent.get("extra_headers")
|
||||
|
||||
create_data: Dict[str, Any] = {
|
||||
"agent_name": agent_name,
|
||||
"litellm_params": litellm_params,
|
||||
|
|
@ -137,6 +145,10 @@ class AgentRegistry:
|
|||
"created_at": datetime.now(timezone.utc),
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
}
|
||||
if static_headers_val is not None:
|
||||
create_data["static_headers"] = static_headers_val
|
||||
if extra_headers_val is not None:
|
||||
create_data["extra_headers"] = extra_headers_val
|
||||
if object_permission_id is not None:
|
||||
create_data["object_permission_id"] = object_permission_id
|
||||
|
||||
|
|
@ -214,6 +226,16 @@ class AgentRegistry:
|
|||
update_data["agent_card_params"] = safe_dumps(
|
||||
augment_agent.get("agent_card_params")
|
||||
)
|
||||
if "static_headers" in agent:
|
||||
headers_value = agent.get("static_headers")
|
||||
update_data["static_headers"] = safe_dumps(
|
||||
dict(headers_value) if headers_value is not None else {}
|
||||
)
|
||||
if "extra_headers" in agent:
|
||||
extra_headers_value = agent.get("extra_headers")
|
||||
update_data["extra_headers"] = (
|
||||
extra_headers_value if extra_headers_value is not None else []
|
||||
)
|
||||
if agent.get("object_permission") is not None:
|
||||
agent_copy = dict(augment_agent)
|
||||
existing_object_permission_id = existing_agent.get(
|
||||
|
|
@ -281,10 +303,21 @@ class AgentRegistry:
|
|||
)
|
||||
agent_card_params: str = safe_dumps(agent_card_params_dict)
|
||||
|
||||
# Serialize static_headers for update
|
||||
static_headers_obj_u = agent.get("static_headers")
|
||||
static_headers_val_u: str = (
|
||||
safe_dumps(dict(static_headers_obj_u))
|
||||
if static_headers_obj_u is not None
|
||||
else safe_dumps({})
|
||||
)
|
||||
extra_headers_val_u: List[str] = agent.get("extra_headers") or []
|
||||
|
||||
update_data: Dict[str, Any] = {
|
||||
"agent_name": agent_name,
|
||||
"litellm_params": litellm_params,
|
||||
"agent_card_params": agent_card_params,
|
||||
"static_headers": static_headers_val_u,
|
||||
"extra_headers": extra_headers_val_u,
|
||||
"updated_by": updated_by,
|
||||
"updated_at": datetime.now(timezone.utc),
|
||||
}
|
||||
|
|
|
|||
27
litellm/proxy/agent_endpoints/utils.py
Normal file
27
litellm/proxy/agent_endpoints/utils.py
Normal file
|
|
@ -0,0 +1,27 @@
|
|||
"""Utility helpers for A2A agent endpoints."""
|
||||
|
||||
from typing import Dict, Mapping, Optional
|
||||
|
||||
|
||||
def merge_agent_headers(
|
||||
*,
|
||||
dynamic_headers: Optional[Mapping[str, str]] = None,
|
||||
static_headers: Optional[Mapping[str, str]] = None,
|
||||
) -> Optional[Dict[str, str]]:
|
||||
"""Merge outbound HTTP headers for A2A agent calls.
|
||||
|
||||
Merge rules:
|
||||
- Start with ``dynamic_headers`` (values extracted from the incoming client request).
|
||||
- Overlay ``static_headers`` (admin-configured per agent).
|
||||
|
||||
If both contain the same key, ``static_headers`` wins.
|
||||
"""
|
||||
merged: Dict[str, str] = {}
|
||||
|
||||
if dynamic_headers:
|
||||
merged.update({str(k): str(v) for k, v in dynamic_headers.items()})
|
||||
|
||||
if static_headers:
|
||||
merged.update({str(k): str(v) for k, v in static_headers.items()})
|
||||
|
||||
return merged or None
|
||||
|
|
@ -63,6 +63,8 @@ model LiteLLM_AgentsTable {
|
|||
agent_name String @unique
|
||||
litellm_params Json?
|
||||
agent_card_params Json
|
||||
static_headers Json? @default("{}")
|
||||
extra_headers String[] @default([])
|
||||
agent_access_groups String[] @default([])
|
||||
object_permission_id String?
|
||||
object_permission LiteLLM_ObjectPermissionTable? @relation(fields: [object_permission_id], references: [object_permission_id])
|
||||
|
|
|
|||
|
|
@ -179,6 +179,8 @@ class AgentConfig(TypedDict, total=False):
|
|||
agent_card_params: Required[AgentCard]
|
||||
litellm_params: Dict[str, Any] # allow for any future litellm params
|
||||
object_permission: AgentObjectPermission
|
||||
static_headers: Optional[Dict[str, str]]
|
||||
extra_headers: Optional[List[str]]
|
||||
|
||||
|
||||
class PatchAgentRequest(TypedDict, total=False):
|
||||
|
|
@ -186,6 +188,8 @@ class PatchAgentRequest(TypedDict, total=False):
|
|||
agent_card_params: AgentCard
|
||||
litellm_params: Dict[str, Any]
|
||||
object_permission: AgentObjectPermission
|
||||
static_headers: Optional[Dict[str, str]]
|
||||
extra_headers: Optional[List[str]]
|
||||
|
||||
|
||||
# Request/Response models for CRUD endpoints
|
||||
|
|
@ -197,6 +201,8 @@ class AgentResponse(BaseModel):
|
|||
litellm_params: Optional[Dict[str, Any]] = None
|
||||
agent_card_params: Dict[str, Any]
|
||||
object_permission: Optional[Dict[str, Any]] = None
|
||||
static_headers: Optional[Dict[str, str]] = None
|
||||
extra_headers: Optional[List[str]] = None
|
||||
created_at: Optional[datetime] = None
|
||||
updated_at: Optional[datetime] = None
|
||||
created_by: Optional[str] = None
|
||||
|
|
|
|||
|
|
@ -0,0 +1,248 @@
|
|||
"""
|
||||
Tests that prove header isolation between agents.
|
||||
|
||||
Before the fix these tests FAIL — agent A's headers bleed into agent B
|
||||
because create_a2a_client mutates a globally cached httpx client.
|
||||
After the fix they pass.
|
||||
"""
|
||||
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, call, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helpers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_agent(agent_id, agent_name, static_headers=None, extra_headers=None, url="http://0.0.0.0:9999"):
|
||||
a = MagicMock()
|
||||
a.agent_id = agent_id
|
||||
a.agent_name = agent_name
|
||||
a.agent_card_params = {"url": url, "name": agent_name}
|
||||
a.litellm_params = {}
|
||||
a.static_headers = static_headers or {}
|
||||
a.extra_headers = extra_headers or []
|
||||
return a
|
||||
|
||||
|
||||
def _make_request(method="message/send", extra_headers=None):
|
||||
mock_request = MagicMock()
|
||||
headers = {"content-type": "application/json"}
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
mock_request.headers = headers
|
||||
mock_request.json = AsyncMock(
|
||||
return_value={
|
||||
"jsonrpc": "2.0",
|
||||
"id": "test-id",
|
||||
"method": method,
|
||||
"params": {
|
||||
"message": {
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hello"}],
|
||||
"messageId": "msg-1",
|
||||
}
|
||||
},
|
||||
}
|
||||
)
|
||||
return mock_request
|
||||
|
||||
|
||||
def _a2a_types_module():
|
||||
try:
|
||||
from a2a.types import MessageSendParams, SendMessageRequest, SendStreamingMessageRequest
|
||||
m = MagicMock()
|
||||
m.MessageSendParams = MessageSendParams
|
||||
m.SendMessageRequest = SendMessageRequest
|
||||
m.SendStreamingMessageRequest = SendStreamingMessageRequest
|
||||
return m
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
def _cls(name):
|
||||
class C:
|
||||
def __init__(self, **kw):
|
||||
self.__dict__.update(kw)
|
||||
self._kw = kw
|
||||
def model_dump(self, mode="json", exclude_none=False):
|
||||
return dict(self._kw)
|
||||
C.__name__ = name
|
||||
return C
|
||||
|
||||
m = MagicMock()
|
||||
m.MessageSendParams = _cls("MessageSendParams")
|
||||
m.SendMessageRequest = _cls("SendMessageRequest")
|
||||
m.SendStreamingMessageRequest = _cls("SendStreamingMessageRequest")
|
||||
return m
|
||||
|
||||
|
||||
async def _invoke_agent(agent, request):
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1")
|
||||
fastapi_response = MagicMock()
|
||||
mock_response = MagicMock()
|
||||
mock_response.model_dump.return_value = {"jsonrpc": "2.0", "id": "test-id", "result": {}}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.a2a_endpoints._get_agent",
|
||||
return_value=agent,
|
||||
), patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed",
|
||||
new_callable=AsyncMock,
|
||||
return_value=True,
|
||||
), patch(
|
||||
"litellm.proxy.common_request_processing.add_litellm_data_to_request",
|
||||
side_effect=lambda data, **kw: data,
|
||||
), patch(
|
||||
"litellm.a2a_protocol.asend_message",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
) as mock_asend, patch(
|
||||
"litellm.a2a_protocol.create_a2a_client",
|
||||
new_callable=AsyncMock,
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.general_settings", {}
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.proxy_config", MagicMock()
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.version", "1.0.0"
|
||||
), patch.dict(
|
||||
sys.modules,
|
||||
{"a2a": MagicMock(), "a2a.types": _a2a_types_module()},
|
||||
), patch(
|
||||
"litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True
|
||||
):
|
||||
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
|
||||
|
||||
await invoke_agent_a2a(
|
||||
agent_id=agent.agent_id,
|
||||
request=request,
|
||||
fastapi_response=fastapi_response,
|
||||
user_api_key_dict=user_api_key_dict,
|
||||
)
|
||||
return mock_asend.call_args.kwargs.get("agent_extra_headers")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_static_headers_do_not_leak_between_agents():
|
||||
"""
|
||||
Agent A has static_headers={"X-Agent-A-Token": "secret-a"}.
|
||||
Agent B has no headers.
|
||||
After invoking A then B, B must NOT receive X-Agent-A-Token.
|
||||
"""
|
||||
agent_a = _make_agent("id-a", "agent-a", static_headers={"X-Agent-A-Token": "secret-a"})
|
||||
agent_b = _make_agent("id-b", "agent-b")
|
||||
|
||||
headers_a = await _invoke_agent(agent_a, _make_request())
|
||||
headers_b = await _invoke_agent(agent_b, _make_request())
|
||||
|
||||
assert headers_a is not None
|
||||
assert headers_a.get("X-Agent-A-Token") == "secret-a"
|
||||
|
||||
# Agent B must not have agent A's header
|
||||
assert headers_b is None or "X-Agent-A-Token" not in headers_b
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_convention_header_only_matches_own_agent():
|
||||
"""
|
||||
Client sends x-a2a-agent-a-authorization: Bearer for-a.
|
||||
When invoking agent-b, that header must NOT be forwarded.
|
||||
"""
|
||||
agent_b = _make_agent("id-b", "agent-b")
|
||||
# Request carries a header scoped to agent-a, not agent-b
|
||||
req = _make_request(extra_headers={"x-a2a-agent-a-authorization": "Bearer for-a"})
|
||||
|
||||
headers_b = await _invoke_agent(agent_b, req)
|
||||
|
||||
assert headers_b is None or "authorization" not in (headers_b or {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_convention_header_matches_own_agent():
|
||||
"""
|
||||
Client sends x-a2a-agent-b-authorization: Bearer for-b.
|
||||
When invoking agent-b, that header IS forwarded.
|
||||
"""
|
||||
agent_b = _make_agent("id-b", "agent-b")
|
||||
req = _make_request(extra_headers={"x-a2a-agent-b-authorization": "Bearer for-b"})
|
||||
|
||||
headers_b = await _invoke_agent(agent_b, req)
|
||||
|
||||
assert headers_b is not None
|
||||
assert headers_b.get("authorization") == "Bearer for-b"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_each_agent_gets_only_its_own_static_headers():
|
||||
"""
|
||||
Agent A: static_headers={"X-Token": "a"}
|
||||
Agent B: static_headers={"X-Token": "b"}
|
||||
Each must receive only their own value.
|
||||
"""
|
||||
agent_a = _make_agent("id-a", "agent-a", static_headers={"X-Token": "a"})
|
||||
agent_b = _make_agent("id-b", "agent-b", static_headers={"X-Token": "b"})
|
||||
|
||||
headers_a = await _invoke_agent(agent_a, _make_request())
|
||||
headers_b = await _invoke_agent(agent_b, _make_request())
|
||||
|
||||
assert (headers_a or {}).get("X-Token") == "a"
|
||||
assert (headers_b or {}).get("X-Token") == "b"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Unit test: create_a2a_client uses a fresh httpx client per call
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_a2a_client_uses_fresh_httpx_client():
|
||||
"""
|
||||
Two calls to create_a2a_client with different extra_headers must NOT
|
||||
share the same underlying httpx.AsyncClient instance.
|
||||
"""
|
||||
import httpx
|
||||
|
||||
from litellm.a2a_protocol.main import create_a2a_client
|
||||
|
||||
created_clients = []
|
||||
|
||||
fake_agent_card = MagicMock()
|
||||
fake_agent_card.name = "test-agent"
|
||||
|
||||
class FakeResolver:
|
||||
def __init__(self, **kw):
|
||||
created_clients.append(kw.get("httpx_client"))
|
||||
async def get_agent_card(self):
|
||||
return fake_agent_card
|
||||
|
||||
class FakeA2AClient:
|
||||
def __init__(self, httpx_client, agent_card):
|
||||
self._client = httpx_client
|
||||
self._litellm_agent_card = agent_card
|
||||
|
||||
with patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True), patch(
|
||||
"litellm.a2a_protocol.main.A2ACardResolver", FakeResolver
|
||||
), patch("litellm.a2a_protocol.main._A2AClient", FakeA2AClient):
|
||||
await create_a2a_client(
|
||||
base_url="http://agent-a:9999",
|
||||
extra_headers={"Authorization": "Bearer a"},
|
||||
)
|
||||
await create_a2a_client(
|
||||
base_url="http://agent-b:9999",
|
||||
extra_headers={"Authorization": "Bearer b"},
|
||||
)
|
||||
|
||||
assert len(created_clients) == 2
|
||||
# Must be distinct objects
|
||||
assert created_clients[0] is not created_clients[1], (
|
||||
"create_a2a_client reused a cached httpx client — headers will bleed between agents"
|
||||
)
|
||||
339
tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py
Normal file
339
tests/test_litellm/proxy/agent_endpoints/test_agent_headers.py
Normal file
|
|
@ -0,0 +1,339 @@
|
|||
"""
|
||||
Unit tests for A2A agent custom header forwarding.
|
||||
|
||||
Tests cover:
|
||||
- Static headers forwarded to backend agent
|
||||
- Dynamic headers extracted from client request and forwarded
|
||||
- Static headers win over dynamic on conflict
|
||||
- No headers configured — existing behavior unchanged
|
||||
- merge_agent_headers utility
|
||||
"""
|
||||
|
||||
import sys
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Helper: build a minimal mock agent
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def _make_mock_agent(
|
||||
static_headers=None,
|
||||
extra_headers=None,
|
||||
url="http://backend-agent:10001",
|
||||
):
|
||||
mock_agent = MagicMock()
|
||||
mock_agent.agent_id = "agent-123"
|
||||
mock_agent.agent_card_params = {"url": url, "name": "Test Agent"}
|
||||
mock_agent.litellm_params = {}
|
||||
mock_agent.static_headers = static_headers or {}
|
||||
mock_agent.extra_headers = extra_headers or []
|
||||
return mock_agent
|
||||
|
||||
|
||||
def _make_mock_request(extra_headers=None, method="message/send"):
|
||||
"""Build a mock FastAPI Request with configurable headers."""
|
||||
mock_request = MagicMock()
|
||||
headers = {"content-type": "application/json"}
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
mock_request.headers = headers
|
||||
mock_request.json = AsyncMock(
|
||||
return_value={
|
||||
"jsonrpc": "2.0",
|
||||
"id": "test-id",
|
||||
"method": method,
|
||||
"params": {
|
||||
"message": {
|
||||
"role": "user",
|
||||
"parts": [{"kind": "text", "text": "Hello"}],
|
||||
"messageId": "msg-123",
|
||||
}
|
||||
},
|
||||
}
|
||||
)
|
||||
return mock_request
|
||||
|
||||
|
||||
def _make_a2a_types_module():
|
||||
"""Return (module, MessageSendParams, SendMessageRequest, SendStreamingMessageRequest)."""
|
||||
try:
|
||||
from a2a.types import (
|
||||
MessageSendParams,
|
||||
SendMessageRequest,
|
||||
SendStreamingMessageRequest,
|
||||
)
|
||||
mock_a2a_types = MagicMock()
|
||||
mock_a2a_types.MessageSendParams = MessageSendParams
|
||||
mock_a2a_types.SendMessageRequest = SendMessageRequest
|
||||
mock_a2a_types.SendStreamingMessageRequest = SendStreamingMessageRequest
|
||||
return mock_a2a_types
|
||||
except ImportError:
|
||||
pass
|
||||
|
||||
def _make_cls(name):
|
||||
class MockCls:
|
||||
def __init__(self, **kwargs):
|
||||
self.__dict__.update(kwargs)
|
||||
self._kwargs = kwargs
|
||||
|
||||
def model_dump(self, mode="json", exclude_none=False):
|
||||
result = dict(self._kwargs)
|
||||
if exclude_none:
|
||||
result = {k: v for k, v in result.items() if v is not None}
|
||||
return result
|
||||
|
||||
MockCls.__name__ = name
|
||||
return MockCls
|
||||
|
||||
mock_a2a_types = MagicMock()
|
||||
mock_a2a_types.MessageSendParams = _make_cls("MessageSendParams")
|
||||
mock_a2a_types.SendMessageRequest = _make_cls("SendMessageRequest")
|
||||
mock_a2a_types.SendStreamingMessageRequest = _make_cls(
|
||||
"SendStreamingMessageRequest"
|
||||
)
|
||||
return mock_a2a_types
|
||||
|
||||
|
||||
async def _invoke(mock_agent, mock_request, mock_asend_message):
|
||||
"""Run invoke_agent_a2a with standard patches applied."""
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
|
||||
mock_user_api_key_dict = UserAPIKeyAuth(api_key="sk-test", user_id="u1")
|
||||
mock_fastapi_response = MagicMock()
|
||||
mock_a2a_types = _make_a2a_types_module()
|
||||
|
||||
mock_response = MagicMock()
|
||||
mock_response.model_dump.return_value = {
|
||||
"jsonrpc": "2.0",
|
||||
"id": "test-id",
|
||||
"result": {"status": "success"},
|
||||
}
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.agent_endpoints.a2a_endpoints._get_agent",
|
||||
return_value=mock_agent,
|
||||
), patch(
|
||||
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed",
|
||||
new_callable=AsyncMock,
|
||||
return_value=True,
|
||||
), patch(
|
||||
"litellm.proxy.common_request_processing.add_litellm_data_to_request",
|
||||
side_effect=lambda data, **kw: data,
|
||||
), patch(
|
||||
"litellm.a2a_protocol.asend_message",
|
||||
new_callable=AsyncMock,
|
||||
return_value=mock_response,
|
||||
) as mock_asend, patch(
|
||||
"litellm.a2a_protocol.create_a2a_client",
|
||||
new_callable=AsyncMock,
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.general_settings",
|
||||
{},
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.proxy_config",
|
||||
MagicMock(),
|
||||
), patch(
|
||||
"litellm.proxy.proxy_server.version",
|
||||
"1.0.0",
|
||||
), patch.dict(
|
||||
sys.modules,
|
||||
{"a2a": MagicMock(), "a2a.types": mock_a2a_types},
|
||||
), patch(
|
||||
"litellm.a2a_protocol.main.A2A_SDK_AVAILABLE",
|
||||
True,
|
||||
):
|
||||
from litellm.proxy.agent_endpoints.a2a_endpoints import invoke_agent_a2a
|
||||
|
||||
await invoke_agent_a2a(
|
||||
agent_id="test-agent",
|
||||
request=mock_request,
|
||||
fastapi_response=mock_fastapi_response,
|
||||
user_api_key_dict=mock_user_api_key_dict,
|
||||
)
|
||||
return mock_asend
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_static_headers_forwarded():
|
||||
"""Static headers configured on the agent are passed to asend_message."""
|
||||
mock_agent = _make_mock_agent(
|
||||
static_headers={"Authorization": "Bearer token123"}
|
||||
)
|
||||
mock_request = _make_mock_request()
|
||||
|
||||
mock_asend = await _invoke(mock_agent, mock_request, None)
|
||||
|
||||
call_kwargs = mock_asend.call_args.kwargs
|
||||
headers = call_kwargs.get("agent_extra_headers")
|
||||
assert headers is not None, "agent_extra_headers should not be None"
|
||||
assert headers.get("Authorization") == "Bearer token123"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_dynamic_headers_forwarded():
|
||||
"""Dynamic headers listed in extra_headers are extracted from the client request."""
|
||||
mock_agent = _make_mock_agent(extra_headers=["x-api-key"])
|
||||
mock_request = _make_mock_request(extra_headers={"x-api-key": "secret"})
|
||||
|
||||
mock_asend = await _invoke(mock_agent, mock_request, None)
|
||||
|
||||
call_kwargs = mock_asend.call_args.kwargs
|
||||
headers = call_kwargs.get("agent_extra_headers")
|
||||
assert headers is not None
|
||||
assert headers.get("x-api-key") == "secret"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_static_overrides_dynamic():
|
||||
"""When the same header appears in both static and dynamic, static wins."""
|
||||
mock_agent = _make_mock_agent(
|
||||
static_headers={"Authorization": "Bearer static-token"},
|
||||
extra_headers=["Authorization"],
|
||||
)
|
||||
# Client sends a different value for Authorization
|
||||
mock_request = _make_mock_request(
|
||||
extra_headers={"Authorization": "Bearer dynamic-token"}
|
||||
)
|
||||
|
||||
mock_asend = await _invoke(mock_agent, mock_request, None)
|
||||
|
||||
call_kwargs = mock_asend.call_args.kwargs
|
||||
headers = call_kwargs.get("agent_extra_headers")
|
||||
assert headers is not None
|
||||
assert headers.get("Authorization") == "Bearer static-token"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_no_headers():
|
||||
"""When no headers are configured, agent_extra_headers is None and behaviour is unchanged."""
|
||||
mock_agent = _make_mock_agent() # no static_headers or extra_headers
|
||||
mock_request = _make_mock_request()
|
||||
|
||||
mock_asend = await _invoke(mock_agent, mock_request, None)
|
||||
|
||||
call_kwargs = mock_asend.call_args.kwargs
|
||||
headers = call_kwargs.get("agent_extra_headers")
|
||||
assert headers is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Convention-based x-a2a-{agent_id/name}-{header_name} tests
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_convention_header_by_agent_name():
|
||||
"""x-a2a-{agent_name}-{header} is forwarded using the agent name alias."""
|
||||
mock_agent = _make_mock_agent()
|
||||
mock_agent.agent_name = "my-agent"
|
||||
mock_request = _make_mock_request(
|
||||
extra_headers={"x-a2a-my-agent-authorization": "Bearer conv-token"}
|
||||
)
|
||||
|
||||
mock_asend = await _invoke(mock_agent, mock_request, None)
|
||||
|
||||
headers = mock_asend.call_args.kwargs.get("agent_extra_headers")
|
||||
assert headers is not None
|
||||
assert headers.get("authorization") == "Bearer conv-token"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_convention_header_by_agent_id():
|
||||
"""x-a2a-{agent_id}-{header} is forwarded using the agent UUID."""
|
||||
mock_agent = _make_mock_agent()
|
||||
mock_agent.agent_id = "abc-123"
|
||||
mock_agent.agent_name = "other-name"
|
||||
mock_request = _make_mock_request(
|
||||
extra_headers={"x-a2a-abc-123-x-api-key": "id-secret"}
|
||||
)
|
||||
|
||||
mock_asend = await _invoke(mock_agent, mock_request, None)
|
||||
|
||||
headers = mock_asend.call_args.kwargs.get("agent_extra_headers")
|
||||
assert headers is not None
|
||||
assert headers.get("x-api-key") == "id-secret"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_convention_header_static_still_wins():
|
||||
"""Static headers still override convention-based dynamic headers."""
|
||||
mock_agent = _make_mock_agent(
|
||||
static_headers={"authorization": "Bearer static-wins"}
|
||||
)
|
||||
mock_agent.agent_name = "my-agent"
|
||||
mock_request = _make_mock_request(
|
||||
extra_headers={"x-a2a-my-agent-authorization": "Bearer conv-value"}
|
||||
)
|
||||
|
||||
mock_asend = await _invoke(mock_agent, mock_request, None)
|
||||
|
||||
headers = mock_asend.call_args.kwargs.get("agent_extra_headers")
|
||||
assert headers is not None
|
||||
assert headers.get("authorization") == "Bearer static-wins"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_convention_unrelated_prefix_not_forwarded():
|
||||
"""Headers that start with x-a2a- but target a different agent are ignored."""
|
||||
mock_agent = _make_mock_agent()
|
||||
mock_agent.agent_id = "agent-abc"
|
||||
mock_agent.agent_name = "my-agent"
|
||||
mock_request = _make_mock_request(
|
||||
extra_headers={"x-a2a-other-agent-authorization": "Bearer wrong"}
|
||||
)
|
||||
|
||||
mock_asend = await _invoke(mock_agent, mock_request, None)
|
||||
|
||||
headers = mock_asend.call_args.kwargs.get("agent_extra_headers")
|
||||
assert headers is None
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Direct unit test for the merge utility
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_merge_agent_headers_util_dynamic_only():
|
||||
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
|
||||
|
||||
result = merge_agent_headers(dynamic_headers={"x-key": "val"})
|
||||
assert result == {"x-key": "val"}
|
||||
|
||||
|
||||
def test_merge_agent_headers_util_static_only():
|
||||
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
|
||||
|
||||
result = merge_agent_headers(static_headers={"Authorization": "Bearer tok"})
|
||||
assert result == {"Authorization": "Bearer tok"}
|
||||
|
||||
|
||||
def test_merge_agent_headers_util_static_wins():
|
||||
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
|
||||
|
||||
result = merge_agent_headers(
|
||||
dynamic_headers={"Authorization": "dynamic", "x-extra": "d"},
|
||||
static_headers={"Authorization": "static"},
|
||||
)
|
||||
assert result == {"Authorization": "static", "x-extra": "d"}
|
||||
|
||||
|
||||
def test_merge_agent_headers_util_none_returns_none():
|
||||
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
|
||||
|
||||
result = merge_agent_headers()
|
||||
assert result is None
|
||||
|
||||
|
||||
def test_merge_agent_headers_util_empty_dicts_returns_none():
|
||||
from litellm.proxy.agent_endpoints.utils import merge_agent_headers
|
||||
|
||||
result = merge_agent_headers(dynamic_headers={}, static_headers={})
|
||||
assert result is None
|
||||
114
tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py
Normal file
114
tests/test_litellm/proxy/agent_endpoints/test_agent_registry.py
Normal file
|
|
@ -0,0 +1,114 @@
|
|||
"""Unit tests for AgentRegistry DB operations."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.agent_endpoints.agent_registry import AgentRegistry
|
||||
|
||||
|
||||
def _sample_agent_card_params() -> dict:
|
||||
return {
|
||||
"protocolVersion": "1.0",
|
||||
"name": "Test Agent",
|
||||
"description": "desc",
|
||||
"url": "http://localhost",
|
||||
"version": "1.0.0",
|
||||
"capabilities": {"streaming": True},
|
||||
"defaultInputModes": ["text"],
|
||||
"defaultOutputModes": ["text"],
|
||||
"skills": [],
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_agent_in_db_clears_static_headers_and_extra_headers_when_omitted():
|
||||
"""
|
||||
PUT (full-replace) should clear static_headers and extra_headers when omitted.
|
||||
Previously, omitting these fields left stale DB values intact.
|
||||
"""
|
||||
registry = AgentRegistry()
|
||||
mock_prisma = MagicMock()
|
||||
|
||||
# Simulate existing agent that had headers set
|
||||
updated_agent = MagicMock()
|
||||
updated_agent.model_dump.return_value = {
|
||||
"agent_id": "agent-123",
|
||||
"agent_name": "Updated Agent",
|
||||
"agent_card_params": _sample_agent_card_params(),
|
||||
"litellm_params": {},
|
||||
"static_headers": {},
|
||||
"extra_headers": [],
|
||||
"object_permission": None,
|
||||
}
|
||||
updated_agent.object_permission = None
|
||||
|
||||
mock_update = AsyncMock(return_value=updated_agent)
|
||||
mock_prisma.db.litellm_agentstable.update = mock_update
|
||||
|
||||
# Agent config WITHOUT static_headers or extra_headers (omitted)
|
||||
agent_config = {
|
||||
"agent_name": "Updated Agent",
|
||||
"agent_card_params": _sample_agent_card_params(),
|
||||
"litellm_params": {},
|
||||
}
|
||||
|
||||
await registry.update_agent_in_db(
|
||||
agent_id="agent-123",
|
||||
agent=agent_config,
|
||||
prisma_client=mock_prisma,
|
||||
updated_by="test-user",
|
||||
)
|
||||
|
||||
mock_update.assert_awaited_once()
|
||||
call_kwargs = mock_update.call_args.kwargs
|
||||
update_data = call_kwargs["data"]
|
||||
|
||||
# Should include static_headers and extra_headers with empty defaults
|
||||
assert "static_headers" in update_data
|
||||
assert update_data["static_headers"] == "{}"
|
||||
assert "extra_headers" in update_data
|
||||
assert update_data["extra_headers"] == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_agent_in_db_preserves_explicit_static_headers_and_extra_headers():
|
||||
"""PUT with explicit values should still work correctly."""
|
||||
registry = AgentRegistry()
|
||||
mock_prisma = MagicMock()
|
||||
|
||||
updated_agent = MagicMock()
|
||||
updated_agent.model_dump.return_value = {
|
||||
"agent_id": "agent-123",
|
||||
"agent_name": "Updated Agent",
|
||||
"agent_card_params": _sample_agent_card_params(),
|
||||
"litellm_params": {},
|
||||
"static_headers": {"Authorization": "Bearer xyz"},
|
||||
"extra_headers": ["X-Custom-Header"],
|
||||
"object_permission": None,
|
||||
}
|
||||
updated_agent.object_permission = None
|
||||
|
||||
mock_update = AsyncMock(return_value=updated_agent)
|
||||
mock_prisma.db.litellm_agentstable.update = mock_update
|
||||
|
||||
agent_config = {
|
||||
"agent_name": "Updated Agent",
|
||||
"agent_card_params": _sample_agent_card_params(),
|
||||
"litellm_params": {},
|
||||
"static_headers": {"Authorization": "Bearer xyz"},
|
||||
"extra_headers": ["X-Custom-Header"],
|
||||
}
|
||||
|
||||
await registry.update_agent_in_db(
|
||||
agent_id="agent-123",
|
||||
agent=agent_config,
|
||||
prisma_client=mock_prisma,
|
||||
updated_by="test-user",
|
||||
)
|
||||
|
||||
call_kwargs = mock_update.call_args.kwargs
|
||||
update_data = call_kwargs["data"]
|
||||
|
||||
assert update_data["static_headers"] == '{"Authorization": "Bearer xyz"}'
|
||||
assert update_data["extra_headers"] == ["X-Custom-Header"]
|
||||
|
|
@ -269,6 +269,23 @@ export const buildAgentDataFromForm = (values: any, existingAgent?: any) => {
|
|||
agentData.litellm_params = params;
|
||||
}
|
||||
|
||||
// static_headers: convert [{header, value}, ...] → {header: value, ...}
|
||||
if (Array.isArray(values.static_headers) && values.static_headers.length > 0) {
|
||||
const staticHeaders: Record<string, string> = {};
|
||||
values.static_headers.forEach((entry: { header?: string; value?: string }) => {
|
||||
const key = entry?.header?.trim();
|
||||
if (key) staticHeaders[key] = entry?.value ?? "";
|
||||
});
|
||||
if (Object.keys(staticHeaders).length > 0) {
|
||||
agentData.static_headers = staticHeaders;
|
||||
}
|
||||
}
|
||||
|
||||
// extra_headers: already an array of strings from Select tags
|
||||
if (Array.isArray(values.extra_headers) && values.extra_headers.length > 0) {
|
||||
agentData.extra_headers = values.extra_headers;
|
||||
}
|
||||
|
||||
return agentData;
|
||||
};
|
||||
|
||||
|
|
@ -302,5 +319,14 @@ export const parseAgentForForm = (agent: any) => {
|
|||
cost_per_query: agent.litellm_params?.cost_per_query,
|
||||
input_cost_per_token: agent.litellm_params?.input_cost_per_token,
|
||||
output_cost_per_token: agent.litellm_params?.output_cost_per_token,
|
||||
// static_headers: {key: value} → [{header, value}, ...]
|
||||
static_headers: agent.static_headers
|
||||
? Object.entries(agent.static_headers as Record<string, string>).map(([header, value]) => ({
|
||||
header,
|
||||
value,
|
||||
}))
|
||||
: [],
|
||||
// extra_headers: already an array of strings
|
||||
extra_headers: agent.extra_headers ?? [],
|
||||
};
|
||||
};
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import React from "react";
|
||||
import { Form, Input, Switch, Collapse } from "antd";
|
||||
import { Form, Input, Switch, Collapse, Select, Space, Tooltip } from "antd";
|
||||
import { Button as AntButton } from "antd";
|
||||
import { PlusOutlined, MinusCircleOutlined } from "@ant-design/icons";
|
||||
import { PlusOutlined, MinusCircleOutlined, InfoCircleOutlined } from "@ant-design/icons";
|
||||
import { AGENT_FORM_CONFIG, SKILL_FIELD_CONFIG } from "./agent_config";
|
||||
|
||||
import CostConfigFields from "./cost_config_fields";
|
||||
|
|
@ -188,6 +188,72 @@ const AgentFormFields: React.FC<AgentFormFieldsProps> = ({ showAgentName = true,
|
|||
))}
|
||||
</Panel>
|
||||
)}
|
||||
|
||||
{/* Authentication Headers */}
|
||||
{shouldShow("auth_headers") && (
|
||||
<Panel header="Authentication Headers" key="auth_headers">
|
||||
{/* Static Headers */}
|
||||
<Form.Item
|
||||
label={
|
||||
<span>
|
||||
Static Headers{" "}
|
||||
<Tooltip title="Headers always sent to the backend agent, regardless of the client request. Admin-configured, static wins on conflict.">
|
||||
<InfoCircleOutlined style={{ color: "#8c8c8c" }} />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
>
|
||||
<Form.List name="static_headers">
|
||||
{(fields, { add, remove }) => (
|
||||
<>
|
||||
{fields.map(({ key, name, ...restField }) => (
|
||||
<Space key={key} style={{ display: "flex", marginBottom: 8 }} align="baseline">
|
||||
<Form.Item
|
||||
{...restField}
|
||||
name={[name, "header"]}
|
||||
rules={[{ required: true, message: "Header name required" }]}
|
||||
>
|
||||
<Input placeholder="Header name (e.g. Authorization)" style={{ width: 220 }} />
|
||||
</Form.Item>
|
||||
<Form.Item
|
||||
{...restField}
|
||||
name={[name, "value"]}
|
||||
rules={[{ required: true, message: "Value required" }]}
|
||||
>
|
||||
<Input placeholder="Value (e.g. Bearer token123)" style={{ width: 260 }} />
|
||||
</Form.Item>
|
||||
<MinusCircleOutlined onClick={() => remove(name)} style={{ color: "#ff4d4f" }} />
|
||||
</Space>
|
||||
))}
|
||||
<AntButton type="dashed" onClick={() => add()} icon={<PlusOutlined />} style={{ width: "100%" }}>
|
||||
Add Static Header
|
||||
</AntButton>
|
||||
</>
|
||||
)}
|
||||
</Form.List>
|
||||
</Form.Item>
|
||||
|
||||
{/* Extra Headers (dynamic forwarding) */}
|
||||
<Form.Item
|
||||
label={
|
||||
<span>
|
||||
Forward Client Headers{" "}
|
||||
<Tooltip title="Header names to extract from the client's request and forward to the agent. Type a name and press Enter.">
|
||||
<InfoCircleOutlined style={{ color: "#8c8c8c" }} />
|
||||
</Tooltip>
|
||||
</span>
|
||||
}
|
||||
name="extra_headers"
|
||||
>
|
||||
<Select
|
||||
mode="tags"
|
||||
style={{ width: "100%" }}
|
||||
placeholder="e.g. x-api-key, Authorization"
|
||||
tokenSeparators={[","]}
|
||||
/>
|
||||
</Form.Item>
|
||||
</Panel>
|
||||
)}
|
||||
</Collapse>
|
||||
</>
|
||||
);
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue