mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat: add Amazon Bedrock AgentCore Web Search as a native search provider
Adds 'agentcore' to SearchProviders, backed by an AgentCore Gateway web-search connector target (MCP tools/call over Streamable HTTP). Web Search on Amazon Bedrock AgentCore is an AWS-managed web index (GA June 2026). Exposing it as a native search provider lets Bedrock users enable Claude Code / Anthropic-native WebSearch through websearch_interception with a pure-YAML config and AWS-native auth, keeping the whole search path inside AWS. Implementation: - New AgentCoreSearchConfig (litellm/llms/bedrock/search/) reusing BaseAWSLLM credential resolution. Auth follows the gateway's inbound authorizer type: AWS_IAM gateways get a SigV4-signed request (explicit aws_access_key_id/aws_secret_access_key params or the default credential chain); CUSTOM_JWT gateways get an OAuth2 bearer token via api_key / AGENTCORE_GATEWAY_TOKEN - SigV4 signing region is derived from the gateway URL so callers don't need aws_region_name to match their default region - Adds an optional sign_request() hook to BaseSearchConfig (no-op by default) and teaches the search HTTP handler to send a signed body verbatim, mirroring the existing anthropic_messages/chat pattern - Handles both plain-JSON and SSE-framed MCP responses, propagates MCP errors, truncates queries to the 200-char gateway limit Tested: - 13 unit tests: payload/signing, explicit AKSK passthrough, bearer token via api_key and env, query truncation, SSE frames, MCP error propagation, region derivation - Verified end-to-end against real AWS_IAM and CUSTOM_JWT gateways, including full Claude Code CLI WebSearch round-trips through the proxy with websearch_interception
This commit is contained in:
parent
5d4c4d0fce
commit
b61484e6c9
8 changed files with 586 additions and 0 deletions
|
|
@ -178,6 +178,29 @@ class BaseSearchConfig:
|
|||
"""
|
||||
return headers
|
||||
|
||||
def sign_request(
|
||||
self,
|
||||
headers: dict,
|
||||
optional_params: dict,
|
||||
request_data: Union[dict, list[dict]],
|
||||
api_base: str,
|
||||
api_key: str | None = None,
|
||||
) -> tuple[dict, bytes | None]:
|
||||
"""
|
||||
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: Optional[str],
|
||||
|
|
|
|||
0
litellm/llms/bedrock/search/__init__.py
Normal file
0
litellm/llms/bedrock/search/__init__.py
Normal file
256
litellm/llms/bedrock/search/transformation.py
Normal file
256
litellm/llms/bedrock/search/transformation.py
Normal file
|
|
@ -0,0 +1,256 @@
|
|||
"""
|
||||
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 typing import Union
|
||||
|
||||
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 = 200
|
||||
|
||||
# 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 = "web-search-tool___WebSearch"
|
||||
|
||||
|
||||
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,
|
||||
api_key: str | None = None,
|
||||
api_base: str | None = None,
|
||||
**kwargs,
|
||||
) -> 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.
|
||||
"""
|
||||
headers["Content-Type"] = "application/json"
|
||||
headers["Accept"] = "application/json, text/event-stream"
|
||||
return headers
|
||||
|
||||
def get_complete_url(
|
||||
self,
|
||||
api_base: str | None,
|
||||
optional_params: dict,
|
||||
data: Union[dict, list[dict]] | None = None,
|
||||
**kwargs,
|
||||
) -> str:
|
||||
api_base = api_base or get_secret_str("AGENTCORE_GATEWAY_URL")
|
||||
if not api_base:
|
||||
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 api_base
|
||||
|
||||
def transform_search_request(
|
||||
self,
|
||||
query: Union[str, list[str]],
|
||||
optional_params: dict,
|
||||
**kwargs,
|
||||
) -> dict:
|
||||
"""
|
||||
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
|
||||
"""
|
||||
if isinstance(query, list):
|
||||
query = " ".join(query)
|
||||
query = query[:AGENTCORE_MAX_QUERY_LENGTH]
|
||||
|
||||
tool_name = (
|
||||
optional_params.get("tool_name")
|
||||
or get_secret_str("AGENTCORE_SEARCH_TOOL_NAME")
|
||||
or AGENTCORE_DEFAULT_TOOL_NAME
|
||||
)
|
||||
|
||||
arguments: dict[str, Union[str, int]] = {"query": query}
|
||||
if "max_results" in optional_params:
|
||||
arguments["maxResults"] = optional_params["max_results"]
|
||||
|
||||
return {
|
||||
"jsonrpc": "2.0",
|
||||
"id": 1,
|
||||
"method": "tools/call",
|
||||
"params": {"name": tool_name, "arguments": arguments},
|
||||
}
|
||||
|
||||
def sign_request(
|
||||
self,
|
||||
headers: dict,
|
||||
optional_params: dict,
|
||||
request_data: Union[dict, list[dict]],
|
||||
api_base: str,
|
||||
api_key: str | None = None,
|
||||
) -> tuple[dict, bytes | None]:
|
||||
"""
|
||||
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 ValueError("AgentCore search expects a single dict request body")
|
||||
|
||||
bearer_token = api_key or get_secret_str("AGENTCORE_GATEWAY_TOKEN")
|
||||
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.
|
||||
signing_params = dict(optional_params)
|
||||
if signing_params.get("aws_region_name") is None:
|
||||
match = re.search(
|
||||
r"\.gateway\.bedrock-agentcore\.([a-z0-9-]+)\.amazonaws\.com",
|
||||
api_base,
|
||||
)
|
||||
if match:
|
||||
signing_params["aws_region_name"] = match.group(1)
|
||||
|
||||
return self._sign_request(
|
||||
service_name="bedrock-agentcore",
|
||||
headers=headers,
|
||||
optional_params=signing_params,
|
||||
request_data=request_data,
|
||||
api_base=api_base,
|
||||
)
|
||||
|
||||
def transform_search_response(
|
||||
self,
|
||||
raw_response: httpx.Response,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
**kwargs,
|
||||
) -> 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 = self._parse_mcp_body(raw_response)
|
||||
|
||||
if "error" in response_json:
|
||||
raise BedrockError(
|
||||
status_code=raw_response.status_code if raw_response.status_code >= 400 else 502,
|
||||
message=f"AgentCore gateway MCP error: {response_json['error']}",
|
||||
)
|
||||
|
||||
results: list[SearchResult] = []
|
||||
for block in response_json.get("result", {}).get("content", []):
|
||||
if block.get("type") != "text":
|
||||
continue
|
||||
try:
|
||||
parsed = json.loads(block["text"])
|
||||
except (json.JSONDecodeError, TypeError):
|
||||
continue
|
||||
items = parsed.get("results", []) if isinstance(parsed, dict) else parsed
|
||||
for item in items:
|
||||
if not isinstance(item, dict):
|
||||
continue
|
||||
results.append(
|
||||
SearchResult(
|
||||
title=item.get("title") or "",
|
||||
url=item.get("url") or "",
|
||||
snippet=item.get("text") or item.get("snippet") or "",
|
||||
date=item.get("publishedDate") or item.get("date"),
|
||||
last_updated=None,
|
||||
)
|
||||
)
|
||||
|
||||
return SearchResponse(results=results, object="search")
|
||||
|
||||
@staticmethod
|
||||
def _parse_mcp_body(raw_response: httpx.Response) -> dict:
|
||||
"""Parse a JSON or SSE-framed (Streamable HTTP transport) MCP response."""
|
||||
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()
|
||||
|
||||
def get_error_class(
|
||||
self,
|
||||
error_message: str,
|
||||
status_code: int,
|
||||
headers: dict,
|
||||
) -> Exception:
|
||||
return BaseLLMException(
|
||||
status_code=status_code,
|
||||
message=error_message,
|
||||
headers=headers,
|
||||
)
|
||||
|
|
@ -1737,6 +1737,15 @@ class BaseLLMHTTPHandler:
|
|||
api_key=api_key,
|
||||
)
|
||||
|
||||
# Sign the request if the provider requires it (e.g. AWS SigV4)
|
||||
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),
|
||||
|
|
@ -1762,6 +1771,14 @@ class BaseLLMHTTPHandler:
|
|||
url=complete_url,
|
||||
headers=headers,
|
||||
)
|
||||
elif signed_json_body is not None:
|
||||
# Send the signed body verbatim — re-serializing would break the signature
|
||||
response = client.post(
|
||||
url=complete_url,
|
||||
headers=headers,
|
||||
data=signed_json_body,
|
||||
timeout=timeout,
|
||||
)
|
||||
else:
|
||||
# Make POST request with JSON data
|
||||
response = client.post(
|
||||
|
|
@ -1821,6 +1838,15 @@ class BaseLLMHTTPHandler:
|
|||
api_key=api_key,
|
||||
)
|
||||
|
||||
# Sign the request if the provider requires it (e.g. AWS SigV4)
|
||||
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),
|
||||
|
|
@ -1851,6 +1877,14 @@ class BaseLLMHTTPHandler:
|
|||
url=complete_url,
|
||||
headers=headers,
|
||||
)
|
||||
elif signed_json_body is not None:
|
||||
# Send the signed body verbatim — re-serializing would break the signature
|
||||
response = await async_httpx_client.post(
|
||||
url=complete_url,
|
||||
headers=headers,
|
||||
data=signed_json_body,
|
||||
timeout=timeout,
|
||||
)
|
||||
else:
|
||||
# Make async POST request with JSON data
|
||||
response = await async_httpx_client.post(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -3489,6 +3489,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
|
||||
|
|
|
|||
|
|
@ -8860,6 +8860,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
|
||||
|
|
@ -8897,6 +8898,7 @@ class ProviderConfigManager:
|
|||
SearchProviders.YOU_COM: YouComSearchConfig,
|
||||
SearchProviders.APISERPENT: APISerpentSearchConfig,
|
||||
SearchProviders.TINYFISH: TinyfishSearchConfig,
|
||||
SearchProviders.AGENTCORE: AgentCoreSearchConfig,
|
||||
}
|
||||
config_class = PROVIDER_TO_CONFIG_MAP.get(provider, None)
|
||||
if config_class is None:
|
||||
|
|
|
|||
231
tests/search_tests/test_agentcore_search.py
Normal file
231
tests/search_tests/test_agentcore_search.py
Normal file
|
|
@ -0,0 +1,231 @@
|
|||
"""
|
||||
Tests for Amazon Bedrock AgentCore Web Search integration.
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, patch, MagicMock
|
||||
|
||||
sys.path.insert(0, os.path.abspath("../.."))
|
||||
|
||||
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 "json" not in call_kwargs
|
||||
|
||||
# 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_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_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):
|
||||
config = AgentCoreSearchConfig()
|
||||
os.environ["AGENTCORE_GATEWAY_TOKEN"] = "env-jwt-token"
|
||||
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)
|
||||
|
||||
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"
|
||||
Loading…
Add table
Reference in a new issue