Merge pull request #22888 from BerriAI/litellm_a2a-custom-headers

[Feat] Add a2a custom headers
This commit is contained in:
Sameer Kankute 2026-03-06 18:24:21 +05:30 • committed by GitHub
commit 8b0375f99c
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
14 changed files with 1183 additions and 13 deletions

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

View file

@ -539,6 +539,7 @@ const sidebars = {
items: [
"a2a",
"a2a_invoking_agents",
"a2a_agent_headers",
"a2a_cost_tracking",
"a2a_agent_permissions"
],

View file

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

View file

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

View file

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

View file

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

View 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

View file

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

View file

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

View file

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

View 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

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

View file

@ -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 ?? [],
};
};

View file

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