diff --git a/litellm/a2a_protocol/litellm_completion_bridge/handler.py b/litellm/a2a_protocol/litellm_completion_bridge/handler.py index c3d2e41523..53aac1d3e6 100644 --- a/litellm/a2a_protocol/litellm_completion_bridge/handler.py +++ b/litellm/a2a_protocol/litellm_completion_bridge/handler.py @@ -48,20 +48,19 @@ class A2ACompletionBridgeHandler: # Get provider config for custom_llm_provider custom_llm_provider = litellm_params.get("custom_llm_provider") a2a_provider_config = A2AProviderConfigManager.get_provider_config( - custom_llm_provider=custom_llm_provider + custom_llm_provider=custom_llm_provider, + model=litellm_params.get("model"), ) # If provider config exists, use it if a2a_provider_config is not None: - if api_base is None: - raise ValueError(f"api_base is required for {custom_llm_provider}") - verbose_logger.info(f"A2A: Using provider config for {custom_llm_provider}") response_data = await a2a_provider_config.handle_non_streaming( request_id=request_id, params=params, api_base=api_base, + litellm_params=litellm_params, ) return response_data @@ -147,14 +146,12 @@ class A2ACompletionBridgeHandler: # Get provider config for custom_llm_provider custom_llm_provider = litellm_params.get("custom_llm_provider") a2a_provider_config = A2AProviderConfigManager.get_provider_config( - custom_llm_provider=custom_llm_provider + custom_llm_provider=custom_llm_provider, + model=litellm_params.get("model"), ) # If provider config exists, use it if a2a_provider_config is not None: - if api_base is None: - raise ValueError(f"api_base is required for {custom_llm_provider}") - verbose_logger.info( f"A2A: Using provider config for {custom_llm_provider} (streaming)" ) @@ -163,6 +160,7 @@ class A2ACompletionBridgeHandler: request_id=request_id, params=params, api_base=api_base, + litellm_params=litellm_params, ): yield chunk diff --git a/litellm/a2a_protocol/providers/base.py b/litellm/a2a_protocol/providers/base.py index a2354b3495..3ac1cb47fc 100644 --- a/litellm/a2a_protocol/providers/base.py +++ b/litellm/a2a_protocol/providers/base.py @@ -3,7 +3,7 @@ Base configuration for A2A protocol providers. """ from abc import ABC, abstractmethod -from typing import Any, AsyncIterator, Dict +from typing import Any, AsyncIterator, Dict, Optional class BaseA2AProviderConfig(ABC): @@ -19,7 +19,7 @@ class BaseA2AProviderConfig(ABC): self, request_id: str, params: Dict[str, Any], - api_base: str, + api_base: Optional[str] = None, **kwargs, ) -> Dict[str, Any]: """ @@ -41,7 +41,7 @@ class BaseA2AProviderConfig(ABC): self, request_id: str, params: Dict[str, Any], - api_base: str, + api_base: Optional[str] = None, **kwargs, ) -> AsyncIterator[Dict[str, Any]]: """ diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/__init__.py b/litellm/a2a_protocol/providers/bedrock_agentcore/__init__.py new file mode 100644 index 0000000000..a61d8f98b3 --- /dev/null +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/__init__.py @@ -0,0 +1,22 @@ +""" +Bedrock AgentCore A2A provider. + +Preserves JSON-RPC envelopes for AgentCore agents that speak A2A natively, +bypassing the completion bridge that would otherwise strip the envelope. +""" + +from litellm.a2a_protocol.providers.bedrock_agentcore.config import ( + BedrockAgentCoreA2AConfig, +) +from litellm.a2a_protocol.providers.bedrock_agentcore.handler import ( + BedrockAgentCoreA2AHandler, +) +from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( + BedrockAgentCoreA2ATransformation, +) + +__all__ = [ + "BedrockAgentCoreA2AConfig", + "BedrockAgentCoreA2AHandler", + "BedrockAgentCoreA2ATransformation", +] diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/config.py b/litellm/a2a_protocol/providers/bedrock_agentcore/config.py new file mode 100644 index 0000000000..679e19c23c --- /dev/null +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/config.py @@ -0,0 +1,61 @@ +""" +Bedrock AgentCore A2A provider configuration. +""" + +from typing import Any, AsyncIterator, Dict, Optional + +from litellm.a2a_protocol.providers.base import BaseA2AProviderConfig +from litellm.a2a_protocol.providers.bedrock_agentcore.handler import ( + BedrockAgentCoreA2AHandler, +) + + +class BedrockAgentCoreA2AConfig(BaseA2AProviderConfig): + """ + Provider configuration for Bedrock AgentCore A2A-native agents. + + AgentCore agents that speak A2A natively expect the full JSON-RPC envelope. + This config bypasses the completion bridge and forwards requests directly, + deriving the endpoint URL from the model ARN and signing with SigV4/JWT. + """ + + async def handle_non_streaming( + self, + request_id: str, + params: Dict[str, Any], + api_base: Optional[str] = None, + **kwargs, + ) -> Dict[str, Any]: + """Handle non-streaming request to AgentCore A2A agent.""" + litellm_params = kwargs.get("litellm_params") + if not litellm_params: + raise ValueError( + "litellm_params is required for BedrockAgentCoreA2AConfig " + "(must contain model with AgentCore ARN)" + ) + return await BedrockAgentCoreA2AHandler.handle_non_streaming( + request_id=request_id, + params=params, + litellm_params=litellm_params, + ) + + async def handle_streaming( + self, + request_id: str, + params: Dict[str, Any], + api_base: Optional[str] = None, + **kwargs, + ) -> AsyncIterator[Dict[str, Any]]: + """Handle streaming request to AgentCore A2A agent.""" + litellm_params = kwargs.get("litellm_params") + if not litellm_params: + raise ValueError( + "litellm_params is required for BedrockAgentCoreA2AConfig " + "(must contain model with AgentCore ARN)" + ) + async for chunk in BedrockAgentCoreA2AHandler.handle_streaming( + request_id=request_id, + params=params, + litellm_params=litellm_params, + ): + yield chunk diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py b/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py new file mode 100644 index 0000000000..d7445dfc25 --- /dev/null +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py @@ -0,0 +1,134 @@ +""" +Handler for Bedrock AgentCore A2A-native agents. + +Sends JSON-RPC envelopes directly to AgentCore endpoints, bypassing the +completion bridge that would otherwise strip the envelope. +""" + +import json +from typing import Any, AsyncIterator, Dict, cast + +from litellm._logging import verbose_logger +from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( + BedrockAgentCoreA2ATransformation, +) +from litellm.llms.custom_httpx.http_handler import get_async_httpx_client +from litellm.types.llms.custom_http import httpxSpecialProvider + + +class BedrockAgentCoreA2AHandler: + """ + Handler for Bedrock AgentCore A2A requests. + + Constructs JSON-RPC envelopes, signs them via AmazonAgentCoreConfig, + and POSTs directly to the AgentCore endpoint. + """ + + @staticmethod + async def handle_non_streaming( + request_id: str, + params: Dict[str, Any], + litellm_params: Dict[str, Any], + ) -> Dict[str, Any]: + """ + Handle non-streaming A2A request to AgentCore. + + Args: + request_id: A2A JSON-RPC request ID + params: A2A MessageSendParams containing the message + litellm_params: Agent's litellm_params (model, api_key, etc.) + + Returns: + A2A JSON-RPC response dict from the AgentCore agent + """ + url, headers, body = ( + BedrockAgentCoreA2ATransformation.get_url_and_signed_request( + request_id=request_id, + params=params, + litellm_params=litellm_params, + method="message/send", + ) + ) + + verbose_logger.info( + f"BedrockAgentCore A2A: Sending non-streaming request to {url}" + ) + + client = get_async_httpx_client( + llm_provider=cast(Any, httpxSpecialProvider.A2AProvider), + ) + response = await client.post( + url, + headers=headers, + data=body, + ) + response.raise_for_status() + response_data = response.json() + + if "error" in response_data: + verbose_logger.warning( + f"BedrockAgentCore A2A: Agent returned error: {response_data['error']}" + ) + + return response_data + + @staticmethod + async def handle_streaming( + request_id: str, + params: Dict[str, Any], + litellm_params: Dict[str, Any], + ) -> AsyncIterator[Dict[str, Any]]: + """ + Handle streaming A2A request to AgentCore. + + Args: + request_id: A2A JSON-RPC request ID + params: A2A MessageSendParams containing the message + litellm_params: Agent's litellm_params (model, api_key, etc.) + + Yields: + A2A streaming response events from the AgentCore agent + """ + url, headers, body = ( + BedrockAgentCoreA2ATransformation.get_url_and_signed_request( + request_id=request_id, + params=params, + litellm_params=litellm_params, + method="message/send", + stream=True, + ) + ) + + verbose_logger.info( + f"BedrockAgentCore A2A: Sending streaming request to {url}" + ) + + client = get_async_httpx_client( + llm_provider=cast(Any, httpxSpecialProvider.A2AProvider), + ) + response = await client.post( + url, + headers=headers, + data=body, + stream=True, + ) + response.raise_for_status() + + # Check content type — AgentCore may return JSON instead of SSE + content_type = response.headers.get("content-type", "").lower() + + if "application/json" in content_type: + # Single JSON response fallback (not SSE) + verbose_logger.debug( + "BedrockAgentCore A2A streaming: received JSON instead of SSE, " + "yielding as single event" + ) + response_body = await response.aread() + response_data = json.loads(response_body) + yield response_data + else: + # SSE stream — parse data: lines + async for event in BedrockAgentCoreA2ATransformation.parse_sse_events( + response + ): + yield event diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py b/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py new file mode 100644 index 0000000000..44dc10fe2b --- /dev/null +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/transformation.py @@ -0,0 +1,134 @@ +""" +Transformation layer for Bedrock AgentCore A2A provider. + +Constructs JSON-RPC envelopes, derives AgentCore URLs from model ARNs, +and signs requests via AmazonAgentCoreConfig (SigV4 or JWT). +""" + +import json +from typing import Any, AsyncIterator, Dict, Tuple + +from litellm._logging import verbose_logger +from litellm.llms.bedrock.chat.agentcore.transformation import AmazonAgentCoreConfig + + +class BedrockAgentCoreA2ATransformation: + """ + Request/response transformation for Bedrock AgentCore A2A agents. + + Reuses AmazonAgentCoreConfig for URL construction, ARN parsing, + and request signing. No logic is duplicated. + """ + + @staticmethod + def get_url_and_signed_request( + request_id: str, + params: Dict[str, Any], + litellm_params: Dict[str, Any], + method: str = "message/send", + stream: bool = False, + ) -> Tuple[str, dict, bytes]: + """ + Build the AgentCore URL, construct a JSON-RPC envelope, and sign the request. + + Args: + request_id: A2A JSON-RPC request ID + params: A2A MessageSendParams + litellm_params: Agent's litellm_params (model, api_key, etc.) + method: JSON-RPC method name (default: "message/send") + stream: Whether this is a streaming request + + Returns: + Tuple of (url, signed_headers, signed_body_bytes) + """ + # Extract model and strip the "bedrock/" prefix + # "bedrock/agentcore/arn:aws:..." → "agentcore/arn:aws:..." + model = litellm_params.get("model", "") + if model.startswith("bedrock/"): + agentcore_model = model[len("bedrock/") :] + else: + agentcore_model = model + + # Build optional_params from litellm_params (everything except model and custom_llm_provider) + optional_params = { + k: v + for k, v in litellm_params.items() + if k not in ("model", "custom_llm_provider") + } + + agentcore_config = AmazonAgentCoreConfig() + + # Derive URL from ARN + url = agentcore_config.get_complete_url( + api_base=optional_params.get("api_base"), + api_key=optional_params.get("api_key"), + model=agentcore_model, + optional_params=optional_params, + litellm_params=litellm_params, + stream=stream, + ) + + # Construct JSON-RPC 2.0 envelope + json_rpc_body = { + "jsonrpc": "2.0", + "method": method, + "id": request_id, + "params": params, + } + + # Set required AgentCore session headers (normally set by transform_request, + # which we skip because it also builds {"prompt": "..."}) + headers: dict = {} + session_id = agentcore_config._get_runtime_session_id(optional_params) + headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] = session_id + runtime_user_id = agentcore_config._get_runtime_user_id(optional_params) + if runtime_user_id: + headers["X-Amzn-Bedrock-AgentCore-Runtime-User-Id"] = runtime_user_id + + # Sign the request (SigV4 or JWT depending on api_key presence) + signed_headers, signed_body = agentcore_config.sign_request( + headers=headers, + optional_params=optional_params, + request_data=json_rpc_body, + api_base=url, + api_key=optional_params.get("api_key"), + model=agentcore_model, + stream=stream, + ) + + # sign_request returns Optional[bytes] — ensure we have bytes + if signed_body is None: + signed_body = json.dumps(json_rpc_body).encode() + + return url, signed_headers, signed_body + + @staticmethod + async def parse_sse_events(response: Any) -> AsyncIterator[Dict[str, Any]]: + """ + Parse SSE events from an httpx streaming response. + + Reads line-by-line, parses `data:` lines as JSON, and yields each parsed dict. + + Args: + response: httpx streaming response + + Yields: + Parsed JSON dicts from SSE data lines + """ + async for line in response.aiter_lines(): + line = line.strip() + if not line: + continue + + if line.startswith("data:"): + data_str = line[len("data:") :].strip() + if not data_str: + continue + try: + event = json.loads(data_str) + yield event + except json.JSONDecodeError: + verbose_logger.debug( + f"BedrockAgentCore A2A: Skipping non-JSON SSE line: {data_str[:100]}" + ) + continue diff --git a/litellm/a2a_protocol/providers/config_manager.py b/litellm/a2a_protocol/providers/config_manager.py index a8b9566c17..d684efd475 100644 --- a/litellm/a2a_protocol/providers/config_manager.py +++ b/litellm/a2a_protocol/providers/config_manager.py @@ -19,12 +19,14 @@ class A2AProviderConfigManager: @staticmethod def get_provider_config( custom_llm_provider: Optional[str], + model: Optional[str] = None, ) -> Optional[BaseA2AProviderConfig]: """ Get the provider configuration for a given custom_llm_provider. Args: custom_llm_provider: The provider identifier (e.g., "pydantic_ai_agents") + model: The model string (used to distinguish sub-providers, e.g. agentcore vs other bedrock) Returns: Provider configuration instance or None if not found @@ -39,9 +41,11 @@ class A2AProviderConfigManager: return PydanticAIProviderConfig() - # Add more providers here as needed - # elif custom_llm_provider == "another_provider": - # from litellm.a2a_protocol.providers.another_provider.config import AnotherProviderConfig - # return AnotherProviderConfig() + if custom_llm_provider == "bedrock" and model and "agentcore" in model: + from litellm.a2a_protocol.providers.bedrock_agentcore.config import ( + BedrockAgentCoreA2AConfig, + ) + + return BedrockAgentCoreA2AConfig() return None diff --git a/tests/test_litellm/a2a_protocol/providers/__init__.py b/tests/test_litellm/a2a_protocol/providers/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/__init__.py b/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/__init__.py new file mode 100644 index 0000000000..e69de29bb2 diff --git a/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py b/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py new file mode 100644 index 0000000000..f21faecaa2 --- /dev/null +++ b/tests/test_litellm/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py @@ -0,0 +1,327 @@ +""" +Tests for Bedrock AgentCore A2A provider. + +Verifies that: +- JSON-RPC envelopes are preserved (not stripped by the completion bridge) +- URLs are derived from the model ARN +- Auth uses JWT Bearer or SigV4 +- Config manager routes "bedrock" correctly +- Handler passes litellm_params and allows api_base=None +""" + +import json + +import pytest +from unittest.mock import AsyncMock, MagicMock, patch + + +SAMPLE_ARN = "arn:aws:bedrock-agentcore:us-west-2:123456789:runtime/my_agent" +SAMPLE_MODEL = f"bedrock/agentcore/{SAMPLE_ARN}" +SAMPLE_PARAMS = { + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "what is 1+1?"}], + "messageId": "msg-001", + } +} +SAMPLE_LITELLM_PARAMS = { + "model": SAMPLE_MODEL, + "custom_llm_provider": "bedrock", + "api_key": "test-jwt-token", +} + + +class TestTransformation: + """Test URL construction and JSON-RPC envelope building.""" + + def test_json_rpc_envelope_structure(self): + """Verify JSON-RPC body has jsonrpc, method, id, and params.""" + from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( + BedrockAgentCoreA2ATransformation, + ) + + url, headers, body = ( + BedrockAgentCoreA2ATransformation.get_url_and_signed_request( + request_id="req-001", + params=SAMPLE_PARAMS, + litellm_params=SAMPLE_LITELLM_PARAMS, + method="message/send", + ) + ) + body_dict = json.loads(body) + assert body_dict["jsonrpc"] == "2.0" + assert body_dict["method"] == "message/send" + assert body_dict["id"] == "req-001" + assert body_dict["params"] == SAMPLE_PARAMS + + def test_url_derived_from_arn(self): + """Verify URL is constructed from the ARN, not from api_base.""" + from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( + BedrockAgentCoreA2ATransformation, + ) + + url, _, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request( + request_id="req-001", + params=SAMPLE_PARAMS, + litellm_params=SAMPLE_LITELLM_PARAMS, + ) + assert "bedrock-agentcore.us-west-2.amazonaws.com" in url + assert "/runtimes/" in url + assert "/invocations" in url + + def test_jwt_auth_uses_bearer_header(self): + """When api_key is set, Authorization header uses Bearer token.""" + from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( + BedrockAgentCoreA2ATransformation, + ) + + _, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request( + request_id="req-001", + params=SAMPLE_PARAMS, + litellm_params=SAMPLE_LITELLM_PARAMS, + ) + assert headers["Authorization"] == "Bearer test-jwt-token" + + def test_session_id_header_set(self): + """Verify X-Amzn-Bedrock-AgentCore-Runtime-Session-Id is set.""" + from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( + BedrockAgentCoreA2ATransformation, + ) + + _, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request( + request_id="req-001", + params=SAMPLE_PARAMS, + litellm_params=SAMPLE_LITELLM_PARAMS, + ) + session_id = headers.get("X-Amzn-Bedrock-AgentCore-Runtime-Session-Id", "") + assert len(session_id) >= 33 + + def test_custom_session_id_header(self): + """Verify custom runtimeSessionId is used when provided.""" + from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( + BedrockAgentCoreA2ATransformation, + ) + + params_with_session = {**SAMPLE_LITELLM_PARAMS, "runtimeSessionId": "a" * 40} + _, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request( + request_id="req-001", + params=SAMPLE_PARAMS, + litellm_params=params_with_session, + ) + assert headers["X-Amzn-Bedrock-AgentCore-Runtime-Session-Id"] == "a" * 40 + + def test_sigv4_auth_when_no_api_key(self): + """When no api_key, falls through to SigV4 signing.""" + from litellm.a2a_protocol.providers.bedrock_agentcore.transformation import ( + BedrockAgentCoreA2ATransformation, + ) + + litellm_params_no_key = { + "model": SAMPLE_MODEL, + "custom_llm_provider": "bedrock", + "aws_access_key_id": "AKIAIOSFODNN7EXAMPLE", + "aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG/bPxRfiCYEXAMPLEKEY", + "aws_region_name": "us-west-2", + } + + # Mock _sign_request to avoid hitting real botocore credential resolution + fake_sigv4_headers = { + "Authorization": "AWS4-HMAC-SHA256 Credential=AKIA.../bedrock-agentcore/aws4_request", + "Content-Type": "application/json", + "Accept": "application/json, text/event-stream", + } + fake_body = b'{"jsonrpc":"2.0"}' + + with patch( + "litellm.llms.bedrock.chat.agentcore.transformation.AmazonAgentCoreConfig._sign_request", + return_value=(fake_sigv4_headers, fake_body), + ): + _, headers, _ = BedrockAgentCoreA2ATransformation.get_url_and_signed_request( + request_id="req-001", + params=SAMPLE_PARAMS, + litellm_params=litellm_params_no_key, + ) + # SigV4 produces an Authorization header starting with "AWS4-HMAC-SHA256" + assert "Authorization" in headers + assert headers["Authorization"].startswith("AWS4-HMAC-SHA256") + + +class TestNonStreaming: + """Test end-to-end non-streaming flow.""" + + @pytest.mark.asyncio + async def test_json_rpc_body_sent_to_agentcore(self): + """Verify the full JSON-RPC envelope is POSTed, not {"prompt": "..."}.""" + from litellm.a2a_protocol.providers.bedrock_agentcore.config import ( + BedrockAgentCoreA2AConfig, + ) + + mock_response = MagicMock() + mock_response.json.return_value = { + "jsonrpc": "2.0", + "id": "req-001", + "result": { + "message": { + "role": "agent", + "parts": [{"kind": "text", "text": "2"}], + "messageId": "resp-001", + } + }, + } + mock_response.raise_for_status = MagicMock() + + with patch( + "litellm.a2a_protocol.providers.bedrock_agentcore.handler.get_async_httpx_client" + ) as mock_get_client: + mock_client = AsyncMock() + mock_client.post = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_client + + config = BedrockAgentCoreA2AConfig() + result = await config.handle_non_streaming( + request_id="req-001", + params=SAMPLE_PARAMS, + litellm_params=SAMPLE_LITELLM_PARAMS, + ) + + # Verify the POST was called + mock_client.post.assert_called_once() + call_kwargs = mock_client.post.call_args + + # Verify sent body is JSON-RPC, not {"prompt": "..."} + sent_body = json.loads(call_kwargs.kwargs["data"]) + assert "jsonrpc" in sent_body + assert "method" in sent_body + assert sent_body["method"] == "message/send" + assert sent_body["params"]["message"]["parts"][0]["text"] == "what is 1+1?" + + # Verify response is passed through + assert result["result"]["message"]["parts"][0]["text"] == "2" + + @pytest.mark.asyncio + async def test_a2a_error_response_passthrough(self): + """JSON-RPC error responses from the agent are returned as-is.""" + from litellm.a2a_protocol.providers.bedrock_agentcore.config import ( + BedrockAgentCoreA2AConfig, + ) + + error_response = { + "jsonrpc": "2.0", + "id": "req-001", + "error": {"code": -32600, "message": "Bad request"}, + } + mock_response = MagicMock() + mock_response.json.return_value = error_response + mock_response.raise_for_status = MagicMock() + + with patch( + "litellm.a2a_protocol.providers.bedrock_agentcore.handler.get_async_httpx_client" + ) as mock_get_client: + mock_client = AsyncMock() + mock_client.post = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_client + + config = BedrockAgentCoreA2AConfig() + result = await config.handle_non_streaming( + request_id="req-001", + params=SAMPLE_PARAMS, + litellm_params=SAMPLE_LITELLM_PARAMS, + ) + + assert result["error"]["code"] == -32600 + assert result["error"]["message"] == "Bad request" + + +class TestConfigManager: + """Test that config manager routes 'bedrock' correctly.""" + + def test_bedrock_returns_config(self): + from litellm.a2a_protocol.providers.bedrock_agentcore.config import ( + BedrockAgentCoreA2AConfig, + ) + from litellm.a2a_protocol.providers.config_manager import ( + A2AProviderConfigManager, + ) + + config = A2AProviderConfigManager.get_provider_config( + "bedrock", model=SAMPLE_MODEL + ) + assert config is not None + assert isinstance(config, BedrockAgentCoreA2AConfig) + + def test_bedrock_non_agentcore_returns_none(self): + """Non-agentcore bedrock models should fall through to completion bridge.""" + from litellm.a2a_protocol.providers.config_manager import ( + A2AProviderConfigManager, + ) + + config = A2AProviderConfigManager.get_provider_config( + "bedrock", model="bedrock/anthropic.claude-3-sonnet" + ) + assert config is None + + def test_unknown_provider_returns_none(self): + from litellm.a2a_protocol.providers.config_manager import ( + A2AProviderConfigManager, + ) + + assert A2AProviderConfigManager.get_provider_config("unknown") is None + + +class TestHandlerIntegration: + """Test handler.py changes — litellm_params passed through, api_base not required.""" + + @pytest.mark.asyncio + async def test_provider_config_receives_litellm_params(self): + """Verify handler passes litellm_params to provider config via kwargs.""" + from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2ACompletionBridgeHandler, + ) + + mock_config = AsyncMock() + mock_config.handle_non_streaming = AsyncMock( + return_value={"jsonrpc": "2.0", "id": "req-001", "result": {}} + ) + + with patch( + "litellm.a2a_protocol.litellm_completion_bridge.handler.A2AProviderConfigManager.get_provider_config", + return_value=mock_config, + ): + await A2ACompletionBridgeHandler.handle_non_streaming( + request_id="req-001", + params=SAMPLE_PARAMS, + litellm_params=SAMPLE_LITELLM_PARAMS, + api_base=None, + ) + + mock_config.handle_non_streaming.assert_called_once_with( + request_id="req-001", + params=SAMPLE_PARAMS, + api_base=None, + litellm_params=SAMPLE_LITELLM_PARAMS, + ) + + @pytest.mark.asyncio + async def test_api_base_none_allowed_with_provider_config(self): + """api_base=None no longer raises when a provider config is registered.""" + from litellm.a2a_protocol.litellm_completion_bridge.handler import ( + A2ACompletionBridgeHandler, + ) + + mock_config = AsyncMock() + mock_config.handle_non_streaming = AsyncMock( + return_value={"jsonrpc": "2.0", "id": "req-001", "result": {}} + ) + + with patch( + "litellm.a2a_protocol.litellm_completion_bridge.handler.A2AProviderConfigManager.get_provider_config", + return_value=mock_config, + ): + # Should NOT raise ValueError + result = await A2ACompletionBridgeHandler.handle_non_streaming( + request_id="req-001", + params=SAMPLE_PARAMS, + litellm_params=SAMPLE_LITELLM_PARAMS, + api_base=None, + ) + assert result is not None