diff --git a/litellm/llms/a2a/chat/streaming_iterator.py b/litellm/llms/a2a/chat/streaming_iterator.py index 1983c18a6b3..0cbff95b9d3 100644 --- a/litellm/llms/a2a/chat/streaming_iterator.py +++ b/litellm/llms/a2a/chat/streaming_iterator.py @@ -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", {}) diff --git a/litellm/proxy/agent_endpoints/a2a_routing.py b/litellm/proxy/agent_endpoints/a2a_routing.py index 0373033ee10..2963240f41a 100644 --- a/litellm/proxy/agent_endpoints/a2a_routing.py +++ b/litellm/proxy/agent_endpoints/a2a_routing.py @@ -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 diff --git a/litellm/proxy/agent_endpoints/model_list_helpers.py b/litellm/proxy/agent_endpoints/model_list_helpers.py index 624136c020b..67d721f9faa 100644 --- a/litellm/proxy/agent_endpoints/model_list_helpers.py +++ b/litellm/proxy/agent_endpoints/model_list_helpers.py @@ -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}", diff --git a/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py b/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py new file mode 100644 index 00000000000..579d6f8efef --- /dev/null +++ b/tests/test_litellm/llms/a2a/chat/test_a2a_streaming_iterator.py @@ -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" diff --git a/tests/test_litellm/proxy/agent_endpoints/test_model_list_helpers.py b/tests/test_litellm/proxy/agent_endpoints/test_model_list_helpers.py index b82cd1dd231..6ce22ed1a3f 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_model_list_helpers.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_model_list_helpers.py @@ -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""" diff --git a/tests/test_litellm/proxy/test_route_a2a_models.py b/tests/test_litellm/proxy/test_route_a2a_models.py index 6273fdc0df0..1dd2edc8bf7 100644 --- a/tests/test_litellm/proxy/test_route_a2a_models.py +++ b/tests/test_litellm/proxy/test_route_a2a_models.py @@ -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