feat(search): add Amazon Bedrock AgentCore web search provider

Adds 'agentcore' as a search provider backed by an AgentCore Gateway MCP web-search target, usable from litellm.search()/`/search` and as a websearch_interception backend. Supports SigV4 (AWS_IAM gateways) and bearer tokens (CUSTOM_JWT gateways) via a new BaseSearchConfig.sign_request hook.

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
Devin AI 2026-08-09 01:32:21 +00:00
commit 07f14617f6
10 changed files with 902 additions and 10 deletions

View file

@ -178,6 +178,29 @@ class BaseSearchConfig:
"""
return headers
def sign_request(
self,
headers: dict, # mutable-ok: matches the request header dict every other hook on this base takes
optional_params: dict, # mutable-ok: matches the optional params dict every other hook on this base takes
request_data: dict | list[dict], # mutable-ok: matches transform_search_request's JSON body return type
api_base: str,
api_key: str | None = None,
) -> tuple[dict, bytes | None]: # mutable-ok: the handler passes these headers straight to httpx
"""
OPTIONAL
Sign the request. Providers like Bedrock AgentCore need to SigV4-sign
the request before sending it to the API.
For all other providers, this is a no-op and we just return the headers.
Returns:
Tuple of (headers, signed_json_body). When signed_json_body is not
None, the handler MUST send it verbatim as the request body —
re-serializing the payload would invalidate the signature.
"""
return headers, None
def get_complete_url(
self,
api_base: str | None,

View file

View file

@ -0,0 +1,388 @@
"""
Calls an Amazon Bedrock AgentCore Gateway web-search target (MCP protocol) to search the web.
Web Search on Amazon Bedrock AgentCore exposes Amazon's managed web index through
an AgentCore Gateway MCP endpoint.
AWS docs: https://docs.aws.amazon.com/bedrock-agentcore/latest/devguide/gateway-target-connector-web-search-tool.html
Authentication (matches the gateway's inbound authorizer type):
- AWS_IAM gateway: the request is SigV4-signed. Credentials come from explicit
params (aws_access_key_id / aws_secret_access_key / aws_session_token /
aws_region_name, also settable in a proxy search_tools entry) or the
standard AWS credential chain (env / profile / IRSA / assumed role)
- CUSTOM_JWT gateway: pass the OAuth2 bearer token (e.g. Cognito
client_credentials) as api_key, or set AGENTCORE_GATEWAY_TOKEN
Setup:
1. Create an AgentCore Gateway with a web-search connector target
2. Set AGENTCORE_GATEWAY_URL (or pass api_base) to the gateway MCP endpoint, e.g.
https://<gateway-id>.gateway.bedrock-agentcore.<region>.amazonaws.com/mcp
3. AWS_IAM: ensure the credentials allow bedrock-agentcore:InvokeGateway
CUSTOM_JWT: set AGENTCORE_GATEWAY_TOKEN (or pass api_key)
Usage:
response = litellm.search(
query="latest AI developments",
search_provider="agentcore",
max_results=5,
aws_access_key_id="...", # optional, omit to use the default chain
aws_secret_access_key="...",
)
"""
import json
import re
from collections.abc import Iterator, Mapping, Sequence
from typing import Final
import httpx
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.llms.base_llm.search.transformation import (
BaseSearchConfig,
SearchResponse,
SearchResult,
)
from litellm.llms.bedrock.base_aws_llm import BaseAWSLLM
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.secret_managers.main import get_secret_str
# AgentCore web-search rejects queries longer than 200 characters
AGENTCORE_MAX_QUERY_LENGTH: Final = 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: Final = 10
# Default MCP tool name for a gateway web-search connector target:
# "<target-name>___<tool-name>". Override with AGENTCORE_SEARCH_TOOL_NAME
# or optional_params["tool_name"] when the target uses a custom name.
AGENTCORE_DEFAULT_TOOL_NAME: Final = "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: Final = "___WebSearch"
_GATEWAY_REGION_PATTERN: Final = re.compile(r"\.gateway\.bedrock-agentcore\.([a-z0-9-]+)\.amazonaws\.com")
_SSE_EVENT_SEPARATOR: Final = re.compile(r"\n[ \t]*\n")
_SSE_LINE_PREFIXES: Final = ("event:", "data:", ":", "id:", "retry:")
def _string_field(item: Mapping[str, object], *keys: str) -> str | None:
return next(
(value for key in keys if isinstance(value := item.get(key), str) and value),
None,
)
def _to_search_result(item: Mapping[str, object]) -> SearchResult:
return SearchResult(
title=_string_field(item, "title") or "",
url=_string_field(item, "url") or "",
snippet=_string_field(item, "text", "snippet") or "",
date=_string_field(item, "publishedDate", "date"),
last_updated=None,
)
def _parse_result_items(raw_text: object) -> tuple[Mapping[str, object], ...]:
"""
Parse one MCP text block into the search result objects it carries.
A block holds either a JSON list of results or a {"results": [...]} object;
anything unparseable is skipped rather than failing the whole response.
"""
if not isinstance(raw_text, str):
return ()
try:
parsed: Final = json.loads(raw_text)
except json.JSONDecodeError:
return ()
items: Final = parsed.get("results", ()) if isinstance(parsed, dict) else parsed
if not isinstance(items, Sequence) or isinstance(items, (str, bytes)):
return ()
return tuple(item for item in items if isinstance(item, dict))
def _iter_sse_events(text: str) -> Iterator[Mapping[str, object]]:
"""
Yield the JSON payload of each SSE event in a Streamable HTTP 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.
"""
for chunk in _SSE_EVENT_SEPARATOR.split(text):
payload = "\n".join(
line[len("data:") :].lstrip() for line in chunk.splitlines() if line.startswith("data:")
)
if not payload:
continue
try:
parsed = json.loads(payload)
except json.JSONDecodeError:
continue
if isinstance(parsed, dict):
yield parsed
class AgentCoreSearchConfig(BaseSearchConfig, BaseAWSLLM):
def __init__(self) -> None:
BaseSearchConfig.__init__(self)
BaseAWSLLM.__init__(self)
@staticmethod
def ui_friendly_name() -> str:
return "Web Search on Amazon Bedrock"
def validate_environment(
self,
headers: dict, # mutable-ok: BaseSearchConfig hands providers the mutable request header dict
api_key: str | None = None,
api_base: str | None = None,
**kwargs: object, # kwargs-ok: BaseSearchConfig.validate_environment forwards provider-specific extras
) -> dict: # mutable-ok: the handler passes these headers straight to httpx, which wants a dict
"""
Set MCP transport headers. Per the MCP Streamable HTTP transport spec,
the client MUST accept both application/json and text/event-stream.
Authentication itself happens in sign_request(): bearer token for
CUSTOM_JWT gateways, AWS SigV4 for AWS_IAM gateways.
"""
return { # mutable-ok: httpx request headers are a dict
**headers,
"Content-Type": "application/json",
"Accept": "application/json, text/event-stream",
}
def get_complete_url(
self,
api_base: str | None,
optional_params: dict, # mutable-ok: BaseSearchConfig passes optional params as a dict
data: dict | list[dict] | None = None, # mutable-ok: BaseSearchConfig request bodies are JSON dicts
**kwargs: object, # kwargs-ok: BaseSearchConfig.get_complete_url forwards provider-specific extras
) -> str:
gateway_url: Final = api_base or get_secret_str("AGENTCORE_GATEWAY_URL")
if not gateway_url:
raise ValueError(
"AGENTCORE_GATEWAY_URL is not set. Set it to your AgentCore Gateway MCP "
"endpoint (https://<gateway-id>.gateway.bedrock-agentcore.<region>"
".amazonaws.com/mcp) or pass api_base."
)
return gateway_url
def transform_search_request(
self,
query: str | list[str], # mutable-ok: BaseSearchConfig accepts a list of queries
optional_params: dict, # mutable-ok: BaseSearchConfig passes optional params as a dict
**kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_request forwards provider-specific extras
) -> dict: # mutable-ok: the JSON-RPC body is serialized as a JSON object
"""
Transform Search request to an MCP tools/call request.
Args:
query: Search query (string or list of strings). AgentCore only
supports single string queries; lists are joined with spaces.
optional_params: Optional parameters for the request
- max_results: Maximum number of results (1-25), default 10
- tool_name: Override the MCP tool name of the gateway target
Returns:
Dict with the JSON-RPC 2.0 request body
"""
joined_query: Final = " ".join(query) if isinstance(query, list) else query
tool_name: Final = (
optional_params.get("tool_name")
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."
)
return { # mutable-ok: JSON-RPC request bodies are JSON objects
"jsonrpc": "2.0",
"id": 1,
"method": "tools/call",
"params": { # mutable-ok: JSON-RPC request bodies are JSON objects
"name": tool_name,
"arguments": { # mutable-ok: JSON-RPC request bodies are JSON objects
"query": joined_query[:AGENTCORE_MAX_QUERY_LENGTH],
"maxResults": optional_params.get("max_results", AGENTCORE_DEFAULT_MAX_RESULTS),
},
},
}
def sign_request(
self,
headers: dict, # mutable-ok: BaseSearchConfig hands providers the mutable request header dict
optional_params: dict, # mutable-ok: BaseSearchConfig passes optional params as a dict
request_data: dict | list[dict], # mutable-ok: BaseSearchConfig request bodies are JSON dicts
api_base: str,
api_key: str | None = None,
) -> tuple[dict, bytes | None]: # mutable-ok: BaseSearchConfig.sign_request returns httpx headers
"""
Authenticate the MCP request.
CUSTOM_JWT gateways: attach the caller's OAuth2 bearer token (api_key
or AGENTCORE_GATEWAY_TOKEN), no AWS credentials involved.
AWS_IAM gateways: SigV4-sign with the bedrock-agentcore service name.
"""
if not isinstance(request_data, dict):
raise TypeError("AgentCore search expects a single dict request body")
# 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: Final = 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:
bearer_headers: Final = { # mutable-ok: httpx request headers are a dict
**headers,
"Authorization": f"Bearer {bearer_token}",
}
return bearer_headers, json.dumps(request_data).encode()
signing_params: Final = (
optional_params
if optional_params.get("aws_region_name") is not None
else { # mutable-ok: BaseAWSLLM._sign_request takes optional params as a dict
**optional_params,
"aws_region_name": self._signing_region(api_base),
}
)
# 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="",
)
@staticmethod
def _signing_region(api_base: str) -> str:
"""
Resolve the SigV4 signing region, which must match the gateway's region.
Standard gateway hostnames carry it, so callers don't have to set
aws_region_name to a region different from their default. For custom or
private hostnames, defer to the AWS configuration chain (env vars and
the shared config / profile region), and error out when that yields
nothing rather than silently signing for a guessed region the gateway
would reject with a confusing auth error.
"""
match: Final = _GATEWAY_REGION_PATTERN.search(api_base)
if match:
return match.group(1)
# boto3's session resolution covers env vars AND the AWS shared config
# (profile region), unlike BaseAWSLLM's helper, which silently defaults
# to us-west-2 when nothing is configured.
import boto3
configured_region: Final = boto3.Session().region_name
if configured_region:
return configured_region
raise ValueError(
f"Cannot derive the SigV4 signing region from api_base '{api_base}' "
"or the AWS configuration chain. Set aws_region_name (or AWS_REGION / "
"a profile region) to the gateway's region when using a custom hostname."
)
def transform_search_response(
self,
raw_response: httpx.Response,
logging_obj: LiteLLMLoggingObj,
**kwargs: object, # kwargs-ok: BaseSearchConfig.transform_search_response forwards provider-specific extras
) -> SearchResponse:
"""
Transform an MCP tools/call response to LiteLLM unified SearchResponse.
The gateway returns JSON-RPC (as plain JSON or a single-message SSE
stream) whose result.content[] text blocks contain a JSON list of
{title, url, date/publishedDate, text} entries.
"""
response_json: Final = self._parse_mcp_body(raw_response)
error: Final = response_json.get("error")
if error is not None:
raise BedrockError(
status_code=raw_response.status_code if raw_response.status_code >= 400 else 502,
message=f"AgentCore gateway MCP error: {error}",
)
return SearchResponse(
results=[ # mutable-ok: SearchResponse.results is a pydantic list field
_to_search_result(item)
for block in self._text_blocks(response_json)
for item in _parse_result_items(block.get("text"))
],
object="search",
)
@staticmethod
def _text_blocks(response_json: Mapping[str, object]) -> tuple[Mapping[str, object], ...]:
result: Final = response_json.get("result")
content: Final = result.get("content") if isinstance(result, dict) else None
if not isinstance(content, Sequence) or isinstance(content, (str, bytes)):
return ()
return tuple(block for block in content if isinstance(block, dict) and block.get("type") == "text")
@staticmethod
def _parse_mcp_body(raw_response: httpx.Response) -> Mapping[str, object]:
"""
Parse a JSON or SSE-framed (Streamable HTTP transport) MCP response.
Return the event whose payload carries the JSON-RPC response, i.e. one
containing ``result`` or ``error``, falling back to the last event when
the stream carries only notifications.
"""
text: Final = raw_response.text
if not text.lstrip().startswith(_SSE_LINE_PREFIXES):
return raw_response.json()
events: Final = tuple(_iter_sse_events(text))
response_event: Final = next(
(event for event in events if "result" in event or "error" in event),
None,
)
if response_event is not None:
return response_event
if events:
return events[-1]
raise BedrockError(
status_code=502,
message=f"AgentCore gateway returned SSE without a JSON data frame: {text[:200]}",
)
def get_error_class(
self,
error_message: str,
status_code: int,
headers: dict, # mutable-ok: BaseSearchConfig.get_error_class takes the response headers as a dict
) -> Exception:
return BaseLLMException(
status_code=status_code,
message=error_message,
headers=headers,
)

View file

@ -1740,6 +1740,14 @@ class BaseLLMHTTPHandler:
api_key=api_key,
)
signed_headers, signed_json_body = provider_config.sign_request(
headers=headers,
optional_params=optional_params,
request_data=data,
api_base=complete_url,
api_key=api_key,
)
## LOGGING
logging_obj.pre_call(
input=query if isinstance(query, str) else str(query),
@ -1747,7 +1755,7 @@ class BaseLLMHTTPHandler:
additional_args={
"complete_input_dict": data,
"api_base": complete_url,
"headers": headers,
"headers": signed_headers,
},
)
@ -1763,14 +1771,15 @@ class BaseLLMHTTPHandler:
# Note: timeout is set on the client itself, not per-request for GET
response = client.get(
url=complete_url,
headers=headers,
headers=signed_headers,
)
else:
# Make POST request with JSON data
# A signed body must be sent verbatim, re-serializing it would break the signature
response = client.post(
url=complete_url,
headers=headers,
json=data,
headers=signed_headers,
data=signed_json_body,
json=data if signed_json_body is None else None,
timeout=timeout,
)
except Exception as e:
@ -1824,6 +1833,14 @@ class BaseLLMHTTPHandler:
api_key=api_key,
)
signed_headers, signed_json_body = provider_config.sign_request(
headers=headers,
optional_params=optional_params,
request_data=data,
api_base=complete_url,
api_key=api_key,
)
## LOGGING
logging_obj.pre_call(
input=query if isinstance(query, str) else str(query),
@ -1831,7 +1848,7 @@ class BaseLLMHTTPHandler:
additional_args={
"complete_input_dict": data,
"api_base": complete_url,
"headers": headers,
"headers": signed_headers,
},
)
@ -1852,14 +1869,15 @@ class BaseLLMHTTPHandler:
# Note: timeout is set on the client itself, not per-request for GET
response = await async_httpx_client.get(
url=complete_url,
headers=headers,
headers=signed_headers,
)
else:
# Make async POST request with JSON data
# A signed body must be sent verbatim, re-serializing it would break the signature
response = await async_httpx_client.post(
url=complete_url,
headers=headers,
json=data,
headers=signed_headers,
data=signed_json_body,
json=data if signed_json_body is None else None,
timeout=timeout,
)
except Exception as e:

View file

@ -15721,6 +15721,14 @@
"notes": "APISerpent deep search (/api/search), multi-engine (Google, Bing, Yahoo, DuckDuckGo). Pricing: $0.60/1k searches."
}
},
"agentcore/search": {
"input_cost_per_query": 0.0,
"litellm_provider": "agentcore",
"mode": "search",
"metadata": {
"notes": "Web Search on Amazon Bedrock AgentCore, billed by AWS on the gateway"
}
},
"tinyfish/search": {
"input_cost_per_query": 0.0,
"litellm_provider": "tinyfish",

View file

@ -0,0 +1,39 @@
# Claude Code / Anthropic-native web search on Bedrock, backed by
# Amazon Bedrock AgentCore Web Search (AWS-managed web index, no third-party
# search API). See litellm/llms/bedrock/search/transformation.py for details.
model_list:
- model_name: claude-sonnet
litellm_params:
model: bedrock/us.anthropic.claude-sonnet-4-5-20250929-v1:0
aws_region_name: us-east-1
search_tools:
- search_tool_name: agentcore-search
litellm_params:
search_provider: agentcore
# Your AgentCore Gateway MCP endpoint (gateway must have a `web-search`
# connector target). Alternatively set the AGENTCORE_GATEWAY_URL env var.
api_base: https://<gateway-id>.gateway.bedrock-agentcore.us-east-1.amazonaws.com/mcp
# The gateway exposes the connector as "<target-name>___WebSearch".
# Default is "web-search-tool___WebSearch", matching the target name used
# in the AWS docs' boto3/CLI setup examples. Set this ONLY if your target
# was created with a different name (misconfiguration surfaces as an MCP
# "tool not found" error):
# tool_name: MyWebSearchTarget___WebSearch
# AWS_IAM gateway (default): SigV4-signed. Omit keys to use the standard
# AWS credential chain (env / profile / IRSA / instance role), or set them
# explicitly:
# aws_access_key_id: os.environ/AWS_ACCESS_KEY_ID
# aws_secret_access_key: os.environ/AWS_SECRET_ACCESS_KEY
# CUSTOM_JWT gateway alternative — OAuth2 bearer token instead of SigV4:
# api_key: os.environ/AGENTCORE_GATEWAY_TOKEN
litellm_settings:
callbacks: ["websearch_interception"]
websearch_interception_params:
enabled_providers: ["bedrock"]
search_tool_name: agentcore-search

View file

@ -3749,6 +3749,7 @@ class SearchProviders(str, Enum):
YOU_COM = "you_com"
APISERPENT = "apiserpent"
TINYFISH = "tinyfish"
AGENTCORE = "agentcore"
# Create a set of all search provider values for quick lookup

View file

@ -8935,6 +8935,7 @@ class ProviderConfigManager:
from litellm.llms.apiserpent.search.transformation import (
APISerpentSearchConfig,
)
from litellm.llms.bedrock.search.transformation import AgentCoreSearchConfig
from litellm.llms.brave.search.transformation import BraveSearchConfig
from litellm.llms.dataforseo.search.transformation import DataForSEOSearchConfig
from litellm.llms.duckduckgo.search.transformation import DuckDuckGoSearchConfig
@ -8972,6 +8973,7 @@ class ProviderConfigManager:
SearchProviders.YOU_COM: YouComSearchConfig,
SearchProviders.APISERPENT: APISerpentSearchConfig,
SearchProviders.TINYFISH: TinyfishSearchConfig,
SearchProviders.AGENTCORE: AgentCoreSearchConfig,
}
config_class: Final = PROVIDER_TO_CONFIG_MAP.get(provider, None)
if config_class is None:

View file

@ -15721,6 +15721,14 @@
"notes": "APISerpent deep search (/api/search), multi-engine (Google, Bing, Yahoo, DuckDuckGo). Pricing: $0.60/1k searches."
}
},
"agentcore/search": {
"input_cost_per_query": 0.0,
"litellm_provider": "agentcore",
"mode": "search",
"metadata": {
"notes": "Web Search on Amazon Bedrock AgentCore, billed by AWS on the gateway"
}
},
"tinyfish/search": {
"input_cost_per_query": 0.0,
"litellm_provider": "tinyfish",

View file

@ -0,0 +1,405 @@
"""
Tests for Amazon Bedrock AgentCore Web Search integration.
Mirror of tests/search_tests/test_agentcore_search.py placed in the
test_litellm tree so the AgentCoreSearchConfig transformation is exercised by
the sharded CI (coverage collection runs against this tree).
"""
import json
import os
import pytest
from unittest.mock import AsyncMock, patch, MagicMock
import litellm
from litellm.llms.bedrock.search.transformation import AgentCoreSearchConfig
GATEWAY_URL = "https://testgateway-abc123.gateway.bedrock-agentcore.us-east-1.amazonaws.com/mcp"
MCP_RESULTS = [
{
"title": "Test Result 1",
"url": "https://example.com/1",
"text": "Snippet for result 1",
"publishedDate": "2026-06-16",
},
{
"title": "Test Result 2",
"url": "https://example.com/2",
"text": "Snippet for result 2",
},
]
def _mcp_response_body() -> dict:
return {
"jsonrpc": "2.0",
"id": 1,
"result": {"content": [{"type": "text", "text": json.dumps(MCP_RESULTS)}]},
}
def _make_mock_response(json_body: dict = None, text: str = None) -> MagicMock:
mock_response = MagicMock()
mock_response.status_code = 200
if text is not None:
mock_response.text = text
else:
mock_response.text = json.dumps(json_body)
mock_response.json.return_value = json_body
return mock_response
class TestAgentCoreSearch:
"""
Tests for AgentCore Web Search functionality with mocked network/signing.
"""
@pytest.mark.asyncio
async def test_agentcore_search_request_payload(self):
"""Validates the MCP tools/call payload and SigV4 signing without real AWS calls."""
os.environ["AGENTCORE_GATEWAY_URL"] = GATEWAY_URL
mock_response = _make_mock_response(_mcp_response_body())
with (
patch(
"litellm.llms.custom_httpx.http_handler.AsyncHTTPHandler.post",
new_callable=AsyncMock,
) as mock_post,
patch.object(
AgentCoreSearchConfig,
"_sign_request",
return_value=(
{"Authorization": "AWS4-HMAC-SHA256 test", "Content-Type": "application/json"},
json.dumps({"signed": True}).encode(),
),
) as mock_sign,
):
mock_post.return_value = mock_response
response = await litellm.asearch(
query="latest developments in AI",
search_provider="agentcore",
max_results=5,
)
mock_post.assert_called_once()
call_kwargs = mock_post.call_args.kwargs
assert call_kwargs["url"] == GATEWAY_URL
# Signed body must be sent verbatim
assert call_kwargs["data"] == json.dumps({"signed": True}).encode()
assert call_kwargs["json"] is None
# Signing was invoked with the MCP request
mock_sign.assert_called_once()
sign_kwargs = mock_sign.call_args.kwargs
request_data = sign_kwargs["request_data"]
assert request_data["method"] == "tools/call"
assert request_data["params"]["name"] == "web-search-tool___WebSearch"
assert request_data["params"]["arguments"]["query"] == "latest developments in AI"
assert request_data["params"]["arguments"]["maxResults"] == 5
assert sign_kwargs["service_name"] == "bedrock-agentcore"
assert len(response.results) == 2
assert response.results[0].title == "Test Result 1"
assert response.results[0].url == "https://example.com/1"
assert response.results[0].snippet == "Snippet for result 1"
assert response.results[0].date == "2026-06-16"
def test_transform_search_request_query_truncation(self):
"""AgentCore rejects queries > 200 chars; the request must truncate."""
config = AgentCoreSearchConfig()
long_query = "a" * 300
data = config.transform_search_request(query=long_query, optional_params={})
assert len(data["params"]["arguments"]["query"]) == 200
def test_transform_search_request_joins_list_queries(self):
config = AgentCoreSearchConfig()
data = config.transform_search_request(query=["foo", "bar"], optional_params={})
assert data["params"]["arguments"]["query"] == "foo bar"
def test_transform_search_request_custom_tool_name(self):
config = AgentCoreSearchConfig()
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)
with pytest.raises(ValueError, match="AGENTCORE_GATEWAY_URL"):
config.get_complete_url(api_base=None, optional_params={})
def test_get_complete_url_prefers_api_base(self):
config = AgentCoreSearchConfig()
assert config.get_complete_url(api_base=GATEWAY_URL, optional_params={}) == GATEWAY_URL
def test_validate_environment_sets_mcp_headers(self):
"""MCP Streamable HTTP requires accepting both JSON and SSE."""
config = AgentCoreSearchConfig()
headers = config.validate_environment(headers={})
assert headers["Accept"] == "application/json, text/event-stream"
assert headers["Content-Type"] == "application/json"
def test_transform_search_response_parses_sse_frame(self):
"""Gateway may answer with an SSE-framed JSON-RPC message."""
config = AgentCoreSearchConfig()
body = _mcp_response_body()
sse_text = f"event: message\ndata: {json.dumps(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
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(
{"jsonrpc": "2.0", "id": 1, "error": {"code": -32601, "message": "tool not found"}}
)
with pytest.raises(Exception, match="tool not found"):
config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock())
def test_sign_request_uses_bearer_token_when_api_key_set(self):
"""CUSTOM_JWT gateways: api_key is sent as a bearer token, no SigV4."""
config = AgentCoreSearchConfig()
request_data = {"jsonrpc": "2.0", "id": 1}
headers, signed_body = config.sign_request(
headers={"Content-Type": "application/json"},
optional_params={},
request_data=request_data,
api_base=GATEWAY_URL,
api_key="test-jwt-token",
)
assert headers["Authorization"] == "Bearer test-jwt-token"
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={},
optional_params={},
request_data={"jsonrpc": "2.0"},
api_base=GATEWAY_URL,
)
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):
"""Custom hostname + empty AWS config chain → clear error, no guessed region."""
config = AgentCoreSearchConfig()
mock_session = MagicMock()
mock_session.region_name = None # nothing configured anywhere
with patch("boto3.Session", return_value=mock_session):
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",
)
def test_sign_request_custom_hostname_uses_shared_config_region(self):
"""Custom hostname + region from AWS shared config (profile) must be honored."""
config = AgentCoreSearchConfig()
mock_session = MagicMock()
mock_session.region_name = "eu-west-1" # e.g. from ~/.aws/config profile
with (
patch("boto3.Session", return_value=mock_session),
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="https://gateway.internal.example.com/mcp",
)
assert mock_base_sign.call_args.kwargs["optional_params"]["aws_region_name"] == "eu-west-1"
def test_sign_request_passes_explicit_aws_credentials(self):
"""Explicit aws_* params (e.g. from a proxy search_tools entry) reach the signer."""
config = AgentCoreSearchConfig()
with patch.object(
AgentCoreSearchConfig.__mro__[2], # BaseAWSLLM
"_sign_request",
return_value=({}, b"{}"),
) as mock_base_sign:
config.sign_request(
headers={},
optional_params={
"aws_access_key_id": "AKIATEST",
"aws_secret_access_key": "secret",
"aws_session_token": "token",
},
request_data={"jsonrpc": "2.0"},
api_base=GATEWAY_URL,
)
passed = mock_base_sign.call_args.kwargs["optional_params"]
assert passed["aws_access_key_id"] == "AKIATEST"
assert passed["aws_secret_access_key"] == "secret"
assert passed["aws_session_token"] == "token"
def test_sign_request_derives_region_from_gateway_url(self):
"""Signing region must come from the gateway URL, not the caller's default region."""
config = AgentCoreSearchConfig()
eu_url = "https://gw-x.gateway.bedrock-agentcore.eu-central-1.amazonaws.com/mcp"
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=eu_url,
)
assert mock_base_sign.call_args.kwargs["optional_params"]["aws_region_name"] == "eu-central-1"
class TestAgentCoreSearchEdgeCases:
"""Branch coverage for response parsing and error mapping."""
def test_transform_search_response_skips_non_text_and_bad_json_blocks(self):
"""Non-text blocks and unparseable text blocks are skipped, not fatal."""
config = AgentCoreSearchConfig()
body = {
"jsonrpc": "2.0",
"id": 1,
"result": {
"content": [
{"type": "image", "data": "..."},
{"type": "text", "text": "not-json"},
{"type": "text", "text": json.dumps(["scalar", {"title": "T", "url": "u", "text": "s"}])},
]
},
}
mock_response = _make_mock_response(body)
response = config.transform_search_response(raw_response=mock_response, logging_obj=MagicMock())
# only the one dict item survives; non-dict list entries are skipped
assert len(response.results) == 1
assert response.results[0].title == "T"
def test_parse_mcp_body_sse_without_json_frame_raises(self):
"""An SSE stream carrying no parseable JSON object is a 502."""
config = AgentCoreSearchConfig()
mock_response = _make_mock_response(text="event: ping\ndata: not-json\n\n")
with pytest.raises(Exception, match="SSE without a JSON data frame"):
config._parse_mcp_body(mock_response)
def test_parse_mcp_body_returns_last_event_when_no_result_frame(self):
"""A stream of only notifications returns the last parsed event."""
config = AgentCoreSearchConfig()
note = {"jsonrpc": "2.0", "method": "notifications/progress"}
mock_response = _make_mock_response(text=f"data: {json.dumps(note)}\n\n")
assert config._parse_mcp_body(mock_response) == note
def test_sign_request_rejects_list_request_body(self):
config = AgentCoreSearchConfig()
with pytest.raises(TypeError, match="single dict"):
config.sign_request(
headers={},
optional_params={},
request_data=[{"jsonrpc": "2.0"}],
api_base=GATEWAY_URL,
)
def test_get_error_class_maps_status_and_message(self):
config = AgentCoreSearchConfig()
err = config.get_error_class(error_message="boom", status_code=503, headers={})
assert getattr(err, "status_code", None) == 503
assert "boom" in str(err)
def test_search_cost_lookup_is_mapped(self):
from litellm.search.cost_calculator import search_provider_cost_per_query
assert search_provider_cost_per_query(model="agentcore/search", custom_llm_provider="agentcore") == (0.0, 0.0)