mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: [BREAKING]: Introducing Options as TypedDict and Generic (#3140)
* WIP typeddict for options * updated all clients and ChatAgents * updated everything * added ADR * fix mypy * proper typevar imports * fixed import * fixed other imports * slight update in the sample * updated from feedback * fixes * fixed missing covariants and test fixes * fixed typing * updated anthropic thinking config * ruff fixes * fixed int tests * fix tests and mypy * updated integration tests * updated docstring and test fix * improved options handling in obser * mypy fix * updated a host of integration tests * fix tests * bedrock fix
This commit is contained in:
@@ -2,7 +2,7 @@
|
||||
|
||||
import importlib.metadata
|
||||
|
||||
from ._chat_client import BedrockChatClient
|
||||
from ._chat_client import BedrockChatClient, BedrockChatOptions, BedrockGuardrailConfig, BedrockSettings
|
||||
|
||||
try:
|
||||
__version__ = importlib.metadata.version(__name__)
|
||||
@@ -11,5 +11,8 @@ except importlib.metadata.PackageNotFoundError:
|
||||
|
||||
__all__ = [
|
||||
"BedrockChatClient",
|
||||
"BedrockChatOptions",
|
||||
"BedrockGuardrailConfig",
|
||||
"BedrockSettings",
|
||||
"__version__",
|
||||
]
|
||||
|
||||
@@ -2,9 +2,10 @@
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import sys
|
||||
from collections import deque
|
||||
from collections.abc import AsyncIterable, MutableMapping, MutableSequence, Sequence
|
||||
from typing import Any, ClassVar
|
||||
from typing import Any, ClassVar, Generic, Literal, TypedDict
|
||||
from uuid import uuid4
|
||||
|
||||
from agent_framework import (
|
||||
@@ -28,6 +29,7 @@ from agent_framework import (
|
||||
prepare_function_call_results,
|
||||
use_chat_middleware,
|
||||
use_function_invocation,
|
||||
validate_tool_mode,
|
||||
)
|
||||
from agent_framework._pydantic import AFBaseSettings
|
||||
from agent_framework.exceptions import ServiceInitializationError, ServiceInvalidResponseError
|
||||
@@ -37,11 +39,151 @@ from botocore.client import BaseClient
|
||||
from botocore.config import Config as BotoConfig
|
||||
from pydantic import SecretStr, ValidationError
|
||||
|
||||
if sys.version_info >= (3, 13):
|
||||
from typing import TypeVar
|
||||
else:
|
||||
from typing_extensions import TypeVar
|
||||
|
||||
if sys.version_info >= (3, 12):
|
||||
from typing import override # type: ignore # pragma: no cover
|
||||
else:
|
||||
from typing_extensions import override # type: ignore[import] # pragma: no cover
|
||||
|
||||
logger = get_logger("agent_framework.bedrock")
|
||||
|
||||
|
||||
__all__ = [
|
||||
"BedrockChatClient",
|
||||
"BedrockChatOptions",
|
||||
"BedrockGuardrailConfig",
|
||||
"BedrockSettings",
|
||||
]
|
||||
|
||||
|
||||
# region Bedrock Chat Options TypedDict
|
||||
|
||||
|
||||
DEFAULT_REGION = "us-east-1"
|
||||
DEFAULT_MAX_TOKENS = 1024
|
||||
|
||||
|
||||
class BedrockGuardrailConfig(TypedDict, total=False):
|
||||
"""Amazon Bedrock Guardrails configuration.
|
||||
|
||||
See: https://docs.aws.amazon.com/bedrock/latest/userguide/guardrails.html
|
||||
"""
|
||||
|
||||
guardrailIdentifier: str
|
||||
"""The identifier of the guardrail to apply."""
|
||||
|
||||
guardrailVersion: str
|
||||
"""The version of the guardrail to use."""
|
||||
|
||||
trace: Literal["enabled", "disabled"]
|
||||
"""Whether to include guardrail trace information in the response."""
|
||||
|
||||
streamProcessingMode: Literal["sync", "async"]
|
||||
"""How to process guardrails during streaming (sync blocks, async does not)."""
|
||||
|
||||
|
||||
class BedrockChatOptions(ChatOptions, total=False):
|
||||
"""Amazon Bedrock Converse API-specific chat options dict.
|
||||
|
||||
Extends base ChatOptions with Bedrock-specific parameters.
|
||||
Bedrock uses a unified Converse API that works across multiple
|
||||
foundation models (Claude, Titan, Llama, etc.).
|
||||
|
||||
See: https://docs.aws.amazon.com/bedrock/latest/APIReference/API_runtime_Converse.html
|
||||
|
||||
Keys:
|
||||
# Inherited from ChatOptions (mapped to Bedrock):
|
||||
model_id: The Bedrock model identifier,
|
||||
translates to ``modelId`` in Bedrock API.
|
||||
temperature: Sampling temperature,
|
||||
translates to ``inferenceConfig.temperature``.
|
||||
top_p: Nucleus sampling parameter,
|
||||
translates to ``inferenceConfig.topP``.
|
||||
max_tokens: Maximum number of tokens to generate,
|
||||
translates to ``inferenceConfig.maxTokens``.
|
||||
stop: Stop sequences,
|
||||
translates to ``inferenceConfig.stopSequences``.
|
||||
tools: List of tools available to the model,
|
||||
translates to ``toolConfig.tools``.
|
||||
tool_choice: How the model should use tools,
|
||||
translates to ``toolConfig.toolChoice``.
|
||||
|
||||
# Options not supported in Bedrock Converse API:
|
||||
seed: Not supported.
|
||||
frequency_penalty: Not supported.
|
||||
presence_penalty: Not supported.
|
||||
allow_multiple_tool_calls: Not supported (models handle parallel calls automatically).
|
||||
response_format: Not directly supported (use model-specific prompting).
|
||||
user: Not supported.
|
||||
store: Not supported.
|
||||
logit_bias: Not supported.
|
||||
metadata: Not supported (use additional_properties for additionalModelRequestFields).
|
||||
|
||||
# Bedrock-specific options:
|
||||
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
|
||||
guardrailConfig: BedrockGuardrailConfig
|
||||
"""Guardrails configuration for content filtering and safety."""
|
||||
|
||||
performanceConfig: dict[str, Any]
|
||||
"""Performance optimization settings (e.g., latency optimization).
|
||||
See: https://docs.aws.amazon.com/bedrock/latest/userguide/inference-performance.html"""
|
||||
|
||||
requestMetadata: dict[str, str]
|
||||
"""Key-value metadata for the request (max 2048 characters total)."""
|
||||
|
||||
promptVariables: dict[str, dict[str, str]]
|
||||
"""Variables for prompt management when using managed prompts."""
|
||||
|
||||
# ChatOptions fields not supported in Bedrock
|
||||
seed: None # type: ignore[misc]
|
||||
"""Not supported in Bedrock Converse API."""
|
||||
|
||||
frequency_penalty: None # type: ignore[misc]
|
||||
"""Not supported in Bedrock Converse API."""
|
||||
|
||||
presence_penalty: None # type: ignore[misc]
|
||||
"""Not supported in Bedrock Converse API."""
|
||||
|
||||
allow_multiple_tool_calls: None # type: ignore[misc]
|
||||
"""Not supported. Bedrock models handle parallel tool calls automatically."""
|
||||
|
||||
response_format: None # type: ignore[misc]
|
||||
"""Not directly supported. Use model-specific prompting for JSON output."""
|
||||
|
||||
user: None # type: ignore[misc]
|
||||
"""Not supported in Bedrock Converse API."""
|
||||
|
||||
store: None # type: ignore[misc]
|
||||
"""Not supported in Bedrock Converse API."""
|
||||
|
||||
logit_bias: None # type: ignore[misc]
|
||||
"""Not supported in Bedrock Converse API."""
|
||||
|
||||
|
||||
BEDROCK_OPTION_TRANSLATIONS: dict[str, str] = {
|
||||
"model_id": "modelId",
|
||||
"max_tokens": "maxTokens",
|
||||
"top_p": "topP",
|
||||
"stop": "stopSequences",
|
||||
}
|
||||
"""Maps ChatOptions keys to Bedrock Converse API parameter names."""
|
||||
|
||||
TBedrockChatOptions = TypeVar("TBedrockChatOptions", bound=TypedDict, default="BedrockChatOptions", covariant=True) # type: ignore[valid-type]
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
|
||||
ROLE_MAP: dict[Role, str] = {
|
||||
Role.USER: "user",
|
||||
Role.ASSISTANT: "assistant",
|
||||
@@ -74,7 +216,7 @@ class BedrockSettings(AFBaseSettings):
|
||||
@use_function_invocation
|
||||
@use_instrumentation
|
||||
@use_chat_middleware
|
||||
class BedrockChatClient(BaseChatClient):
|
||||
class BedrockChatClient(BaseChatClient[TBedrockChatOptions], Generic[TBedrockChatOptions]):
|
||||
"""Async chat client for Amazon Bedrock's Converse API."""
|
||||
|
||||
OTEL_PROVIDER_NAME: ClassVar[str] = "aws.bedrock" # type: ignore[reportIncompatibleVariableOverride, misc]
|
||||
@@ -106,6 +248,26 @@ class BedrockChatClient(BaseChatClient):
|
||||
env_file_path: Optional .env file path used by ``BedrockSettings`` to load defaults.
|
||||
env_file_encoding: Encoding for the optional .env file.
|
||||
kwargs: Additional arguments forwarded to ``BaseChatClient``.
|
||||
|
||||
Examples:
|
||||
.. code-block:: python
|
||||
|
||||
from agent_framework.bedrock import BedrockChatClient
|
||||
|
||||
# Basic usage with default credentials
|
||||
client = BedrockChatClient(model_id="<model name>")
|
||||
|
||||
# Using custom ChatOptions with type safety:
|
||||
from typing import TypedDict
|
||||
from agent_framework_bedrock import BedrockChatOptions
|
||||
|
||||
|
||||
class MyOptions(BedrockChatOptions, total=False):
|
||||
my_custom_option: str
|
||||
|
||||
|
||||
client = BedrockChatClient[MyOptions](model_id="<model name>")
|
||||
response = await client.get_response("Hello", options={"my_custom_option": "value"})
|
||||
"""
|
||||
try:
|
||||
settings = BedrockSettings(
|
||||
@@ -143,25 +305,27 @@ class BedrockChatClient(BaseChatClient):
|
||||
session_kwargs["aws_session_token"] = settings.session_token.get_secret_value()
|
||||
return Boto3Session(**session_kwargs)
|
||||
|
||||
@override
|
||||
async def _inner_get_response(
|
||||
self,
|
||||
*,
|
||||
messages: MutableSequence[ChatMessage],
|
||||
chat_options: ChatOptions,
|
||||
options: dict[str, Any],
|
||||
**kwargs: Any,
|
||||
) -> ChatResponse:
|
||||
request = self._build_converse_request(messages, chat_options, **kwargs)
|
||||
request = self._prepare_options(messages, options, **kwargs)
|
||||
raw_response = await asyncio.to_thread(self._bedrock_client.converse, **request)
|
||||
return self._process_converse_response(raw_response)
|
||||
|
||||
@override
|
||||
async def _inner_get_streaming_response(
|
||||
self,
|
||||
*,
|
||||
messages: MutableSequence[ChatMessage],
|
||||
chat_options: ChatOptions,
|
||||
options: dict[str, Any],
|
||||
**kwargs: Any,
|
||||
) -> AsyncIterable[ChatResponseUpdate]:
|
||||
response = await self._inner_get_response(messages=messages, chat_options=chat_options, **kwargs)
|
||||
response = await self._inner_get_response(messages=messages, options=options, **kwargs)
|
||||
contents = list(response.messages[0].contents if response.messages else [])
|
||||
if response.usage_details:
|
||||
contents.append(UsageContent(details=response.usage_details))
|
||||
@@ -173,13 +337,13 @@ class BedrockChatClient(BaseChatClient):
|
||||
raw_representation=response.raw_representation,
|
||||
)
|
||||
|
||||
def _build_converse_request(
|
||||
def _prepare_options(
|
||||
self,
|
||||
messages: MutableSequence[ChatMessage],
|
||||
chat_options: ChatOptions,
|
||||
options: dict[str, Any],
|
||||
**kwargs: Any,
|
||||
) -> dict[str, Any]:
|
||||
model_id = chat_options.model_id or self.model_id
|
||||
model_id = options.get("model_id") or self.model_id
|
||||
if not model_id:
|
||||
raise ServiceInitializationError(
|
||||
"Bedrock model_id is required. Set via chat options or BEDROCK_CHAT_MODEL_ID environment variable."
|
||||
@@ -188,40 +352,42 @@ class BedrockChatClient(BaseChatClient):
|
||||
system_prompts, conversation = self._prepare_bedrock_messages(messages)
|
||||
if not conversation:
|
||||
raise ServiceInitializationError("At least one non-system message is required for Bedrock requests.")
|
||||
# Prepend instructions from options if they exist
|
||||
if instructions := options.get("instructions"):
|
||||
system_prompts = [{"text": instructions}, *system_prompts]
|
||||
|
||||
payload: dict[str, Any] = {
|
||||
run_options: dict[str, Any] = {
|
||||
"modelId": model_id,
|
||||
"messages": conversation,
|
||||
"inferenceConfig": {"maxTokens": options.get("max_tokens", DEFAULT_MAX_TOKENS)},
|
||||
}
|
||||
if system_prompts:
|
||||
payload["system"] = system_prompts
|
||||
run_options["system"] = system_prompts
|
||||
|
||||
inference_config: dict[str, Any] = {}
|
||||
inference_config["maxTokens"] = (
|
||||
chat_options.max_tokens if chat_options.max_tokens is not None else DEFAULT_MAX_TOKENS
|
||||
)
|
||||
if chat_options.temperature is not None:
|
||||
inference_config["temperature"] = chat_options.temperature
|
||||
if chat_options.top_p is not None:
|
||||
inference_config["topP"] = chat_options.top_p
|
||||
if chat_options.stop is not None:
|
||||
inference_config["stopSequences"] = chat_options.stop
|
||||
if inference_config:
|
||||
payload["inferenceConfig"] = inference_config
|
||||
if (temperature := options.get("temperature")) is not None:
|
||||
run_options["inferenceConfig"]["temperature"] = temperature
|
||||
if (top_p := options.get("top_p")) is not None:
|
||||
run_options["inferenceConfig"]["topP"] = top_p
|
||||
if (stop := options.get("stop")) is not None:
|
||||
run_options["inferenceConfig"]["stopSequences"] = stop
|
||||
|
||||
tool_config = self._convert_tools_to_bedrock_config(chat_options.tools)
|
||||
if tool_choice := self._convert_tool_choice(chat_options.tool_choice):
|
||||
if tool_config is None:
|
||||
tool_config = {}
|
||||
tool_config["toolChoice"] = tool_choice
|
||||
tool_config = self._prepare_tools(options.get("tools"))
|
||||
if tool_mode := validate_tool_mode(options.get("tool_choice")):
|
||||
tool_config = tool_config or {}
|
||||
match tool_mode.get("mode"):
|
||||
case "auto" | "none":
|
||||
tool_config["toolChoice"] = {tool_mode.get("mode"): {}}
|
||||
case "required":
|
||||
if required_name := tool_mode.get("required_function_name"):
|
||||
tool_config["toolChoice"] = {"tool": {"name": required_name}}
|
||||
else:
|
||||
tool_config["toolChoice"] = {"any": {}}
|
||||
case _:
|
||||
raise ServiceInitializationError(f"Unsupported tool mode for Bedrock: {tool_mode.get('mode')}")
|
||||
if tool_config:
|
||||
payload["toolConfig"] = tool_config
|
||||
run_options["toolConfig"] = tool_config
|
||||
|
||||
if chat_options.additional_properties:
|
||||
payload.update(chat_options.additional_properties)
|
||||
if kwargs:
|
||||
payload.update(kwargs)
|
||||
return payload
|
||||
return run_options
|
||||
|
||||
def _prepare_bedrock_messages(
|
||||
self, messages: Sequence[ChatMessage]
|
||||
@@ -374,12 +540,10 @@ class BedrockChatClient(BaseChatClient):
|
||||
return {"text": str(value)}
|
||||
return {"text": str(value)}
|
||||
|
||||
def _convert_tools_to_bedrock_config(
|
||||
self, tools: list[ToolProtocol | MutableMapping[str, Any]] | None
|
||||
) -> dict[str, Any] | None:
|
||||
def _prepare_tools(self, tools: list[ToolProtocol | MutableMapping[str, Any]] | None) -> dict[str, Any] | None:
|
||||
converted: list[dict[str, Any]] = []
|
||||
if not tools:
|
||||
return None
|
||||
converted: list[dict[str, Any]] = []
|
||||
for tool in tools:
|
||||
if isinstance(tool, MutableMapping):
|
||||
converted.append(dict(tool))
|
||||
@@ -396,24 +560,6 @@ class BedrockChatClient(BaseChatClient):
|
||||
logger.debug("Ignoring unsupported tool type for Bedrock: %s", type(tool))
|
||||
return {"tools": converted} if converted else None
|
||||
|
||||
def _convert_tool_choice(self, tool_choice: Any) -> dict[str, Any] | None:
|
||||
if not tool_choice:
|
||||
return None
|
||||
mode = tool_choice.mode if hasattr(tool_choice, "mode") else str(tool_choice)
|
||||
required_name = getattr(tool_choice, "required_function_name", None)
|
||||
match mode:
|
||||
case "auto":
|
||||
return {"auto": {}}
|
||||
case "none":
|
||||
return {"none": {}}
|
||||
case "required":
|
||||
if required_name:
|
||||
return {"tool": {"name": required_name}}
|
||||
return {"any": {}}
|
||||
case _:
|
||||
logger.debug("Unsupported tool choice mode for Bedrock: %s", mode)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _generate_tool_call_id() -> str:
|
||||
return f"tool-call-{uuid4().hex}"
|
||||
|
||||
Reference in New Issue
Block a user