diff --git a/python/packages/bedrock/AGENTS.md b/python/packages/bedrock/AGENTS.md index 00245229f52..ec3110c2fc1 100644 --- a/python/packages/bedrock/AGENTS.md +++ b/python/packages/bedrock/AGENTS.md @@ -8,6 +8,8 @@ Integration with AWS Bedrock for LLM inference. - **`BedrockChatOptions`** - Options TypedDict for Bedrock-specific parameters - **`BedrockGuardrailConfig`** - Configuration for Bedrock guardrails - **`BedrockSettings`** - Pydantic settings for Bedrock configuration +- **`BedrockKnowledgeBaseTool`** - `FunctionTool` for retrieving from an Amazon Bedrock Knowledge Base (agentic retrieval with fallback to standard Retrieve) +- **`BedrockKnowledgeBaseProvider`** - `ContextProvider` that injects Knowledge Base passages before each agent run ## Usage diff --git a/python/packages/bedrock/BEDROCK_MANAGED_KB.md b/python/packages/bedrock/BEDROCK_MANAGED_KB.md new file mode 100644 index 00000000000..8542388a7b2 --- /dev/null +++ b/python/packages/bedrock/BEDROCK_MANAGED_KB.md @@ -0,0 +1,67 @@ +# Bedrock Managed Knowledge Base Support + +## Overview +Adds an Agent Framework tool that queries Amazon Bedrock Knowledge Bases for managed retrieval within agent pipelines. + +## Usage +```python +from agent_framework import Agent +from agent_framework_bedrock import BedrockKnowledgeBaseTool, BedrockChatClient, BedrockChatOptions + +tool = BedrockKnowledgeBaseTool( + knowledge_base_id="YOUR_KB_ID", + region_name="us-east-1", +) + +# As a FunctionTool, pass directly to an Agent: +agent = Agent(client=BedrockChatClient(options=BedrockChatOptions(model_id="...")), tools=[tool]) + +# Or invoke directly for testing: +import asyncio +result = asyncio.run(tool.invoke(arguments={"query": "What are the compliance requirements?"})) +print(result) # List of Content items with retrieval results +``` + +## Configuration + +All configuration is via constructor parameters: + +| Parameter | Description | Default | +|---|---|---| +| `knowledge_base_id` | Bedrock Knowledge Base ID (required) | — | +| `region_name` | AWS region for the KB | `us-east-1` | +| `number_of_results` | Maximum retrieval results | `5` | +| `use_agentic_retrieval` | Enable agentic multi-hop retrieval | `True` | +| `client` | Pre-configured boto3 client (optional) | Auto-created | + +## Features +- Managed search (no vector store needed) +- **BedrockKnowledgeBaseTool**: Agentic retrieval with query decomposition + reranking, automatic fallback to standard Retrieve +- **BedrockKnowledgeBaseProvider**: Standard managed retrieval injected as context before each agent run +- Multi-source support (S3, Web, Confluence, SharePoint) +- Compatible with Agent Framework FunctionTool and ContextProvider interfaces + +## SDK Requirements +- boto3 >= 1.43.32 + +## Required IAM Permissions +```json +{ + "Version": "2012-10-17", + "Statement": [ + { + "Effect": "Allow", + "Action": [ + "bedrock:Retrieve", + "bedrock:AgenticRetrieveStream" + ], + "Resource": "arn:aws:bedrock:::knowledge-base/" + } + ] +} +``` + +## References +- [Build a Managed Knowledge Base](https://docs.aws.amazon.com/bedrock/latest/userguide/kb-build-managed.html) +- [Retrieve API](https://docs.aws.amazon.com/bedrock/latest/userguide/kb-test-retrieve.html) +- [Agentic Retrieval](https://docs.aws.amazon.com/bedrock/latest/userguide/kb-test-agentic.html) diff --git a/python/packages/bedrock/agent_framework_bedrock/__init__.py b/python/packages/bedrock/agent_framework_bedrock/__init__.py index 3fbf5c15cf5..a1948280e63 100644 --- a/python/packages/bedrock/agent_framework_bedrock/__init__.py +++ b/python/packages/bedrock/agent_framework_bedrock/__init__.py @@ -4,6 +4,8 @@ from ._chat_client import BedrockChatClient, BedrockChatOptions, BedrockGuardrailConfig, BedrockSettings from ._embedding_client import BedrockEmbeddingClient, BedrockEmbeddingOptions, BedrockEmbeddingSettings +from ._knowledge_base import BedrockKnowledgeBaseTool +from ._knowledge_base_provider import BedrockKnowledgeBaseProvider try: __version__ = importlib.metadata.version(__name__) @@ -17,6 +19,8 @@ "BedrockEmbeddingOptions", "BedrockEmbeddingSettings", "BedrockGuardrailConfig", + "BedrockKnowledgeBaseProvider", + "BedrockKnowledgeBaseTool", "BedrockSettings", "__version__", ] diff --git a/python/packages/bedrock/agent_framework_bedrock/_chat_client.py b/python/packages/bedrock/agent_framework_bedrock/_chat_client.py index c38b813fdae..b88dbcfa540 100644 --- a/python/packages/bedrock/agent_framework_bedrock/_chat_client.py +++ b/python/packages/bedrock/agent_framework_bedrock/_chat_client.py @@ -469,6 +469,10 @@ def _prepare_bedrock_messages( prompts: list[dict[str, str]] = [] conversation: list[dict[str, Any]] = [] pending_tool_use_ids: deque[str] = deque() + # Track the original role of the last appended conversation turn so we only + # coalesce genuine user-role messages (see below), never tool/system turns + # that merely map to the Bedrock "user" role. + last_appended_role: str | None = None for message in messages: if message.role == "system": text_value = message.text @@ -495,7 +499,23 @@ def _prepare_bedrock_messages( else: pending_tool_use_ids.clear() - conversation.append({"role": role, "content": content_blocks}) + # Coalesce adjacent genuine user-role turns only. Context providers + # (e.g. the Bedrock Knowledge Base provider) inject retrieved passages as + # separate user messages that would otherwise sit next to the real user + # input and violate Bedrock's role-alternation requirement. We restrict + # this to messages whose ORIGINAL role is "user" so that tool-result turns + # (message.role == "tool", which also map to the Bedrock "user" role) are + # never merged — preserving function-call/tool-result serialization. + if ( + message.role == "user" + and last_appended_role == "user" + and conversation + and conversation[-1]["role"] == "user" + ): + conversation[-1]["content"].extend(content_blocks) + else: + conversation.append({"role": role, "content": content_blocks}) + last_appended_role = message.role return prompts, conversation diff --git a/python/packages/bedrock/agent_framework_bedrock/_knowledge_base.py b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base.py new file mode 100644 index 00000000000..8976ccb5db6 --- /dev/null +++ b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base.py @@ -0,0 +1,205 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Amazon Bedrock Knowledge Base retrieval tool for Agent Framework.""" + +from __future__ import annotations + +import asyncio +import logging +from typing import TYPE_CHECKING, Annotated, Any, Optional + +from agent_framework import FunctionTool +from agent_framework._telemetry import get_user_agent, mark_feature_used +from pydantic import BaseModel, Field + +from ._feature_usage import FeatureIndex + +if TYPE_CHECKING: + from botocore.client import BaseClient + +try: + import boto3 + from botocore.config import Config as BotoConfig +except ImportError as e: + raise ImportError( + "boto3 is required for BedrockKnowledgeBaseTool. " + "Install it with: pip install boto3>=1.43.32" + ) from e + +logger = logging.getLogger("agent_framework.bedrock") + + +def _get_source_uri(result: dict[str, Any]) -> str: + """Extract source URI from a retrieval result.""" + location = result.get("location", {}) + if "s3Location" in location: + return location["s3Location"].get("uri", "") + if "webLocation" in location: + return location["webLocation"].get("url", "") + if "confluenceLocation" in location: + return location["confluenceLocation"].get("url", "") + if "sharePointLocation" in location: + return location["sharePointLocation"].get("url", "") + if "customDocumentLocation" in location: + return location["customDocumentLocation"].get("id", "") + return "" + + +class _BedrockKBQueryInput(BaseModel): + """Input schema for the Bedrock Knowledge Base tool.""" + + query: Annotated[str, Field(description="The search query to find relevant documents in the knowledge base.")] + + +class BedrockKnowledgeBaseTool(FunctionTool): + """Tool that retrieves documents from Amazon Bedrock Knowledge Bases. + + Subclasses FunctionTool so it can be passed directly to any Agent or ChatClient. + + Usage: + from agent_framework_bedrock import BedrockKnowledgeBaseTool, BedrockChatClient, BedrockChatOptions + from agent_framework import Agent + + tool = BedrockKnowledgeBaseTool(knowledge_base_id="YOUR_KB_ID") + agent = Agent(client=BedrockChatClient(options=BedrockChatOptions(model_id="...")), tools=[tool]) + """ + + def __init__( + self, + *, + knowledge_base_id: str, + region_name: str = "us-east-1", + number_of_results: int = 5, + use_agentic_retrieval: bool = True, + client: Optional[BaseClient] = None, + name: str = "bedrock_knowledge_base", + description: str = ( + "Retrieves relevant documents from an Amazon Bedrock Knowledge Base. " + "Use this to answer questions that require specific knowledge or context." + ), + ) -> None: + """Create a Bedrock Knowledge Base tool. + + Args: + knowledge_base_id: The Bedrock Knowledge Base ID. + region_name: AWS region name. + number_of_results: Maximum number of results to return. + use_agentic_retrieval: Use AgenticRetrieveStream for query decomposition + reranking. + client: Pre-configured bedrock-agent-runtime client. If not provided, one is created. + name: Tool name for model registration. + description: Tool description for model context. + """ + self.knowledge_base_id = knowledge_base_id + self.region_name = region_name + self.number_of_results = number_of_results + self.use_agentic_retrieval = use_agentic_retrieval + + if client is not None: + self._client = client + else: + self._client = boto3.client( + "bedrock-agent-runtime", + region_name=self.region_name, + config=BotoConfig(user_agent_extra=f"{get_user_agent()} bedrock-kb"), + ) + + super().__init__( + name=name, + description=description, + func=self._retrieve, + input_model=_BedrockKBQueryInput, + ) + + async def _retrieve(self, query: str) -> str: + """Retrieve documents from the knowledge base. + + Args: + query: The search query. + + Returns: + Formatted string of retrieval results. + """ + mark_feature_used(FeatureIndex.BEDROCK) + + if self.use_agentic_retrieval: + try: + results = await asyncio.to_thread(self._agentic_retrieve, query) + if results: + return self._format_results(results) + except asyncio.CancelledError: + raise + except Exception as e: + logger.debug("Agentic retrieval failed, falling back: %s", e) + + results = await asyncio.to_thread(self._standard_retrieve, query) + return self._format_results(results) + + def _agentic_retrieve(self, query: str) -> list[dict[str, Any]]: + """Use AgenticRetrieveStream for query decomposition + managed reranking.""" + response = self._client.agentic_retrieve_stream( + messages=[{"content": {"text": query}, "role": "user"}], + # This tool returns retrieval passages only; the agent's own model + # generates the final answer. AgenticRetrieveStream defaults to + # generating a response (streamed responseEvents we would discard), + # so disable it explicitly to avoid unnecessary generation latency/cost. + generateResponse=False, + retrievers=[{ + "configuration": { + "knowledgeBase": { + "knowledgeBaseId": self.knowledge_base_id, + "retrievalOverrides": {"maxNumberOfResults": self.number_of_results}, + } + } + }], + agenticRetrieveConfiguration={ + "foundationModelType": "MANAGED", + "rerankingModelType": "MANAGED", + }, + ) + results = [] + for event in response.get("stream", []): + if "result" in event and "results" in event["result"]: + for r in event["result"]["results"]: + # AgenticRetrieveStream results use a different schema than standard + # Retrieve: they expose `content`/`metadata`/`sourceRetriever` and do + # NOT include `score` or `location`. The source URI lives in metadata, + # and managed reranking orders results without exposing a numeric score. + metadata = r.get("metadata", {}) or {} + results.append({ + "content": r.get("content", {}).get("text", ""), + "source": metadata.get("_source_uri", ""), + "score": None, + }) + return results + + def _standard_retrieve(self, query: str) -> list[dict[str, Any]]: + """Use standard Retrieve API with managed search configuration.""" + response = self._client.retrieve( + knowledgeBaseId=self.knowledge_base_id, + retrievalQuery={"text": query}, + retrievalConfiguration={"managedSearchConfiguration": {"numberOfResults": self.number_of_results}}, + ) + results = [] + for r in response.get("retrievalResults", []): + results.append({ + "content": r.get("content", {}).get("text", ""), + "source": _get_source_uri(r), + "score": r.get("score", 0), + }) + return results + + @staticmethod + def _format_results(results: list[dict[str, Any]]) -> str: + """Format retrieval results as a readable string.""" + if not results: + return "No relevant documents found." + parts = [] + for i, r in enumerate(results, 1): + source = r.get("source", "") + content = r.get("content", "") + score = r.get("score") + # Standard Retrieve results carry a numeric relevance score; agentic + # (managed reranking) results do not, so only render it when present. + header = f"[{i}] (score: {score:.3f})" if isinstance(score, (int, float)) else f"[{i}]" + parts.append(f"{header} {content}\n Source: {source}") + return "\n\n".join(parts) diff --git a/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py new file mode 100644 index 00000000000..a47bfbc1a87 --- /dev/null +++ b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py @@ -0,0 +1,159 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Amazon Bedrock Knowledge Base context provider for Agent Framework.""" + +from __future__ import annotations + +import asyncio +import logging +from typing import TYPE_CHECKING, Any, Optional + +from agent_framework import AgentSession, ContextProvider, Message, SessionContext +from agent_framework._telemetry import get_user_agent, mark_feature_used + +if TYPE_CHECKING: + from agent_framework._agents import SupportsAgentRun + from botocore.client import BaseClient + +try: + import boto3 + from botocore.config import Config as BotoConfig +except ImportError as e: + raise ImportError( + "boto3 is required for BedrockKnowledgeBaseProvider. " + "Install it with: pip install boto3>=1.43.32" + ) from e + +from ._feature_usage import FeatureIndex +from ._knowledge_base import _get_source_uri + +logger = logging.getLogger("agent_framework.bedrock") + + +class BedrockKnowledgeBaseProvider(ContextProvider): + """Context provider that injects Bedrock Knowledge Base results before agent runs. + + Subclasses ContextProvider and implements before_run() to automatically + retrieve relevant context from a Bedrock Knowledge Base on every agent invocation. + + Usage: + from agent_framework_bedrock import BedrockKnowledgeBaseProvider + + provider = BedrockKnowledgeBaseProvider(knowledge_base_id="YOUR_KB_ID") + agent = Agent(context_providers=[provider]) + """ + + DEFAULT_CONTEXT_PROMPT = ( + "## Knowledge Base Context\n" + "The following passages were retrieved from the knowledge base. " + "Use them to answer the user's question:" + ) + + def __init__( + self, + *, + knowledge_base_id: str, + region_name: str = "us-east-1", + number_of_results: int = 5, + min_score: float = 0.0, + source_id: str = "bedrock-kb", + context_prompt: str | None = None, + client: Optional[BaseClient] = None, + ) -> None: + """Create a Bedrock Knowledge Base context provider. + + Args: + knowledge_base_id: The Bedrock Knowledge Base ID. + region_name: AWS region name. + number_of_results: Maximum number of results to inject as context. + min_score: Minimum relevance score threshold. + source_id: Identifier for this context source. + context_prompt: Custom prompt to prepend to retrieved context. + client: Pre-configured bedrock-agent-runtime client. If not provided, one is created. + """ + super().__init__(source_id) + self.knowledge_base_id = knowledge_base_id + self.region_name = region_name + self.number_of_results = number_of_results + self.min_score = min_score + self.context_prompt = context_prompt or self.DEFAULT_CONTEXT_PROMPT + + if client is not None: + self._client = client + else: + self._client = boto3.client( + "bedrock-agent-runtime", + region_name=self.region_name, + config=BotoConfig(user_agent_extra=f"{get_user_agent()} bedrock-kb"), + ) + + async def before_run( + self, + *, + agent: SupportsAgentRun, + session: AgentSession, + context: SessionContext, + state: dict[str, Any], + ) -> None: + """Retrieve relevant KB context and inject it into the session context. + + Called automatically before each model invocation. Extracts the user's + query from input messages, retrieves relevant passages, and adds them + as a delimited user-role message (untrusted external content). + + Args: + agent: The agent running this invocation. + session: The current session. + context: The invocation context - add messages here. + state: The provider-scoped mutable state dict. + """ + # Extract query from input messages + input_text = "\n".join( + msg.text for msg in context.input_messages if msg and msg.text and msg.text.strip() + ) + if not input_text.strip(): + return + + # Retrieve from knowledge base (non-fatal — agent continues without context on failure) + mark_feature_used(FeatureIndex.BEDROCK) + try: + retrieved_context = await self._retrieve(input_text) + except asyncio.CancelledError: + raise + except Exception: + logger.debug("KB retrieval failed, continuing without context", exc_info=True) + return + + if not retrieved_context: + return + + # Inject as a user-role message (untrusted external content), consistent with + # other context providers in this repo (e.g. azure-cosmos-memory), which keep + # retrieved/generated content in the untrusted user channel rather than elevating + # it to system instructions (avoids stored prompt-injection). Bedrock's + # role-alternation requirement is handled by coalescing adjacent same-role + # messages in BedrockChatClient._prepare_bedrock_messages. + context.extend_messages( + self.source_id, + [Message(role="user", contents=[f"{self.context_prompt}\n\n{retrieved_context}"])], + ) + + async def _retrieve(self, query: str) -> str: + """Retrieve and format context from the knowledge base.""" + response = await asyncio.to_thread( + lambda: self._client.retrieve( + knowledgeBaseId=self.knowledge_base_id, + retrievalQuery={"text": query}, + retrievalConfiguration={"managedSearchConfiguration": {"numberOfResults": self.number_of_results}}, + ) + ) + + passages = [] + for r in response.get("retrievalResults", []): + score = r.get("score", 0) + if score >= self.min_score: + content = r.get("content", {}).get("text", "") + source = _get_source_uri(r) + passages.append(f"[Source: {source}]\n{content}") + + return "\n\n---\n\n".join(passages) if passages else "" diff --git a/python/packages/bedrock/pyproject.toml b/python/packages/bedrock/pyproject.toml index 41d278e4458..c096e29b391 100644 --- a/python/packages/bedrock/pyproject.toml +++ b/python/packages/bedrock/pyproject.toml @@ -24,8 +24,8 @@ classifiers = [ ] dependencies = [ "agent-framework-core>=1.13.0,<2", - "boto3>=1.35.0,<2.0.0", - "botocore>=1.35.0,<2.0.0", + "boto3>=1.43.32,<2.0.0", + "botocore>=1.43.32,<2.0.0", ] [tool.uv] diff --git a/python/packages/bedrock/samples/README.md b/python/packages/bedrock/samples/README.md new file mode 100644 index 00000000000..6dba5132e1e --- /dev/null +++ b/python/packages/bedrock/samples/README.md @@ -0,0 +1,38 @@ +# Bedrock Knowledge Base Examples + +This folder contains examples demonstrating how to use Amazon Bedrock Knowledge Bases with the Agent Framework. + +## Examples + +| File | Description | +|------|-------------| +| [`bedrock_kb_tool.py`](bedrock_kb_tool.py) | Using `BedrockKnowledgeBaseTool` as a FunctionTool — agent calls it on-demand when it needs knowledge base context. | +| [`bedrock_kb_context_provider.py`](bedrock_kb_context_provider.py) | Using `BedrockKnowledgeBaseProvider` as a ContextProvider — automatically injects KB context before every agent invocation. | + +## When to use each pattern + +- **Tool pattern** (`BedrockKnowledgeBaseTool`): When the agent should decide *when* to search the KB. Best for multi-tool agents where KB retrieval is one of several capabilities. +- **Provider pattern** (`BedrockKnowledgeBaseProvider`): When KB context should *always* be available. Best for single-purpose assistants that always need domain knowledge. + +## Environment Variables + +- `AWS_DEFAULT_REGION`: AWS region where your Knowledge Base is deployed +- AWS credentials: Configure via environment variables, IAM role, or AWS profiles + +## Required IAM Permissions + +```json +{ + "Version": "2012-10-17", + "Statement": [ + { + "Effect": "Allow", + "Action": [ + "bedrock:Retrieve", + "bedrock:AgenticRetrieveStream" + ], + "Resource": "arn:aws:bedrock:*:*:knowledge-base/*" + } + ] +} +``` diff --git a/python/packages/bedrock/samples/__init__.py b/python/packages/bedrock/samples/__init__.py new file mode 100644 index 00000000000..2a50eae8941 --- /dev/null +++ b/python/packages/bedrock/samples/__init__.py @@ -0,0 +1 @@ +# Copyright (c) Microsoft. All rights reserved. diff --git a/python/packages/bedrock/samples/bedrock_kb_context_provider.py b/python/packages/bedrock/samples/bedrock_kb_context_provider.py new file mode 100644 index 00000000000..09f6df4d5f9 --- /dev/null +++ b/python/packages/bedrock/samples/bedrock_kb_context_provider.py @@ -0,0 +1,50 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Sample: Using BedrockKnowledgeBaseProvider for automatic context injection. + +This demonstrates the ContextProvider pattern where KB context is automatically +retrieved and injected before every agent invocation — no explicit tool calling needed. + +Prerequisites: + pip install agent-framework-bedrock + export AWS_DEFAULT_REGION=us-west-2 + # AWS credentials configured (IAM role with bedrock:Retrieve) +""" + +import asyncio + +from agent_framework import Agent +from agent_framework_bedrock import BedrockChatClient, BedrockChatOptions, BedrockKnowledgeBaseProvider + + +async def main() -> None: + # Create the Knowledge Base context provider — subclasses ContextProvider + kb_provider = BedrockKnowledgeBaseProvider( + knowledge_base_id="YOUR_KB_ID", # Replace with your managed KB ID + region_name="us-west-2", + number_of_results=3, + min_score=0.3, # Only include results above this relevance threshold + source_id="company-docs", # Unique ID for this context source + ) + + # Create a Bedrock chat client + chat_client = BedrockChatClient( + options=BedrockChatOptions(model_id="us.anthropic.claude-sonnet-4-20250514-v1:0") + ) + + # Create an agent with the context provider — context is injected automatically + agent = Agent( + client=chat_client, + name="ContextualAssistant", + instructions="You are a helpful assistant that answers based on provided context.", + context_providers=[kb_provider], # ContextProvider subclass, injects context on every run + ) + + # Run the agent — KB context is retrieved and injected automatically via before_run() + session = agent.create_session() + response = await agent.run("What data sources does Bedrock support?", session=session) + print(f"Agent response: {response.text}") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/python/packages/bedrock/samples/bedrock_kb_tool.py b/python/packages/bedrock/samples/bedrock_kb_tool.py new file mode 100644 index 00000000000..a4de08cbe1f --- /dev/null +++ b/python/packages/bedrock/samples/bedrock_kb_tool.py @@ -0,0 +1,50 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Sample: Using BedrockKnowledgeBaseTool with an Agent. + +This demonstrates how the Bedrock Knowledge Base tool integrates with +Agent Framework primitives. The tool subclasses FunctionTool and can be +passed directly to any Agent or ChatClient. + +Prerequisites: + pip install agent-framework-bedrock + export AWS_DEFAULT_REGION=us-west-2 + # AWS credentials configured (IAM role with bedrock:Retrieve and bedrock:AgenticRetrieveStream) +""" + +import asyncio + +from agent_framework import Agent +from agent_framework_bedrock import BedrockChatClient, BedrockChatOptions, BedrockKnowledgeBaseTool + + +async def main() -> None: + # Create the Knowledge Base tool — subclasses FunctionTool, pass directly to Agent + kb_tool = BedrockKnowledgeBaseTool( + knowledge_base_id="YOUR_KB_ID", # Replace with your managed KB ID + region_name="us-west-2", + number_of_results=5, + use_agentic_retrieval=True, # Uses query decomposition + managed reranking + ) + + # Create a Bedrock chat client + chat_client = BedrockChatClient( + options=BedrockChatOptions(model_id="us.anthropic.claude-sonnet-4-20250514-v1:0") + ) + + # Create an agent with the KB tool — Agent will call it when it needs context + agent = Agent( + client=chat_client, + name="KnowledgeAssistant", + instructions="You are a helpful assistant. Use the knowledge base tool to answer questions about the company.", + tools=[kb_tool], # FunctionTool subclass, works with any ChatClient + ) + + # Run the agent + session = agent.create_session() + response = await agent.run("What is our return policy for electronics?", session=session) + print(f"Agent response: {response.text}") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/python/packages/bedrock/tests/test_bedrock_client.py b/python/packages/bedrock/tests/test_bedrock_client.py index 6d339ae3c59..4239f40f169 100644 --- a/python/packages/bedrock/tests/test_bedrock_client.py +++ b/python/packages/bedrock/tests/test_bedrock_client.py @@ -381,6 +381,66 @@ def test_prepare_bedrock_messages_skips_unsupported_content_and_unmatched_tool_r assert conversation == [{"role": "user", "content": [{"text": "hello"}]}] +def test_prepare_bedrock_messages_coalesces_adjacent_user_turns() -> None: + """Adjacent user-role messages (e.g. injected KB context + real input) must be + merged into a single user turn so Bedrock's role-alternation rule is satisfied.""" + client = _make_client() + messages = [ + Message(role="user", contents=[Content.from_text(text="[KB context] policy is 30 days")]), + Message(role="user", contents=[Content.from_text(text="What is the policy?")]), + ] + + prompts, conversation = client._prepare_bedrock_messages(messages) + + assert prompts == [] + assert conversation == [ + { + "role": "user", + "content": [ + {"text": "[KB context] policy is 30 days"}, + {"text": "What is the policy?"}, + ], + } + ] + + +def test_prepare_bedrock_messages_does_not_coalesce_across_assistant() -> None: + """User turns separated by an assistant turn must remain distinct.""" + client = _make_client() + messages = [ + Message(role="user", contents=[Content.from_text(text="first")]), + Message(role="assistant", contents=[Content.from_text(text="reply")]), + Message(role="user", contents=[Content.from_text(text="second")]), + ] + + _, conversation = client._prepare_bedrock_messages(messages) + + assert [m["role"] for m in conversation] == ["user", "assistant", "user"] + + +def test_prepare_bedrock_messages_does_not_coalesce_tool_results_into_user_text() -> None: + """A tool-result turn (role='tool' -> Bedrock 'user') must NOT be merged into a + preceding genuine user text turn; function-call/tool-result serialization is preserved.""" + client = _make_client() + messages = [ + Message(role="user", contents=[Content.from_text(text="run the tool")]), + Message( + role="assistant", + contents=[Content.from_function_call(call_id="call-1", name="do_it", arguments={})], + ), + Message(role="tool", contents=[Content.from_function_result(call_id="call-1", result={"ok": True})]), + ] + + _, conversation = client._prepare_bedrock_messages(messages) + + # user text, assistant toolUse, then a SEPARATE user turn holding the toolResult + assert [m["role"] for m in conversation] == ["user", "assistant", "user"] + # the tool-result turn must contain the toolResult block, not be merged with "run the tool" + last = conversation[-1] + assert any(isinstance(b, dict) and "toolResult" in b for b in last["content"]) + assert not any(isinstance(b, dict) and b.get("text") == "run the tool" for b in last["content"]) + + def test_align_tool_results_handles_pending_edge_cases() -> None: """Tool result alignment should preserve valid blocks and drop invalid or extra results.""" client = _make_client() diff --git a/python/packages/bedrock/tests/test_bedrock_knowledge_base.py b/python/packages/bedrock/tests/test_bedrock_knowledge_base.py new file mode 100644 index 00000000000..79f74299588 --- /dev/null +++ b/python/packages/bedrock/tests/test_bedrock_knowledge_base.py @@ -0,0 +1,279 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Tests for Bedrock Knowledge Base tool and provider.""" + +import asyncio +from unittest.mock import MagicMock, patch + +from agent_framework import FunctionTool, ContextProvider + + +class TestBedrockKnowledgeBaseTool: + def test_is_function_tool_subclass(self): + from agent_framework_bedrock._knowledge_base import BedrockKnowledgeBaseTool + + mock_client = MagicMock() + tool = BedrockKnowledgeBaseTool(knowledge_base_id="TEST_KB", client=mock_client) + assert isinstance(tool, FunctionTool) + + def test_tool_has_correct_name_and_description(self): + from agent_framework_bedrock._knowledge_base import BedrockKnowledgeBaseTool + + mock_client = MagicMock() + tool = BedrockKnowledgeBaseTool(knowledge_base_id="TEST_KB", client=mock_client) + assert tool.name == "bedrock_knowledge_base" + assert "knowledge" in tool.description.lower() + + def test_retrieve_returns_formatted_results(self): + from agent_framework_bedrock._knowledge_base import BedrockKnowledgeBaseTool + + mock_client = MagicMock() + mock_client.retrieve.return_value = { + "retrievalResults": [ + {"content": {"text": "Result 1"}, "score": 0.95, "location": {"s3Location": {"uri": "s3://b/k"}}}, + {"content": {"text": "Result 2"}, "score": 0.80, "location": {"webLocation": {"url": "https://example.com"}}}, + ] + } + + tool = BedrockKnowledgeBaseTool( + knowledge_base_id="TEST_KB", + region_name="us-west-2", + use_agentic_retrieval=False, + client=mock_client, + ) + + result = asyncio.run(tool._retrieve(query="test query")) + assert "Result 1" in result + assert "Result 2" in result + assert "s3://b/k" in result + assert "0.950" in result + + def test_agentic_with_fallback(self): + from agent_framework_bedrock._knowledge_base import BedrockKnowledgeBaseTool + + mock_client = MagicMock() + mock_client.agentic_retrieve_stream.side_effect = Exception("Not available") + mock_client.retrieve.return_value = {"retrievalResults": [ + {"content": {"text": "Fallback"}, "score": 0.7, "location": {}}, + ]} + + tool = BedrockKnowledgeBaseTool( + knowledge_base_id="TEST_KB", + use_agentic_retrieval=True, + client=mock_client, + ) + + result = asyncio.run(tool._retrieve(query="test")) + assert "Fallback" in result + mock_client.agentic_retrieve_stream.assert_called_once() + mock_client.retrieve.assert_called_once() + + def test_agentic_retrieve_success(self): + from agent_framework_bedrock._knowledge_base import BedrockKnowledgeBaseTool + + mock_client = MagicMock() + mock_client.agentic_retrieve_stream.return_value = { + "stream": [ + {"result": {"results": [ + # AgenticRetrieveStream schema: content/metadata/sourceRetriever + # (no score, no location). Source URI comes from metadata._source_uri. + { + "content": {"mimeType": "text/plain", "text": "Agentic result"}, + "metadata": {"_source_uri": "s3://b/doc", "_document_title": "Doc"}, + "sourceRetriever": {"identifier": "TEST_KB"}, + }, + ]}} + ] + } + + tool = BedrockKnowledgeBaseTool( + knowledge_base_id="TEST_KB", + use_agentic_retrieval=True, + client=mock_client, + ) + + result = asyncio.run(tool._retrieve(query="complex question")) + assert "Agentic result" in result + assert "s3://b/doc" in result + # Agentic results must not fabricate a numeric score + assert "score:" not in result + # Response generation must be disabled (tool returns passages only) + assert mock_client.agentic_retrieve_stream.call_args.kwargs["generateResponse"] is False + mock_client.retrieve.assert_not_called() + + def test_client_uses_get_user_agent(self): + from agent_framework_bedrock._knowledge_base import BedrockKnowledgeBaseTool + + with patch("agent_framework_bedrock._knowledge_base.boto3.client") as mock_boto: + mock_boto.return_value = MagicMock() + _ = BedrockKnowledgeBaseTool(knowledge_base_id="TEST_KB", region_name="us-west-2") + config = mock_boto.call_args.kwargs["config"] + ua = getattr(config, "user_agent_extra", "") + assert "bedrock-kb" in ua + + def test_no_results_returns_message(self): + from agent_framework_bedrock._knowledge_base import BedrockKnowledgeBaseTool + + mock_client = MagicMock() + mock_client.retrieve.return_value = {"retrievalResults": []} + + tool = BedrockKnowledgeBaseTool( + knowledge_base_id="TEST_KB", + use_agentic_retrieval=False, + client=mock_client, + ) + + result = asyncio.run(tool._retrieve(query="unknown")) + assert "No relevant documents found" in result + + def test_invoke_end_to_end(self): + """Test the public FunctionTool.invoke() path with argument validation.""" + from agent_framework_bedrock._knowledge_base import BedrockKnowledgeBaseTool + + mock_client = MagicMock() + mock_client.retrieve.return_value = { + "retrievalResults": [ + {"content": {"text": "Invoked result"}, "score": 0.88, "location": {"s3Location": {"uri": "s3://b/invoke"}}} + ] + } + + tool = BedrockKnowledgeBaseTool( + knowledge_base_id="TEST_KB", + use_agentic_retrieval=False, + client=mock_client, + ) + + # Call via the public invoke() API — exercises argument validation + Content parsing + result = asyncio.run(tool.invoke(arguments={"query": "test invoke"})) + # invoke() returns list[Content] by default + assert len(result) > 0 + assert "Invoked result" in result[0].text + + +class TestBedrockKnowledgeBaseProvider: + def test_is_context_provider_subclass(self): + from agent_framework_bedrock._knowledge_base_provider import BedrockKnowledgeBaseProvider + + mock_client = MagicMock() + provider = BedrockKnowledgeBaseProvider(knowledge_base_id="TEST_KB", client=mock_client) + assert isinstance(provider, ContextProvider) + + def test_has_source_id(self): + from agent_framework_bedrock._knowledge_base_provider import BedrockKnowledgeBaseProvider + + mock_client = MagicMock() + provider = BedrockKnowledgeBaseProvider( + knowledge_base_id="TEST_KB", source_id="my-kb", client=mock_client + ) + assert provider.source_id == "my-kb" + + def test_retrieve_returns_formatted_context(self): + from agent_framework_bedrock._knowledge_base_provider import BedrockKnowledgeBaseProvider + + mock_client = MagicMock() + mock_client.retrieve.return_value = { + "retrievalResults": [ + {"content": {"text": "Passage 1"}, "score": 0.9, "location": {"s3Location": {"uri": "s3://b/doc.pdf"}}}, + {"content": {"text": "Passage 2"}, "score": 0.5, "location": {}}, + ] + } + + provider = BedrockKnowledgeBaseProvider( + knowledge_base_id="TEST_KB", + client=mock_client, + ) + + context = asyncio.run(provider._retrieve("test query")) + assert "Passage 1" in context + assert "s3://b/doc.pdf" in context + + def test_min_score_filtering(self): + from agent_framework_bedrock._knowledge_base_provider import BedrockKnowledgeBaseProvider + + mock_client = MagicMock() + mock_client.retrieve.return_value = { + "retrievalResults": [ + {"content": {"text": "High"}, "score": 0.9, "location": {}}, + {"content": {"text": "Low"}, "score": 0.2, "location": {}}, + ] + } + + provider = BedrockKnowledgeBaseProvider( + knowledge_base_id="TEST_KB", + min_score=0.5, + client=mock_client, + ) + + context = asyncio.run(provider._retrieve("test")) + assert "High" in context + assert "Low" not in context + + def test_has_before_run_method(self): + from agent_framework_bedrock._knowledge_base_provider import BedrockKnowledgeBaseProvider + + mock_client = MagicMock() + provider = BedrockKnowledgeBaseProvider(knowledge_base_id="TEST_KB", client=mock_client) + assert hasattr(provider, "before_run") + assert asyncio.iscoroutinefunction(provider.before_run) + + def test_before_run_injects_context(self): + from agent_framework import Message, SessionContext + from agent_framework_bedrock._knowledge_base_provider import BedrockKnowledgeBaseProvider + + mock_client = MagicMock() + mock_client.retrieve.return_value = { + "retrievalResults": [ + {"content": {"text": "Relevant passage"}, "score": 0.9, "location": {"s3Location": {"uri": "s3://b/doc"}}}, + ] + } + + provider = BedrockKnowledgeBaseProvider( + knowledge_base_id="TEST_KB", + client=mock_client, + ) + + # Create a SessionContext with an input message + context = SessionContext( + input_messages=[Message(role="user", contents=["What is our policy?"])], + ) + + # Verify context_messages is empty before + assert len(context.context_messages) == 0 + + # Run before_run + asyncio.run(provider.before_run( + agent=MagicMock(), + session=MagicMock(), + context=context, + state={}, + )) + + # Verify context injected as an untrusted user-role message (matches repo + # convention; role alternation is handled by _prepare_bedrock_messages coalescing) + assert "bedrock-kb" in context.context_messages + injected = context.context_messages["bedrock-kb"] + assert len(injected) == 1 + assert injected[0].role == "user" + assert "Relevant passage" in injected[0].text + assert "s3://b/doc" in injected[0].text + + def test_before_run_skips_empty_input(self): + from agent_framework import SessionContext + from agent_framework_bedrock._knowledge_base_provider import BedrockKnowledgeBaseProvider + + mock_client = MagicMock() + provider = BedrockKnowledgeBaseProvider(knowledge_base_id="TEST_KB", client=mock_client) + + # Empty input messages + context = SessionContext(input_messages=[]) + + asyncio.run(provider.before_run( + agent=MagicMock(), + session=MagicMock(), + context=context, + state={}, + )) + + # Should not call retrieve + mock_client.retrieve.assert_not_called() + assert len(context.context_messages) == 0 diff --git a/python/uv.lock b/python/uv.lock index 3ed405101ea..30a275cbeab 100644 --- a/python/uv.lock +++ b/python/uv.lock @@ -357,8 +357,8 @@ dependencies = [ [package.metadata] requires-dist = [ { name = "agent-framework-core", editable = "packages/core" }, - { name = "boto3", specifier = ">=1.35.0,<2.0.0" }, - { name = "botocore", specifier = ">=1.35.0,<2.0.0" }, + { name = "boto3", specifier = ">=1.43.32,<2.0.0" }, + { name = "botocore", specifier = ">=1.43.32,<2.0.0" }, ] [[package]]