mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
fix: preserve registered A2A request context
This commit is contained in:
parent
c3d599dd27
commit
02c5c74b37
6 changed files with 253 additions and 19 deletions
|
|
@ -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", {})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue