diff --git a/litellm/llms/bedrock/search/transformation.py b/litellm/llms/bedrock/search/transformation.py index 16d671f26b0..3bae84f0026 100644 --- a/litellm/llms/bedrock/search/transformation.py +++ b/litellm/llms/bedrock/search/transformation.py @@ -51,11 +51,20 @@ from litellm.secret_managers.main import get_secret_str # AgentCore web-search rejects queries longer than 200 characters AGENTCORE_MAX_QUERY_LENGTH = 200 +# The provider contract documents a default of 10 results — send it explicitly +# so the gateway can't silently apply a different default. +AGENTCORE_DEFAULT_MAX_RESULTS = 10 + # Default MCP tool name for a gateway web-search connector target: # "___". Override with AGENTCORE_SEARCH_TOOL_NAME # or optional_params["tool_name"] when the target uses a custom name. AGENTCORE_DEFAULT_TOOL_NAME = "web-search-tool___WebSearch" +# All web-search connector tools share this suffix; rejecting other names keeps +# a caller-supplied tool_name from invoking unrelated tools on the same gateway +# with the proxy's credentials. +AGENTCORE_TOOL_NAME_SUFFIX = "___WebSearch" + class AgentCoreSearchConfig(BaseSearchConfig, BaseAWSLLM): def __init__(self) -> None: @@ -128,10 +137,15 @@ class AgentCoreSearchConfig(BaseSearchConfig, BaseAWSLLM): or get_secret_str("AGENTCORE_SEARCH_TOOL_NAME") or AGENTCORE_DEFAULT_TOOL_NAME ) + if not tool_name.endswith(AGENTCORE_TOOL_NAME_SUFFIX): + raise ValueError( + f"Invalid AgentCore search tool_name '{tool_name}': must end with " + f"'{AGENTCORE_TOOL_NAME_SUFFIX}' (a web-search connector tool). " + "Other gateway tools cannot be invoked through this provider." + ) arguments: dict[str, Union[str, int]] = {"query": query} - if "max_results" in optional_params: - arguments["maxResults"] = optional_params["max_results"] + arguments["maxResults"] = optional_params.get("max_results", AGENTCORE_DEFAULT_MAX_RESULTS) return { "jsonrpc": "2.0", @@ -159,14 +173,28 @@ class AgentCoreSearchConfig(BaseSearchConfig, BaseAWSLLM): if not isinstance(request_data, dict): raise ValueError("AgentCore search expects a single dict request body") - bearer_token = api_key or get_secret_str("AGENTCORE_GATEWAY_TOKEN") + # Server-managed token fallback is gated on the request targeting the + # operator-configured gateway host — otherwise an authenticated caller + # could point api_base at their own server (e.g. via + # /search_tools/test_connection) and receive AGENTCORE_GATEWAY_TOKEN. + bearer_token = self.resolve_server_api_key( + caller_api_key=api_key, + caller_api_base=api_base, + key_env_vars=("AGENTCORE_GATEWAY_TOKEN",), + base_env_var="AGENTCORE_GATEWAY_URL", + default_api_base=None, + ) if bearer_token: headers["Authorization"] = f"Bearer {bearer_token}" return headers, json.dumps(request_data).encode() # The signing region must match the gateway's region — derive it from - # the gateway URL so callers don't have to set aws_region_name to a - # region different from their default. + # standard gateway hostnames so callers don't have to set + # aws_region_name to a region different from their default. Custom or + # private hostnames can't be parsed: fall back to an explicitly + # configured region (param or AWS env vars), and error out rather than + # silently signing for a guessed region the gateway would reject with + # a confusing auth error. signing_params = dict(optional_params) if signing_params.get("aws_region_name") is None: match = re.search( @@ -175,13 +203,23 @@ class AgentCoreSearchConfig(BaseSearchConfig, BaseAWSLLM): ) if match: signing_params["aws_region_name"] = match.group(1) + elif not any(get_secret_str(var) for var in ("AWS_REGION", "AWS_REGION_NAME", "AWS_DEFAULT_REGION")): + raise ValueError( + f"Cannot derive the SigV4 signing region from api_base '{api_base}'. " + "Set aws_region_name (or the AWS_REGION env var) to the gateway's " + "region when using a custom hostname." + ) + # api_key="" (not None, but falsy) disables BaseAWSLLM's fallback to the + # AWS_BEARER_TOKEN_BEDROCK env var: that token is a Bedrock Runtime + # credential and must not be sent to an AgentCore gateway. return self._sign_request( service_name="bedrock-agentcore", headers=headers, optional_params=signing_params, request_data=request_data, api_base=api_base, + api_key="", ) def transform_search_response( @@ -231,17 +269,42 @@ class AgentCoreSearchConfig(BaseSearchConfig, BaseAWSLLM): @staticmethod def _parse_mcp_body(raw_response: httpx.Response) -> dict: - """Parse a JSON or SSE-framed (Streamable HTTP transport) MCP response.""" + """ + Parse a JSON or SSE-framed (Streamable HTTP transport) MCP response. + + Per the SSE spec, an event's data is the concatenation of all its + ``data:`` lines (joined with newlines), and a stream may carry several + events (e.g. progress notifications before the JSON-RPC response). + Return the event whose payload carries the ``id``-matched JSON-RPC + response — i.e. one containing ``result`` or ``error``. + """ text = raw_response.text - if text.lstrip().startswith(("event:", "data:")): - for line in text.splitlines(): - if line.startswith("data:"): - return json.loads(line[len("data:") :].strip()) - raise BedrockError( - status_code=502, - message=f"AgentCore gateway returned SSE without a data frame: {text[:200]}", - ) - return raw_response.json() + if not text.lstrip().startswith(("event:", "data:", ":", "id:", "retry:")): + return raw_response.json() + + last_parsed: dict | None = None + data_lines: list[str] = [] + # Trailing sentinel flushes the final event even without a blank line + for line in text.splitlines() + [""]: + if line.startswith("data:"): + data_lines.append(line[len("data:") :].lstrip()) + continue + if line == "" and data_lines: + try: + parsed = json.loads("\n".join(data_lines)) + except json.JSONDecodeError: + parsed = None + data_lines = [] + if isinstance(parsed, dict): + last_parsed = parsed + if "result" in parsed or "error" in parsed: + return parsed + if last_parsed is not None: + return last_parsed + raise BedrockError( + status_code=502, + message=f"AgentCore gateway returned SSE without a JSON data frame: {text[:200]}", + ) def get_error_class( self, diff --git a/tests/search_tests/test_agentcore_search.py b/tests/search_tests/test_agentcore_search.py index 5d2e4d1c3cd..d5f6ab3e82f 100644 --- a/tests/search_tests/test_agentcore_search.py +++ b/tests/search_tests/test_agentcore_search.py @@ -123,6 +123,18 @@ class TestAgentCoreSearch: data = config.transform_search_request(query="q", optional_params={"tool_name": "my-target___WebSearch"}) assert data["params"]["name"] == "my-target___WebSearch" + def test_transform_search_request_rejects_non_websearch_tool_name(self): + """A caller-supplied tool_name must not reach other tools on the gateway.""" + config = AgentCoreSearchConfig() + with pytest.raises(ValueError, match="must end with"): + config.transform_search_request(query="q", optional_params={"tool_name": "admin-target___DeleteUser"}) + + def test_transform_search_request_sends_documented_default_max_results(self): + """The documented default of 10 is sent explicitly, not left to the gateway.""" + config = AgentCoreSearchConfig() + data = config.transform_search_request(query="q", optional_params={}) + assert data["params"]["arguments"]["maxResults"] == 10 + def test_get_complete_url_requires_gateway_url(self): config = AgentCoreSearchConfig() os.environ.pop("AGENTCORE_GATEWAY_URL", None) @@ -151,6 +163,29 @@ class TestAgentCoreSearch: assert len(response.results) == 2 assert response.results[1].url == "https://example.com/2" + def test_transform_search_response_parses_multiline_sse_data(self): + """SSE data may be split across several data: lines (joined per spec).""" + config = AgentCoreSearchConfig() + pretty = json.dumps(_mcp_response_body(), indent=2) + sse_text = "event: message\n" + "\n".join(f"data: {line}" for line in pretty.splitlines()) + "\n\n" + mock_response = _make_mock_response(text=sse_text) + + response = config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock()) + assert len(response.results) == 2 + + def test_transform_search_response_skips_progress_events(self): + """A progress notification before the JSON-RPC result must not shadow it.""" + config = AgentCoreSearchConfig() + progress = {"jsonrpc": "2.0", "method": "notifications/progress", "params": {"progress": 1}} + sse_text = ( + f"event: message\ndata: {json.dumps(progress)}\n\n" + f"event: message\ndata: {json.dumps(_mcp_response_body())}\n\n" + ) + mock_response = _make_mock_response(text=sse_text) + + response = config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock()) + assert len(response.results) == 2 + def test_transform_search_response_raises_on_mcp_error(self): config = AgentCoreSearchConfig() mock_response = _make_mock_response( @@ -175,8 +210,10 @@ class TestAgentCoreSearch: assert signed_body == json.dumps(request_data).encode() def test_sign_request_uses_bearer_token_from_env(self): + """Server token is attached when the request targets the configured gateway host.""" config = AgentCoreSearchConfig() os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token" + os.environ["AGENTCORE_GATEWAY_URL"] = GATEWAY_URL try: headers, _ = config.sign_request( headers={}, @@ -187,6 +224,61 @@ class TestAgentCoreSearch: assert headers["Authorization"] == "Bearer env-jwt-token" finally: os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None) + os.environ.pop("AGENTCORE_GATEWAY_URL", None) + + def test_sign_request_refuses_server_token_to_untrusted_host(self): + """Server-managed token must not be sent to a caller-chosen api_base.""" + config = AgentCoreSearchConfig() + os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token" + os.environ["AGENTCORE_GATEWAY_URL"] = GATEWAY_URL + try: + with pytest.raises(ValueError, match="Refusing to send"): + config.sign_request( + headers={}, + optional_params={}, + request_data={"jsonrpc": "2.0"}, + api_base="https://attacker.example.com/mcp", + ) + finally: + os.environ.pop("AGENTCORE_GATEWAY_TOKEN", None) + os.environ.pop("AGENTCORE_GATEWAY_URL", None) + + def test_sign_request_does_not_leak_bedrock_bearer_token(self): + """AWS_BEARER_TOKEN_BEDROCK is a Bedrock Runtime credential — it must not + replace SigV4 on requests to an AgentCore gateway.""" + config = AgentCoreSearchConfig() + + with patch.object( + AgentCoreSearchConfig.__mro__[2], # BaseAWSLLM + "_sign_request", + return_value=({}, b"{}"), + ) as mock_base_sign: + config.sign_request( + headers={}, + optional_params={}, + request_data={"jsonrpc": "2.0"}, + api_base=GATEWAY_URL, + ) + # api_key="" (falsy, not None) disables the base class's + # AWS_BEARER_TOKEN_BEDROCK env fallback. + assert mock_base_sign.call_args.kwargs["api_key"] == "" + + def test_sign_request_custom_hostname_requires_region(self): + """Non-standard hostnames can't yield a signing region — require it explicitly.""" + config = AgentCoreSearchConfig() + saved = {var: os.environ.pop(var, None) for var in ("AWS_REGION", "AWS_REGION_NAME", "AWS_DEFAULT_REGION")} + try: + with pytest.raises(ValueError, match="signing region"): + config.sign_request( + headers={}, + optional_params={}, + request_data={"jsonrpc": "2.0"}, + api_base="https://gateway.internal.example.com/mcp", + ) + finally: + for var, val in saved.items(): + if val is not None: + os.environ[var] = val def test_sign_request_passes_explicit_aws_credentials(self): """Explicit aws_* params (e.g. from a proxy search_tools entry) reach the signer."""