From 5d812862de1d1ec1e8924ef7d8fdec734fa94366 Mon Sep 17 00:00:00 2001 From: PVidyadhar Date: Sat, 11 Jul 2026 02:30:02 +0000 Subject: [PATCH] feat: add Amazon Bedrock Knowledge Base tool and context provider - Created BedrockKnowledgeBaseTool with async run() + get_tool_definition() - Created BedrockKnowledgeBaseProvider (ContextProvider subclass) with before_run() - Two integration points: standalone tool + automatic context injection - Supports managed search and agentic retrieval with fallback - Unit tests included - Added BEDROCK_MANAGED_KB.md design doc --- python/packages/bedrock/BEDROCK_MANAGED_KB.md | 62 +++++ .../agent_framework_bedrock/__init__.py | 4 + .../_knowledge_base.py | 185 +++++++++++++ .../_knowledge_base_provider.py | 141 ++++++++++ python/packages/bedrock/pyproject.toml | 6 +- python/packages/bedrock/samples/README.md | 33 +++ python/packages/bedrock/samples/__init__.py | 1 + .../samples/bedrock_kb_context_provider.py | 53 ++++ .../bedrock/samples/bedrock_kb_tool.py | 53 ++++ .../tests/test_bedrock_knowledge_base.py | 246 ++++++++++++++++++ 10 files changed, 781 insertions(+), 3 deletions(-) create mode 100644 python/packages/bedrock/BEDROCK_MANAGED_KB.md create mode 100644 python/packages/bedrock/agent_framework_bedrock/_knowledge_base.py create mode 100644 python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py create mode 100644 python/packages/bedrock/samples/README.md create mode 100644 python/packages/bedrock/samples/__init__.py create mode 100644 python/packages/bedrock/samples/bedrock_kb_context_provider.py create mode 100644 python/packages/bedrock/samples/bedrock_kb_tool.py create mode 100644 python/packages/bedrock/tests/test_bedrock_knowledge_base.py diff --git a/python/packages/bedrock/BEDROCK_MANAGED_KB.md b/python/packages/bedrock/BEDROCK_MANAGED_KB.md new file mode 100644 index 00000000000..57e3fa61018 --- /dev/null +++ b/python/packages/bedrock/BEDROCK_MANAGED_KB.md @@ -0,0 +1,62 @@ +# 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 + +tool = BedrockKnowledgeBaseTool( + knowledge_base_id="YOUR_KB_ID", + region_name="us-east-1", +) + +# As a FunctionTool, pass directly to an Agent: +agent = Agent(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 +{ + "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..b40d756f00e 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__) @@ -18,5 +20,7 @@ "BedrockEmbeddingSettings", "BedrockGuardrailConfig", "BedrockSettings", + "BedrockKnowledgeBaseTool", + "BedrockKnowledgeBaseProvider", "__version__", ] 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..4eb14c190ab --- /dev/null +++ b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base.py @@ -0,0 +1,185 @@ +# 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 +from pydantic import BaseModel, Field + +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(__name__) + + +def _get_source_uri(result: dict) -> 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 + + tool = BedrockKnowledgeBaseTool(knowledge_base_id="YOUR_KB_ID") + agent = Agent(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: + 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. + """ + if self.use_agentic_retrieval: + try: + results = await asyncio.to_thread(self._agentic_retrieve, query) + if results: + return self._format_results(results) + 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"}], + 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"]: + results.append({ + "content": r.get("content", {}).get("text", ""), + "source": _get_source_uri(r), + "score": r.get("score", 0), + }) + 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", 0) + parts.append(f"[{i}] (score: {score:.3f}) {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..977b556d1a6 --- /dev/null +++ b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py @@ -0,0 +1,141 @@ +# Copyright (c) Microsoft. All rights reserved. + +"""Amazon Bedrock Knowledge Base context provider for Agent Framework.""" + +from __future__ import annotations + +import asyncio +from typing import TYPE_CHECKING, Any, Optional + +from agent_framework import Message +from agent_framework._sessions import AgentSession, ContextProvider, SessionContext +from agent_framework._telemetry import get_user_agent + +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 agent_framework_bedrock._knowledge_base import _get_source_uri + + +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: + 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 system message to the context. + + 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 + retrieved_context = await self._retrieve(input_text) + if not retrieved_context: + return + + # Inject as a system message via extend_messages + context_message = Message(role="system", contents=[f"{self.context_prompt}\n\n{retrieved_context}"]) + context.extend_messages(self, [context_message]) + + 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 3b3570e8250..060829ca5d6 100644 --- a/python/packages/bedrock/pyproject.toml +++ b/python/packages/bedrock/pyproject.toml @@ -23,9 +23,9 @@ classifiers = [ "Typing :: Typed", ] dependencies = [ - "agent-framework-core>=1.10.0,<2", - "boto3>=1.35.0,<2.0.0", - "botocore>=1.35.0,<2.0.0", + "agent-framework-core>=1.13.0,<2", + "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..546efd39d46 --- /dev/null +++ b/python/packages/bedrock/samples/README.md @@ -0,0 +1,33 @@ +# 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 +{ + "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..7987fba03b6 --- /dev/null +++ b/python/packages/bedrock/samples/bedrock_kb_context_provider.py @@ -0,0 +1,53 @@ +# 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( + name="ContextualAssistant", + instructions="You are a helpful assistant that answers based on provided context.", + chat_client=chat_client, + 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.invoke( + session=session, + input_message="What data sources does Bedrock support?", + ) + 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..4b7e4f0e322 --- /dev/null +++ b/python/packages/bedrock/samples/bedrock_kb_tool.py @@ -0,0 +1,53 @@ +# 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( + name="KnowledgeAssistant", + instructions="You are a helpful assistant. Use the knowledge base tool to answer questions about the company.", + chat_client=chat_client, + tools=[kb_tool], # FunctionTool subclass, works with any ChatClient + ) + + # Run the agent + session = agent.create_session() + response = await agent.invoke( + session=session, + input_message="What is our return policy for electronics?", + ) + print(f"Agent response: {response.text}") + + +if __name__ == "__main__": + asyncio.run(main()) 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..f7885f40cf8 --- /dev/null +++ b/python/packages/bedrock/tests/test_bedrock_knowledge_base.py @@ -0,0 +1,246 @@ +# 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 +from agent_framework._sessions import 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": [ + {"content": {"text": "Agentic result"}, "score": 0.99, "location": {"s3Location": {"uri": "s3://b/doc"}}}, + ]}} + ] + } + + 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 + 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 + + +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 + from agent_framework._sessions import 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 was injected via extend_messages + assert "bedrock-kb" in context.context_messages + injected = context.context_messages["bedrock-kb"] + assert len(injected) == 1 + 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._sessions 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