From 5d812862de1d1ec1e8924ef7d8fdec734fa94366 Mon Sep 17 00:00:00 2001 From: PVidyadhar Date: Sat, 11 Jul 2026 02:30:02 +0000 Subject: [PATCH 1/5] 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 From 3d1ae25443aad3eb5a4612e64553c26fc8d718fb Mon Sep 17 00:00:00 2001 From: PVidyadhar Date: Wed, 9 Sep 2026 07:44:47 +0000 Subject: [PATCH 2/5] fix: inject KB context as instructions to avoid consecutive user roles Addresses reviewer feedback (@moonbox3): when the provider is used with BedrockChatClient, injecting retrieved context as a separate user message produced two consecutive user turns in _prepare_bedrock_messages (which does not coalesce same-role messages). Route the retrieved context through extend_instructions() so it lands in Bedrock's system field, separate from the conversation array. This is model-agnostic and also avoids adding untrusted content as a system conversation message. - provider uses context.extend_instructions(self.source_id, ...) - removed unused Message import - updated tests to assert on context.instructions - 65 tests pass, verified E2E via agent.run() with BedrockChatClient + live KB --- .../_knowledge_base_provider.py | 13 ++++++++----- .../tests/test_bedrock_knowledge_base.py | 17 ++++++++--------- 2 files changed, 16 insertions(+), 14 deletions(-) diff --git a/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py index 32d309d3fd7..f408df3c266 100644 --- a/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py +++ b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py @@ -8,7 +8,7 @@ import logging from typing import TYPE_CHECKING, Any, Optional -from agent_framework import AgentSession, ContextProvider, Message, SessionContext +from agent_framework import AgentSession, ContextProvider, SessionContext from agent_framework._telemetry import get_user_agent, mark_feature_used if TYPE_CHECKING: @@ -99,7 +99,7 @@ async def before_run( Called automatically before each model invocation. Extracts the user's query from input messages, retrieves relevant passages, and adds them - as a user message to the context. + as instructions to the context (prepended to the system prompt). Args: agent: The agent running this invocation. @@ -127,9 +127,12 @@ async def before_run( if not retrieved_context: return - # Inject as a user message (untrusted external content) via extend_messages - context_message = Message(role="user", contents=[f"{self.context_prompt}\n\n{retrieved_context}"]) - context.extend_messages(self, [context_message]) + # Inject retrieved context as instructions rather than a message. + # Adding it as a separate user message would produce consecutive user + # roles when the agent appends the real input (SessionContext.get_messages + # with include_input=True), which Bedrock Converse rejects since roles must + # alternate. Instructions are prepended to the system context and avoid this. + context.extend_instructions(self.source_id, f"{self.context_prompt}\n\n{retrieved_context}") async def _retrieve(self, query: str) -> str: """Retrieve and format context from the knowledge base.""" diff --git a/python/packages/bedrock/tests/test_bedrock_knowledge_base.py b/python/packages/bedrock/tests/test_bedrock_knowledge_base.py index 0a682f0fe37..41b048b4f27 100644 --- a/python/packages/bedrock/tests/test_bedrock_knowledge_base.py +++ b/python/packages/bedrock/tests/test_bedrock_knowledge_base.py @@ -227,8 +227,8 @@ def test_before_run_injects_context(self): input_messages=[Message(role="user", contents=["What is our policy?"])], ) - # Verify context_messages is empty before - assert len(context.context_messages) == 0 + # Verify instructions are empty before + assert len(context.instructions) == 0 # Run before_run asyncio.run(provider.before_run( @@ -238,12 +238,11 @@ def test_before_run_injects_context(self): 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 + # Verify context was injected as instructions (avoids consecutive user-role + # issue with BedrockChatClient which requires alternating roles) + assert len(context.instructions) == 1 + assert "Relevant passage" in context.instructions[0] + assert "s3://b/doc" in context.instructions[0] def test_before_run_skips_empty_input(self): from agent_framework import SessionContext @@ -264,4 +263,4 @@ def test_before_run_skips_empty_input(self): # Should not call retrieve mock_client.retrieve.assert_not_called() - assert len(context.context_messages) == 0 + assert len(context.instructions) == 0 From 9f8ea758fcda7207c6b7026bedf9f8acd6605eaa Mon Sep 17 00:00:00 2001 From: PVidyadhar Date: Wed, 9 Sep 2026 08:03:42 +0000 Subject: [PATCH 3/5] fix: address Copilot review on #8173 - Keep retrieved KB passages as untrusted user-role context instead of elevating to system instructions (matches azure-cosmos-memory convention; avoids stored prompt-injection). Solve Bedrock role alternation by coalescing adjacent user-role messages in _prepare_bedrock_messages (assistant turns left untouched to preserve tool-use/tool-result pairing). - Regenerate python/uv.lock for the boto3/botocore >=1.43.32 floor. - Add BedrockKnowledgeBaseTool/Provider to bedrock AGENTS.md class list. - Tests: coalescing + no-coalesce-across-assistant cases; 67 pass. Verified E2E via agent.run() with BedrockChatClient + live KB. --- python/packages/bedrock/AGENTS.md | 2 + .../agent_framework_bedrock/_chat_client.py | 11 +++++- .../_knowledge_base_provider.py | 20 ++++++---- .../bedrock/tests/test_bedrock_client.py | 37 +++++++++++++++++++ .../tests/test_bedrock_knowledge_base.py | 19 ++++++---- python/uv.lock | 4 +- 6 files changed, 74 insertions(+), 19 deletions(-) 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/agent_framework_bedrock/_chat_client.py b/python/packages/bedrock/agent_framework_bedrock/_chat_client.py index c38b813fdae..13a922394a5 100644 --- a/python/packages/bedrock/agent_framework_bedrock/_chat_client.py +++ b/python/packages/bedrock/agent_framework_bedrock/_chat_client.py @@ -495,7 +495,16 @@ def _prepare_bedrock_messages( else: pending_tool_use_ids.clear() - conversation.append({"role": role, "content": content_blocks}) + # Coalesce adjacent user-role turns. Context providers (e.g. the Bedrock + # Knowledge Base provider) inject retrieved passages as separate user + # messages, which would otherwise sit next to the real user input and + # violate Bedrock's role-alternation requirement. Merging their content + # blocks into a single user turn keeps the conversation valid. Assistant + # turns are intentionally not merged to preserve tool-use/tool-result pairing. + if role == "user" and conversation and conversation[-1]["role"] == "user": + conversation[-1]["content"].extend(content_blocks) + else: + conversation.append({"role": role, "content": content_blocks}) return prompts, conversation diff --git a/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py index f408df3c266..a47bfbc1a87 100644 --- a/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py +++ b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base_provider.py @@ -8,7 +8,7 @@ import logging from typing import TYPE_CHECKING, Any, Optional -from agent_framework import AgentSession, ContextProvider, SessionContext +from agent_framework import AgentSession, ContextProvider, Message, SessionContext from agent_framework._telemetry import get_user_agent, mark_feature_used if TYPE_CHECKING: @@ -99,7 +99,7 @@ async def before_run( Called automatically before each model invocation. Extracts the user's query from input messages, retrieves relevant passages, and adds them - as instructions to the context (prepended to the system prompt). + as a delimited user-role message (untrusted external content). Args: agent: The agent running this invocation. @@ -127,12 +127,16 @@ async def before_run( if not retrieved_context: return - # Inject retrieved context as instructions rather than a message. - # Adding it as a separate user message would produce consecutive user - # roles when the agent appends the real input (SessionContext.get_messages - # with include_input=True), which Bedrock Converse rejects since roles must - # alternate. Instructions are prepended to the system context and avoid this. - context.extend_instructions(self.source_id, f"{self.context_prompt}\n\n{retrieved_context}") + # 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.""" diff --git a/python/packages/bedrock/tests/test_bedrock_client.py b/python/packages/bedrock/tests/test_bedrock_client.py index 6d339ae3c59..f6e4d21a758 100644 --- a/python/packages/bedrock/tests/test_bedrock_client.py +++ b/python/packages/bedrock/tests/test_bedrock_client.py @@ -381,6 +381,43 @@ 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_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 index 41b048b4f27..016342dbb36 100644 --- a/python/packages/bedrock/tests/test_bedrock_knowledge_base.py +++ b/python/packages/bedrock/tests/test_bedrock_knowledge_base.py @@ -227,8 +227,8 @@ def test_before_run_injects_context(self): input_messages=[Message(role="user", contents=["What is our policy?"])], ) - # Verify instructions are empty before - assert len(context.instructions) == 0 + # Verify context_messages is empty before + assert len(context.context_messages) == 0 # Run before_run asyncio.run(provider.before_run( @@ -238,11 +238,14 @@ def test_before_run_injects_context(self): state={}, )) - # Verify context was injected as instructions (avoids consecutive user-role - # issue with BedrockChatClient which requires alternating roles) - assert len(context.instructions) == 1 - assert "Relevant passage" in context.instructions[0] - assert "s3://b/doc" in context.instructions[0] + # 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 @@ -263,4 +266,4 @@ def test_before_run_skips_empty_input(self): # Should not call retrieve mock_client.retrieve.assert_not_called() - assert len(context.instructions) == 0 + 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]] From dfd0d2fa14031e6bdb1e09bfaf9d0f41256d8eed Mon Sep 17 00:00:00 2001 From: PVidyadhar Date: Wed, 9 Sep 2026 08:27:43 +0000 Subject: [PATCH 4/5] fix: correct agentic result parsing + scope user-message coalescing Addresses second Copilot review on #8173: 1. AgenticRetrieveStream results use a different schema (content/metadata/ sourceRetriever) than standard Retrieve (score/location). Previously every agentic result was normalized to score 0 with a blank source. Now parse the source URI from metadata._source_uri and omit the score (managed reranking does not expose one); the formatter only renders a score when present. Updated the agentic test mock to the real SDK schema. 2. Restrict _prepare_bedrock_messages coalescing to messages whose ORIGINAL role is 'user', so tool-result turns (role='tool', which map to Bedrock 'user') are never merged into a preceding user text turn. This keeps function-call/tool-result serialization unchanged. Added a regression test for the tool-call/tool-result path. Verified E2E against live KB: agentic results show real source URLs and no fabricated scores. 68 unit tests pass. --- .../agent_framework_bedrock/_chat_client.py | 25 +++++++++++++------ .../_knowledge_base.py | 16 +++++++++--- .../bedrock/tests/test_bedrock_client.py | 23 +++++++++++++++++ .../tests/test_bedrock_knowledge_base.py | 10 +++++++- 4 files changed, 62 insertions(+), 12 deletions(-) diff --git a/python/packages/bedrock/agent_framework_bedrock/_chat_client.py b/python/packages/bedrock/agent_framework_bedrock/_chat_client.py index 13a922394a5..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,16 +499,23 @@ def _prepare_bedrock_messages( else: pending_tool_use_ids.clear() - # Coalesce adjacent user-role turns. Context providers (e.g. the Bedrock - # Knowledge Base provider) inject retrieved passages as separate user - # messages, which would otherwise sit next to the real user input and - # violate Bedrock's role-alternation requirement. Merging their content - # blocks into a single user turn keeps the conversation valid. Assistant - # turns are intentionally not merged to preserve tool-use/tool-result pairing. - if role == "user" and conversation and conversation[-1]["role"] == "user": + # 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 index 5c7f546db4b..febd10e1a02 100644 --- a/python/packages/bedrock/agent_framework_bedrock/_knowledge_base.py +++ b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base.py @@ -155,10 +155,15 @@ def _agentic_retrieve(self, query: str) -> list[dict[str, Any]]: 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": _get_source_uri(r), - "score": r.get("score", 0), + "source": metadata.get("_source_uri", ""), + "score": None, }) return results @@ -187,6 +192,9 @@ def _format_results(results: list[dict[str, Any]]) -> str: 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}") + 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/tests/test_bedrock_client.py b/python/packages/bedrock/tests/test_bedrock_client.py index f6e4d21a758..4239f40f169 100644 --- a/python/packages/bedrock/tests/test_bedrock_client.py +++ b/python/packages/bedrock/tests/test_bedrock_client.py @@ -418,6 +418,29 @@ def test_prepare_bedrock_messages_does_not_coalesce_across_assistant() -> None: 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 index 016342dbb36..57a32e78254 100644 --- a/python/packages/bedrock/tests/test_bedrock_knowledge_base.py +++ b/python/packages/bedrock/tests/test_bedrock_knowledge_base.py @@ -75,7 +75,13 @@ def test_agentic_retrieve_success(self): mock_client.agentic_retrieve_stream.return_value = { "stream": [ {"result": {"results": [ - {"content": {"text": "Agentic result"}, "score": 0.99, "location": {"s3Location": {"uri": "s3://b/doc"}}}, + # 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"}, + }, ]}} ] } @@ -89,6 +95,8 @@ def test_agentic_retrieve_success(self): 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 mock_client.retrieve.assert_not_called() def test_client_uses_get_user_agent(self): From 8da7c321d54f3e407a3eeb8f15c550d0aadb502e Mon Sep 17 00:00:00 2001 From: PVidyadhar Date: Wed, 9 Sep 2026 08:39:20 +0000 Subject: [PATCH 5/5] fix: disable response generation in agentic retrieve Addresses third Copilot review on #8173: - AgenticRetrieveStream defaults to generating a response (verified: 331 streamed responseEvents when omitted vs 0 with generateResponse=False). The tool only formats retrieval passages and discards generation, so pass generateResponse=False to avoid unnecessary model generation latency/cost. - Added a test asserting generateResponse=False is sent. - PR description updated separately to match the actual implementation (user-role injection + serializer coalescing, not extend_instructions). 68 unit tests pass; verified generateResponse behavior against live API. --- .../bedrock/agent_framework_bedrock/_knowledge_base.py | 5 +++++ python/packages/bedrock/tests/test_bedrock_knowledge_base.py | 2 ++ 2 files changed, 7 insertions(+) diff --git a/python/packages/bedrock/agent_framework_bedrock/_knowledge_base.py b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base.py index febd10e1a02..8976ccb5db6 100644 --- a/python/packages/bedrock/agent_framework_bedrock/_knowledge_base.py +++ b/python/packages/bedrock/agent_framework_bedrock/_knowledge_base.py @@ -138,6 +138,11 @@ 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": { diff --git a/python/packages/bedrock/tests/test_bedrock_knowledge_base.py b/python/packages/bedrock/tests/test_bedrock_knowledge_base.py index 57a32e78254..79f74299588 100644 --- a/python/packages/bedrock/tests/test_bedrock_knowledge_base.py +++ b/python/packages/bedrock/tests/test_bedrock_knowledge_base.py @@ -97,6 +97,8 @@ def test_agentic_retrieve_success(self): 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):