Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -137,16 +137,21 @@ 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.
promptVariables: Variables for prompt management (if using managed prompts).
"""

# 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."""

Expand Down Expand Up @@ -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)
Expand Down Expand Up @@ -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",
Comment thread
kimnamu marked this conversation as resolved.
"performanceConfig",
"requestMetadata",
"promptVariables",
Comment thread
kimnamu marked this conversation as resolved.
):
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(
Expand Down
74 changes: 73 additions & 1 deletion python/packages/bedrock/tests/test_bedrock_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand Down
Loading