fix: preserve registered A2A request context

This commit is contained in:
aiedwardyi 2026-08-24 14:18:12 +09:00
parent c3d599dd27
commit 02c5c74b37
No known key found for this signature in database
6 changed files with 253 additions and 19 deletions

View file

@ -83,6 +83,13 @@ class A2AModelResponseIterator(BaseModelResponseIterator):
tool_use=None,
)
def _handle_string_chunk(
self, str_line: str | dict
) -> GenericStreamingChunk | ModelResponseStream:
if isinstance(str_line, dict):
return self.chunk_parser(chunk=str_line)
return super()._handle_string_chunk(str_line=str_line)
def _get_finish_reason(self, chunk: dict) -> str | None:
"""Extract finish reason from A2A chunk"""
result: Final = chunk.get("result", {})

View file

@ -7,6 +7,7 @@ Looks up agents in the registry and injects their API base URL.
from __future__ import annotations
import asyncio
from collections.abc import Awaitable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Final
@ -18,6 +19,7 @@ from typing_extensions import ReadOnly, TypedDict
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.interactions.agents.utils import merge_agent_headers
from litellm.llms.a2a.common_utils import A2AError, convert_messages_to_prompt, extract_text_from_a2a_response
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
from litellm.types.llms.openai import AllMessageValues
@ -29,6 +31,41 @@ if TYPE_CHECKING:
_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
_HEADERS_ADAPTER: Final = TypeAdapter(dict[str, str])
_MESSAGES_ADAPTER: Final = TypeAdapter(list[AllMessageValues])
_FORWARDED_REQUEST_PARAMS: Final = frozenset(
{
"audio",
"frequency_penalty",
"functions",
"function_call",
"include_server_side_tool_invocations",
"logit_bias",
"logprobs",
"guardrails",
"max_completion_tokens",
"max_tokens",
"modalities",
"n",
"parallel_tool_calls",
"prediction",
"presence_penalty",
"reasoning_effort",
"response_format",
"seed",
"service_tier",
"stop",
"store",
"temperature",
"thinking",
"timeout",
"tool_choice",
"tools",
"top_logprobs",
"top_p",
"user",
"verbosity",
"web_search_options",
}
)
class _A2ATextPart(TypedDict):
@ -49,8 +86,9 @@ class _A2AParams(TypedDict):
async def _route_registered_provider(
data: Mapping[str, object],
model_name: str,
api_base: str,
api_base: str | None,
litellm_params: Mapping[str, object],
static_headers: Mapping[str, str] | None,
) -> ModelResponse | CustomStreamWrapper:
from litellm.a2a_protocol.litellm_completion_bridge.handler import (
A2ACompletionBridgeHandler,
@ -70,12 +108,25 @@ async def _route_registered_provider(
"messageId": str(uuid4()),
}
}
provider_params: Final = _OBJECT_DICT_ADAPTER.validate_python(litellm_params)
provider_params: Final = {
**_OBJECT_DICT_ADAPTER.validate_python(litellm_params),
**{
key: data[key]
for key in _FORWARDED_REQUEST_PARAMS
if key in data and data[key] is not None
},
}
bridge_params: Final = _OBJECT_DICT_ADAPTER.validate_python(params)
configured_headers: Final = litellm_params.get("extra_headers") or litellm_params.get("headers")
agent_extra_headers: Final = (
configured_headers_dict: Final = (
_HEADERS_ADAPTER.validate_python(configured_headers) if isinstance(configured_headers, dict) else None
)
agent_extra_headers: Final = merge_agent_headers(
dynamic_headers=configured_headers_dict,
static_headers=static_headers,
)
if agent_extra_headers:
provider_params["extra_headers"] = agent_extra_headers
if stream:
streaming_response: Final = A2ACompletionBridgeHandler.handle_streaming(
@ -125,9 +176,63 @@ async def _route_registered_provider(
Choices(finish_reason="stop", index=0, message=Message(content=text, role="assistant"))
],
)
usage: Final = response.get("usage")
if usage is not None:
setattr(model_response, "usage", usage)
logging_obj: Final = data.get("litellm_logging_obj")
if isinstance(logging_obj, Logging):
def _enqueue_logging() -> None:
asyncio.create_task(
logging_obj.dispatch_success_handlers(
model_response,
cache_hit=False,
prefer_async_handlers=True,
)
)
logging_obj._enqueue_deferred_logging = _enqueue_logging
return model_response
def _merge_agent_guardrails(
data: Mapping[str, object],
agent_guardrails: object,
) -> Mapping[str, object]:
if not agent_guardrails:
return data
configured_guardrails: list[object] = (
agent_guardrails if isinstance(agent_guardrails, list) else [agent_guardrails]
)
metadata = data.get("metadata")
metadata_guardrails = metadata.get("guardrails") if isinstance(metadata, dict) else None
root_guardrails = data.get("guardrails")
existing_guardrails: list[object] = []
for value in (metadata_guardrails, root_guardrails):
if isinstance(value, list):
existing_guardrails.extend(value)
elif value:
existing_guardrails.append(value)
merged_guardrails = existing_guardrails + [
guardrail for guardrail in configured_guardrails if guardrail not in existing_guardrails
]
if isinstance(data, dict):
data["guardrails"] = merged_guardrails
if isinstance(metadata, dict):
metadata["guardrails"] = merged_guardrails
return data
merged_data = dict(data)
merged_data["guardrails"] = merged_guardrails
if isinstance(metadata, dict):
merged_data["metadata"] = {**metadata, "guardrails": merged_guardrails}
return merged_data
async def route_a2a_agent_request(
data: Mapping[str, object],
route_type: str,
@ -185,11 +290,6 @@ async def route_a2a_agent_request(
# Get API base URL from agent config
agent_card_params: Final = agent.agent_card_params
agent_url: Final = agent_card_params.get("url") if agent_card_params else None
if not isinstance(agent_url, str) or not agent_url:
verbose_proxy_logger.error("[A2A] Agent '%s' has no URL configured", agent_name)
route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type)
raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False)
registered_params_value: Final = agent.litellm_params
registered_provider_value: Final = (
registered_params_value.get("custom_llm_provider") if registered_params_value else None
@ -197,6 +297,19 @@ async def route_a2a_agent_request(
registered_provider: Final = registered_provider_value if isinstance(registered_provider_value, str) else None
configured_api_base: Final = registered_params_value.get("api_base") if registered_params_value else None
api_base: Final = configured_api_base if isinstance(configured_api_base, str) else agent_url
registered_model: Final = registered_params_value.get("model") if registered_params_value else None
cardless_provider: Final = (
registered_provider == "bedrock" and isinstance(registered_model, str) and "agentcore" in registered_model
)
if (not isinstance(agent_url, str) or not agent_url) and not cardless_provider:
verbose_proxy_logger.error("[A2A] Agent '%s' has no URL configured", agent_name)
route_name = ROUTE_ENDPOINT_MAPPING.get(route_type, route_type)
raise ProxyModelNotFoundError(route=route_name, model_name=model_name, retryable_with_model_read_through=False)
routed_data: Final = _merge_agent_guardrails(
data=data,
agent_guardrails=registered_params_value.get("guardrails") if registered_params_value else None,
)
if (
registered_provider
and registered_provider != "a2a"
@ -205,12 +318,13 @@ async def route_a2a_agent_request(
):
verbose_proxy_logger.debug("[A2A] Routing %s through %s", model_name, registered_provider)
return _route_registered_provider(
data=data,
data=routed_data,
model_name=model_name,
api_base=api_base,
litellm_params=registered_params_value,
static_headers=agent.static_headers,
)
completion_data: Final = MappingProxyType({**data, "api_base": api_base})
completion_data: Final = MappingProxyType({**routed_data, "api_base": api_base})
verbose_proxy_logger.debug("[A2A] Routing %s to %s", model_name, api_base)
return getattr(litellm, f"{route_type}")(**completion_data) # pyright: ignore[reportAny] # dynamic SDK route

View file

@ -6,16 +6,12 @@ Used by proxy model endpoints to make agents appear in UI alongside models.
from typing import Final
from pydantic import TypeAdapter
from litellm._logging import verbose_proxy_logger
from litellm.proxy._types import UserAPIKeyAuth
from litellm.types.proxy.management_endpoints.model_management_endpoints import (
ModelGroupInfoProxy,
)
_OBJECT_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
async def append_agents_to_model_group(
model_groups: list[ModelGroupInfoProxy],
@ -39,11 +35,16 @@ async def append_agents_to_model_group(
for agent_id in allowed_agent_ids:
agent = global_agent_registry.get_agent_by_id(agent_id)
if agent is not None:
agent_params: Final = agent.litellm_params
provider_value: Final = agent_params.get("custom_llm_provider") if agent_params else None
custom_llm_provider: Final = (
provider_value if isinstance(provider_value, str) else "a2a"
)
model_groups.append(
ModelGroupInfoProxy(
model_group=f"a2a/{agent.agent_name}",
mode="chat",
providers=["a2a"],
providers=[custom_llm_provider],
)
)
case _:
@ -76,9 +77,11 @@ async def append_agents_to_model_info(
for agent_id in allowed_agent_ids:
agent = global_agent_registry.get_agent_by_id(agent_id)
if agent is not None:
agent_params = agent.litellm_params
provider_value = agent_params.get("custom_llm_provider") if agent_params else None
custom_llm_provider = provider_value if isinstance(provider_value, str) else "a2a"
agent_params: Final = agent.litellm_params
provider_value: Final = agent_params.get("custom_llm_provider") if agent_params else None
custom_llm_provider: Final = (
provider_value if isinstance(provider_value, str) else "a2a"
)
models.append(
{
"model_name": f"a2a/{agent.agent_name}",

View file

@ -0,0 +1,26 @@
"""Tests for the A2A chat streaming iterator."""
import pytest
from litellm.llms.a2a.chat.streaming_iterator import A2AModelResponseIterator
@pytest.mark.asyncio
async def test_async_iterator_accepts_decoded_a2a_events():
async def _events():
yield {
"jsonrpc": "2.0",
"result": {
"kind": "artifact-update",
"artifact": {"parts": [{"kind": "text", "text": "Hello"}]},
},
}
iterator = A2AModelResponseIterator(
streaming_response=_events(),
sync_stream=False,
)
chunk = await iterator.__aiter__().__anext__()
assert chunk["text"] == "Hello"

View file

@ -64,6 +64,32 @@ async def test_append_agents_to_model_group():
assert result[0].providers == ["a2a"]
@pytest.mark.asyncio
async def test_append_agents_to_model_group_preserves_registered_provider():
agent = AgentResponse(
agent_id="agent-123",
agent_name="test-agent",
agent_card_params={"url": "http://example.com"},
litellm_params={"custom_llm_provider": "pydantic_ai_agents"},
)
registry = Mock()
registry.get_agent_by_id = Mock(return_value=agent)
with (
patch(
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.resolve_agent_access",
AsyncMock(return_value=RestrictedAgentAccess(frozenset({"agent-123"}))),
),
patch("litellm.proxy.agent_endpoints.agent_registry.global_agent_registry", registry),
):
result = await append_agents_to_model_group(
model_groups=[],
user_api_key_dict=Mock(spec=UserAPIKeyAuth),
)
assert result[0].providers == ["pydantic_ai_agents"]
@pytest.mark.asyncio
async def test_append_agents_to_model_info():
"""Test agents are converted to model info format with a2a/ prefix"""

View file

@ -78,11 +78,20 @@ async def test_route_a2a_model_uses_registered_provider():
agent_id="test-agent-id",
agent_name="test-agent",
agent_card_params={"url": "http://agent.example.com"},
litellm_params={"custom_llm_provider": "pydantic_ai_agents"},
litellm_params={
"custom_llm_provider": "pydantic_ai_agents",
"guardrails": ["agent-guardrail"],
},
static_headers={"Authorization": "Bearer static"},
)
data = {
"model": "a2a/test-agent",
"messages": [{"role": "user", "content": "Hello"}],
"guardrails": ["request-guardrail"],
"max_tokens": 32,
"temperature": 0.2,
"timeout": 12.0,
"tools": [{"type": "function", "function": {"name": "lookup"}}],
}
bridge_response = {
"jsonrpc": "2.0",
@ -118,6 +127,55 @@ async def test_route_a2a_model_uses_registered_provider():
bridge.assert_awaited_once()
generic_completion.assert_not_called()
assert response.choices[0].message.content == "Hello back"
bridge_kwargs = bridge.await_args.kwargs
assert bridge_kwargs["litellm_params"]["max_tokens"] == 32
assert bridge_kwargs["litellm_params"]["temperature"] == 0.2
assert bridge_kwargs["litellm_params"]["timeout"] == 12.0
assert bridge_kwargs["litellm_params"]["tools"] == data["tools"]
assert bridge_kwargs["litellm_params"]["guardrails"] == ["request-guardrail", "agent-guardrail"]
assert bridge_kwargs["litellm_params"]["extra_headers"] == {"Authorization": "Bearer static"}
@pytest.mark.asyncio
async def test_route_a2a_cardless_bedrock_agentcore_uses_registered_model():
from litellm.types.agents import AgentResponse
agent = AgentResponse(
agent_id="test-agent-id",
agent_name="test-agent",
agent_card_params={},
litellm_params={
"custom_llm_provider": "bedrock",
"model": "bedrock/agentcore/arn:aws:bedrock-agentcore:us-west-2:123:runtime/test",
},
)
bridge_response = {
"jsonrpc": "2.0",
"id": "request-id",
"result": {"kind": "message", "parts": [{"kind": "text", "text": "Hello back"}]},
}
with (
patch(
"litellm.proxy.common_utils.registry_read_through.get_agent_with_read_through",
AsyncMock(return_value=agent),
),
patch(
"litellm.proxy.agent_endpoints.auth.agent_permission_handler.AgentRequestHandler.is_agent_allowed",
AsyncMock(return_value=True),
),
patch(
"litellm.a2a_protocol.litellm_completion_bridge.handler.A2ACompletionBridgeHandler.handle_non_streaming",
AsyncMock(return_value=bridge_response),
) as bridge,
):
call = await route_a2a_agent_request(
{"model": "a2a/test-agent", "messages": [{"role": "user", "content": "Hello"}]},
"acompletion",
)
await call
assert bridge.await_args.kwargs["api_base"] is None
@pytest.mark.asyncio