diff --git a/python/packages/bedrock/agent_framework_bedrock/_chat_client.py b/python/packages/bedrock/agent_framework_bedrock/_chat_client.py index 75362c019ce..ea67ecdb3d4 100644 --- a/python/packages/bedrock/agent_framework_bedrock/_chat_client.py +++ b/python/packages/bedrock/agent_framework_bedrock/_chat_client.py @@ -137,9 +137,10 @@ class BedrockChatOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], t user: Not supported. store: Not supported. logit_bias: Not supported. - metadata: Not supported (use additional_properties for additionalModelRequestFields). + metadata: Not supported (use requestMetadata). # Bedrock-specific options: + additionalModelRequestFields: Model-specific request fields not covered by the Converse API. guardrailConfig: Guardrails configuration for content filtering. performanceConfig: Performance optimization settings. requestMetadata: Key-value metadata for the request. @@ -147,6 +148,10 @@ class BedrockChatOptions(ChatOptions[ResponseModelT], Generic[ResponseModelT], t """ # Bedrock-specific options + additionalModelRequestFields: dict[str, Any] + """Model-specific request fields passed through as ``additionalModelRequestFields`` + (e.g. ``{"reasoning": {"effort": "low"}}``).""" + guardrailConfig: BedrockGuardrailConfig """Guardrails configuration for content filtering and safety.""" @@ -387,6 +392,10 @@ async def _stream() -> AsyncIterable[ChatResponseUpdate]: return self._build_response_stream(_stream(), response_format=options.get("response_format")) # Non-streaming mode + if guardrail_config := request.get("guardrailConfig"): + # streamProcessingMode is only valid for ConverseStream; Converse rejects requests that include it. + request["guardrailConfig"] = {k: v for k, v in guardrail_config.items() if k != "streamProcessingMode"} + async def _get_response() -> ChatResponse: raw_response = await asyncio.to_thread(self._invoke_converse, request) return self._process_converse_response(raw_response, options) @@ -458,6 +467,30 @@ def _prepare_options( if output_config := self._prepare_output_config(options.get("response_format")): run_options["outputConfig"] = output_config + for key in ( + "additionalModelRequestFields", + "guardrailConfig", + "performanceConfig", + "requestMetadata", + "promptVariables", + ): + if (value := options.get(key)) is not None: + run_options[key] = value + if ":prompt/" in model: + # A Prompt Management ARN takes these fields from the prompt, and Converse rejects requests that set them. + if run_options["inferenceConfig"] == {"maxTokens": DEFAULT_MAX_TOKENS}: + del run_options["inferenceConfig"] # only the client default, nothing the caller asked for + if omitted := [ + key + for key in ("inferenceConfig", "system", "toolConfig", "additionalModelRequestFields") + if run_options.pop(key, None) is not None + ]: + logger.warning( + "Converse does not accept %s with a Prompt Management prompt; they are omitted from the request. " + "Define them on the prompt in Prompt Management instead.", + ", ".join(omitted), + ) + return run_options def _prepare_bedrock_messages( diff --git a/python/packages/bedrock/tests/test_bedrock_client.py b/python/packages/bedrock/tests/test_bedrock_client.py index c1667075104..5e25cad9692 100644 --- a/python/packages/bedrock/tests/test_bedrock_client.py +++ b/python/packages/bedrock/tests/test_bedrock_client.py @@ -15,7 +15,7 @@ from boto3.session import Session as Boto3Session from botocore.client import BaseClient -from agent_framework_bedrock import BedrockChatClient, BedrockEmbeddingClient +from agent_framework_bedrock import BedrockChatClient, BedrockChatOptions, BedrockEmbeddingClient from agent_framework_bedrock._chat_client import BedrockSettings from agent_framework_bedrock._feature_usage import FeatureIndex @@ -465,6 +465,78 @@ def test_prepare_options_single_stop_string_becomes_list() -> None: assert request["inferenceConfig"]["stopSequences"] == ["DONE"] +async def test_get_response_forwards_bedrock_specific_options() -> None: + """Bedrock-specific options should reach the Converse request as top-level fields.""" + stub = _StubBedrockRuntime() + client = BedrockChatClient( + model="us.openai.gpt-6-sol", + region="us-east-1", + client=stub, # pyrefly: ignore[bad-argument-type] # ty: ignore[invalid-argument-type] # pyright: ignore[reportArgumentType] + ) + bedrock_options: BedrockChatOptions = { + "additionalModelRequestFields": {"reasoning": {"effort": "low"}}, + "guardrailConfig": {"guardrailIdentifier": "gr-123", "guardrailVersion": "1"}, + "performanceConfig": {"latency": "optimized"}, + "requestMetadata": {"tenant": "contoso"}, + "promptVariables": {"topic": {"text": "hash maps"}}, + } + + await client.get_response( + [Message(role="user", contents=[Content.from_text(text="hello")])], options=bedrock_options + ) + + payload = stub.calls[0] + assert {key: payload.get(key) for key in bedrock_options} == bedrock_options + + +async def test_guardrail_stream_processing_mode_is_sent_only_to_converse_stream() -> None: + """ConverseStream accepts streamProcessingMode and Converse rejects it, so only the Converse request drops it.""" + stub = _StubBedrockStreamRuntime([{"messageStop": {"stopReason": "end_turn"}}]) + client = BedrockChatClient( + model="us.openai.gpt-6-sol", + region="us-east-1", + client=stub, # pyrefly: ignore[bad-argument-type] # ty: ignore[invalid-argument-type] # pyright: ignore[reportArgumentType] + ) + messages = [Message(role="user", contents=[Content.from_text(text="hello")])] + options: BedrockChatOptions = { + "guardrailConfig": {"guardrailIdentifier": "gr-123", "guardrailVersion": "1", "streamProcessingMode": "async"} + } + + await client.get_response(messages, options=options) + stream = client._inner_get_response(messages=messages, options=options, stream=True) + assert isinstance(stream, ResponseStream) + _ = [update async for update in stream] + + assert stub.calls[0]["guardrailConfig"] == {"guardrailIdentifier": "gr-123", "guardrailVersion": "1"} + assert stub.calls[1]["guardrailConfig"] == options["guardrailConfig"] + assert options["guardrailConfig"]["streamProcessingMode"] == "async" + + +def test_prepare_options_prompt_management_arn_omits_fields_converse_rejects(caplog: pytest.LogCaptureFixture) -> None: + """Converse rejects inferenceConfig, system, toolConfig and additionalModelRequestFields with a prompt ARN.""" + client = _make_client() + client.model = "arn:aws:bedrock:us-east-1:123456789012:prompt/PROMPT1234:1" + messages = [Message(role="user", contents=[Content.from_text(text="hello")])] + variables: BedrockChatOptions = {"promptVariables": {"topic": {"text": "hash maps"}}} + + with caplog.at_level("WARNING", logger="agent_framework.bedrock"): + request = client._prepare_options(messages, variables) + assert set(request) == {"modelId", "messages", "promptVariables"} + assert not caplog.records # the client's default maxTokens is dropped without a warning + + caller_set: BedrockChatOptions = { + **variables, + "instructions": "Be brief.", + "temperature": 0.2, + "tools": [{"toolSpec": {"name": "get_weather", "description": "Get weather", "inputSchema": {"json": {}}}}], + "additionalModelRequestFields": {"reasoning": {"effort": "low"}}, + } + with caplog.at_level("WARNING", logger="agent_framework.bedrock"): + request = client._prepare_options(messages, caller_set) + assert set(request) == {"modelId", "messages", "promptVariables"} + assert "inferenceConfig, system, toolConfig, additionalModelRequestFields" in caplog.text + + def test_prepare_options_unsupported_tool_mode_raises(monkeypatch: pytest.MonkeyPatch) -> None: """Unexpected tool modes should raise a clear error.""" from agent_framework_bedrock import _chat_client as chat_client_module