diff --git a/.github/workflows/integration-testing.yml b/.github/workflows/integration-testing.yml index 23f0c09a..569842e2 100644 --- a/.github/workflows/integration-testing.yml +++ b/.github/workflows/integration-testing.yml @@ -124,8 +124,8 @@ jobs: ignore: "" # TODO: expand to full tests_integ/memory once test stability is addressed - group: memory - path: tests_integ/memory/test_controlplane.py tests_integ/memory/test_memory_client.py tests_integ/memory/integrations/test_session_manager.py - timeout: 15 + path: tests_integ/memory/test_controlplane.py tests_integ/memory/test_memory_client.py tests_integ/memory/integrations/strands/memorysessionmanager/test_session_manager.py tests_integ/memory/integrations/strands/memorystore/test_memory_store.py + timeout: 30 extra-deps: "" ignore: "" - group: evaluation @@ -189,7 +189,7 @@ jobs: EXTRA_DEPS: ${{ matrix.extra-deps }} run: | pip install -e . - pip install --no-cache-dir pytest pytest-xdist pytest-order pytest-rerunfailures requests strands-agents uvicorn httpx starlette websockets $EXTRA_DEPS + pip install --no-cache-dir pytest pytest-asyncio pytest-xdist pytest-order pytest-rerunfailures requests strands-agents uvicorn httpx starlette websockets $EXTRA_DEPS - name: Run integration tests env: diff --git a/pyproject.toml b/pyproject.toml index d430d87e..c168435d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -26,8 +26,8 @@ classifiers = [ "Topic :: Software Development :: Libraries :: Python Modules", ] dependencies = [ - "boto3>=1.43.31", - "botocore>=1.43.31", + "boto3>=1.43.35", + "botocore>=1.43.35", "pydantic>=2.0.0,<2.41.3", "urllib3>=1.26.0", "starlette>=0.46.2", @@ -149,7 +149,7 @@ dev = [ "ruff>=0.12.0", "websockets>=14.1", "wheel>=0.45.1", - "strands-agents>=1.20.0", + "strands-agents>=1.46.0", "strands-agents-evals>=0.1.0", "langchain>=1.0.0", "langgraph>=1.0.0", @@ -163,7 +163,7 @@ dev = [ a2a = ["a2a-sdk[http-server]>=0.3,<1.0"] ag-ui = ["ag-ui-protocol>=0.1.10"] strands-agents = [ - "strands-agents>=1.20.0", + "strands-agents>=1.46.0", "mcp>=1.23.0,<2.0.0", ] langgraph = [ diff --git a/src/bedrock_agentcore/memory/integrations/strands/README.md b/src/bedrock_agentcore/memory/integrations/strands/README.md index 72ebe2db..e5f2ab73 100644 --- a/src/bedrock_agentcore/memory/integrations/strands/README.md +++ b/src/bedrock_agentcore/memory/integrations/strands/README.md @@ -1,375 +1,6 @@ -# Strands AgentCore Memory Examples +# Strands AgentCore Memory integrations -This directory contains comprehensive examples demonstrating how to use the Strands AgentCoreMemorySessionManager with Amazon Bedrock AgentCore Memory for persistent conversation storage and intelligent retrieval (Supports STM and LTM). +Choose the guide for the Strands memory integration you want to use: -## Quick Setup - -```bash -pip install 'bedrock-agentcore[strands-agents]' -``` - -or to develop locally: -```bash -git clone https://github.com/aws/bedrock-agentcore-sdk-python.git -cd bedrock-agentcore-sdk-python -uv sync -source .venv/bin/activate -``` - -## Examples Overview - -### 1. Short-Term Memory (STM) -Basic memory functionality for conversation persistence within a session. - -### 2. Long-Term Memory (LTM) -Advanced memory with multiple strategies for user preferences, facts, and session summaries. - ---- - -## Short-Term Memory Example - -### Basic Setup - -```python -import uuid -import boto3 -from datetime import date -from strands import Agent -from bedrock_agentcore.memory import MemoryClient -from bedrock_agentcore.memory.integrations.strands.config import AgentCoreMemoryConfig, RetrievalConfig -from bedrock_agentcore.memory.integrations.strands.session_manager import AgentCoreMemorySessionManager -``` - -### Create a Basic Memory - -```python -client = MemoryClient(region_name="us-east-1") -basic_memory = client.create_memory( - name="BasicTestMemory", - description="Basic memory for testing short-term functionality" -) -print(basic_memory.get('id')) -``` - -### Configure and Use Agent - -```python -MEM_ID = basic_memory.get('id') -ACTOR_ID = "actor_id_test_%s" % datetime.now().strftime("%Y%m%d%H%M%S") -SESSION_ID = "testing_session_id_%s" % datetime.now().strftime("%Y%m%d%H%M%S") - - -# Configure memory -agentcore_memory_config = AgentCoreMemoryConfig( - memory_id=MEM_ID, - session_id=SESSION_ID, - actor_id=ACTOR_ID -) - -# Create session manager -session_manager = AgentCoreMemorySessionManager( - agentcore_memory_config=agentcore_memory_config, - region_name="us-east-1" -) - -# Create agent -agent = Agent( - system_prompt="You are a helpful assistant. Use all you know about the user to provide helpful responses.", - session_manager=session_manager, -) -``` - -### Example Conversation - -```python -agent("I like sushi with tuna") -# Agent remembers this preference - -agent("I like pizza") -# Agent acknowledges both preferences - -agent("What should I buy for lunch today?") -# Agent suggests options based on remembered preferences -``` - ---- - -## Long-Term Memory Example - -### Create LTM Memory with Strategies - -```python -from bedrock_agentcore.memory.integrations.strands.config import AgentCoreMemoryConfig, RetrievalConfig -from bedrock_agentcore.memory.integrations.strands.session_manager import AgentCoreMemorySessionManager -from datetime import datetime - -# Create comprehensive memory with all built-in strategies -client = MemoryClient(region_name="us-east-1") -comprehensive_memory = client.create_memory_and_wait( - name="ComprehensiveAgentMemory", - description="Full-featured memory with all built-in strategies", - strategies=[ - { - "summaryMemoryStrategy": { - "name": "SessionSummarizer", - "namespaceTemplates": ["/summaries/{actorId}/{sessionId}/"] - } - }, - { - "userPreferenceMemoryStrategy": { - "name": "PreferenceLearner", - "namespaceTemplates": ["/preferences/{actorId}/"] - } - }, - { - "semanticMemoryStrategy": { - "name": "FactExtractor", - "namespaceTemplates": ["/facts/{actorId}/"] - } - } - ] -) -MEM_ID = comprehensive_memory.get('id') -ACTOR_ID = "actor_id_test_%s" % datetime.now().strftime("%Y%m%d%H%M%S") -SESSION_ID = "testing_session_id_%s" % datetime.now().strftime("%Y%m%d%H%M%S") - -``` - -### Single Namespace Retrieval - -```python -config = AgentCoreMemoryConfig( - memory_id=MEM_ID, - session_id=SESSION_ID, - actor_id=ACTOR_ID, - retrieval_config={ - "/preferences/{actorId}/": RetrievalConfig( - top_k=5, - relevance_score=0.7 - ) - } -) -session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') -ltm_agent = Agent(session_manager=session_manager) -``` - -### Multiple Namespace Retrieval - -```python -config = AgentCoreMemoryConfig( - memory_id=MEM_ID, - session_id=SESSION_ID, - actor_id=ACTOR_ID, - retrieval_config={ - "/preferences/{actorId}/": RetrievalConfig( - top_k=5, - relevance_score=0.7 - ), - "/facts/{actorId}/": RetrievalConfig( - top_k=10, - relevance_score=0.3 - ), - "/summaries/{actorId}/{sessionId}/": RetrievalConfig( - top_k=5, - relevance_score=0.5 - ) - } -) -session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') -agent_with_multiple_namespaces = Agent(session_manager=session_manager) -``` - ---- - -## Large Payload example processing an Image using the [strands_tools](https://github.com/strands-agents/tools) library - -### Agent with Image Processing - -```python -from strands import Agent, tool -from strands_tools import generate_image, image_reader - -ACTOR_ID = "actor_id_test_%s" % datetime.now().strftime("%Y%m%d%H%M%S") -SESSION_ID = "testing_session_id_%s" % datetime.now().strftime("%Y%m%d%H%M%S") - -config = AgentCoreMemoryConfig( - memory_id=MEM_ID, - session_id=SESSION_ID, - actor_id=ACTOR_ID, -) -session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') -agent_with_tools = Agent( - tools=[image_reader], - system_prompt="You will be provided with a filesystem path to an image. Describe the image in detail.", - session_manager=session_manager, - agent_id='my_test_agent_id' -) -# Use with image -result = agent_with_tools("/path/to/image.png") -``` - ---- - -## Key Configuration Options - -### AgentCoreMemoryConfig Parameters - -- `memory_id`: ID of the Bedrock AgentCore Memory resource -- `session_id`: Unique identifier for the conversation session -- `actor_id`: Unique identifier for the user/actor -- `retrieval_config`: Dictionary mapping namespaces to RetrievalConfig objects -- `batch_size`: Number of messages to buffer before sending to AgentCore Memory (1-100, default: 1). A value of 1 sends immediately (no batching). -- `default_metadata`: Optional dictionary of key-value metadata to attach to every message event. Maximum 15 total keys per event (including internal keys). Example: `{"location": {"stringValue": "NYC"}}` -- `metadata_provider`: Optional callable returning a metadata dictionary. Called at each event creation for dynamic values (e.g., traceId). Merged after `default_metadata`. - -### RetrievalConfig Parameters - -- `top_k`: Number of top results to retrieve (default: 5) -- `relevance_score`: Minimum relevance threshold (0.0-1.0) - -### Memory Strategies -https://docs.aws.amazon.com/bedrock-agentcore/latest/devguide/memory-strategies.html - -1. **summaryMemoryStrategy**: Summarizes conversation sessions -2. **userPreferenceMemoryStrategy**: Learns and stores user preferences -3. **semanticMemoryStrategy**: Extracts and stores factual information - -### Namespace Patterns - -- `/preferences/{actorId}/`: User-specific preferences -- `/facts/{actorId}/`: User-specific facts -- `/summaries/{actorId}/{sessionId}/`: Session-specific summaries - - ---- - -## Event Metadata - -You can attach custom key-value metadata to every message event. This is useful for tagging -conversations with contextual information (e.g., location, project, case type) that can later -be used to filter events with `list_events`. - -### Default Metadata (applied to all messages) - -```python -config = AgentCoreMemoryConfig( - memory_id=MEM_ID, - session_id=SESSION_ID, - actor_id=ACTOR_ID, - default_metadata={ - "project": "atlas", - "env": "production", - }, -) -session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') -agent = Agent(session_manager=session_manager) -agent("Hello!") # This event will have project=atlas and env=production metadata -``` - -> Plain strings are auto-wrapped to `{"stringValue": "..."}`. The explicit form -> `{"project": {"stringValue": "atlas"}}` also works. - -### Dynamic Metadata (metadata_provider) - -For values that change per invocation (e.g., traceId for Langfuse), use `metadata_provider` — -a callable invoked at each event creation: - -```python -from langfuse.decorators import langfuse_context - -def get_trace_metadata(): - return {"traceId": langfuse_context.get_current_trace_id() or ""} - -config = AgentCoreMemoryConfig( - memory_id=MEM_ID, - session_id=SESSION_ID, - actor_id=ACTOR_ID, - metadata_provider=get_trace_metadata, -) -session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') -agent = Agent(session_manager=session_manager) -agent("Hello!") # Event gets the current traceId automatically -``` - -### Per-call Metadata - -You can also pass metadata on individual `create_message` calls. Per-call metadata is merged -with `default_metadata` and `metadata_provider` (per-call values override both for the same key): - -```python -session_manager.create_message( - session_id, agent_id, message, - metadata={"priority": "high"}, -) -``` - -> **Note:** The API allows a maximum of 15 metadata key-value pairs per event. -> The keys `stateType` and `agentId` are reserved for internal use. - ---- - -## Message Batching - -When `batch_size` is greater than 1, messages are buffered in memory and sent to AgentCore Memory -in a single API call once the buffer reaches the configured size. This reduces the number of API -requests in high-throughput conversations. - -> **Important:** When using `batch_size > 1`, you **must** use a `with` block or call `close()` -> when the session is complete. Otherwise, any buffered messages that have not yet reached the -> batch threshold will be lost. - -### Recommended: Context Manager - -```python -config = AgentCoreMemoryConfig( - memory_id=MEM_ID, - session_id=SESSION_ID, - actor_id=ACTOR_ID, - batch_size=10, # Buffer up to 10 messages before sending -) - -# The `with` block guarantees all buffered messages are flushed on exit -with AgentCoreMemorySessionManager(config, region_name='us-east-1') as session_manager: - agent = Agent( - system_prompt="You are a helpful assistant.", - session_manager=session_manager, - ) - agent("Hello!") - agent("Tell me about AWS") -# All remaining buffered messages are automatically flushed here -``` - -### Alternative: Explicit close() - -If you cannot use a `with` block, call `close()` manually: - -```python -session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') -try: - agent = Agent( - system_prompt="You are a helpful assistant.", - session_manager=session_manager, - ) - agent("Hello!") -finally: - session_manager.close() # Flush any remaining buffered messages -``` - ---- - -## Important Notes - -### Session Management -- Only **one** agent per session is currently supported -- Creating multiple agents with the same session will show a warning - -### Memory Types -- **STM (Short-Term Memory)**: Basic conversation persistence within a session -- **LTM (Long-Term Memory)**: Advanced memory with multiple strategies for learning user preferences, facts, and summaries - -### Best Practices -- Use unique `session_id` for each conversation -- Use consistent `actor_id` for the same user across sessions -- Configure appropriate `relevance_score` thresholds for your use case -- Test with different `top_k` values to optimize retrieval performance -- When using `batch_size > 1`, always use a `with` block or call `close()` to ensure buffered messages are flushed before the session ends +- [Session manager](memorysessionmanager/README.md) — persist Strands sessions in AgentCore Memory and restore short-term conversation state, with optional long-term memory retrieval. +- [Memory store](memorystore/README.md) — connect AgentCore long-term memory to Strands' native `MemoryManager` for recall and extraction. diff --git a/src/bedrock_agentcore/memory/integrations/strands/__init__.py b/src/bedrock_agentcore/memory/integrations/strands/__init__.py index 26a9d043..41a9c9b2 100644 --- a/src/bedrock_agentcore/memory/integrations/strands/__init__.py +++ b/src/bedrock_agentcore/memory/integrations/strands/__init__.py @@ -1,5 +1,5 @@ -"""Strands integration for Bedrock AgentCore Memory.""" +"""Strands integrations for Bedrock AgentCore Memory.""" -from .converters import MemoryConverter, OpenAIConverseConverter +from .memorysessionmanager.converters import MemoryConverter, OpenAIConverseConverter __all__ = ["MemoryConverter", "OpenAIConverseConverter"] diff --git a/src/bedrock_agentcore/memory/integrations/strands/bedrock_converter.py b/src/bedrock_agentcore/memory/integrations/strands/bedrock_converter.py index 90a26617..14f733db 100644 --- a/src/bedrock_agentcore/memory/integrations/strands/bedrock_converter.py +++ b/src/bedrock_agentcore/memory/integrations/strands/bedrock_converter.py @@ -1,108 +1,12 @@ -"""Bedrock AgentCore Memory conversion utilities.""" +"""Compatibility alias for the Strands AgentCore Memory converter.""" -import json -import logging -from typing import Any, Tuple +import sys +from typing import TYPE_CHECKING -from strands.types.session import SessionMessage +from .memorysessionmanager import bedrock_converter as _canonical_module -logger = logging.getLogger(__name__) +if TYPE_CHECKING: + from .memorysessionmanager.bedrock_converter import CONVERSATIONAL_MAX_SIZE as CONVERSATIONAL_MAX_SIZE + from .memorysessionmanager.bedrock_converter import AgentCoreMemoryConverter as AgentCoreMemoryConverter -# Bedrock AgentCore Data Plane conversational payload text max is 100000 chars. -# Ref: https://docs.aws.amazon.com/cli/latest/reference/bedrock-agentcore/create-event.html -CONVERSATIONAL_MAX_SIZE = 100000 - - -class AgentCoreMemoryConverter: - """Handles conversion between Strands and Bedrock AgentCore Memory formats.""" - - @staticmethod - def _filter_empty_text(message: dict) -> dict: - """The Bedrock Converse API can't take empty text as input. So we need to filter out empty text.""" - content = message.get("content", []) - filtered_content = [item for item in content if "text" not in item or item.get("text", "").strip() != ""] - return {**message, "content": filtered_content} - - @staticmethod - def message_to_payload(session_message: SessionMessage) -> list[Tuple[str, str]]: - """Convert a SessionMessage to Bedrock AgentCore Memory message format. - - Args: - session_message (SessionMessage): The session message to convert. - - Returns: - list[Tuple[str, str]]: list of (text, role) tuples for Bedrock AgentCore Memory. - Returns empty list if message has no content after filtering. - """ - # First convert to dict (which encodes bytes to base64), - # then filter empty text on the encoded version - session_dict = session_message.to_dict() - filtered_message = AgentCoreMemoryConverter._filter_empty_text(session_dict["message"]) - if not filtered_message.get("content"): - logger.debug("Skipping message with no content after filtering empty text") - return [] - session_dict["message"] = filtered_message - return [(json.dumps(session_dict), filtered_message["role"])] - - @staticmethod - def events_to_messages(events: list[dict[str, Any]]) -> list[SessionMessage]: - """Convert Bedrock AgentCore Memory events to SessionMessages. - - Args: - events (list[dict[str, Any]]): list of events from Bedrock AgentCore Memory. - Each individual event looks as follows: - ``` - { - "memoryId": "unique_mem_id", - "actorId": "actor_id", - "sessionId": "session_id", - "eventId": "0000001756147154000#ffa53e54", - "eventTimestamp": datetime.datetime(2025, 8, 25, 15, 12, 34, tzinfo=tzlocal()), - "payload": [ - { - "conversational": { - "content": {"text": "What is the weather?"}, - "role": "USER", - } - } - ], - "branch": {"name": "main"}, - } - ``` - - Returns: - list[SessionMessage]: list of SessionMessage objects. - """ - messages = [] - for event in reversed(events): - for payload_item in event.get("payload", []): - if "conversational" in payload_item: - conv = payload_item["conversational"] - session_msg = SessionMessage.from_dict(json.loads(conv["content"]["text"])) - session_msg.message = AgentCoreMemoryConverter._filter_empty_text(session_msg.message) - if session_msg.message.get("content"): - messages.append(session_msg) - elif "blob" in payload_item: - try: - blob_data = json.loads(payload_item["blob"]) - if isinstance(blob_data, (tuple, list)) and len(blob_data) == 2: - try: - session_msg = SessionMessage.from_dict(json.loads(blob_data[0])) - session_msg.message = AgentCoreMemoryConverter._filter_empty_text(session_msg.message) - if session_msg.message.get("content"): - messages.append(session_msg) - except (json.JSONDecodeError, ValueError): - logger.error("This is not a SessionMessage but just a blob message. Ignoring") - except (json.JSONDecodeError, ValueError): - logger.error("Failed to parse blob content: %s", payload_item) - return messages - - @staticmethod - def total_length(message: tuple[str, str]) -> int: - """Calculate total length of a message tuple.""" - return sum(len(text) for text in message) - - @staticmethod - def exceeds_conversational_limit(message: tuple[str, str]) -> bool: - """Check if message exceeds conversational size limit.""" - return AgentCoreMemoryConverter.total_length(message) >= CONVERSATIONAL_MAX_SIZE +sys.modules[__name__] = _canonical_module diff --git a/src/bedrock_agentcore/memory/integrations/strands/config.py b/src/bedrock_agentcore/memory/integrations/strands/config.py index a42e824e..07ef5811 100644 --- a/src/bedrock_agentcore/memory/integrations/strands/config.py +++ b/src/bedrock_agentcore/memory/integrations/strands/config.py @@ -1,107 +1,14 @@ -"""Configuration for AgentCore Memory Session Manager.""" +"""Compatibility alias for the Strands AgentCore Memory session manager configuration.""" -from enum import Enum -from typing import Any, Callable, Dict, Optional +import sys +from typing import TYPE_CHECKING -from pydantic import BaseModel, Field, field_validator +from .memorysessionmanager import config as _canonical_module +if TYPE_CHECKING: + from .memorysessionmanager.config import AgentCoreMemoryConfig as AgentCoreMemoryConfig + from .memorysessionmanager.config import PersistenceMode as PersistenceMode + from .memorysessionmanager.config import RetrievalConfig as RetrievalConfig + from .memorysessionmanager.config import normalize_metadata as normalize_metadata -def normalize_metadata(raw: Dict[str, Any]) -> Dict[str, Any]: - """Normalize metadata values: plain strings become {"stringValue": value}.""" - return {k: {"stringValue": v} if isinstance(v, str) else v for k, v in raw.items()} - - -class PersistenceMode(str, Enum): - """Controls what gets persisted to AgentCore Memory. - - Attributes: - FULL: Persist everything (session, agent state, messages) to AgentCore Memory. Default behavior. - NONE: Disable all persistence. Local session/agent state management and memory injection - (LTM retrieval) still work, but no create_event calls are made to AgentCore Memory. - """ - - FULL = "FULL" - NONE = "NONE" - - -class RetrievalConfig(BaseModel): - """Configuration for memory retrieval operations. - - Attributes: - top_k: Number of top-scoring records to return from semantic search (default: 10) - relevance_score: Relevance score to filter responses from semantic search (default: 0.2) - strategy_id: Optional parameter to filter memory strategies (default: None) - initialization_query: Optional custom query for initialization retrieval (default: None) - """ - - top_k: int = Field(default=10, gt=0, le=1000) - relevance_score: float = Field(default=0.2, ge=0.0, le=1.0) - strategy_id: Optional[str] = None - initialization_query: Optional[str] = None - - -class AgentCoreMemoryConfig(BaseModel): - """Configuration for AgentCore Memory Session Manager. - - Attributes: - memory_id: Required Bedrock AgentCore Memory ID - session_id: Required unique ID for the session - actor_id: Required unique ID for the agent instance/user - retrieval_config: Optional dictionary mapping namespaces to retrieval configurations - batch_size: Number of messages to batch before sending to AgentCore Memory. - Default of 1 means immediate sending (no batching). Max 100. - flush_interval_seconds: Optional interval in seconds for automatic buffer flushing. - Useful for long-running agents to ensure messages are persisted regularly. - Default is None (disabled). - context_tag: XML tag name used to wrap retrieved memory context injected into messages. - Default is "user_context". - filter_restored_tool_context: When True, strip historical toolUse/toolResult blocks from - restored messages before loading them into Strands runtime memory. Default is False. - default_metadata: Optional default metadata key-value pairs to attach to every message event. - Merged with any per-call metadata. Maximum 15 total keys per event (including internal keys). - Accepts plain strings (auto-wrapped) or explicit MetadataValue dicts. - Example: {"location": "NYC"} or {"location": {"stringValue": "NYC"}} - metadata_provider: Optional callable that returns metadata key-value pairs. Called at each - event creation, so it can return dynamic values (e.g. current traceId). The returned - dict is merged after default_metadata but before per-call metadata. - Accepts plain strings (auto-wrapped) or explicit MetadataValue dicts. - persistence_mode: Controls what gets persisted to AgentCore Memory. - FULL (default): persist everything. NONE: disable all persistence while keeping - local state management and memory injection working. - async_mode: When True, the session manager registers async hook callbacks that - offload the per-turn boto3 calls (append_message, sync_agent, - retrieve_customer_context, and buffer flushes) to a thread via - asyncio.to_thread, keeping the asyncio event loop unblocked. Intended for - async agent runtimes (e.g. Agent.stream_async() in a WebSocket server). - Default is False (existing synchronous behavior, unchanged). - - Requires async invocation (stream_async / invoke_async). Sync agent() calls - will raise RuntimeError from Strands' hook registry because it refuses to - dispatch coroutine callbacks through the sync path. - - Note: this does NOT cover agent initialization. Strands disallows async - callbacks for AgentInitializedEvent, so the read_session / read_agent / - list_messages calls that run during Agent(...) construction still block - the calling thread. If that matters, construct the Agent off-loop - (e.g. `await asyncio.to_thread(Agent, ...)`). - """ - - memory_id: str = Field(min_length=1) - session_id: str = Field(min_length=1) - actor_id: str = Field(min_length=1) - retrieval_config: Optional[Dict[str, RetrievalConfig]] = None - batch_size: int = Field(default=1, ge=1, le=100) - flush_interval_seconds: Optional[float] = Field(default=None, gt=0) - context_tag: str = Field(default="user_context", min_length=1) - filter_restored_tool_context: bool = Field(default=False) - default_metadata: Optional[Dict[str, Any]] = None - metadata_provider: Optional[Callable[[], Dict[str, Any]]] = None - persistence_mode: PersistenceMode = Field(default=PersistenceMode.FULL) - async_mode: bool = Field(default=False) - - @field_validator("default_metadata", mode="before") - @classmethod - def _normalize_default_metadata(cls, v: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]: - if v is None: - return None - return normalize_metadata(v) +sys.modules[__name__] = _canonical_module diff --git a/src/bedrock_agentcore/memory/integrations/strands/converters/__init__.py b/src/bedrock_agentcore/memory/integrations/strands/converters/__init__.py index 56d3093e..69e59693 100644 --- a/src/bedrock_agentcore/memory/integrations/strands/converters/__init__.py +++ b/src/bedrock_agentcore/memory/integrations/strands/converters/__init__.py @@ -1,9 +1,16 @@ -"""Converters for Strands <-> STM message formats.""" +"""Compatibility aliases for Strands session manager converters.""" -from .openai import OpenAIConverseConverter -from .protocol import MemoryConverter +import sys +from typing import TYPE_CHECKING -__all__ = [ - "OpenAIConverseConverter", - "MemoryConverter", -] +from ..memorysessionmanager import converters as _canonical_module +from ..memorysessionmanager.converters import openai as _openai_module +from ..memorysessionmanager.converters import protocol as _protocol_module + +if TYPE_CHECKING: + from ..memorysessionmanager.converters import MemoryConverter as MemoryConverter + from ..memorysessionmanager.converters import OpenAIConverseConverter as OpenAIConverseConverter + +sys.modules[f"{__name__}.openai"] = _openai_module +sys.modules[f"{__name__}.protocol"] = _protocol_module +sys.modules[__name__] = _canonical_module diff --git a/src/bedrock_agentcore/memory/integrations/strands/converters/openai.py b/src/bedrock_agentcore/memory/integrations/strands/converters/openai.py index e5acfc4f..ecf10e08 100644 --- a/src/bedrock_agentcore/memory/integrations/strands/converters/openai.py +++ b/src/bedrock_agentcore/memory/integrations/strands/converters/openai.py @@ -1,191 +1,11 @@ -"""OpenAI-format converter for AgentCore Memory. +"""Compatibility alias for the OpenAI Converse session manager converter.""" -Converts between Strands SessionMessages (Strands-native message shape) and OpenAI message format -stored in AgentCore Memory STM events. -""" +import sys +from typing import TYPE_CHECKING -import json -import logging -from typing import Any, Tuple +from ..memorysessionmanager.converters import openai as _canonical_module -from strands.types.session import SessionMessage +if TYPE_CHECKING: + from ..memorysessionmanager.converters.openai import OpenAIConverseConverter as OpenAIConverseConverter -from .protocol import exceeds_conversational_limit - -logger = logging.getLogger(__name__) - - -def _bedrock_to_openai(message: dict) -> dict: - """Convert a Strands-native message dict to OpenAI message format.""" - role = message.get("role", "user") - content = message.get("content", []) - - if content and "toolResult" in content[0]: - tool_result = content[0]["toolResult"] - text_parts = [c.get("text", "") for c in tool_result.get("content", []) if "text" in c] - result = { - "role": "tool", - "tool_call_id": tool_result["toolUseId"], - "content": "\n".join(text_parts), - } - if "status" in tool_result: - result["status"] = tool_result["status"] - return result - - text_parts = [] - tool_calls = [] - reasoning_blocks: list[dict[str, Any]] = [] - for item in content: - if "text" in item: - text_value = item.get("text") - if isinstance(text_value, str): - text = text_value.strip() - if text: - text_parts.append(text) - elif "reasoningContent" in item: - # OpenAI message shape does not have a stable multi-turn reasoning block field. - # Preserve original block(s) in storage-only extension field for lossless restore. - reasoning_blocks.append(item) - elif "toolUse" in item: - tu = item["toolUse"] - tool_calls.append( - { - "id": tu["toolUseId"], - "type": "function", - "function": { - "name": tu["name"], - "arguments": json.dumps(tu.get("input", {})), - }, - } - ) - - result: dict[str, Any] = {"role": role} - - if tool_calls: - result["content"] = "\n".join(text_parts) if text_parts else None - result["tool_calls"] = tool_calls - else: - result["content"] = "\n".join(text_parts) if text_parts else "" - - if reasoning_blocks: - result["_strands_reasoning_content"] = reasoning_blocks - - return result - - -def _openai_to_bedrock(openai_msg: dict) -> dict: - """Convert an OpenAI message dict to Strands-native message shape.""" - role = openai_msg.get("role", "user") - content_items: list[dict[str, Any]] = [] - reasoning_items: list[dict[str, Any]] = [] - - if role == "tool": - tool_result: dict[str, Any] = { - "toolUseId": openai_msg["tool_call_id"], - "content": [{"text": openai_msg.get("content", "")}], - } - if "status" in openai_msg: - tool_result["status"] = openai_msg["status"] - return { - "role": "user", - "content": [{"toolResult": tool_result}], - } - - if role == "system": - return { - "role": "user", - "content": [{"text": openai_msg.get("content", "")}], - } - - text_content = openai_msg.get("content") - if text_content and isinstance(text_content, str): - content_items.append({"text": text_content}) - - for tc in openai_msg.get("tool_calls", []): - fn = tc.get("function", {}) - args_str = fn.get("arguments", "{}") - try: - args = json.loads(args_str) - except (json.JSONDecodeError, ValueError): - args = {} - content_items.append( - { - "toolUse": { - "toolUseId": tc["id"], - "name": fn["name"], - "input": args, - } - } - ) - - for rc in openai_msg.get("_strands_reasoning_content", []): - if isinstance(rc, dict) and "reasoningContent" in rc: - reasoning_items.append(rc) - - bedrock_role = "assistant" if role == "assistant" else "user" - - # Reasoning blocks MUST come first per Bedrock API: - # "If an assistant message contains any thinking blocks, the first block must be thinking." - return {"role": bedrock_role, "content": reasoning_items + content_items} - - -class OpenAIConverseConverter: - """Converts between Strands SessionMessages and OpenAI message format in STM.""" - - @staticmethod - def message_to_payload(session_message: SessionMessage) -> list[Tuple[str, str]]: - """Convert a SessionMessage (Strands-native shape) to OpenAI-format STM payload.""" - message = session_message.message - content = message.get("content", []) - if not content: - return [] - - has_non_empty = any( - (isinstance(item.get("text"), str) and item["text"].strip()) or "toolUse" in item or "toolResult" in item - for item in content - ) - if not has_non_empty: - return [] - - openai_msg = _bedrock_to_openai(message) - role = openai_msg.get("role", "user") - return [(json.dumps(openai_msg), role)] - - @staticmethod - def events_to_messages(events: list[dict[str, Any]]) -> list[SessionMessage]: - """Convert STM events containing OpenAI-format messages to SessionMessages.""" - messages: list[SessionMessage] = [] - - for event in reversed(events): - for payload_item in event.get("payload", []): - openai_msg = None - - if "conversational" in payload_item: - conv = payload_item["conversational"] - try: - openai_msg = json.loads(conv["content"]["text"]) - except (json.JSONDecodeError, KeyError, ValueError): - logger.error("Failed to parse conversational payload as OpenAI message") - continue - - elif "blob" in payload_item: - try: - blob_data = json.loads(payload_item["blob"]) - if isinstance(blob_data, (tuple, list)) and len(blob_data) == 2: - openai_msg = json.loads(blob_data[0]) - except (json.JSONDecodeError, ValueError): - logger.error("Failed to parse blob payload: %s", payload_item) - continue - - if openai_msg and isinstance(openai_msg, dict): - bedrock_msg = _openai_to_bedrock(openai_msg) - if bedrock_msg.get("content"): - session_msg = SessionMessage(message=bedrock_msg, message_id=0) - messages.append(session_msg) - - return messages - - @staticmethod - def exceeds_conversational_limit(message: tuple[str, str]) -> bool: - """Check if message exceeds conversational payload size limit.""" - return exceeds_conversational_limit(message) +sys.modules[__name__] = _canonical_module diff --git a/src/bedrock_agentcore/memory/integrations/strands/converters/protocol.py b/src/bedrock_agentcore/memory/integrations/strands/converters/protocol.py index 2ae5943e..7a4e92e6 100644 --- a/src/bedrock_agentcore/memory/integrations/strands/converters/protocol.py +++ b/src/bedrock_agentcore/memory/integrations/strands/converters/protocol.py @@ -1,28 +1,13 @@ -"""Shared protocol and utilities for memory converters.""" +"""Compatibility alias for the session manager converter protocol.""" -from typing import Any, Protocol, Tuple +import sys +from typing import TYPE_CHECKING -from strands.types.session import SessionMessage +from ..memorysessionmanager.converters import protocol as _canonical_module -CONVERSATIONAL_MAX_SIZE = 100000 +if TYPE_CHECKING: + from ..memorysessionmanager.converters.protocol import CONVERSATIONAL_MAX_SIZE as CONVERSATIONAL_MAX_SIZE + from ..memorysessionmanager.converters.protocol import MemoryConverter as MemoryConverter + from ..memorysessionmanager.converters.protocol import exceeds_conversational_limit as exceeds_conversational_limit - -class MemoryConverter(Protocol): - """Protocol for converting between Strands messages and STM event payloads.""" - - @staticmethod - def message_to_payload(session_message: SessionMessage) -> list[Tuple[str, str]]: - """Convert SessionMessage to STM event payload format.""" - - @staticmethod - def events_to_messages(events: list[dict[str, Any]]) -> list[SessionMessage]: - """Convert STM events to SessionMessages.""" - - @staticmethod - def exceeds_conversational_limit(message: tuple[str, str]) -> bool: - """Check if message exceeds conversational payload size limit.""" - - -def exceeds_conversational_limit(message: tuple[str, str]) -> bool: - """Check if message exceeds the conversational payload size limit.""" - return sum(len(text) for text in message) >= CONVERSATIONAL_MAX_SIZE +sys.modules[__name__] = _canonical_module diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/README.md b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/README.md new file mode 100644 index 00000000..92a55e50 --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/README.md @@ -0,0 +1,375 @@ +# Strands AgentCore Memory Examples + +This directory contains comprehensive examples demonstrating how to use the Strands AgentCoreMemorySessionManager with Amazon Bedrock AgentCore Memory for persistent conversation storage and intelligent retrieval (Supports STM and LTM). + +## Quick Setup + +```bash +pip install 'bedrock-agentcore[strands-agents]' +``` + +or to develop locally: +```bash +git clone https://github.com/aws/bedrock-agentcore-sdk-python.git +cd bedrock-agentcore-sdk-python +uv sync +source .venv/bin/activate +``` + +## Examples Overview + +### 1. Short-Term Memory (STM) +Basic memory functionality for conversation persistence within a session. + +### 2. Long-Term Memory (LTM) +Advanced memory with multiple strategies for user preferences, facts, and session summaries. + +--- + +## Short-Term Memory Example + +### Basic Setup + +```python +import uuid +import boto3 +from datetime import date +from strands import Agent +from bedrock_agentcore.memory import MemoryClient +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.config import AgentCoreMemoryConfig, RetrievalConfig +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager import AgentCoreMemorySessionManager +``` + +### Create a Basic Memory + +```python +client = MemoryClient(region_name="us-east-1") +basic_memory = client.create_memory( + name="BasicTestMemory", + description="Basic memory for testing short-term functionality" +) +print(basic_memory.get('id')) +``` + +### Configure and Use Agent + +```python +MEM_ID = basic_memory.get('id') +ACTOR_ID = "actor_id_test_%s" % datetime.now().strftime("%Y%m%d%H%M%S") +SESSION_ID = "testing_session_id_%s" % datetime.now().strftime("%Y%m%d%H%M%S") + + +# Configure memory +agentcore_memory_config = AgentCoreMemoryConfig( + memory_id=MEM_ID, + session_id=SESSION_ID, + actor_id=ACTOR_ID +) + +# Create session manager +session_manager = AgentCoreMemorySessionManager( + agentcore_memory_config=agentcore_memory_config, + region_name="us-east-1" +) + +# Create agent +agent = Agent( + system_prompt="You are a helpful assistant. Use all you know about the user to provide helpful responses.", + session_manager=session_manager, +) +``` + +### Example Conversation + +```python +agent("I like sushi with tuna") +# Agent remembers this preference + +agent("I like pizza") +# Agent acknowledges both preferences + +agent("What should I buy for lunch today?") +# Agent suggests options based on remembered preferences +``` + +--- + +## Long-Term Memory Example + +### Create LTM Memory with Strategies + +```python +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.config import AgentCoreMemoryConfig, RetrievalConfig +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager import AgentCoreMemorySessionManager +from datetime import datetime + +# Create comprehensive memory with all built-in strategies +client = MemoryClient(region_name="us-east-1") +comprehensive_memory = client.create_memory_and_wait( + name="ComprehensiveAgentMemory", + description="Full-featured memory with all built-in strategies", + strategies=[ + { + "summaryMemoryStrategy": { + "name": "SessionSummarizer", + "namespaceTemplates": ["/summaries/{actorId}/{sessionId}/"] + } + }, + { + "userPreferenceMemoryStrategy": { + "name": "PreferenceLearner", + "namespaceTemplates": ["/preferences/{actorId}/"] + } + }, + { + "semanticMemoryStrategy": { + "name": "FactExtractor", + "namespaceTemplates": ["/facts/{actorId}/"] + } + } + ] +) +MEM_ID = comprehensive_memory.get('id') +ACTOR_ID = "actor_id_test_%s" % datetime.now().strftime("%Y%m%d%H%M%S") +SESSION_ID = "testing_session_id_%s" % datetime.now().strftime("%Y%m%d%H%M%S") + +``` + +### Single Namespace Retrieval + +```python +config = AgentCoreMemoryConfig( + memory_id=MEM_ID, + session_id=SESSION_ID, + actor_id=ACTOR_ID, + retrieval_config={ + "/preferences/{actorId}/": RetrievalConfig( + top_k=5, + relevance_score=0.7 + ) + } +) +session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') +ltm_agent = Agent(session_manager=session_manager) +``` + +### Multiple Namespace Retrieval + +```python +config = AgentCoreMemoryConfig( + memory_id=MEM_ID, + session_id=SESSION_ID, + actor_id=ACTOR_ID, + retrieval_config={ + "/preferences/{actorId}/": RetrievalConfig( + top_k=5, + relevance_score=0.7 + ), + "/facts/{actorId}/": RetrievalConfig( + top_k=10, + relevance_score=0.3 + ), + "/summaries/{actorId}/{sessionId}/": RetrievalConfig( + top_k=5, + relevance_score=0.5 + ) + } +) +session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') +agent_with_multiple_namespaces = Agent(session_manager=session_manager) +``` + +--- + +## Large Payload example processing an Image using the [strands_tools](https://github.com/strands-agents/tools) library + +### Agent with Image Processing + +```python +from strands import Agent, tool +from strands_tools import generate_image, image_reader + +ACTOR_ID = "actor_id_test_%s" % datetime.now().strftime("%Y%m%d%H%M%S") +SESSION_ID = "testing_session_id_%s" % datetime.now().strftime("%Y%m%d%H%M%S") + +config = AgentCoreMemoryConfig( + memory_id=MEM_ID, + session_id=SESSION_ID, + actor_id=ACTOR_ID, +) +session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') +agent_with_tools = Agent( + tools=[image_reader], + system_prompt="You will be provided with a filesystem path to an image. Describe the image in detail.", + session_manager=session_manager, + agent_id='my_test_agent_id' +) +# Use with image +result = agent_with_tools("/path/to/image.png") +``` + +--- + +## Key Configuration Options + +### AgentCoreMemoryConfig Parameters + +- `memory_id`: ID of the Bedrock AgentCore Memory resource +- `session_id`: Unique identifier for the conversation session +- `actor_id`: Unique identifier for the user/actor +- `retrieval_config`: Dictionary mapping namespaces to RetrievalConfig objects +- `batch_size`: Number of messages to buffer before sending to AgentCore Memory (1-100, default: 1). A value of 1 sends immediately (no batching). +- `default_metadata`: Optional dictionary of key-value metadata to attach to every message event. Maximum 15 total keys per event (including internal keys). Example: `{"location": {"stringValue": "NYC"}}` +- `metadata_provider`: Optional callable returning a metadata dictionary. Called at each event creation for dynamic values (e.g., traceId). Merged after `default_metadata`. + +### RetrievalConfig Parameters + +- `top_k`: Number of top results to retrieve (default: 5) +- `relevance_score`: Minimum relevance threshold (0.0-1.0) + +### Memory Strategies +https://docs.aws.amazon.com/bedrock-agentcore/latest/devguide/memory-strategies.html + +1. **summaryMemoryStrategy**: Summarizes conversation sessions +2. **userPreferenceMemoryStrategy**: Learns and stores user preferences +3. **semanticMemoryStrategy**: Extracts and stores factual information + +### Namespace Patterns + +- `/preferences/{actorId}/`: User-specific preferences +- `/facts/{actorId}/`: User-specific facts +- `/summaries/{actorId}/{sessionId}/`: Session-specific summaries + + +--- + +## Event Metadata + +You can attach custom key-value metadata to every message event. This is useful for tagging +conversations with contextual information (e.g., location, project, case type) that can later +be used to filter events with `list_events`. + +### Default Metadata (applied to all messages) + +```python +config = AgentCoreMemoryConfig( + memory_id=MEM_ID, + session_id=SESSION_ID, + actor_id=ACTOR_ID, + default_metadata={ + "project": "atlas", + "env": "production", + }, +) +session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') +agent = Agent(session_manager=session_manager) +agent("Hello!") # This event will have project=atlas and env=production metadata +``` + +> Plain strings are auto-wrapped to `{"stringValue": "..."}`. The explicit form +> `{"project": {"stringValue": "atlas"}}` also works. + +### Dynamic Metadata (metadata_provider) + +For values that change per invocation (e.g., traceId for Langfuse), use `metadata_provider` — +a callable invoked at each event creation: + +```python +from langfuse.decorators import langfuse_context + +def get_trace_metadata(): + return {"traceId": langfuse_context.get_current_trace_id() or ""} + +config = AgentCoreMemoryConfig( + memory_id=MEM_ID, + session_id=SESSION_ID, + actor_id=ACTOR_ID, + metadata_provider=get_trace_metadata, +) +session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') +agent = Agent(session_manager=session_manager) +agent("Hello!") # Event gets the current traceId automatically +``` + +### Per-call Metadata + +You can also pass metadata on individual `create_message` calls. Per-call metadata is merged +with `default_metadata` and `metadata_provider` (per-call values override both for the same key): + +```python +session_manager.create_message( + session_id, agent_id, message, + metadata={"priority": "high"}, +) +``` + +> **Note:** The API allows a maximum of 15 metadata key-value pairs per event. +> The keys `stateType` and `agentId` are reserved for internal use. + +--- + +## Message Batching + +When `batch_size` is greater than 1, messages are buffered in memory and sent to AgentCore Memory +in a single API call once the buffer reaches the configured size. This reduces the number of API +requests in high-throughput conversations. + +> **Important:** When using `batch_size > 1`, you **must** use a `with` block or call `close()` +> when the session is complete. Otherwise, any buffered messages that have not yet reached the +> batch threshold will be lost. + +### Recommended: Context Manager + +```python +config = AgentCoreMemoryConfig( + memory_id=MEM_ID, + session_id=SESSION_ID, + actor_id=ACTOR_ID, + batch_size=10, # Buffer up to 10 messages before sending +) + +# The `with` block guarantees all buffered messages are flushed on exit +with AgentCoreMemorySessionManager(config, region_name='us-east-1') as session_manager: + agent = Agent( + system_prompt="You are a helpful assistant.", + session_manager=session_manager, + ) + agent("Hello!") + agent("Tell me about AWS") +# All remaining buffered messages are automatically flushed here +``` + +### Alternative: Explicit close() + +If you cannot use a `with` block, call `close()` manually: + +```python +session_manager = AgentCoreMemorySessionManager(config, region_name='us-east-1') +try: + agent = Agent( + system_prompt="You are a helpful assistant.", + session_manager=session_manager, + ) + agent("Hello!") +finally: + session_manager.close() # Flush any remaining buffered messages +``` + +--- + +## Important Notes + +### Session Management +- Only **one** agent per session is currently supported +- Creating multiple agents with the same session will show a warning + +### Memory Types +- **STM (Short-Term Memory)**: Basic conversation persistence within a session +- **LTM (Long-Term Memory)**: Advanced memory with multiple strategies for learning user preferences, facts, and summaries + +### Best Practices +- Use unique `session_id` for each conversation +- Use consistent `actor_id` for the same user across sessions +- Configure appropriate `relevance_score` thresholds for your use case +- Test with different `top_k` values to optimize retrieval performance +- When using `batch_size > 1`, always use a `with` block or call `close()` to ensure buffered messages are flushed before the session ends diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/__init__.py b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/__init__.py new file mode 100644 index 00000000..df4c1226 --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/__init__.py @@ -0,0 +1,16 @@ +"""AgentCore Memory session manager integration for Strands.""" + +from .bedrock_converter import AgentCoreMemoryConverter +from .config import AgentCoreMemoryConfig, PersistenceMode, RetrievalConfig +from .converters import MemoryConverter, OpenAIConverseConverter +from .session_manager import AgentCoreMemorySessionManager + +__all__ = [ + "AgentCoreMemoryConfig", + "AgentCoreMemoryConverter", + "AgentCoreMemorySessionManager", + "MemoryConverter", + "OpenAIConverseConverter", + "PersistenceMode", + "RetrievalConfig", +] diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/bedrock_converter.py b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/bedrock_converter.py new file mode 100644 index 00000000..90a26617 --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/bedrock_converter.py @@ -0,0 +1,108 @@ +"""Bedrock AgentCore Memory conversion utilities.""" + +import json +import logging +from typing import Any, Tuple + +from strands.types.session import SessionMessage + +logger = logging.getLogger(__name__) + +# Bedrock AgentCore Data Plane conversational payload text max is 100000 chars. +# Ref: https://docs.aws.amazon.com/cli/latest/reference/bedrock-agentcore/create-event.html +CONVERSATIONAL_MAX_SIZE = 100000 + + +class AgentCoreMemoryConverter: + """Handles conversion between Strands and Bedrock AgentCore Memory formats.""" + + @staticmethod + def _filter_empty_text(message: dict) -> dict: + """The Bedrock Converse API can't take empty text as input. So we need to filter out empty text.""" + content = message.get("content", []) + filtered_content = [item for item in content if "text" not in item or item.get("text", "").strip() != ""] + return {**message, "content": filtered_content} + + @staticmethod + def message_to_payload(session_message: SessionMessage) -> list[Tuple[str, str]]: + """Convert a SessionMessage to Bedrock AgentCore Memory message format. + + Args: + session_message (SessionMessage): The session message to convert. + + Returns: + list[Tuple[str, str]]: list of (text, role) tuples for Bedrock AgentCore Memory. + Returns empty list if message has no content after filtering. + """ + # First convert to dict (which encodes bytes to base64), + # then filter empty text on the encoded version + session_dict = session_message.to_dict() + filtered_message = AgentCoreMemoryConverter._filter_empty_text(session_dict["message"]) + if not filtered_message.get("content"): + logger.debug("Skipping message with no content after filtering empty text") + return [] + session_dict["message"] = filtered_message + return [(json.dumps(session_dict), filtered_message["role"])] + + @staticmethod + def events_to_messages(events: list[dict[str, Any]]) -> list[SessionMessage]: + """Convert Bedrock AgentCore Memory events to SessionMessages. + + Args: + events (list[dict[str, Any]]): list of events from Bedrock AgentCore Memory. + Each individual event looks as follows: + ``` + { + "memoryId": "unique_mem_id", + "actorId": "actor_id", + "sessionId": "session_id", + "eventId": "0000001756147154000#ffa53e54", + "eventTimestamp": datetime.datetime(2025, 8, 25, 15, 12, 34, tzinfo=tzlocal()), + "payload": [ + { + "conversational": { + "content": {"text": "What is the weather?"}, + "role": "USER", + } + } + ], + "branch": {"name": "main"}, + } + ``` + + Returns: + list[SessionMessage]: list of SessionMessage objects. + """ + messages = [] + for event in reversed(events): + for payload_item in event.get("payload", []): + if "conversational" in payload_item: + conv = payload_item["conversational"] + session_msg = SessionMessage.from_dict(json.loads(conv["content"]["text"])) + session_msg.message = AgentCoreMemoryConverter._filter_empty_text(session_msg.message) + if session_msg.message.get("content"): + messages.append(session_msg) + elif "blob" in payload_item: + try: + blob_data = json.loads(payload_item["blob"]) + if isinstance(blob_data, (tuple, list)) and len(blob_data) == 2: + try: + session_msg = SessionMessage.from_dict(json.loads(blob_data[0])) + session_msg.message = AgentCoreMemoryConverter._filter_empty_text(session_msg.message) + if session_msg.message.get("content"): + messages.append(session_msg) + except (json.JSONDecodeError, ValueError): + logger.error("This is not a SessionMessage but just a blob message. Ignoring") + except (json.JSONDecodeError, ValueError): + logger.error("Failed to parse blob content: %s", payload_item) + return messages + + @staticmethod + def total_length(message: tuple[str, str]) -> int: + """Calculate total length of a message tuple.""" + return sum(len(text) for text in message) + + @staticmethod + def exceeds_conversational_limit(message: tuple[str, str]) -> bool: + """Check if message exceeds conversational size limit.""" + return AgentCoreMemoryConverter.total_length(message) >= CONVERSATIONAL_MAX_SIZE diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/config.py b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/config.py new file mode 100644 index 00000000..a42e824e --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/config.py @@ -0,0 +1,107 @@ +"""Configuration for AgentCore Memory Session Manager.""" + +from enum import Enum +from typing import Any, Callable, Dict, Optional + +from pydantic import BaseModel, Field, field_validator + + +def normalize_metadata(raw: Dict[str, Any]) -> Dict[str, Any]: + """Normalize metadata values: plain strings become {"stringValue": value}.""" + return {k: {"stringValue": v} if isinstance(v, str) else v for k, v in raw.items()} + + +class PersistenceMode(str, Enum): + """Controls what gets persisted to AgentCore Memory. + + Attributes: + FULL: Persist everything (session, agent state, messages) to AgentCore Memory. Default behavior. + NONE: Disable all persistence. Local session/agent state management and memory injection + (LTM retrieval) still work, but no create_event calls are made to AgentCore Memory. + """ + + FULL = "FULL" + NONE = "NONE" + + +class RetrievalConfig(BaseModel): + """Configuration for memory retrieval operations. + + Attributes: + top_k: Number of top-scoring records to return from semantic search (default: 10) + relevance_score: Relevance score to filter responses from semantic search (default: 0.2) + strategy_id: Optional parameter to filter memory strategies (default: None) + initialization_query: Optional custom query for initialization retrieval (default: None) + """ + + top_k: int = Field(default=10, gt=0, le=1000) + relevance_score: float = Field(default=0.2, ge=0.0, le=1.0) + strategy_id: Optional[str] = None + initialization_query: Optional[str] = None + + +class AgentCoreMemoryConfig(BaseModel): + """Configuration for AgentCore Memory Session Manager. + + Attributes: + memory_id: Required Bedrock AgentCore Memory ID + session_id: Required unique ID for the session + actor_id: Required unique ID for the agent instance/user + retrieval_config: Optional dictionary mapping namespaces to retrieval configurations + batch_size: Number of messages to batch before sending to AgentCore Memory. + Default of 1 means immediate sending (no batching). Max 100. + flush_interval_seconds: Optional interval in seconds for automatic buffer flushing. + Useful for long-running agents to ensure messages are persisted regularly. + Default is None (disabled). + context_tag: XML tag name used to wrap retrieved memory context injected into messages. + Default is "user_context". + filter_restored_tool_context: When True, strip historical toolUse/toolResult blocks from + restored messages before loading them into Strands runtime memory. Default is False. + default_metadata: Optional default metadata key-value pairs to attach to every message event. + Merged with any per-call metadata. Maximum 15 total keys per event (including internal keys). + Accepts plain strings (auto-wrapped) or explicit MetadataValue dicts. + Example: {"location": "NYC"} or {"location": {"stringValue": "NYC"}} + metadata_provider: Optional callable that returns metadata key-value pairs. Called at each + event creation, so it can return dynamic values (e.g. current traceId). The returned + dict is merged after default_metadata but before per-call metadata. + Accepts plain strings (auto-wrapped) or explicit MetadataValue dicts. + persistence_mode: Controls what gets persisted to AgentCore Memory. + FULL (default): persist everything. NONE: disable all persistence while keeping + local state management and memory injection working. + async_mode: When True, the session manager registers async hook callbacks that + offload the per-turn boto3 calls (append_message, sync_agent, + retrieve_customer_context, and buffer flushes) to a thread via + asyncio.to_thread, keeping the asyncio event loop unblocked. Intended for + async agent runtimes (e.g. Agent.stream_async() in a WebSocket server). + Default is False (existing synchronous behavior, unchanged). + + Requires async invocation (stream_async / invoke_async). Sync agent() calls + will raise RuntimeError from Strands' hook registry because it refuses to + dispatch coroutine callbacks through the sync path. + + Note: this does NOT cover agent initialization. Strands disallows async + callbacks for AgentInitializedEvent, so the read_session / read_agent / + list_messages calls that run during Agent(...) construction still block + the calling thread. If that matters, construct the Agent off-loop + (e.g. `await asyncio.to_thread(Agent, ...)`). + """ + + memory_id: str = Field(min_length=1) + session_id: str = Field(min_length=1) + actor_id: str = Field(min_length=1) + retrieval_config: Optional[Dict[str, RetrievalConfig]] = None + batch_size: int = Field(default=1, ge=1, le=100) + flush_interval_seconds: Optional[float] = Field(default=None, gt=0) + context_tag: str = Field(default="user_context", min_length=1) + filter_restored_tool_context: bool = Field(default=False) + default_metadata: Optional[Dict[str, Any]] = None + metadata_provider: Optional[Callable[[], Dict[str, Any]]] = None + persistence_mode: PersistenceMode = Field(default=PersistenceMode.FULL) + async_mode: bool = Field(default=False) + + @field_validator("default_metadata", mode="before") + @classmethod + def _normalize_default_metadata(cls, v: Optional[Dict[str, Any]]) -> Optional[Dict[str, Any]]: + if v is None: + return None + return normalize_metadata(v) diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/converters/__init__.py b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/converters/__init__.py new file mode 100644 index 00000000..b9f216a8 --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/converters/__init__.py @@ -0,0 +1,6 @@ +"""Converters for Strands and AgentCore short-term memory message formats.""" + +from .openai import OpenAIConverseConverter +from .protocol import MemoryConverter + +__all__ = ["MemoryConverter", "OpenAIConverseConverter"] diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/converters/openai.py b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/converters/openai.py new file mode 100644 index 00000000..e5acfc4f --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/converters/openai.py @@ -0,0 +1,191 @@ +"""OpenAI-format converter for AgentCore Memory. + +Converts between Strands SessionMessages (Strands-native message shape) and OpenAI message format +stored in AgentCore Memory STM events. +""" + +import json +import logging +from typing import Any, Tuple + +from strands.types.session import SessionMessage + +from .protocol import exceeds_conversational_limit + +logger = logging.getLogger(__name__) + + +def _bedrock_to_openai(message: dict) -> dict: + """Convert a Strands-native message dict to OpenAI message format.""" + role = message.get("role", "user") + content = message.get("content", []) + + if content and "toolResult" in content[0]: + tool_result = content[0]["toolResult"] + text_parts = [c.get("text", "") for c in tool_result.get("content", []) if "text" in c] + result = { + "role": "tool", + "tool_call_id": tool_result["toolUseId"], + "content": "\n".join(text_parts), + } + if "status" in tool_result: + result["status"] = tool_result["status"] + return result + + text_parts = [] + tool_calls = [] + reasoning_blocks: list[dict[str, Any]] = [] + for item in content: + if "text" in item: + text_value = item.get("text") + if isinstance(text_value, str): + text = text_value.strip() + if text: + text_parts.append(text) + elif "reasoningContent" in item: + # OpenAI message shape does not have a stable multi-turn reasoning block field. + # Preserve original block(s) in storage-only extension field for lossless restore. + reasoning_blocks.append(item) + elif "toolUse" in item: + tu = item["toolUse"] + tool_calls.append( + { + "id": tu["toolUseId"], + "type": "function", + "function": { + "name": tu["name"], + "arguments": json.dumps(tu.get("input", {})), + }, + } + ) + + result: dict[str, Any] = {"role": role} + + if tool_calls: + result["content"] = "\n".join(text_parts) if text_parts else None + result["tool_calls"] = tool_calls + else: + result["content"] = "\n".join(text_parts) if text_parts else "" + + if reasoning_blocks: + result["_strands_reasoning_content"] = reasoning_blocks + + return result + + +def _openai_to_bedrock(openai_msg: dict) -> dict: + """Convert an OpenAI message dict to Strands-native message shape.""" + role = openai_msg.get("role", "user") + content_items: list[dict[str, Any]] = [] + reasoning_items: list[dict[str, Any]] = [] + + if role == "tool": + tool_result: dict[str, Any] = { + "toolUseId": openai_msg["tool_call_id"], + "content": [{"text": openai_msg.get("content", "")}], + } + if "status" in openai_msg: + tool_result["status"] = openai_msg["status"] + return { + "role": "user", + "content": [{"toolResult": tool_result}], + } + + if role == "system": + return { + "role": "user", + "content": [{"text": openai_msg.get("content", "")}], + } + + text_content = openai_msg.get("content") + if text_content and isinstance(text_content, str): + content_items.append({"text": text_content}) + + for tc in openai_msg.get("tool_calls", []): + fn = tc.get("function", {}) + args_str = fn.get("arguments", "{}") + try: + args = json.loads(args_str) + except (json.JSONDecodeError, ValueError): + args = {} + content_items.append( + { + "toolUse": { + "toolUseId": tc["id"], + "name": fn["name"], + "input": args, + } + } + ) + + for rc in openai_msg.get("_strands_reasoning_content", []): + if isinstance(rc, dict) and "reasoningContent" in rc: + reasoning_items.append(rc) + + bedrock_role = "assistant" if role == "assistant" else "user" + + # Reasoning blocks MUST come first per Bedrock API: + # "If an assistant message contains any thinking blocks, the first block must be thinking." + return {"role": bedrock_role, "content": reasoning_items + content_items} + + +class OpenAIConverseConverter: + """Converts between Strands SessionMessages and OpenAI message format in STM.""" + + @staticmethod + def message_to_payload(session_message: SessionMessage) -> list[Tuple[str, str]]: + """Convert a SessionMessage (Strands-native shape) to OpenAI-format STM payload.""" + message = session_message.message + content = message.get("content", []) + if not content: + return [] + + has_non_empty = any( + (isinstance(item.get("text"), str) and item["text"].strip()) or "toolUse" in item or "toolResult" in item + for item in content + ) + if not has_non_empty: + return [] + + openai_msg = _bedrock_to_openai(message) + role = openai_msg.get("role", "user") + return [(json.dumps(openai_msg), role)] + + @staticmethod + def events_to_messages(events: list[dict[str, Any]]) -> list[SessionMessage]: + """Convert STM events containing OpenAI-format messages to SessionMessages.""" + messages: list[SessionMessage] = [] + + for event in reversed(events): + for payload_item in event.get("payload", []): + openai_msg = None + + if "conversational" in payload_item: + conv = payload_item["conversational"] + try: + openai_msg = json.loads(conv["content"]["text"]) + except (json.JSONDecodeError, KeyError, ValueError): + logger.error("Failed to parse conversational payload as OpenAI message") + continue + + elif "blob" in payload_item: + try: + blob_data = json.loads(payload_item["blob"]) + if isinstance(blob_data, (tuple, list)) and len(blob_data) == 2: + openai_msg = json.loads(blob_data[0]) + except (json.JSONDecodeError, ValueError): + logger.error("Failed to parse blob payload: %s", payload_item) + continue + + if openai_msg and isinstance(openai_msg, dict): + bedrock_msg = _openai_to_bedrock(openai_msg) + if bedrock_msg.get("content"): + session_msg = SessionMessage(message=bedrock_msg, message_id=0) + messages.append(session_msg) + + return messages + + @staticmethod + def exceeds_conversational_limit(message: tuple[str, str]) -> bool: + """Check if message exceeds conversational payload size limit.""" + return exceeds_conversational_limit(message) diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/converters/protocol.py b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/converters/protocol.py new file mode 100644 index 00000000..2ae5943e --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/converters/protocol.py @@ -0,0 +1,28 @@ +"""Shared protocol and utilities for memory converters.""" + +from typing import Any, Protocol, Tuple + +from strands.types.session import SessionMessage + +CONVERSATIONAL_MAX_SIZE = 100000 + + +class MemoryConverter(Protocol): + """Protocol for converting between Strands messages and STM event payloads.""" + + @staticmethod + def message_to_payload(session_message: SessionMessage) -> list[Tuple[str, str]]: + """Convert SessionMessage to STM event payload format.""" + + @staticmethod + def events_to_messages(events: list[dict[str, Any]]) -> list[SessionMessage]: + """Convert STM events to SessionMessages.""" + + @staticmethod + def exceeds_conversational_limit(message: tuple[str, str]) -> bool: + """Check if message exceeds conversational payload size limit.""" + + +def exceeds_conversational_limit(message: tuple[str, str]) -> bool: + """Check if message exceeds the conversational payload size limit.""" + return sum(len(text) for text in message) >= CONVERSATIONAL_MAX_SIZE diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/session_manager.py b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/session_manager.py new file mode 100644 index 00000000..f930b548 --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/session_manager.py @@ -0,0 +1,1343 @@ +"""AgentCore Memory-based session manager for Bedrock AgentCore Memory integration.""" + +import asyncio +import json +import logging +import threading +from concurrent.futures import ThreadPoolExecutor, as_completed +from datetime import datetime, timedelta, timezone +from enum import Enum +from typing import TYPE_CHECKING, Any, Dict, NamedTuple, Optional + +import boto3 +from botocore.config import Config as BotocoreConfig +from strands.experimental.hooks.events import ( + BidiAfterInvocationEvent, + BidiAgentInitializedEvent, + BidiMessageAddedEvent, +) +from strands.experimental.hooks.multiagent.events import ( + AfterMultiAgentInvocationEvent, + AfterNodeCallEvent, + MultiAgentInitializedEvent, +) +from strands.hooks import AfterInvocationEvent, MessageAddedEvent +from strands.hooks.events import AgentInitializedEvent +from strands.hooks.registry import HookRegistry +from strands.session.repository_session_manager import RepositorySessionManager +from strands.session.session_repository import SessionRepository +from strands.types.content import Message +from strands.types.exceptions import SessionException +from strands.types.session import Session, SessionAgent, SessionMessage +from typing_extensions import override + +from bedrock_agentcore.memory.client import MemoryClient +from bedrock_agentcore.memory.models.filters import ( + EventMetadataFilter, + LeftExpression, + MetadataValue, + OperatorType, + RightExpression, +) + +from .bedrock_converter import AgentCoreMemoryConverter +from .config import AgentCoreMemoryConfig, PersistenceMode, RetrievalConfig, normalize_metadata +from .converters import MemoryConverter + +if TYPE_CHECKING: + from strands.agent.agent import Agent + +logger = logging.getLogger(__name__) + +MAX_FETCH_ALL_RESULTS = 10000 + +# Legacy prefixes for backwards compatibility with old events +LEGACY_SESSION_PREFIX = "session_" +LEGACY_AGENT_PREFIX = "agent_" + +# Metadata keys for event identification +STATE_TYPE_KEY = "stateType" +AGENT_ID_KEY = "agentId" + +# Maximum metadata key-value pairs per event (API limit) +MAX_METADATA_KEYS = 15 + +# Reserved internal metadata keys that users cannot override +RESERVED_METADATA_KEYS = frozenset({STATE_TYPE_KEY, AGENT_ID_KEY}) + + +class BufferedMessage(NamedTuple): + """A pre-processed message waiting to be flushed to AgentCore Memory.""" + + session_id: str + messages: list[tuple[str, str]] + is_blob: bool + timestamp: datetime + metadata: Optional[Dict[str, MetadataValue]] = None + + +class StateType(Enum): + """State type for distinguishing session and agent metadata in events.""" + + SESSION = "SESSION" + AGENT = "AGENT" + + +class AgentCoreMemorySessionManager(RepositorySessionManager, SessionRepository): + """AgentCore Memory-based session manager for Bedrock AgentCore Memory integration. + + This session manager integrates Strands agents with Amazon Bedrock AgentCore Memory, + providing seamless synchronization between Strands' session management and Bedrock's + short-term and long-term memory capabilities. + + Key Features: + - Automatic synchronization of conversation messages to Bedrock AgentCore Memory events + - Loading of conversation history from short-term memory during agent initialization + - Integration with long-term memory for context injection into agent state + - Support for custom retrieval configurations per namespace + - Consistent with existing Strands Session managers (such as: FileSessionManager, S3SessionManager) + """ + + def _get_monotonic_timestamp(self, desired_timestamp: Optional[datetime] = None) -> datetime: + """Get a monotonically increasing timestamp for this session. + + Ties are broken at millisecond granularity, which is the resolution + AgentCore Memory stores and orders ``eventTimestamp`` at. + + Args: + desired_timestamp (Optional[datetime]): The desired timestamp. If None, uses current time. + + Returns: + datetime: A timestamp guaranteed to be greater than any previously returned timestamp. + """ + if desired_timestamp is None: + desired_timestamp = datetime.now(timezone.utc) + + # Floor to milliseconds — the resolution AgentCore Memory actually stores + # and orders eventTimestamp at. Comparing at microsecond precision would + # miss two events that fall in the same millisecond but different + # microseconds: they pass the tie check here yet collide once stored. + desired_timestamp = desired_timestamp.replace(microsecond=(desired_timestamp.microsecond // 1000) * 1000) + + with self._timestamp_lock: + if self._last_timestamp is not None and desired_timestamp <= self._last_timestamp: + # Break the tie at millisecond granularity (the service's resolution). + desired_timestamp = self._last_timestamp + timedelta(milliseconds=1) + self._last_timestamp = desired_timestamp + return desired_timestamp + + def __init__( + self, + agentcore_memory_config: AgentCoreMemoryConfig, + region_name: Optional[str] = None, + boto_session: Optional[boto3.Session] = None, + boto_client_config: Optional[BotocoreConfig] = None, + *, + converter: Optional[type[MemoryConverter]] = None, + **kwargs: Any, + ): + """Initialize AgentCoreMemorySessionManager with Bedrock AgentCore Memory. + + Args: + agentcore_memory_config (AgentCoreMemoryConfig): Configuration for AgentCore Memory integration. + region_name (Optional[str], optional): AWS region for Bedrock AgentCore Memory. Defaults to None. + boto_session (Optional[boto3.Session], optional): Optional boto3 session. Defaults to None. + boto_client_config (Optional[BotocoreConfig], optional): Optional boto3 client configuration. + Defaults to None. + converter (Optional[type[MemoryConverter]], optional): Optional custom converter. + If None, native Bedrock/Strands converter is used. + **kwargs (Any): Additional keyword arguments. + """ + self.converter = converter or AgentCoreMemoryConverter + self.config = agentcore_memory_config + self.persistence_mode = agentcore_memory_config.persistence_mode + self.memory_client = MemoryClient(region_name=region_name) + session = boto_session or boto3.Session(region_name=region_name) + self.has_existing_agent = False + + # Instance-scoped monotonic-timestamp state. Per-instance so concurrent + # managers for different sessions in one process do not perturb each + # other's ordering. + self._timestamp_lock = threading.Lock() + self._last_timestamp: Optional[datetime] = None + + # Batching support - stores pre-processed messages + self._message_buffer: list[BufferedMessage] = [] + self._message_lock = threading.Lock() + + # Agent state buffering - stores all agent state updates: (session_id, agent) + self._agent_state_buffer: list[tuple[str, SessionAgent]] = [] + self._agent_state_lock = threading.Lock() + + # Cache for agent created_at timestamps to avoid fetching on every update + self._agent_created_at_cache: dict[str, datetime] = {} + + # Track last synced internal state for each agent (required by parent RepositorySessionManager) + self._last_synced_internal_state: dict[str, Any] = {} + + # Track if this is a new session (required by parent RepositorySessionManager) + self._is_new_session: bool = True + + # Interval-based flushing support + self._flush_timer: Optional[threading.Timer] = None + self._timer_lock = threading.Lock() + self._shutdown = False + + # Add strands-agents to the request user agent + if boto_client_config: + existing_user_agent = getattr(boto_client_config, "user_agent_extra", None) + if existing_user_agent: + new_user_agent = f"{existing_user_agent} strands-agents" + else: + new_user_agent = "strands-agents" + client_config = boto_client_config.merge(BotocoreConfig(user_agent_extra=new_user_agent)) + else: + client_config = BotocoreConfig(user_agent_extra="strands-agents") + + # Override the memory client's boto3 clients + self.memory_client.gmcp_client = session.client( + "bedrock-agentcore-control", region_name=region_name or session.region_name, config=client_config + ) + self.memory_client.gmdp_client = session.client( + "bedrock-agentcore", region_name=region_name or session.region_name, config=client_config + ) + super().__init__(session_id=self.config.session_id, session_repository=self) + + # Start interval-based flush timer if configured + if self.config.flush_interval_seconds: + self._start_flush_timer() + + def _build_metadata( + self, + internal_metadata: Optional[Dict[str, MetadataValue]] = None, + per_call_metadata: Optional[Dict[str, MetadataValue]] = None, + ) -> Optional[Dict[str, MetadataValue]]: + """Build merged metadata from config defaults, provider, per-call overrides, and internal keys. + + Merge precedence (highest wins): + 1. internal_metadata (stateType, agentId) — always wins + 2. per_call_metadata (passed via **kwargs) + 3. metadata_provider() (called at event creation time for dynamic values) + 4. self.config.default_metadata (set at config construction time) + + Args: + internal_metadata: System-reserved metadata (e.g. stateType, agentId). + per_call_metadata: Caller-supplied metadata for a single operation. + + Returns: + Merged metadata dict, or None if empty. + + Raises: + ValueError: If user metadata contains reserved keys or total keys exceed MAX_METADATA_KEYS. + """ + merged: Dict[str, MetadataValue] = {} + + if self.config.default_metadata: + merged.update(self.config.default_metadata) + + if self.config.metadata_provider: + merged.update(normalize_metadata(self.config.metadata_provider())) + + if per_call_metadata: + merged.update(per_call_metadata) + + # Validate user-supplied keys before merging internal keys + user_reserved = RESERVED_METADATA_KEYS & merged.keys() + if user_reserved: + raise ValueError( + f"Metadata keys {user_reserved} are reserved for internal use. Reserved keys: {RESERVED_METADATA_KEYS}" + ) + + if internal_metadata: + merged.update(internal_metadata) + + if len(merged) > MAX_METADATA_KEYS: + raise ValueError(f"Combined metadata has {len(merged)} keys, exceeding the maximum of {MAX_METADATA_KEYS}.") + + return merged or None + + # region SessionRepository interface implementation + def create_session(self, session: Session, **kwargs: Any) -> Session: + """Create a new session in AgentCore Memory. + + Note: AgentCore Memory doesn't have explicit session creation, + so we just validate the session and return it. + + Args: + session (Session): The session to create. + **kwargs (Any): Additional keyword arguments. + + Returns: + Session: The created session. + + Raises: + SessionException: If session ID doesn't match configuration. + """ + if session.session_id != self.config.session_id: + raise SessionException(f"Session ID mismatch: expected {self.config.session_id}, got {session.session_id}") + + if self.persistence_mode is not PersistenceMode.NONE: + event = self.memory_client.gmdp_client.create_event( + memoryId=self.config.memory_id, + actorId=self.config.actor_id, + sessionId=self.session_id, + payload=[ + {"blob": json.dumps(session.to_dict())}, + ], + eventTimestamp=self._get_monotonic_timestamp(), + metadata={STATE_TYPE_KEY: {"stringValue": StateType.SESSION.value}}, + ) + logger.info("Created session: %s with event: %s", session.session_id, event.get("event", {}).get("eventId")) + + return session + + def read_session(self, session_id: str, **kwargs: Any) -> Optional[Session]: + """Read session data. + + AgentCore Memory does not have a `get_session` method. + Which is fine as AgentCore Memory is a managed service we therefore do not need to read/update + the session data. We just return the session object. + + Args: + session_id (str): The session ID to read. + **kwargs (Any): Additional keyword arguments. + + Returns: + Optional[Session]: The session if found, None otherwise. + """ + if session_id != self.config.session_id: + return None + + # 1. Try new approach (metadata filter) + event_metadata = [ + EventMetadataFilter.build_expression( + left_operand=LeftExpression.build(STATE_TYPE_KEY), + operator=OperatorType.EQUALS_TO, + right_operand=RightExpression.build(StateType.SESSION.value), + ) + ] + + events = self.memory_client.list_events( + memory_id=self.config.memory_id, + actor_id=self.config.actor_id, + session_id=session_id, + event_metadata=event_metadata, + max_results=1, + ) + if events: + session_data = json.loads(events[0].get("payload", {})[0].get("blob")) + return Session.from_dict(session_data) + + # 2. Fallback: check for legacy event and migrate + legacy_actor_id = f"{LEGACY_SESSION_PREFIX}{session_id}" + events = self.memory_client.list_events( + memory_id=self.config.memory_id, + actor_id=legacy_actor_id, + session_id=session_id, + max_results=1, + ) + if events: + old_event = events[0] + session_data = json.loads(old_event.get("payload", {})[0].get("blob")) + session = Session.from_dict(session_data) + # Migrate: create new event with metadata, delete old + if self.persistence_mode is not PersistenceMode.NONE: + self.create_session(session) + self.memory_client.gmdp_client.delete_event( + memoryId=self.config.memory_id, + actorId=legacy_actor_id, + sessionId=session_id, + eventId=old_event.get("eventId"), + ) + logger.info("Migrated legacy session event for session: %s", session_id) + return session + + return None + + def delete_session(self, session_id: str, **kwargs: Any) -> None: + """Delete session and all associated data. + + Note: AgentCore Memory doesn't support deletion of events, + so this is a no-op operation. + + Args: + session_id (str): The session ID to delete. + **kwargs (Any): Additional keyword arguments. + """ + logger.warning("Session deletion not supported in AgentCore Memory: %s", session_id) + + def create_agent(self, session_id: str, session_agent: SessionAgent, **kwargs: Any) -> None: + """Create a new agent in the session. + + For AgentCore Memory, we don't need to explicitly create agents; we have Implicit Agent Existence + The agent's existence is inferred from the presence of events/messages in the memory system, + but we validate the session_id matches our config. + + Args: + session_id (str): The session ID to create the agent in. + session_agent (SessionAgent): The agent to create. + **kwargs (Any): Additional keyword arguments. + + Raises: + SessionException: If session ID doesn't match configuration. + """ + if session_id != self.config.session_id: + raise SessionException(f"Session ID mismatch: expected {self.config.session_id}, got {session_id}") + + # Cache the created_at timestamp to avoid re-fetching on updates + if session_agent.created_at: + self._agent_created_at_cache[session_agent.agent_id] = session_agent.created_at + + if self.persistence_mode is PersistenceMode.NONE: + return + + if self.config.batch_size > 1: + # Buffer the agent state events + should_flush = False + with self._agent_state_lock: + self._agent_state_buffer.append((session_id, session_agent)) + should_flush = len(self._agent_state_buffer) >= self.config.batch_size + + # Flush only agent states outside the lock to prevent deadlock + if should_flush: + self._flush_agent_states_only() + + logger.info( + "Buffered agent creation: %s in session: %s", + session_agent.agent_id, + session_id, + ) + else: + # Immediate send when batching is disabled + event = self.memory_client.gmdp_client.create_event( + memoryId=self.config.memory_id, + actorId=self.config.actor_id, + sessionId=self.session_id, + payload=[ + {"blob": json.dumps(session_agent.to_dict())}, + ], + eventTimestamp=self._get_monotonic_timestamp(), + metadata={ + STATE_TYPE_KEY: {"stringValue": StateType.AGENT.value}, + AGENT_ID_KEY: {"stringValue": session_agent.agent_id}, + }, + ) + + logger.info( + "Created agent: %s in session: %s with event %s", + session_agent.agent_id, + session_id, + event.get("event", {}).get("eventId"), + ) + + def read_agent(self, session_id: str, agent_id: str, **kwargs: Any) -> Optional[SessionAgent]: + """Read agent data from AgentCore Memory events. + + We reconstruct the agent state from the conversation history. + + Args: + session_id (str): The session ID to read from. + agent_id (str): The agent ID to read. + **kwargs (Any): Additional keyword arguments. + + Returns: + Optional[SessionAgent]: The agent if found, None otherwise. + """ + if session_id != self.config.session_id: + return None + try: + # 1. Try new approach (metadata filter) + event_metadata = [ + EventMetadataFilter.build_expression( + left_operand=LeftExpression.build(STATE_TYPE_KEY), + operator=OperatorType.EQUALS_TO, + right_operand=RightExpression.build(StateType.AGENT.value), + ), + EventMetadataFilter.build_expression( + left_operand=LeftExpression.build(AGENT_ID_KEY), + operator=OperatorType.EQUALS_TO, + right_operand=RightExpression.build(agent_id), + ), + ] + + events = self.memory_client.list_events( + memory_id=self.config.memory_id, + actor_id=self.config.actor_id, + session_id=session_id, + event_metadata=event_metadata, + max_results=1, + ) + + if events: + agent_data = json.loads(events[0].get("payload", {})[0].get("blob")) + agent = SessionAgent.from_dict(agent_data) + # Cache the created_at timestamp to avoid re-fetching on updates + if agent.created_at: + self._agent_created_at_cache[agent_id] = agent.created_at + return agent + + # 2. Fallback: check for legacy event and migrate + legacy_actor_id = f"{LEGACY_AGENT_PREFIX}{agent_id}" + events = self.memory_client.list_events( + memory_id=self.config.memory_id, + actor_id=legacy_actor_id, + session_id=session_id, + max_results=1, + ) + if events: + old_event = events[0] + agent_data = json.loads(old_event.get("payload", {})[0].get("blob")) + agent = SessionAgent.from_dict(agent_data) + # Migrate: create new event with metadata, delete old + if self.persistence_mode is not PersistenceMode.NONE: + self.create_agent(session_id, agent) + self.memory_client.gmdp_client.delete_event( + memoryId=self.config.memory_id, + actorId=legacy_actor_id, + sessionId=session_id, + eventId=old_event.get("eventId"), + ) + logger.info("Migrated legacy agent event for agent: %s", agent_id) + return agent + + return None + except Exception as e: + logger.error("Failed to read agent %s", e) + return None + + def update_agent(self, session_id: str, session_agent: SessionAgent, **kwargs: Any) -> None: + """Update agent data. + + Args: + session_id (str): The session ID containing the agent. + session_agent (SessionAgent): The agent to update. + **kwargs (Any): Additional keyword arguments. + + Raises: + SessionException: If session ID doesn't match configuration. + """ + agent_id = session_agent.agent_id + + # Verify agent exists and get created_at timestamp if not cached + if agent_id not in self._agent_created_at_cache: + previous_agent = self.read_agent(session_id=session_id, agent_id=agent_id) + if previous_agent is None: + raise SessionException(f"Agent {agent_id} in session {session_id} does not exist") + + # Set created_at from cache before creating the update event + session_agent.created_at = self._agent_created_at_cache[agent_id] + + # Create a new agent event (AgentCore Memory is immutable) + # create_agent will handle batching and caching appropriately + self.create_agent(session_id, session_agent) + + def create_message( + self, session_id: str, agent_id: str, session_message: SessionMessage, **kwargs: Any + ) -> Optional[dict[str, Any]]: + """Create a new message in AgentCore Memory. + + If batch_size > 1, the message is buffered and sent when the buffer reaches batch_size. + Use _flush_messages() or close() to send any remaining buffered messages. + + Args: + session_id (str): The session ID to create the message in. + agent_id (str): The agent ID associated with the message (only here for the interface. + We use the actorId for AgentCore). + session_message (SessionMessage): The message to create. + **kwargs (Any): Additional keyword arguments. + + Returns: + Optional[dict[str, Any]]: The created event data from AgentCore Memory. + Returns empty dict if message is buffered (batch_size > 1). + + Raises: + SessionException: If session ID doesn't match configuration or message creation fails. + + Note: + The returned created message `event` looks like: + ```python + { + "memoryId": "my-mem-id", + "actorId": "user_1", + "sessionId": "test_session_id", + "eventId": "0000001752235548000#97f30a6b", + "eventTimestamp": datetime.datetime(2025, 8, 18, 12, 45, 48, tzinfo=tzlocal()), + "branch": {"name": "main"}, + } + ``` + """ + if session_id != self.config.session_id: + raise SessionException(f"Session ID mismatch: expected {self.config.session_id}, got {session_id}") + + # Convert and check size ONCE (not again at flush) + messages = self.converter.message_to_payload(session_message) + if not messages: + return None + + if self.persistence_mode is PersistenceMode.NONE: + return {} + + is_blob = self.converter.exceeds_conversational_limit(messages[0]) + + # Build merged metadata from config defaults + per-call overrides + merged_metadata = self._build_metadata(per_call_metadata=kwargs.get("metadata")) + + # Parse the original timestamp and use it as desired timestamp + original_timestamp = datetime.fromisoformat(session_message.created_at.replace("Z", "+00:00")) + monotonic_timestamp = self._get_monotonic_timestamp(original_timestamp) + + if self.config.batch_size > 1: + # Buffer the pre-processed message + should_flush = False + with self._message_lock: + self._message_buffer.append( + BufferedMessage( + session_id=session_id, + messages=messages, + is_blob=is_blob, + timestamp=monotonic_timestamp, + metadata=merged_metadata, + ) + ) + should_flush = len(self._message_buffer) >= self.config.batch_size + + # Flush only messages outside the lock to prevent deadlock + if should_flush: + self._flush_messages_only() + + return {} # No eventId yet + + # Immediate send (batch_size == 1) + try: + if not is_blob: + event = self.memory_client.create_event( + memory_id=self.config.memory_id, + actor_id=self.config.actor_id, + session_id=session_id, + messages=messages, + event_timestamp=monotonic_timestamp, + metadata=merged_metadata, + ) + else: + create_event_kwargs: dict[str, Any] = { + "memoryId": self.config.memory_id, + "actorId": self.config.actor_id, + "sessionId": session_id, + "payload": [{"blob": json.dumps(messages[0])}], + "eventTimestamp": monotonic_timestamp, + } + if merged_metadata: + create_event_kwargs["metadata"] = merged_metadata + event = self.memory_client.gmdp_client.create_event(**create_event_kwargs) + logger.debug("Created event: %s for message: %s", event.get("eventId"), session_message.message_id) + return event + except Exception as e: + logger.error("Failed to create message in AgentCore Memory: %s", e) + raise SessionException(f"Failed to create message: {e}") from e + + def read_message(self, session_id: str, agent_id: str, message_id: int, **kwargs: Any) -> Optional[SessionMessage]: + """Read a specific message by ID from AgentCore Memory. + + Args: + session_id (str): The session ID to read from. + agent_id (str): The agent ID associated with the message. + message_id (int): The message ID to read. + **kwargs (Any): Additional keyword arguments. + + Returns: + Optional[SessionMessage]: The message if found, None otherwise. + + Note: + This reads a single event by ID from AgentCore Memory. + """ + result = self.memory_client.gmdp_client.get_event( + memoryId=self.config.memory_id, actorId=self.config.actor_id, sessionId=session_id, eventId=message_id + ) + return SessionMessage.from_dict(result) if result else None + + def update_message(self, session_id: str, agent_id: str, session_message: SessionMessage, **kwargs: Any) -> None: + """Update message data in AgentCore Memory. + + Since AgentCore Memory events are immutable, this method performs an update by + creating a new event with the updated content and deleting the old event. + This enables features like guardrail redaction via Strands' redact_latest_message(). + + If the message has not yet been persisted (e.g., still in the message buffer when + batch_size > 1), the buffered message is replaced in-place instead. + + Args: + session_id (str): The session ID containing the message. + agent_id (str): The agent ID associated with the message. + session_message (SessionMessage): The message to update (with updated content + and the original message_id/eventId). + **kwargs (Any): Additional keyword arguments. + + Raises: + SessionException: If session ID doesn't match configuration or update fails. + """ + if session_id != self.config.session_id: + raise SessionException(f"Session ID mismatch: expected {self.config.session_id}, got {session_id}") + + old_message_id = session_message.message_id + + # If message hasn't been persisted yet (still in buffer), update it there + if old_message_id is None: + if self._update_buffered_message(session_message): + logger.debug("Updated buffered message (not yet persisted to AgentCore Memory)") + return + logger.debug("Message has no event ID and was not found in buffer - skipping update") + return + + # Create a new event with the updated message content + try: + updated_message = SessionMessage( + message=session_message.message, + message_id=0, + created_at=session_message.created_at, + ) + new_event = self.create_message(session_id, agent_id, updated_message) + except Exception as e: + logger.error("Failed to update message in AgentCore Memory: %s", e) + raise SessionException(f"Failed to update message: {e}") from e + + new_event_id = new_event.get("eventId") if new_event else None + if not new_event_id: + logger.warning("create_message did not return an eventId — skipping delete of old event %s", old_message_id) + return + + # Delete the old event; if this fails, roll back the newly created event + try: + self.memory_client.gmdp_client.delete_event( + memoryId=self.config.memory_id, + actorId=self.config.actor_id, + sessionId=session_id, + eventId=old_message_id, + ) + except Exception as delete_error: + logger.warning( + "Failed to delete old event %s after creating replacement: %s. Attempting rollback.", + old_message_id, + delete_error, + ) + try: + self.memory_client.gmdp_client.delete_event( + memoryId=self.config.memory_id, + actorId=self.config.actor_id, + sessionId=session_id, + eventId=new_event_id, + ) + logger.info("Rolled back new event %s after failed delete of old event", new_event_id) + except Exception as rollback_error: + logger.error( + "Rollback failed: could not delete new event %s: %s. Both old (%s) and new events may exist.", + new_event_id, + rollback_error, + old_message_id, + ) + raise SessionException( + f"Failed to update message: could not delete old event: {delete_error}" + ) from delete_error + + # Update _latest_agent_message so it doesn't hold a stale eventId + latest_messages = getattr(self, "_latest_agent_message", None) + if latest_messages and agent_id in latest_messages: + old_latest = self._latest_agent_message[agent_id] + if old_latest.message_id == old_message_id: + self._latest_agent_message[agent_id] = SessionMessage( + message=session_message.message, + message_id=new_event_id, + created_at=session_message.created_at, + ) + + logger.info("Updated message in AgentCore Memory: replaced event %s", old_message_id) + + def list_messages( + self, + session_id: str, + agent_id: str, + limit: Optional[int] = None, + offset: int = 0, + **kwargs: Any, + ) -> list[SessionMessage]: + """List messages for an agent from AgentCore Memory with pagination. + + Args: + session_id (str): The session ID to list messages from. + agent_id (str): The agent ID to list messages for. + limit (Optional[int], optional): Maximum number of messages to return. Defaults to None. + offset (int, optional): Number of messages to skip. Defaults to 0. + **kwargs (Any): Additional keyword arguments. + + Returns: + list[SessionMessage]: list of messages for the agent. + + Raises: + SessionException: If session ID doesn't match configuration. + """ + if session_id != self.config.session_id: + raise SessionException(f"Session ID mismatch: expected {self.config.session_id}, got {session_id}") + + try: + max_results = (limit + offset) if limit else MAX_FETCH_ALL_RESULTS + + events = self.memory_client.list_events( + memory_id=self.config.memory_id, + actor_id=self.config.actor_id, + session_id=session_id, + max_results=max_results, + ) + messages = self.converter.events_to_messages(events) + if self.config.filter_restored_tool_context: + messages = self._filter_restored_tool_context(messages) + if limit is not None: + return messages[offset : offset + limit] + else: + return messages[offset:] + + except Exception as e: + logger.error("Failed to list messages from AgentCore Memory: %s", e) + return [] + + def _filter_restored_tool_context(self, messages: list[SessionMessage]) -> list[SessionMessage]: + """Strip historical toolUse/toolResult context from restored messages.""" + filtered_messages: list[SessionMessage] = [] + for session_message in messages: + message = session_message.to_message() + filtered_content = [ + content + for content in message.get("content", []) + if "toolUse" not in content and "toolResult" not in content + ] + + if not filtered_content: + continue + + filtered_message: Message = {"role": message["role"], "content": filtered_content} + filtered_messages.append( + SessionMessage( + message=filtered_message, + message_id=session_message.message_id, + redact_message=session_message.redact_message, + created_at=session_message.created_at, + updated_at=session_message.updated_at, + ) + ) + + return filtered_messages + + # endregion SessionRepository interface implementation + + # region RepositorySessionManager overrides + @override + def append_message(self, message: Message, agent: "Agent", **kwargs: Any) -> None: + """Append a message to the agent's session using AgentCore's eventId as message_id. + + Args: + message: Message to add to the agent in the session + agent: Agent to append the message to + **kwargs: Additional keyword arguments for future extensibility. + """ + created_message = self.create_message(self.session_id, agent.agent_id, SessionMessage.from_message(message, 0)) + if created_message is None: + return + session_message = SessionMessage.from_message(message, created_message.get("eventId")) + self._latest_agent_message[agent.agent_id] = session_message + + def retrieve_customer_context(self, event: MessageAddedEvent) -> None: + """Retrieve customer LTM context before processing support query. + + Args: + event (MessageAddedEvent): The message added event containing the agent and message data. + """ + messages = event.agent.messages + if not messages or messages[-1].get("role") != "user": + return None + content = messages[-1].get("content") + if not content or "text" not in content[0]: + return None + if not self.config.retrieval_config: + # Only retrieve LTM + return None + + user_query = messages[-1]["content"][0]["text"] + + def retrieve_for_namespace(namespace: str, retrieval_config: RetrievalConfig): + """Helper function to retrieve memories for a single namespace.""" + resolved_namespace = namespace.format( + actorId=self.config.actor_id, + sessionId=self.config.session_id, + memoryStrategyId=retrieval_config.strategy_id or "", + ) + + memories = self.memory_client.retrieve_memories( + memory_id=self.config.memory_id, + namespace_path=resolved_namespace, + query=user_query, + top_k=retrieval_config.top_k, + ) + if retrieval_config.relevance_score: + memories = [m for m in memories if m.get("score", 0.0) >= retrieval_config.relevance_score] + context_items = [] + for memory in memories: + if isinstance(memory, dict): + content = memory.get("content", {}) + if isinstance(content, dict): + text = content.get("text", "").strip() + if text: + context_items.append(text) + return context_items + + try: + # Retrieve customer context from all namespaces in parallel + all_context = [] + + with ThreadPoolExecutor() as executor: + future_to_namespace = { + executor.submit(retrieve_for_namespace, namespace, retrieval_config): namespace + for namespace, retrieval_config in self.config.retrieval_config.items() + } + for future in as_completed(future_to_namespace): + try: + context_items = future.result() + all_context.extend(context_items) + except Exception as e: + # Continue processing other futures event if one fails rather than failing the entire operation + namespace = future_to_namespace[future] + logger.error("Failed to retrieve memories for namespace %s: %s", namespace, e) + + # Inject retrieved memory as a content block in the last user message. + # Prepended so the user's query text remains last (avoids assistant-prefill + # errors on Claude 4.6+ and keeps the user request in the position models + # attend to most). + if all_context: + context_text = "\n".join(all_context) + event.agent.messages[-1]["content"].insert( + 0, {"text": f"<{self.config.context_tag}>{context_text}"} + ) + logger.info("Retrieved %s customer context items", len(all_context)) + + except Exception as e: + logger.error("Failed to retrieve customer context: %s", e) + + @override + def register_hooks(self, registry: HookRegistry, **kwargs) -> None: + """Register additional hooks. + + In sync mode (the default), delegates to the base class and adds the + retrieve_customer_context + batching callbacks synchronously, preserving + existing behavior exactly. + + In async mode, registers async callbacks that wrap every per-turn + boto3-backed operation (append_message, sync_agent, buffer flushes, + customer-context retrieval) with asyncio.to_thread, so the asyncio + event loop stays free while boto3 is blocking on the network. + + Note: AgentInitializedEvent cannot be async per Strands' HookRegistry, + so agent restoration (read_session / read_agent / list_messages) still + blocks the calling thread in async mode — see AgentCoreMemoryConfig + docstring for mitigations. + + Args: + registry (HookRegistry): The hook registry to register callbacks with. + **kwargs: Additional keyword arguments. + """ + if not self.config.async_mode: + RepositorySessionManager.register_hooks(self, registry, **kwargs) + registry.add_callback(MessageAddedEvent, lambda event: self.retrieve_customer_context(event)) + + # Only register AfterInvocationEvent hook when batching is enabled + if self.config.batch_size > 1: + registry.add_callback(AfterInvocationEvent, lambda event: self._flush_messages()) + return + + # Async mode: register async callbacks that offload the existing sync + # methods to a worker thread via asyncio.to_thread. AgentInitializedEvent + # and BidiAgentInitializedEvent must stay sync (Strands disallows async + # callbacks for AgentInitializedEvent — see strands/hooks/registry.py:227). + logger.warning( + "AgentCoreMemorySessionManager async_mode=True: the agent must be invoked " + "via the async path (e.g. agent.stream_async(...) or agent.invoke_async(...)). " + "Sync invocation will raise RuntimeError from Strands' hook registry." + ) + + def _offload(method, *event_args): + """Build an async callback that offloads `method(*[a(event) for a in event_args])` to a thread. + + Each entry in `event_args` is a callable that extracts an argument from the event; + pass none for a zero-arg method. + """ + + async def _callback(event): + await asyncio.to_thread(method, *(extract(event) for extract in event_args)) + + return _callback + + registry.add_callback(AgentInitializedEvent, lambda event: self.initialize(event.agent)) + + async def _on_message_added_persist(event: MessageAddedEvent) -> None: + await asyncio.to_thread(self.append_message, event.message, event.agent) + await asyncio.to_thread(self.sync_agent, event.agent) + + registry.add_callback(MessageAddedEvent, _on_message_added_persist) + registry.add_callback(AfterInvocationEvent, _offload(self.sync_agent, lambda e: e.agent)) + registry.add_callback(MessageAddedEvent, _offload(self.retrieve_customer_context, lambda e: e)) + + if self.config.batch_size > 1: + registry.add_callback(AfterInvocationEvent, _offload(self._flush_messages)) + + # Register multi-agent callbacks so async-mode parity matches sync-mode + registry.add_callback(MultiAgentInitializedEvent, _offload(self.initialize_multi_agent, lambda e: e.source)) + registry.add_callback(AfterNodeCallEvent, _offload(self.sync_multi_agent, lambda e: e.source)) + registry.add_callback(AfterMultiAgentInvocationEvent, _offload(self.sync_multi_agent, lambda e: e.source)) + + # Register BidiAgent callbacks so async-mode parity matches sync-mode. + # BidiAgentInitializedEvent dispatches through invoke_callbacks (sync), + # so its callback must stay sync; the other two dispatch through + # invoke_callbacks_async, so async wrappers are safe. + registry.add_callback(BidiAgentInitializedEvent, lambda event: self.initialize_bidi_agent(event.agent)) + + async def _on_bidi_message_added(event: BidiMessageAddedEvent) -> None: + await asyncio.to_thread(self.append_bidi_message, event.message, event.agent) + await asyncio.to_thread(self.sync_bidi_agent, event.agent) + + registry.add_callback(BidiMessageAddedEvent, _on_bidi_message_added) + registry.add_callback(BidiAfterInvocationEvent, _offload(self.sync_bidi_agent, lambda e: e.agent)) + + @override + def initialize(self, agent: "Agent", **kwargs: Any) -> None: + if self.has_existing_agent: + logger.warning( + "An Agent already exists in session %s. We currently support one agent per session.", self.session_id + ) + else: + self.has_existing_agent = True + RepositorySessionManager.initialize(self, agent, **kwargs) + + # endregion RepositorySessionManager overrides + + # region Batching support + + def _update_buffered_message(self, session_message: SessionMessage) -> bool: + """Attempt to update a message that is still in the send buffer. + + When batch_size > 1, messages may not yet be persisted to AgentCore Memory. + This method finds the most recent buffered message matching the session_message's + content role and replaces it with the updated content. + + Args: + session_message (SessionMessage): The message with updated content. + + Returns: + bool: True if a buffered message was found and updated, False otherwise. + """ + updated_messages = self.converter.message_to_payload(session_message) + if not updated_messages: + return False + + is_blob = self.converter.exceeds_conversational_limit(updated_messages[0]) + + with self._message_lock: + # Search from the end (most recent) to find the message to update + for i in range(len(self._message_buffer) - 1, -1, -1): + buf = self._message_buffer[i] + if buf.session_id == self.config.session_id and buf.messages: + # Match by role - the most recent message with the same role + existing_role = buf.messages[0][1] if not buf.is_blob else None + new_role = updated_messages[0][1] if not is_blob else None + if existing_role == new_role: + self._message_buffer[i] = BufferedMessage( + session_id=buf.session_id, + messages=updated_messages, + is_blob=is_blob, + timestamp=buf.timestamp, + metadata=buf.metadata, + ) + return True + return False + + def _flush_messages_only(self) -> list[dict[str, Any]]: + """Flush only buffered messages to AgentCore Memory. + + Call this method to send any remaining buffered messages when batch_size > 1. + This is called when the message buffer reaches batch_size. + Messages are batched by session_id - all conversational messages for the same + session are combined into a single create_event() call to reduce API calls. + Messages that exceed the conversational payload limit are sent as blob events individually + as they require a different API path. + + Returns: + list[dict[str, Any]]: List of created event responses from AgentCore Memory. + + Raises: + SessionException: If message creation fails. On failure, messages remain in the buffer. + """ + if self.persistence_mode is PersistenceMode.NONE: + return [] + + with self._message_lock: + messages_to_send = list(self._message_buffer) + self._message_buffer.clear() + + if not messages_to_send: + return [] + + # Group all messages by session_id, combining conversational and blob messages + # Structure: {session_id: {"payload": [...], "timestamp": latest_timestamp, "metadata": {...}}} + session_groups: dict[str, dict[str, Any]] = {} + + for buffered_msg in messages_to_send: + sid = buffered_msg.session_id + if sid not in session_groups: + session_groups[sid] = {"payload": [], "timestamp": buffered_msg.timestamp, "metadata": {}} + + if buffered_msg.is_blob: + for msg in buffered_msg.messages: + session_groups[sid]["payload"].append({"blob": json.dumps(msg)}) + else: + for text, role in buffered_msg.messages: + session_groups[sid]["payload"].append( + {"conversational": {"content": {"text": text}, "role": role.upper()}} + ) + + # Use the latest timestamp for the combined event + if buffered_msg.timestamp > session_groups[sid]["timestamp"]: + session_groups[sid]["timestamp"] = buffered_msg.timestamp + + # Merge metadata (later entries override earlier for same key) + if buffered_msg.metadata: + session_groups[sid]["metadata"].update(buffered_msg.metadata) + + results = [] + try: + # Send one create_event per session_id with all messages (conversational + blob) + for session_id, group in session_groups.items(): + create_event_kwargs: dict[str, Any] = { + "memoryId": self.config.memory_id, + "actorId": self.config.actor_id, + "sessionId": session_id, + "payload": group["payload"], + "eventTimestamp": group["timestamp"], + } + if group["metadata"]: + create_event_kwargs["metadata"] = group["metadata"] + event = self.memory_client.gmdp_client.create_event(**create_event_kwargs) + results.append(event) + logger.debug( + "Flushed batched event for session %s with %d messages: %s", + session_id, + len(group["payload"]), + event.get("eventId"), + ) + + except Exception as e: + # Restore messages to buffer so they aren't lost + with self._message_lock: + self._message_buffer.extend(messages_to_send) + logger.error("Failed to flush messages to AgentCore Memory: %s", e) + raise SessionException(f"Failed to flush messages: {e}") from e + + logger.info("Flushed %d message events to AgentCore Memory", len(results)) + return results + + def _flush_agent_states_only(self) -> list[dict[str, Any]]: + """Flush only buffered agent states to AgentCore Memory. + + Call this method to send any remaining agent state when batch_size > 1. + Agent states are grouped by agent_id and sent as separate events so that + each event carries the correct agentId metadata for read_agent() lookups. + + Returns: + list[dict[str, Any]]: List of created event responses from AgentCore Memory. + + Raises: + SessionException: If agent state creation fails. On failure, agent states remain in the buffer. + """ + if self.persistence_mode is PersistenceMode.NONE: + return [] + + with self._agent_state_lock: + agent_states_to_send = list(self._agent_state_buffer) + self._agent_state_buffer.clear() + + if not agent_states_to_send: + return [] + + results = [] + try: + # Group agent states by agent_id + agent_groups: dict[str, list[dict]] = {} + for _session_id, session_agent in agent_states_to_send: + agent_id = session_agent.agent_id + if agent_id not in agent_groups: + agent_groups[agent_id] = [] + agent_groups[agent_id].append({"blob": json.dumps(session_agent.to_dict())}) + + # Send one event per agent_id with correct metadata + for agent_id, payloads in agent_groups.items(): + event = self.memory_client.gmdp_client.create_event( + memoryId=self.config.memory_id, + actorId=self.config.actor_id, + sessionId=self.config.session_id, + payload=payloads, + eventTimestamp=self._get_monotonic_timestamp(), + metadata={ + STATE_TYPE_KEY: {"stringValue": StateType.AGENT.value}, + AGENT_ID_KEY: {"stringValue": agent_id}, + }, + ) + results.append(event) + logger.debug("Flushed %d agent states for agent %s: %s", len(payloads), agent_id, event.get("eventId")) + + except Exception as e: + # Restore agent states to buffer so they aren't lost + with self._agent_state_lock: + self._agent_state_buffer.extend(agent_states_to_send) + logger.error("Failed to flush agent states to AgentCore Memory: %s", e) + raise SessionException(f"Failed to flush agent states: {e}") from e + + logger.info("Flushed %d agent state events to AgentCore Memory", len(results)) + return results + + def _flush_messages(self) -> list[dict[str, Any]]: + """Flush all buffered messages and agent state to AgentCore Memory. + + Call this method to send any remaining buffered messages and agent state messages. + This is automatically called when the session is complete (via close() or context manager). + + Returns: + list[dict[str, Any]]: List of created event responses from AgentCore Memory. + + Raises: + SessionException: If any message or agent state creation fails. + """ + results = [] + results.extend(self._flush_messages_only()) + results.extend(self._flush_agent_states_only()) + return results + + def pending_message_count(self) -> int: + """Return the number of messages pending in the buffer. + + Returns: + int: Number of buffered messages waiting to be sent. + """ + with self._message_lock: + return len(self._message_buffer) + + def pending_agent_state_count(self) -> int: + """Return the number of agent states pending in the buffer. + + Returns: + int: Number of buffered agent states waiting to be sent. + """ + with self._agent_state_lock: + return len(self._agent_state_buffer) + + def close(self) -> None: + """Explicitly flush pending messages and close the session manager. + + Call this method when the session is complete to ensure all buffered + messages are sent to AgentCore Memory. Alternatively, use the context + manager protocol (with statement) for automatic cleanup. + """ + self._stop_flush_timer() + self._flush_messages() + + def __enter__(self) -> "AgentCoreMemorySessionManager": + """Enter the context manager. + + Returns: + AgentCoreMemorySessionManager: This session manager instance. + """ + return self + + def __exit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None: + """Exit the context manager and flush any pending messages. + + Args: + exc_type: Exception type if an exception occurred. + exc_val: Exception value if an exception occurred. + exc_tb: Exception traceback if an exception occurred. + """ + try: + self._stop_flush_timer() + self._flush_messages() + except Exception as e: + if exc_type is not None: + logger.error("Failed to flush messages during exception handling: %s", e) + else: + raise + + # endregion Batching support + + # region Interval-based flushing support + + def _start_flush_timer(self) -> None: + """Start the interval-based flush timer. + + This method schedules a recurring timer that flushes the message buffer + at regular intervals if flush_interval_seconds is configured. + """ + with self._timer_lock: + if self._shutdown: + return + + # Cancel existing timer if any + if self._flush_timer is not None: + self._flush_timer.cancel() + + # Schedule next flush + self._flush_timer = threading.Timer( + self.config.flush_interval_seconds, + self._interval_flush_callback, + ) + self._flush_timer.daemon = True + self._flush_timer.start() + logger.debug( + "Scheduled interval flush in %.1f seconds", + self.config.flush_interval_seconds, + ) + + def _interval_flush_callback(self) -> None: + """Callback executed by the flush timer. + + Flushes the buffer if it contains messages or agent states, then reschedules the timer. + """ + try: + # Only flush if there are messages or agent states in the buffer + pending_messages = self.pending_message_count() + pending_agent_states = self.pending_agent_state_count() + if pending_messages > 0 or pending_agent_states > 0: + logger.debug( + "Interval flush triggered: %d message(s) and %d agent state(s) pending", + pending_messages, + pending_agent_states, + ) + self._flush_messages() + else: + logger.debug("Interval flush skipped: buffers are empty") + + # Reschedule the timer (unless shutdown) + if not self._shutdown and self.config.flush_interval_seconds: + self._start_flush_timer() + + except Exception as e: + logger.error("Error during interval flush: %s", e) + # Attempt to reschedule even after error + if not self._shutdown and self.config.flush_interval_seconds: + self._start_flush_timer() + + def _stop_flush_timer(self) -> None: + """Stop the interval-based flush timer. + + This method cancels the timer and prevents it from rescheduling. + Should be called during cleanup (close() or __exit__). + """ + with self._timer_lock: + self._shutdown = True + if self._flush_timer is not None: + self._flush_timer.cancel() + self._flush_timer = None + logger.debug("Stopped interval flush timer") + + # endregion Interval-based flushing support diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorystore/README.md b/src/bedrock_agentcore/memory/integrations/strands/memorystore/README.md new file mode 100644 index 00000000..93e3e818 --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorystore/README.md @@ -0,0 +1,136 @@ +# Strands AgentCore MemoryStore + +`AgentCoreMemoryStore` plugs AgentCore long-term memory directly into Strands' `MemoryManager` for +long-term recall and extraction. It requires `strands-agents>=1.46.0`. + +## One namespace + +A store is recall-only by default. Set `writable=True` on exactly one store when Strands should send +messages to AgentCore for server-side long-term extraction: + +```python +import os + +from strands import Agent +from strands.memory import MemoryManager + +from bedrock_agentcore.memory.integrations.strands.memorystore import AgentCoreMemoryStore + +store = AgentCoreMemoryStore( + memory_id=os.environ["AGENTCORE_MEMORY_ID"], + actor_id="demo-user", + session_id="demo-session", + namespace="/facts/{actorId}/", + writable=True, + extraction=True, + region_name="us-east-1", +) +manager = MemoryManager(stores=[store]) +agent = Agent(memory_manager=manager) +agent("Remember that I prefer window seats.") +``` + +`namespace` performs exact-prefix retrieval. Use `namespace_path` instead to search a namespace +subtree. The integration resolves `{actorId}` and `{sessionId}` client-side; substitute other +placeholders and malformed braces before constructing the store. + +## Multiple namespaces + +`create_agentcore_memory_stores` returns `list[MemoryStore]` for direct `MemoryManager` composition. +Each item is a concrete `AgentCoreMemoryStore`; the factory shares one boto3 client and prevents +duplicate writes by allowing at most one writer: + +```python +import os + +from strands.memory import IntervalTrigger, MemoryManager, MemoryMessageFilter + +from bedrock_agentcore.memory.integrations.strands.memorystore import create_agentcore_memory_stores + +stores = create_agentcore_memory_stores( + memory_id=os.environ["AGENTCORE_MEMORY_ID"], + actor_id="demo-user", + session_id="demo-session", + namespaces=[ + { + "namespace": "/preferences/{actorId}/", + "max_search_results": 5, + "min_score": 0.7, + }, + { + "namespace": "/facts/{actorId}/", + "max_search_results": 10, + "min_score": 0.3, + }, + ], + extraction={ + "cadence": IntervalTrigger(turns=10), + "filter": MemoryMessageFilter(exclude=["toolUse", "toolResult", "image"]), + }, + region_name="us-east-1", +) +manager = MemoryManager(stores=stores) +``` + +With extraction enabled, the first namespace not explicitly marked `writable=False` becomes the +writer. Set `writable=True` on one namespace to choose it explicitly. Omit `extraction` or pass +`False` for recall-only stores. + +## Search and write behavior + +- Search defaults to 5 results. `min_score` enables client-side score filtering and over-fetches by + a factor of 4 (configurable with `over_fetch_factor`); only the over-fetched `topK` is capped at 100. +- Returned metadata uses reserved keys `_id`, `_score`, `_namespaces`, and `_createdAt`. +- Writes preserve user/assistant roles, ignore blank and tool-only messages, and batch up to 50 + consecutive turns per AgentCore event by default. `max_turns_per_event` accepts any positive integer. +- `metadata_provider` returns scalar strings, finite numbers, or booleans. Strings pass through; + other finite scalars use Python `json.dumps` formatting. `None`, arrays, objects, non-finite numbers, + and values outside AgentCore's allowed character set are rejected locally. +- Direct `AgentCoreMemoryStore(...)` construction accepts `extraction_mode="SKIP"` to omit long-term + extraction for its events. The multi-namespace factory intentionally does not expose this option. +- `add_messages()` is the supported write interface. The flat-string Strands `add()` API is not + implemented because it loses role and turn information. + +## Batching, cadence, and flush + +Three separate controls determine write timing and cost: + +1. **Batching is always on.** Each flush packs its role-tagged messages into as few `create_event` + requests as `max_turns_per_event` allows. +2. **Cadence controls when buffered messages are dispatched across turns.** `extraction=True` uses + Strands' default trigger. Pass an extraction config with an `IntervalTrigger` or another Strands + trigger to tune cadence. +3. **`flush()` lets pending write attempts settle; it does not acknowledge durability or server-side + extraction.** Strands 1.46 logs and swallows sender failures, rolls back its high-water mark, and + retains the failed batch for a later retry. + +Synchronous `agent(...)` invocations flush automatically. After async invocation or streaming, call +`await manager.flush()` at a lifecycle or shutdown boundary to let pending writes settle. Monitor logs +or telemetry for failures rather than treating `flush()` as proof that data was persisted. AgentCore's +server-side extraction remains eventually consistent, so newly written records may not be immediately +searchable. + +Reuse one manager per `(actor_id, session_id)` while that session is active. Reuse keeps trigger state +and buffered turns alive, allowing a coarser cadence to reduce calls. The application owns manager +caching and eviction. + +## Namespace and error contract + +Recall works only when the query namespace matches the concrete namespace where AgentCore stored the +extracted record. Writes append to the shared `(memory_id, actor_id, session_id)` stream; the memory +resource's strategies decide which namespaces receive extracted records. That is why a store set must +have at most one writer. + +- AgentCore resolves strategy placeholders at extraction time, but retrieval does not. The store + resolves only `{actorId}` and `{sessionId}` and rejects remaining braces at construction. +- Match the namespace template used when provisioning the strategy. `namespace` queries one exact + prefix; `namespace_path` queries a parent subtree. +- A namespace containing `{sessionId}` is session-scoped. Use a stable session id or actor-only + namespace for cross-session recall. +- The store consumes an existing memory resource; it does not provision strategies or the resource. +- Retrieval failures propagate to `MemoryManager`, which applies its per-store partial-failure behavior. +- Sender failures remain buffered for retry and are logged by Strands rather than propagated by + `flush()`. + +The integration calls the boto3 `bedrock-agentcore` data-plane client directly. AWS credentials use +boto3's normal credential chain; no credentials are stored by the integration. diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorystore/__init__.py b/src/bedrock_agentcore/memory/integrations/strands/memorystore/__init__.py new file mode 100644 index 00000000..8730c1fe --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorystore/__init__.py @@ -0,0 +1,40 @@ +"""Strands-native long-term memory stores backed by Bedrock AgentCore Memory.""" + +from .factory import assert_writable_topology, create_agentcore_memory_stores +from .sender import AgentCoreEventSender +from .store import AgentCoreMemoryStore +from .types import ( + RESERVED_METADATA_PREFIX, + AgentCoreEventSenderConfig, + AgentCoreExactNamespaceStoreConfig, + AgentCoreExtractionConfig, + AgentCoreMemoryStoreConfig, + AgentCoreNamespaceConfig, + AgentCoreSubtreeStoreConfig, + CreateAgentCoreMemoryStoresInput, + ExtractionMode, + MetadataProvider, + MetadataValue, + resolve_namespace, + slugify_namespace, +) + +__all__ = [ + "RESERVED_METADATA_PREFIX", + "AgentCoreEventSender", + "AgentCoreEventSenderConfig", + "AgentCoreExactNamespaceStoreConfig", + "AgentCoreExtractionConfig", + "AgentCoreMemoryStore", + "AgentCoreMemoryStoreConfig", + "AgentCoreNamespaceConfig", + "AgentCoreSubtreeStoreConfig", + "CreateAgentCoreMemoryStoresInput", + "ExtractionMode", + "MetadataProvider", + "MetadataValue", + "assert_writable_topology", + "create_agentcore_memory_stores", + "resolve_namespace", + "slugify_namespace", +] diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorystore/_format.py b/src/bedrock_agentcore/memory/integrations/strands/memorystore/_format.py new file mode 100644 index 00000000..5f42fd54 --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorystore/_format.py @@ -0,0 +1,52 @@ +"""Internal Strands-message formatting helpers.""" + +from typing import Literal + +from strands.types.content import Message + +AgentCoreRole = Literal["USER", "ASSISTANT"] + + +def map_role(message: Message) -> AgentCoreRole: + """Map a Strands role to the AgentCore conversational-role subset. + + Args: + message: Strands message to map. + + Returns: + ``USER`` for a user message, otherwise ``ASSISTANT``. + """ + return "USER" if message["role"] == "user" else "ASSISTANT" + + +def extract_text(message: Message) -> str: + """Join non-empty text blocks and ignore all other block kinds. + + Args: + message: Strands message whose text should be extracted. + + Returns: + Trimmed blocks joined by newlines. + """ + # Drop blank blocks before joining so an empty middle block does not + # leave a stray blank line in the concatenated event text. + parts = [] + for block in message["content"]: + if "text" not in block: + continue + text = block["text"].strip() + if text: + parts.append(text) + return "\n".join(parts) + + +def is_user_or_assistant_with_text(message: Message) -> bool: + """Return whether a user/assistant message contains extractable text. + + Args: + message: Strands message to inspect. + + Returns: + ``True`` only for a supported role with non-blank text. + """ + return message["role"] in ("user", "assistant") and bool(extract_text(message)) diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorystore/factory.py b/src/bedrock_agentcore/memory/integrations/strands/memorystore/factory.py new file mode 100644 index 00000000..1b04ca8e --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorystore/factory.py @@ -0,0 +1,152 @@ +"""Factory helpers for multi-namespace AgentCore memory topologies.""" + +from __future__ import annotations + +from collections.abc import Sequence + +import boto3 +from strands.memory import ExtractionConfig, MemoryStore + +from .store import AgentCoreMemoryStore, _create_data_plane_client +from .types import ( + AgentCoreDataPlaneClient, + AgentCoreExtractionConfig, + AgentCoreNamespaceConfig, + MetadataProvider, +) + + +def assert_writable_topology(stores: Sequence[MemoryStore], expect_extraction: bool = False) -> None: + """Require at most one writable store, and optionally require one writer. + + Args: + stores: AgentCore stores sharing one identity stream. + expect_extraction: Whether a writer is required. + + Raises: + ValueError: If the topology would duplicate writes or cannot extract. + """ + # create_event writes to the identity stream rather than a namespace. Multiple + # writers would therefore duplicate the same conversation events. + writers = [store for store in stores if store.writable] + if len(writers) > 1: + names = ", ".join(f'"{store.name}"' for store in writers) + raise ValueError( + f"AgentCore memory: at most one store may be writable, but {len(writers)} are ({names}). " + "create_event is namespace-free, so multiple writable stores would write duplicate events to the " + "same (memory_id, actor_id, session_id) stream. Mark exactly one namespace writable." + ) + if expect_extraction and not writers: + raise ValueError( + "AgentCore memory: extraction is enabled but no store is writable. Mark one namespace writable " + "(or omit extraction for recall-only)." + ) + + +def create_agentcore_memory_stores( + *, + memory_id: str, + actor_id: str, + session_id: str, + namespaces: list[AgentCoreNamespaceConfig], + extraction: bool | AgentCoreExtractionConfig | None = None, + metadata_provider: MetadataProvider | None = None, + max_turns_per_event: int | None = None, + region_name: str | None = None, + boto3_session: boto3.Session | None = None, + client: AgentCoreDataPlaneClient | None = None, +) -> list[MemoryStore]: + """Build one store per exact namespace with one shared boto3 client. + + Args: + memory_id: AgentCore Memory resource identifier. + actor_id: Actor identifier. + session_id: Session identifier. + namespaces: Per-namespace store configuration dictionaries. + extraction: Recall-only switch, or custom cadence/filter configuration. + metadata_provider: Optional per-message event metadata callback. + max_turns_per_event: Maximum turns packed into one event. + region_name: Region used when constructing the shared client. + boto3_session: Session used when constructing the shared client. + client: Preconstructed shared data-plane client. + + Returns: + One store per namespace. + + Raises: + ValueError: If namespace or writer configuration is invalid. + """ + if not isinstance(namespaces, list) or not namespaces: + raise ValueError("create_agentcore_memory_stores: at least one namespace is required") + for index, namespace_config in enumerate(namespaces): + namespace = namespace_config.get("namespace") if isinstance(namespace_config, dict) else None + if not isinstance(namespace, str) or not namespace.strip(): + raise ValueError( + f"create_agentcore_memory_stores: namespaces[{index}].namespace must be a non-empty string" + ) + if max_turns_per_event is not None and (type(max_turns_per_event) is not int or max_turns_per_event < 1): + raise ValueError( + f"create_agentcore_memory_stores: max_turns_per_event must be a positive integer, got {max_turns_per_event}" + ) + + # ``True`` leaves cadence to MemoryManager; only the object form builds a + # custom Strands extraction configuration. + write_enabled = extraction not in (None, False) + extraction_config: bool | ExtractionConfig | None + if not write_enabled: + extraction_config = None + elif isinstance(extraction, dict) and ("cadence" in extraction or "filter" in extraction): + extraction_config = ExtractionConfig() + cadence = extraction.get("cadence") + message_filter = extraction.get("filter") + if cadence is not None: + extraction_config["trigger"] = cadence + if message_filter is not None: + extraction_config["filter"] = message_filter + else: + extraction_config = True + + # Build one connection and reuse it for every namespace in this identity set. + shared_client = ( + client + if client is not None + else _create_data_plane_client(region_name=region_name, boto3_session=boto3_session) + ) + # The default writer skips explicit opt-outs. Keep multiple explicit writers + # intact so the topology check fails loudly instead of silently choosing one. + any_flagged = any(config.get("writable") is True for config in namespaces) + default_writer_index = -1 + if write_enabled and not any_flagged: + default_writer_index = next( + (index for index, config in enumerate(namespaces) if config.get("writable") is not False), -1 + ) + if write_enabled and not any_flagged and default_writer_index == -1: + raise ValueError( + "create_agentcore_memory_stores: extraction is enabled but every namespace is marked writable: false; " + "leave one namespace un-opted-out (or set writable: true on the intended writer)." + ) + + stores: list[AgentCoreMemoryStore] = [] + for index, namespace_config in enumerate(namespaces): + is_writer = namespace_config.get("writable") is True or index == default_writer_index + stores.append( + AgentCoreMemoryStore( + memory_id=memory_id, + actor_id=actor_id, + session_id=session_id, + namespace=str(namespace_config["namespace"]), + name=namespace_config.get("name"), + description=namespace_config.get("description"), + max_search_results=namespace_config.get("max_search_results"), + min_score=namespace_config.get("min_score"), + over_fetch_factor=namespace_config.get("over_fetch_factor", 4), + writable=is_writer, + extraction=extraction_config if is_writer else None, + metadata_provider=metadata_provider, + max_turns_per_event=max_turns_per_event, + client=shared_client, + ) + ) + # MemoryManager validates store-name uniqueness, so only write topology is checked here. + assert_writable_topology(stores, write_enabled) + return list(stores) diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorystore/sender.py b/src/bedrock_agentcore/memory/integrations/strands/memorystore/sender.py new file mode 100644 index 00000000..91e988bd --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorystore/sender.py @@ -0,0 +1,234 @@ +"""AgentCore event sender used by the Strands long-term memory store.""" + +from __future__ import annotations + +import asyncio +import json +import math +import re +import uuid +from dataclasses import dataclass +from datetime import datetime, timezone + +from strands.memory import AggregateMemoryError +from strands.types.content import Message + +from ._format import extract_text, is_user_or_assistant_with_text, map_role +from .types import ( + DEFAULT_MAX_TURNS_PER_EVENT, + AgentCoreDataPlaneClient, + ExtractionMode, + MetadataProvider, + MetadataValue, +) + +# AgentCore applies this metadata value constraint server-side rather than +# exposing it in boto3's generated types, so validate it with a clear key-specific error. +_METADATA_VALUE_PATTERN = re.compile(r"^[a-zA-Z0-9\s._:/=+@-]*$") + + +@dataclass +class _SeqMessage: + message: Message + sequence_number: int | None + + +@dataclass +class _EventGroup: + items: list[_SeqMessage] + metadata: dict[str, dict[str, str]] | None + + +class AgentCoreEventSender: + """Pack role-tagged Strands messages into AgentCore ``create_event`` calls. + + One flush becomes as few events as metadata boundaries and ``max_turns_per_event`` + allow, so the Strands extraction cadence controls API-call volume. The sender has + no retry layer: failures reach Strands, which retains the batch for retry. + + When Strands supplies sequence numbers, each event gets a deterministic token + derived from its covered range and this sender's run-unique identifier. This + deduplicates a re-fire without colliding when restored sessions reset sequences. + """ + + def __init__( + self, + *, + client: AgentCoreDataPlaneClient, + memory_id: str, + actor_id: str, + session_id: str, + metadata_provider: MetadataProvider | None = None, + run_id: str | None = None, + max_turns_per_event: int = DEFAULT_MAX_TURNS_PER_EVENT, + extraction_mode: ExtractionMode | None = None, + ) -> None: + """Initialize the sender. + + Args: + client: Boto3 AgentCore data-plane client. + memory_id: AgentCore Memory resource identifier. + actor_id: Actor identifier. + session_id: Session identifier. + metadata_provider: Optional per-message event metadata callback. + run_id: Run-unique idempotency-token component. Defaults to a UUID per sender. + max_turns_per_event: Maximum role-tagged turns in one event. + extraction_mode: Optional AgentCore long-term extraction mode. + + Raises: + ValueError: If ``max_turns_per_event`` is not a positive integer. + """ + if type(max_turns_per_event) is not int or max_turns_per_event < 1: + raise ValueError( + f"AgentCoreEventSender: max_turns_per_event must be a positive integer, got {max_turns_per_event}" + ) + self._client = client + self._memory_id = memory_id + self._actor_id = actor_id + self._session_id = session_id + self._metadata_provider = metadata_provider + self._run_id = run_id if run_id is not None else str(uuid.uuid4()) + self._max_turns_per_event = max_turns_per_event + self._extraction_mode = extraction_mode + + async def send_batch(self, messages: list[Message], sequence_numbers: list[int] | None = None) -> None: + """Send all eligible messages, attempting every prepared event concurrently. + + The complete all-event operation is shielded from caller cancellation. If + cancellation arrives, this method waits for every in-flight boto3 call to + settle. A failed write then wins as ``AggregateMemoryError`` so Strands can + roll back the batch; otherwise the original cancellation propagates. + + Args: + messages: Strands messages to write. + sequence_numbers: Optional index-aligned message sequence numbers. + + Raises: + AggregateMemoryError: If one or more AgentCore calls fail. + ValueError: If metadata cannot be represented by AgentCore. + """ + sendable = [ + _SeqMessage( + message, + sequence_numbers[index] if sequence_numbers and index < len(sequence_numbers) else None, + ) + for index, message in enumerate(messages) + if is_user_or_assistant_with_text(message) + ] + if not sendable: + return + + # Validate and map every metadata bag before scheduling any network call. + # A deterministic input error must not allow sibling groups to start writing. + events = self._group_into_events(sendable) + operation = asyncio.create_task(self._send_all(events)) + try: + await asyncio.shield(operation) + except asyncio.CancelledError: + # ``asyncio.to_thread`` cannot stop its worker. Keep cancellation + # attached to this coroutine until every write has a known outcome. + while not operation.done(): + try: + await asyncio.shield(operation) + except asyncio.CancelledError: + continue + # A write failure must reach the coordinator as Exception so its + # high-water mark is rolled back instead of trimming pending data. + operation.result() + raise + + async def _send_all(self, events: list[_EventGroup]) -> None: + results = await asyncio.gather(*(self._send_event(event) for event in events), return_exceptions=True) + failures = [result for result in results if isinstance(result, BaseException)] + if failures: + first = str(failures[0]) + raise AggregateMemoryError( + f"AgentCore create_event failed for {len(failures)} of {len(events)} event(s); first error: {first}", + failures, + ) + + def _group_into_events(self, sendable: list[_SeqMessage]) -> list[_EventGroup]: + # Metadata belongs to the event, so start a new group when per-message metadata + # changes or the current event reaches its configured turn cap. + groups: list[_EventGroup] = [] + current: _EventGroup | None = None + current_signature: str | None = None + for item in sendable: + raw_metadata = dict(self._metadata_provider(item.message)) if self._metadata_provider else None + signature = json.dumps(raw_metadata, sort_keys=True) if raw_metadata is not None else "" + metadata = _to_agentcore_metadata(raw_metadata) if raw_metadata else None + at_cap = current is not None and len(current.items) >= self._max_turns_per_event + if current is None or signature != current_signature or at_cap: + current = _EventGroup(items=[], metadata=metadata) + groups.append(current) + current_signature = signature + current.items.append(item) + return groups + + async def _send_event(self, event: _EventGroup) -> None: + payload = [ + { + "conversational": { + "role": map_role(item.message), + "content": {"text": extract_text(item.message)}, + } + } + for item in event.items + ] + kwargs: dict[str, object] = { + "memoryId": self._memory_id, + "actorId": self._actor_id, + "sessionId": self._session_id, + "eventTimestamp": datetime.now(timezone.utc), + "payload": payload, + } + token = self._token_for_sequence_numbers([item.sequence_number for item in event.items]) + if token is not None: + kwargs["clientToken"] = token + if event.metadata: + kwargs["metadata"] = event.metadata + if self._extraction_mode is not None: + kwargs["extractionMode"] = self._extraction_mode + await asyncio.to_thread(self._client.create_event, **kwargs) + + def _token_for_sequence_numbers(self, sequence_numbers: list[int | None]) -> str | None: + # Without a complete range no safe deterministic token is available; tolerate + # an error-path duplicate rather than inventing a time-based token per call. + if not sequence_numbers or any(number is None for number in sequence_numbers): + return None + first = sequence_numbers[0] + last = sequence_numbers[-1] + return f"{self._memory_id}-{self._actor_id}-{self._run_id}-{first}-{last}" + + +def _metadata_scalar_string(key: str, value: object) -> str: + if value is None or (isinstance(value, float) and not math.isfinite(value)): + raise ValueError( + f'AgentCoreEventSender: metadata value for key "{key}" is {value}, which has no valid string ' + "representation. Provide a finite number, boolean, or a string (omit the key instead of passing " + "None)." + ) + if isinstance(value, str): + return value + if isinstance(value, (int, float, bool)): + return json.dumps(value) + raise ValueError( + f'AgentCoreEventSender: metadata value for key "{key}" must be a scalar string, finite number, or boolean; ' + f"got {type(value).__name__}. Arrays, objects, and None are not supported." + ) + + +def _to_agentcore_metadata(metadata: dict[str, MetadataValue]) -> dict[str, dict[str, str]]: + # AgentCore expects each metadata value in a ``stringValue`` wrapper. Reject + # unusable scalars here instead of surfacing an opaque service validation error. + output: dict[str, dict[str, str]] = {} + for key, value in metadata.items(): + string_value = _metadata_scalar_string(key, value) + if not _METADATA_VALUE_PATTERN.fullmatch(string_value): + raise ValueError( + f'AgentCoreEventSender: metadata value for key "{key}" contains characters AgentCore rejects ' + "(allowed: letters, digits, whitespace, and ._:/=+@-). " + f"Got {string_value!r}. Pass a pre-encoded scalar string using only the allowed characters." + ) + output[key] = {"stringValue": string_value} + return output diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorystore/store.py b/src/bedrock_agentcore/memory/integrations/strands/memorystore/store.py new file mode 100644 index 00000000..8791eebe --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorystore/store.py @@ -0,0 +1,270 @@ +"""AgentCore Memory implementation of the native Strands ``MemoryStore`` contract.""" + +from __future__ import annotations + +import asyncio +import logging +import math +import os +from datetime import datetime, timezone +from typing import TYPE_CHECKING, Any + +import boto3 +from botocore.config import Config as BotocoreConfig +from strands.memory import AddMessagesContext, ExtractionConfig, MemoryEntry, MemoryStore, SearchOptions +from strands.types.content import Message + +from bedrock_agentcore._utils.user_agent import build_user_agent_suffix + +from .sender import AgentCoreEventSender +from .types import ( + DEFAULT_MAX_SEARCH_RESULTS, + DEFAULT_MAX_TURNS_PER_EVENT, + DEFAULT_OVERFETCH_FACTOR, + DEFAULT_REGION, + MAX_TOPK, + RESERVED_METADATA_PREFIX, + AgentCoreDataPlaneClient, + ExtractionMode, + MetadataProvider, + assert_non_empty, + assert_resolved_namespace, + resolve_namespace, + slugify_namespace, +) + +if TYPE_CHECKING: + from typing_extensions import Never + + _MemoryStoreBase = MemoryStore +else: + _MemoryStoreBase = object + + +logger = logging.getLogger(__name__) + + +def _create_data_plane_client( + *, region_name: str | None = None, boto3_session: boto3.Session | None = None +) -> AgentCoreDataPlaneClient: + session = boto3_session if boto3_session is not None else boto3.Session() + region = region_name or session.region_name or os.environ.get("AWS_REGION") or DEFAULT_REGION + config = BotocoreConfig(user_agent_extra=build_user_agent_suffix("strands")) + return session.client("bedrock-agentcore", region_name=region, config=config) # type: ignore[no-any-return] + + +class AgentCoreMemoryStore(_MemoryStoreBase): + """Expose AgentCore long-term memory through Strands' native memory interface. + + Identity and one exact or subtree read target are fixed at construction. Only a + writable store carries ``add_messages`` and extraction configuration. The flat + ``add`` method is intentionally absent because it would discard conversation roles. + """ + + # Strands models optional MemoryStore methods on its Protocol for type checking, + # while runtime capability detection requires absent methods to stay absent. + if TYPE_CHECKING: + add: Never + initialize: Never + get_tools: Never + + def __init__( + self, + *, + memory_id: str, + actor_id: str, + session_id: str, + namespace: str | None = None, + namespace_path: str | None = None, + name: str | None = None, + description: str | None = None, + max_search_results: int | None = None, + writable: bool = False, + extraction: bool | ExtractionConfig | None = None, + min_score: float | None = None, + over_fetch_factor: float = DEFAULT_OVERFETCH_FACTOR, + metadata_provider: MetadataProvider | None = None, + max_turns_per_event: int | None = None, + extraction_mode: ExtractionMode | None = None, + region_name: str | None = None, + boto3_session: boto3.Session | None = None, + client: AgentCoreDataPlaneClient | None = None, + ) -> None: + """Initialize one exact-namespace or subtree store. + + Args: + memory_id: AgentCore Memory resource identifier. + actor_id: Actor identifier. + session_id: Session identifier. + namespace: Exact namespace template to retrieve. + namespace_path: Parent namespace template for subtree retrieval. + name: Strands store name; defaults to a namespace slug. + description: Optional human-readable store description. + max_search_results: Default result count for this store. + writable: Whether ``add_messages`` may write conversation events. + extraction: Strands automatic-extraction configuration. + min_score: Optional client-side relevance threshold. + over_fetch_factor: Retrieval multiplier used when ``min_score`` is set. + metadata_provider: Optional per-message event metadata callback. + max_turns_per_event: Maximum turns packed into one event. + extraction_mode: Optional AgentCore extraction mode, including ``SKIP``. + region_name: Region used when constructing the boto3 client. + boto3_session: Session used when constructing the boto3 client. + client: Preconstructed boto3 AgentCore data-plane client. + + Raises: + ValueError: If identity, read target, or numeric options are invalid. + """ + if (namespace is None) == (namespace_path is None): + raise ValueError("AgentCoreMemoryStore: exactly one of namespace or namespace_path is required") + template = namespace_path if namespace_path is not None else namespace + read_field = "namespace_path" if namespace_path is not None else "namespace" + assert template is not None + + self._memory_id = assert_non_empty(memory_id, "memory_id") + self._actor_id = assert_non_empty(actor_id, "actor_id") + self._session_id = assert_non_empty(session_id, "session_id") + assert_non_empty(template, read_field) + self._resolved_namespace = resolve_namespace(template, self._actor_id, self._session_id) + assert_resolved_namespace(self._resolved_namespace, template) + self._read_mode = "subtree" if namespace_path is not None else "exact" + + explicit_name = name.strip() if name is not None else "" + self.name = explicit_name or slugify_namespace(template) + self.description = description + if max_search_results is not None and (type(max_search_results) is not int or max_search_results < 1): + raise ValueError( + f"AgentCoreMemoryStore: max_search_results must be a positive integer, got {max_search_results}" + ) + self.max_search_results = max_search_results + # Recall-safe default: a store never writes unless explicitly enabled. + self.writable = writable + self.extraction: bool | ExtractionConfig | None = extraction if writable else None + if min_score is not None and ( + isinstance(min_score, bool) + or not isinstance(min_score, (int, float)) + or not math.isfinite(min_score) + or min_score < 0 + or min_score > 1 + ): + raise ValueError( + f"AgentCoreMemoryStore: min_score must be a finite number between 0 and 1, got {min_score}" + ) + if ( + isinstance(over_fetch_factor, bool) + or not isinstance(over_fetch_factor, (int, float)) + or not math.isfinite(over_fetch_factor) + or over_fetch_factor < 1 + ): + raise ValueError(f"AgentCoreMemoryStore: over_fetch_factor must be a number >= 1, got {over_fetch_factor}") + self._min_score = min_score + self._over_fetch_factor = over_fetch_factor + self._client = ( + client + if client is not None + else _create_data_plane_client(region_name=region_name, boto3_session=boto3_session) + ) + self._sender = ( + AgentCoreEventSender( + client=self._client, + memory_id=self._memory_id, + actor_id=self._actor_id, + session_id=self._session_id, + metadata_provider=metadata_provider, + max_turns_per_event=( + max_turns_per_event if max_turns_per_event is not None else DEFAULT_MAX_TURNS_PER_EVENT + ), + extraction_mode=extraction_mode, + ) + if writable + else None + ) + if not writable and extraction not in (None, False): + logger.warning( + '[agentcore-memory] store "%s" has an extraction config but writable is false; extraction will not run', + self.name, + ) + + async def search(self, query: str, options: SearchOptions | None = None) -> list[MemoryEntry]: + """Retrieve relevant AgentCore memory records. + + Args: + query: Semantic search query. + options: Optional Strands per-call result cap. + + Returns: + Records mapped to Strands ``MemoryEntry`` objects. + + Raises: + ValueError: If the effective result cap is invalid. + """ + want = ( + options.get("max_search_results") + if options is not None and "max_search_results" in options + else self.max_search_results or DEFAULT_MAX_SEARCH_RESULTS + ) + if type(want) is not int or want < 1: + raise ValueError(f"AgentCoreMemoryStore.search: max_search_results must be a positive integer, got {want}") + top_k = want + if self._min_score is not None: + # Clamp before converting to int: unlike JavaScript's Math.ceil, Python's + # math.ceil cannot convert an overflowed infinite product. + over_fetch = want * self._over_fetch_factor + top_k = MAX_TOPK if over_fetch >= MAX_TOPK else math.ceil(over_fetch) + kwargs: dict[str, Any] = { + "memoryId": self._memory_id, + "searchCriteria": {"searchQuery": query, "topK": top_k}, + "namespacePath" if self._read_mode == "subtree" else "namespace": self._resolved_namespace, + } + # Let retrieval errors propagate: MemoryManager isolates store failures and + # applies its normal partial-failure behavior. + response = await asyncio.to_thread(self._client.retrieve_memory_records, **kwargs) + records = response.get("memoryRecordSummaries") or [] + filtered = [ + record + for record in records + if self._min_score is None or float(record.get("score", 0) or 0) >= self._min_score + ][:want] + return [self._to_entry(record) for record in filtered] + + async def add_messages(self, messages: list[Message], context: AddMessagesContext | None = None) -> None: + """Write role-preserving conversation messages to AgentCore. + + Args: + messages: Strands messages to ingest. + context: Optional manager-provided sequence numbers. + + Raises: + ValueError: If this store is recall-only. + """ + if self._sender is None: + raise ValueError(f'AgentCoreMemoryStore "{self.name}" is not writable; add_messages is unavailable') + await self._sender.send_batch(messages, context.sequence_numbers if context else None) + + @staticmethod + def _to_entry(record: dict[str, Any]) -> MemoryEntry: + content = record.get("content") + text = content.get("text", "") if isinstance(content, dict) and isinstance(content.get("text"), str) else "" + # Reserved keys prevent store-supplied record fields colliding with user metadata. + prefix = RESERVED_METADATA_PREFIX + metadata: dict[str, Any] = {} + if "memoryRecordId" in record: + metadata[f"{prefix}id"] = record["memoryRecordId"] + if "score" in record: + metadata[f"{prefix}score"] = record["score"] + if "namespaces" in record: + metadata[f"{prefix}namespaces"] = record["namespaces"] + if "createdAt" in record: + created_at = record["createdAt"] + metadata[f"{prefix}createdAt"] = _format_created_at(created_at) + return MemoryEntry(content=text, metadata=metadata) + + +def _format_created_at(value: object) -> str: + if not isinstance(value, datetime): + return str(value) + timestamp = value + if timestamp.tzinfo is None: + timestamp = timestamp.replace(tzinfo=timezone.utc) + utc = timestamp.astimezone(timezone.utc) + return utc.isoformat(timespec="milliseconds").replace("+00:00", "Z") diff --git a/src/bedrock_agentcore/memory/integrations/strands/memorystore/types.py b/src/bedrock_agentcore/memory/integrations/strands/memorystore/types.py new file mode 100644 index 00000000..a718b822 --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorystore/types.py @@ -0,0 +1,221 @@ +"""Public types and namespace helpers for the Strands AgentCore memory store.""" + +import re +from collections.abc import Callable, Mapping +from typing import Any, Literal, Protocol + +import boto3 +from strands.memory import ExtractionConfig, ExtractionTrigger, MemoryMessageFilter +from strands.types.content import Message +from typing_extensions import Never, NotRequired, TypedDict + +ExtractionMode = Literal["SKIP"] +"""Long-term extraction control accepted by AgentCore ``create_event``.""" + +# Defaults apply only when neither call-level nor store-level configuration overrides them. +DEFAULT_REGION = "us-west-2" +DEFAULT_MAX_SEARCH_RESULTS = 5 +DEFAULT_OVERFETCH_FACTOR = 4 +# Bound score-filter over-fetching even when callers configure a large multiplier. +MAX_TOPK = 100 +# Packing turns lets extraction cadence control API volume instead of writing one event per message. +DEFAULT_MAX_TURNS_PER_EVENT = 50 +# Store-supplied record fields use this prefix to avoid colliding with user metadata. +RESERVED_METADATA_PREFIX = "_" + + +class AgentCoreDataPlaneClient(Protocol): + """Structural type for the boto3 AgentCore data-plane client.""" + + def create_event(self, **kwargs: Any) -> dict[str, Any]: + """Create an AgentCore memory event.""" + ... + + def retrieve_memory_records(self, **kwargs: Any) -> dict[str, Any]: + """Retrieve AgentCore long-term memory records.""" + ... + + +MetadataValue = str | int | float | bool +"""Scalar metadata value accepted by AgentCore event metadata.""" + +MetadataProvider = Callable[[Message], Mapping[str, MetadataValue]] +"""Derive event metadata from one message. + +Strings pass through; other finite scalars use :func:`json.dumps` formatting. AgentCore +accepts only letters, digits, whitespace, and ``._:/=+@-`` in the resulting value. +""" + + +class _AgentCoreMemoryConnectionConfig(TypedDict): + """Connection and write identity shared internally by AgentCore memory stores.""" + + memory_id: str + actor_id: str + session_id: str + metadata_provider: NotRequired[MetadataProvider] + max_turns_per_event: NotRequired[int] + extraction_mode: NotRequired[ExtractionMode] + region_name: NotRequired[str] + boto3_session: NotRequired[boto3.Session] + client: NotRequired[AgentCoreDataPlaneClient] + + +class _AgentCoreMemoryStoreOptions(_AgentCoreMemoryConnectionConfig, total=False): + """Fields shared by exact-namespace and subtree store configurations.""" + + name: str + description: str + max_search_results: int + writable: bool + extraction: bool | ExtractionConfig + min_score: float + over_fetch_factor: float + + +class AgentCoreExactNamespaceStoreConfig(_AgentCoreMemoryStoreOptions): + """Read one exact namespace prefix after substituting actor/session placeholders.""" + + namespace: str + namespace_path: NotRequired[Never] + + +class AgentCoreSubtreeStoreConfig(_AgentCoreMemoryStoreOptions): + """Read a parent namespace path and all of its child namespaces.""" + + namespace_path: str + namespace: NotRequired[Never] + + +AgentCoreMemoryStoreConfig = AgentCoreExactNamespaceStoreConfig | AgentCoreSubtreeStoreConfig +"""One flat store config with identity and exactly one read-target shape. + +``writable`` defaults to false for recall-only behavior; a name defaults to a slug of +the namespace template. +""" + + +class AgentCoreEventSenderConfig(TypedDict): + """Configuration for :class:`AgentCoreEventSender`.""" + + client: AgentCoreDataPlaneClient + memory_id: str + actor_id: str + session_id: str + metadata_provider: NotRequired[MetadataProvider] + run_id: NotRequired[str] + max_turns_per_event: NotRequired[int] + extraction_mode: NotRequired[ExtractionMode] + + +class AgentCoreNamespaceConfig(TypedDict): + """Per-namespace read configuration used by the store factory.""" + + namespace: str + name: NotRequired[str] + description: NotRequired[str] + max_search_results: NotRequired[int] + min_score: NotRequired[float] + over_fetch_factor: NotRequired[float] + writable: NotRequired[bool] + + +class AgentCoreExtractionConfig(TypedDict, total=False): + """Writable-store extraction cadence and message filtering.""" + + cadence: ExtractionTrigger | list[ExtractionTrigger] + filter: MemoryMessageFilter + + +class CreateAgentCoreMemoryStoresInput(TypedDict): + """Configuration accepted by :func:`create_agentcore_memory_stores`.""" + + memory_id: str + actor_id: str + session_id: str + namespaces: list[AgentCoreNamespaceConfig] + extraction: NotRequired[bool | AgentCoreExtractionConfig] + metadata_provider: NotRequired[MetadataProvider] + max_turns_per_event: NotRequired[int] + region_name: NotRequired[str] + boto3_session: NotRequired[boto3.Session] + client: NotRequired[AgentCoreDataPlaneClient] + + +_UNRESOLVED_PLACEHOLDER = re.compile(r"\{[^{}]*\}") +_ANY_BRACE = re.compile(r"[{}]") +_NAMESPACE_PLACEHOLDER = re.compile(r"\{[^{}]*\}") +_NON_ALPHANUMERIC = re.compile(r"[^a-zA-Z0-9]+") + + +def resolve_namespace(template: str, actor_id: str, session_id: str) -> str: + """Resolve ``{actorId}`` and then ``{sessionId}`` in a namespace template. + + AgentCore resolves strategy placeholders when extracting records, but retrieval + does not resolve placeholders and rejects braces. The store therefore substitutes + its two known identity placeholders before reading and rejects anything left over. + + Args: + template: Namespace template to resolve. + actor_id: Actor identifier substituted for ``{actorId}``. + session_id: Session identifier substituted for ``{sessionId}``. + + Returns: + The resolved namespace. + """ + return template.replace("{actorId}", actor_id).replace("{sessionId}", session_id) + + +def assert_resolved_namespace(resolved: str, template: str) -> None: + """Reject unresolved placeholders and unmatched braces. + + Args: + resolved: Namespace after supported substitutions. + template: Original namespace template, used in the error message. + + Raises: + ValueError: If a token or brace remains. + """ + token = _UNRESOLVED_PLACEHOLDER.search(resolved) + brace = _ANY_BRACE.search(resolved) + offending = token.group(0) if token else brace.group(0) if brace else None + if offending is not None: + raise ValueError( + f'AgentCoreMemoryStore: namespace "{template}" still contains "{offending}" after substitution. ' + "Only {actorId} and {sessionId} are resolved client-side; the AgentCore retrieve path does not " + 'resolve placeholders and rejects "{"/"}". Provide a namespace whose only placeholders are ' + "{actorId}/{sessionId} (and no stray braces), or pre-substitute the others (for example, a concrete " + "strategy id) before constructing the store." + ) + + +def assert_non_empty(value: object, field: str) -> str: + """Return a non-empty string or raise a field-specific error. + + Args: + value: Value to validate. + field: Public field name used in the error. + + Returns: + The validated string. + + Raises: + ValueError: If ``value`` is not a non-empty string. + """ + if not isinstance(value, str) or not value.strip(): + raise ValueError(f"AgentCoreMemoryStore: {field} must be a non-empty string") + return value + + +def slugify_namespace(namespace: str) -> str: + """Derive a stable store name from a namespace template. + + Args: + namespace: Namespace template. + + Returns: + A hyphenated slug, or ``agentcore-memory`` when no usable text remains. + """ + without_placeholders = _NAMESPACE_PLACEHOLDER.sub("", namespace) + slug = _NON_ALPHANUMERIC.sub("-", without_placeholders).strip("-") + return slug or "agentcore-memory" diff --git a/src/bedrock_agentcore/memory/integrations/strands/session_manager.py b/src/bedrock_agentcore/memory/integrations/strands/session_manager.py index f930b548..2706a16e 100644 --- a/src/bedrock_agentcore/memory/integrations/strands/session_manager.py +++ b/src/bedrock_agentcore/memory/integrations/strands/session_manager.py @@ -1,1343 +1,36 @@ -"""AgentCore Memory-based session manager for Bedrock AgentCore Memory integration.""" +"""Compatibility alias for the Strands AgentCore Memory session manager.""" -import asyncio -import json -import logging -import threading -from concurrent.futures import ThreadPoolExecutor, as_completed -from datetime import datetime, timedelta, timezone -from enum import Enum -from typing import TYPE_CHECKING, Any, Dict, NamedTuple, Optional +import sys +from typing import TYPE_CHECKING -import boto3 -from botocore.config import Config as BotocoreConfig -from strands.experimental.hooks.events import ( - BidiAfterInvocationEvent, - BidiAgentInitializedEvent, - BidiMessageAddedEvent, -) -from strands.experimental.hooks.multiagent.events import ( - AfterMultiAgentInvocationEvent, - AfterNodeCallEvent, - MultiAgentInitializedEvent, -) -from strands.hooks import AfterInvocationEvent, MessageAddedEvent -from strands.hooks.events import AgentInitializedEvent -from strands.hooks.registry import HookRegistry -from strands.session.repository_session_manager import RepositorySessionManager -from strands.session.session_repository import SessionRepository -from strands.types.content import Message -from strands.types.exceptions import SessionException -from strands.types.session import Session, SessionAgent, SessionMessage -from typing_extensions import override - -from bedrock_agentcore.memory.client import MemoryClient -from bedrock_agentcore.memory.models.filters import ( - EventMetadataFilter, - LeftExpression, - MetadataValue, - OperatorType, - RightExpression, -) - -from .bedrock_converter import AgentCoreMemoryConverter -from .config import AgentCoreMemoryConfig, PersistenceMode, RetrievalConfig, normalize_metadata -from .converters import MemoryConverter +from .memorysessionmanager import session_manager as _canonical_module if TYPE_CHECKING: - from strands.agent.agent import Agent - -logger = logging.getLogger(__name__) - -MAX_FETCH_ALL_RESULTS = 10000 - -# Legacy prefixes for backwards compatibility with old events -LEGACY_SESSION_PREFIX = "session_" -LEGACY_AGENT_PREFIX = "agent_" - -# Metadata keys for event identification -STATE_TYPE_KEY = "stateType" -AGENT_ID_KEY = "agentId" - -# Maximum metadata key-value pairs per event (API limit) -MAX_METADATA_KEYS = 15 - -# Reserved internal metadata keys that users cannot override -RESERVED_METADATA_KEYS = frozenset({STATE_TYPE_KEY, AGENT_ID_KEY}) - - -class BufferedMessage(NamedTuple): - """A pre-processed message waiting to be flushed to AgentCore Memory.""" - - session_id: str - messages: list[tuple[str, str]] - is_blob: bool - timestamp: datetime - metadata: Optional[Dict[str, MetadataValue]] = None - - -class StateType(Enum): - """State type for distinguishing session and agent metadata in events.""" - - SESSION = "SESSION" - AGENT = "AGENT" - - -class AgentCoreMemorySessionManager(RepositorySessionManager, SessionRepository): - """AgentCore Memory-based session manager for Bedrock AgentCore Memory integration. - - This session manager integrates Strands agents with Amazon Bedrock AgentCore Memory, - providing seamless synchronization between Strands' session management and Bedrock's - short-term and long-term memory capabilities. - - Key Features: - - Automatic synchronization of conversation messages to Bedrock AgentCore Memory events - - Loading of conversation history from short-term memory during agent initialization - - Integration with long-term memory for context injection into agent state - - Support for custom retrieval configurations per namespace - - Consistent with existing Strands Session managers (such as: FileSessionManager, S3SessionManager) - """ - - def _get_monotonic_timestamp(self, desired_timestamp: Optional[datetime] = None) -> datetime: - """Get a monotonically increasing timestamp for this session. - - Ties are broken at millisecond granularity, which is the resolution - AgentCore Memory stores and orders ``eventTimestamp`` at. - - Args: - desired_timestamp (Optional[datetime]): The desired timestamp. If None, uses current time. - - Returns: - datetime: A timestamp guaranteed to be greater than any previously returned timestamp. - """ - if desired_timestamp is None: - desired_timestamp = datetime.now(timezone.utc) - - # Floor to milliseconds — the resolution AgentCore Memory actually stores - # and orders eventTimestamp at. Comparing at microsecond precision would - # miss two events that fall in the same millisecond but different - # microseconds: they pass the tie check here yet collide once stored. - desired_timestamp = desired_timestamp.replace(microsecond=(desired_timestamp.microsecond // 1000) * 1000) - - with self._timestamp_lock: - if self._last_timestamp is not None and desired_timestamp <= self._last_timestamp: - # Break the tie at millisecond granularity (the service's resolution). - desired_timestamp = self._last_timestamp + timedelta(milliseconds=1) - self._last_timestamp = desired_timestamp - return desired_timestamp - - def __init__( - self, - agentcore_memory_config: AgentCoreMemoryConfig, - region_name: Optional[str] = None, - boto_session: Optional[boto3.Session] = None, - boto_client_config: Optional[BotocoreConfig] = None, - *, - converter: Optional[type[MemoryConverter]] = None, - **kwargs: Any, - ): - """Initialize AgentCoreMemorySessionManager with Bedrock AgentCore Memory. - - Args: - agentcore_memory_config (AgentCoreMemoryConfig): Configuration for AgentCore Memory integration. - region_name (Optional[str], optional): AWS region for Bedrock AgentCore Memory. Defaults to None. - boto_session (Optional[boto3.Session], optional): Optional boto3 session. Defaults to None. - boto_client_config (Optional[BotocoreConfig], optional): Optional boto3 client configuration. - Defaults to None. - converter (Optional[type[MemoryConverter]], optional): Optional custom converter. - If None, native Bedrock/Strands converter is used. - **kwargs (Any): Additional keyword arguments. - """ - self.converter = converter or AgentCoreMemoryConverter - self.config = agentcore_memory_config - self.persistence_mode = agentcore_memory_config.persistence_mode - self.memory_client = MemoryClient(region_name=region_name) - session = boto_session or boto3.Session(region_name=region_name) - self.has_existing_agent = False - - # Instance-scoped monotonic-timestamp state. Per-instance so concurrent - # managers for different sessions in one process do not perturb each - # other's ordering. - self._timestamp_lock = threading.Lock() - self._last_timestamp: Optional[datetime] = None - - # Batching support - stores pre-processed messages - self._message_buffer: list[BufferedMessage] = [] - self._message_lock = threading.Lock() - - # Agent state buffering - stores all agent state updates: (session_id, agent) - self._agent_state_buffer: list[tuple[str, SessionAgent]] = [] - self._agent_state_lock = threading.Lock() - - # Cache for agent created_at timestamps to avoid fetching on every update - self._agent_created_at_cache: dict[str, datetime] = {} - - # Track last synced internal state for each agent (required by parent RepositorySessionManager) - self._last_synced_internal_state: dict[str, Any] = {} - - # Track if this is a new session (required by parent RepositorySessionManager) - self._is_new_session: bool = True - - # Interval-based flushing support - self._flush_timer: Optional[threading.Timer] = None - self._timer_lock = threading.Lock() - self._shutdown = False - - # Add strands-agents to the request user agent - if boto_client_config: - existing_user_agent = getattr(boto_client_config, "user_agent_extra", None) - if existing_user_agent: - new_user_agent = f"{existing_user_agent} strands-agents" - else: - new_user_agent = "strands-agents" - client_config = boto_client_config.merge(BotocoreConfig(user_agent_extra=new_user_agent)) - else: - client_config = BotocoreConfig(user_agent_extra="strands-agents") - - # Override the memory client's boto3 clients - self.memory_client.gmcp_client = session.client( - "bedrock-agentcore-control", region_name=region_name or session.region_name, config=client_config - ) - self.memory_client.gmdp_client = session.client( - "bedrock-agentcore", region_name=region_name or session.region_name, config=client_config - ) - super().__init__(session_id=self.config.session_id, session_repository=self) - - # Start interval-based flush timer if configured - if self.config.flush_interval_seconds: - self._start_flush_timer() - - def _build_metadata( - self, - internal_metadata: Optional[Dict[str, MetadataValue]] = None, - per_call_metadata: Optional[Dict[str, MetadataValue]] = None, - ) -> Optional[Dict[str, MetadataValue]]: - """Build merged metadata from config defaults, provider, per-call overrides, and internal keys. - - Merge precedence (highest wins): - 1. internal_metadata (stateType, agentId) — always wins - 2. per_call_metadata (passed via **kwargs) - 3. metadata_provider() (called at event creation time for dynamic values) - 4. self.config.default_metadata (set at config construction time) - - Args: - internal_metadata: System-reserved metadata (e.g. stateType, agentId). - per_call_metadata: Caller-supplied metadata for a single operation. - - Returns: - Merged metadata dict, or None if empty. - - Raises: - ValueError: If user metadata contains reserved keys or total keys exceed MAX_METADATA_KEYS. - """ - merged: Dict[str, MetadataValue] = {} - - if self.config.default_metadata: - merged.update(self.config.default_metadata) - - if self.config.metadata_provider: - merged.update(normalize_metadata(self.config.metadata_provider())) - - if per_call_metadata: - merged.update(per_call_metadata) - - # Validate user-supplied keys before merging internal keys - user_reserved = RESERVED_METADATA_KEYS & merged.keys() - if user_reserved: - raise ValueError( - f"Metadata keys {user_reserved} are reserved for internal use. Reserved keys: {RESERVED_METADATA_KEYS}" - ) - - if internal_metadata: - merged.update(internal_metadata) - - if len(merged) > MAX_METADATA_KEYS: - raise ValueError(f"Combined metadata has {len(merged)} keys, exceeding the maximum of {MAX_METADATA_KEYS}.") - - return merged or None - - # region SessionRepository interface implementation - def create_session(self, session: Session, **kwargs: Any) -> Session: - """Create a new session in AgentCore Memory. - - Note: AgentCore Memory doesn't have explicit session creation, - so we just validate the session and return it. - - Args: - session (Session): The session to create. - **kwargs (Any): Additional keyword arguments. - - Returns: - Session: The created session. - - Raises: - SessionException: If session ID doesn't match configuration. - """ - if session.session_id != self.config.session_id: - raise SessionException(f"Session ID mismatch: expected {self.config.session_id}, got {session.session_id}") - - if self.persistence_mode is not PersistenceMode.NONE: - event = self.memory_client.gmdp_client.create_event( - memoryId=self.config.memory_id, - actorId=self.config.actor_id, - sessionId=self.session_id, - payload=[ - {"blob": json.dumps(session.to_dict())}, - ], - eventTimestamp=self._get_monotonic_timestamp(), - metadata={STATE_TYPE_KEY: {"stringValue": StateType.SESSION.value}}, - ) - logger.info("Created session: %s with event: %s", session.session_id, event.get("event", {}).get("eventId")) - - return session - - def read_session(self, session_id: str, **kwargs: Any) -> Optional[Session]: - """Read session data. - - AgentCore Memory does not have a `get_session` method. - Which is fine as AgentCore Memory is a managed service we therefore do not need to read/update - the session data. We just return the session object. - - Args: - session_id (str): The session ID to read. - **kwargs (Any): Additional keyword arguments. - - Returns: - Optional[Session]: The session if found, None otherwise. - """ - if session_id != self.config.session_id: - return None - - # 1. Try new approach (metadata filter) - event_metadata = [ - EventMetadataFilter.build_expression( - left_operand=LeftExpression.build(STATE_TYPE_KEY), - operator=OperatorType.EQUALS_TO, - right_operand=RightExpression.build(StateType.SESSION.value), - ) - ] - - events = self.memory_client.list_events( - memory_id=self.config.memory_id, - actor_id=self.config.actor_id, - session_id=session_id, - event_metadata=event_metadata, - max_results=1, - ) - if events: - session_data = json.loads(events[0].get("payload", {})[0].get("blob")) - return Session.from_dict(session_data) - - # 2. Fallback: check for legacy event and migrate - legacy_actor_id = f"{LEGACY_SESSION_PREFIX}{session_id}" - events = self.memory_client.list_events( - memory_id=self.config.memory_id, - actor_id=legacy_actor_id, - session_id=session_id, - max_results=1, - ) - if events: - old_event = events[0] - session_data = json.loads(old_event.get("payload", {})[0].get("blob")) - session = Session.from_dict(session_data) - # Migrate: create new event with metadata, delete old - if self.persistence_mode is not PersistenceMode.NONE: - self.create_session(session) - self.memory_client.gmdp_client.delete_event( - memoryId=self.config.memory_id, - actorId=legacy_actor_id, - sessionId=session_id, - eventId=old_event.get("eventId"), - ) - logger.info("Migrated legacy session event for session: %s", session_id) - return session - - return None - - def delete_session(self, session_id: str, **kwargs: Any) -> None: - """Delete session and all associated data. - - Note: AgentCore Memory doesn't support deletion of events, - so this is a no-op operation. - - Args: - session_id (str): The session ID to delete. - **kwargs (Any): Additional keyword arguments. - """ - logger.warning("Session deletion not supported in AgentCore Memory: %s", session_id) - - def create_agent(self, session_id: str, session_agent: SessionAgent, **kwargs: Any) -> None: - """Create a new agent in the session. - - For AgentCore Memory, we don't need to explicitly create agents; we have Implicit Agent Existence - The agent's existence is inferred from the presence of events/messages in the memory system, - but we validate the session_id matches our config. - - Args: - session_id (str): The session ID to create the agent in. - session_agent (SessionAgent): The agent to create. - **kwargs (Any): Additional keyword arguments. - - Raises: - SessionException: If session ID doesn't match configuration. - """ - if session_id != self.config.session_id: - raise SessionException(f"Session ID mismatch: expected {self.config.session_id}, got {session_id}") - - # Cache the created_at timestamp to avoid re-fetching on updates - if session_agent.created_at: - self._agent_created_at_cache[session_agent.agent_id] = session_agent.created_at - - if self.persistence_mode is PersistenceMode.NONE: - return - - if self.config.batch_size > 1: - # Buffer the agent state events - should_flush = False - with self._agent_state_lock: - self._agent_state_buffer.append((session_id, session_agent)) - should_flush = len(self._agent_state_buffer) >= self.config.batch_size - - # Flush only agent states outside the lock to prevent deadlock - if should_flush: - self._flush_agent_states_only() - - logger.info( - "Buffered agent creation: %s in session: %s", - session_agent.agent_id, - session_id, - ) - else: - # Immediate send when batching is disabled - event = self.memory_client.gmdp_client.create_event( - memoryId=self.config.memory_id, - actorId=self.config.actor_id, - sessionId=self.session_id, - payload=[ - {"blob": json.dumps(session_agent.to_dict())}, - ], - eventTimestamp=self._get_monotonic_timestamp(), - metadata={ - STATE_TYPE_KEY: {"stringValue": StateType.AGENT.value}, - AGENT_ID_KEY: {"stringValue": session_agent.agent_id}, - }, - ) - - logger.info( - "Created agent: %s in session: %s with event %s", - session_agent.agent_id, - session_id, - event.get("event", {}).get("eventId"), - ) - - def read_agent(self, session_id: str, agent_id: str, **kwargs: Any) -> Optional[SessionAgent]: - """Read agent data from AgentCore Memory events. - - We reconstruct the agent state from the conversation history. - - Args: - session_id (str): The session ID to read from. - agent_id (str): The agent ID to read. - **kwargs (Any): Additional keyword arguments. - - Returns: - Optional[SessionAgent]: The agent if found, None otherwise. - """ - if session_id != self.config.session_id: - return None - try: - # 1. Try new approach (metadata filter) - event_metadata = [ - EventMetadataFilter.build_expression( - left_operand=LeftExpression.build(STATE_TYPE_KEY), - operator=OperatorType.EQUALS_TO, - right_operand=RightExpression.build(StateType.AGENT.value), - ), - EventMetadataFilter.build_expression( - left_operand=LeftExpression.build(AGENT_ID_KEY), - operator=OperatorType.EQUALS_TO, - right_operand=RightExpression.build(agent_id), - ), - ] - - events = self.memory_client.list_events( - memory_id=self.config.memory_id, - actor_id=self.config.actor_id, - session_id=session_id, - event_metadata=event_metadata, - max_results=1, - ) - - if events: - agent_data = json.loads(events[0].get("payload", {})[0].get("blob")) - agent = SessionAgent.from_dict(agent_data) - # Cache the created_at timestamp to avoid re-fetching on updates - if agent.created_at: - self._agent_created_at_cache[agent_id] = agent.created_at - return agent - - # 2. Fallback: check for legacy event and migrate - legacy_actor_id = f"{LEGACY_AGENT_PREFIX}{agent_id}" - events = self.memory_client.list_events( - memory_id=self.config.memory_id, - actor_id=legacy_actor_id, - session_id=session_id, - max_results=1, - ) - if events: - old_event = events[0] - agent_data = json.loads(old_event.get("payload", {})[0].get("blob")) - agent = SessionAgent.from_dict(agent_data) - # Migrate: create new event with metadata, delete old - if self.persistence_mode is not PersistenceMode.NONE: - self.create_agent(session_id, agent) - self.memory_client.gmdp_client.delete_event( - memoryId=self.config.memory_id, - actorId=legacy_actor_id, - sessionId=session_id, - eventId=old_event.get("eventId"), - ) - logger.info("Migrated legacy agent event for agent: %s", agent_id) - return agent - - return None - except Exception as e: - logger.error("Failed to read agent %s", e) - return None - - def update_agent(self, session_id: str, session_agent: SessionAgent, **kwargs: Any) -> None: - """Update agent data. - - Args: - session_id (str): The session ID containing the agent. - session_agent (SessionAgent): The agent to update. - **kwargs (Any): Additional keyword arguments. - - Raises: - SessionException: If session ID doesn't match configuration. - """ - agent_id = session_agent.agent_id - - # Verify agent exists and get created_at timestamp if not cached - if agent_id not in self._agent_created_at_cache: - previous_agent = self.read_agent(session_id=session_id, agent_id=agent_id) - if previous_agent is None: - raise SessionException(f"Agent {agent_id} in session {session_id} does not exist") - - # Set created_at from cache before creating the update event - session_agent.created_at = self._agent_created_at_cache[agent_id] - - # Create a new agent event (AgentCore Memory is immutable) - # create_agent will handle batching and caching appropriately - self.create_agent(session_id, session_agent) - - def create_message( - self, session_id: str, agent_id: str, session_message: SessionMessage, **kwargs: Any - ) -> Optional[dict[str, Any]]: - """Create a new message in AgentCore Memory. - - If batch_size > 1, the message is buffered and sent when the buffer reaches batch_size. - Use _flush_messages() or close() to send any remaining buffered messages. - - Args: - session_id (str): The session ID to create the message in. - agent_id (str): The agent ID associated with the message (only here for the interface. - We use the actorId for AgentCore). - session_message (SessionMessage): The message to create. - **kwargs (Any): Additional keyword arguments. - - Returns: - Optional[dict[str, Any]]: The created event data from AgentCore Memory. - Returns empty dict if message is buffered (batch_size > 1). - - Raises: - SessionException: If session ID doesn't match configuration or message creation fails. - - Note: - The returned created message `event` looks like: - ```python - { - "memoryId": "my-mem-id", - "actorId": "user_1", - "sessionId": "test_session_id", - "eventId": "0000001752235548000#97f30a6b", - "eventTimestamp": datetime.datetime(2025, 8, 18, 12, 45, 48, tzinfo=tzlocal()), - "branch": {"name": "main"}, - } - ``` - """ - if session_id != self.config.session_id: - raise SessionException(f"Session ID mismatch: expected {self.config.session_id}, got {session_id}") - - # Convert and check size ONCE (not again at flush) - messages = self.converter.message_to_payload(session_message) - if not messages: - return None - - if self.persistence_mode is PersistenceMode.NONE: - return {} - - is_blob = self.converter.exceeds_conversational_limit(messages[0]) - - # Build merged metadata from config defaults + per-call overrides - merged_metadata = self._build_metadata(per_call_metadata=kwargs.get("metadata")) - - # Parse the original timestamp and use it as desired timestamp - original_timestamp = datetime.fromisoformat(session_message.created_at.replace("Z", "+00:00")) - monotonic_timestamp = self._get_monotonic_timestamp(original_timestamp) - - if self.config.batch_size > 1: - # Buffer the pre-processed message - should_flush = False - with self._message_lock: - self._message_buffer.append( - BufferedMessage( - session_id=session_id, - messages=messages, - is_blob=is_blob, - timestamp=monotonic_timestamp, - metadata=merged_metadata, - ) - ) - should_flush = len(self._message_buffer) >= self.config.batch_size - - # Flush only messages outside the lock to prevent deadlock - if should_flush: - self._flush_messages_only() - - return {} # No eventId yet - - # Immediate send (batch_size == 1) - try: - if not is_blob: - event = self.memory_client.create_event( - memory_id=self.config.memory_id, - actor_id=self.config.actor_id, - session_id=session_id, - messages=messages, - event_timestamp=monotonic_timestamp, - metadata=merged_metadata, - ) - else: - create_event_kwargs: dict[str, Any] = { - "memoryId": self.config.memory_id, - "actorId": self.config.actor_id, - "sessionId": session_id, - "payload": [{"blob": json.dumps(messages[0])}], - "eventTimestamp": monotonic_timestamp, - } - if merged_metadata: - create_event_kwargs["metadata"] = merged_metadata - event = self.memory_client.gmdp_client.create_event(**create_event_kwargs) - logger.debug("Created event: %s for message: %s", event.get("eventId"), session_message.message_id) - return event - except Exception as e: - logger.error("Failed to create message in AgentCore Memory: %s", e) - raise SessionException(f"Failed to create message: {e}") from e - - def read_message(self, session_id: str, agent_id: str, message_id: int, **kwargs: Any) -> Optional[SessionMessage]: - """Read a specific message by ID from AgentCore Memory. - - Args: - session_id (str): The session ID to read from. - agent_id (str): The agent ID associated with the message. - message_id (int): The message ID to read. - **kwargs (Any): Additional keyword arguments. - - Returns: - Optional[SessionMessage]: The message if found, None otherwise. - - Note: - This reads a single event by ID from AgentCore Memory. - """ - result = self.memory_client.gmdp_client.get_event( - memoryId=self.config.memory_id, actorId=self.config.actor_id, sessionId=session_id, eventId=message_id - ) - return SessionMessage.from_dict(result) if result else None - - def update_message(self, session_id: str, agent_id: str, session_message: SessionMessage, **kwargs: Any) -> None: - """Update message data in AgentCore Memory. - - Since AgentCore Memory events are immutable, this method performs an update by - creating a new event with the updated content and deleting the old event. - This enables features like guardrail redaction via Strands' redact_latest_message(). - - If the message has not yet been persisted (e.g., still in the message buffer when - batch_size > 1), the buffered message is replaced in-place instead. - - Args: - session_id (str): The session ID containing the message. - agent_id (str): The agent ID associated with the message. - session_message (SessionMessage): The message to update (with updated content - and the original message_id/eventId). - **kwargs (Any): Additional keyword arguments. - - Raises: - SessionException: If session ID doesn't match configuration or update fails. - """ - if session_id != self.config.session_id: - raise SessionException(f"Session ID mismatch: expected {self.config.session_id}, got {session_id}") - - old_message_id = session_message.message_id - - # If message hasn't been persisted yet (still in buffer), update it there - if old_message_id is None: - if self._update_buffered_message(session_message): - logger.debug("Updated buffered message (not yet persisted to AgentCore Memory)") - return - logger.debug("Message has no event ID and was not found in buffer - skipping update") - return - - # Create a new event with the updated message content - try: - updated_message = SessionMessage( - message=session_message.message, - message_id=0, - created_at=session_message.created_at, - ) - new_event = self.create_message(session_id, agent_id, updated_message) - except Exception as e: - logger.error("Failed to update message in AgentCore Memory: %s", e) - raise SessionException(f"Failed to update message: {e}") from e - - new_event_id = new_event.get("eventId") if new_event else None - if not new_event_id: - logger.warning("create_message did not return an eventId — skipping delete of old event %s", old_message_id) - return - - # Delete the old event; if this fails, roll back the newly created event - try: - self.memory_client.gmdp_client.delete_event( - memoryId=self.config.memory_id, - actorId=self.config.actor_id, - sessionId=session_id, - eventId=old_message_id, - ) - except Exception as delete_error: - logger.warning( - "Failed to delete old event %s after creating replacement: %s. Attempting rollback.", - old_message_id, - delete_error, - ) - try: - self.memory_client.gmdp_client.delete_event( - memoryId=self.config.memory_id, - actorId=self.config.actor_id, - sessionId=session_id, - eventId=new_event_id, - ) - logger.info("Rolled back new event %s after failed delete of old event", new_event_id) - except Exception as rollback_error: - logger.error( - "Rollback failed: could not delete new event %s: %s. Both old (%s) and new events may exist.", - new_event_id, - rollback_error, - old_message_id, - ) - raise SessionException( - f"Failed to update message: could not delete old event: {delete_error}" - ) from delete_error - - # Update _latest_agent_message so it doesn't hold a stale eventId - latest_messages = getattr(self, "_latest_agent_message", None) - if latest_messages and agent_id in latest_messages: - old_latest = self._latest_agent_message[agent_id] - if old_latest.message_id == old_message_id: - self._latest_agent_message[agent_id] = SessionMessage( - message=session_message.message, - message_id=new_event_id, - created_at=session_message.created_at, - ) - - logger.info("Updated message in AgentCore Memory: replaced event %s", old_message_id) - - def list_messages( - self, - session_id: str, - agent_id: str, - limit: Optional[int] = None, - offset: int = 0, - **kwargs: Any, - ) -> list[SessionMessage]: - """List messages for an agent from AgentCore Memory with pagination. - - Args: - session_id (str): The session ID to list messages from. - agent_id (str): The agent ID to list messages for. - limit (Optional[int], optional): Maximum number of messages to return. Defaults to None. - offset (int, optional): Number of messages to skip. Defaults to 0. - **kwargs (Any): Additional keyword arguments. - - Returns: - list[SessionMessage]: list of messages for the agent. - - Raises: - SessionException: If session ID doesn't match configuration. - """ - if session_id != self.config.session_id: - raise SessionException(f"Session ID mismatch: expected {self.config.session_id}, got {session_id}") - - try: - max_results = (limit + offset) if limit else MAX_FETCH_ALL_RESULTS - - events = self.memory_client.list_events( - memory_id=self.config.memory_id, - actor_id=self.config.actor_id, - session_id=session_id, - max_results=max_results, - ) - messages = self.converter.events_to_messages(events) - if self.config.filter_restored_tool_context: - messages = self._filter_restored_tool_context(messages) - if limit is not None: - return messages[offset : offset + limit] - else: - return messages[offset:] - - except Exception as e: - logger.error("Failed to list messages from AgentCore Memory: %s", e) - return [] - - def _filter_restored_tool_context(self, messages: list[SessionMessage]) -> list[SessionMessage]: - """Strip historical toolUse/toolResult context from restored messages.""" - filtered_messages: list[SessionMessage] = [] - for session_message in messages: - message = session_message.to_message() - filtered_content = [ - content - for content in message.get("content", []) - if "toolUse" not in content and "toolResult" not in content - ] - - if not filtered_content: - continue - - filtered_message: Message = {"role": message["role"], "content": filtered_content} - filtered_messages.append( - SessionMessage( - message=filtered_message, - message_id=session_message.message_id, - redact_message=session_message.redact_message, - created_at=session_message.created_at, - updated_at=session_message.updated_at, - ) - ) - - return filtered_messages - - # endregion SessionRepository interface implementation - - # region RepositorySessionManager overrides - @override - def append_message(self, message: Message, agent: "Agent", **kwargs: Any) -> None: - """Append a message to the agent's session using AgentCore's eventId as message_id. - - Args: - message: Message to add to the agent in the session - agent: Agent to append the message to - **kwargs: Additional keyword arguments for future extensibility. - """ - created_message = self.create_message(self.session_id, agent.agent_id, SessionMessage.from_message(message, 0)) - if created_message is None: - return - session_message = SessionMessage.from_message(message, created_message.get("eventId")) - self._latest_agent_message[agent.agent_id] = session_message - - def retrieve_customer_context(self, event: MessageAddedEvent) -> None: - """Retrieve customer LTM context before processing support query. - - Args: - event (MessageAddedEvent): The message added event containing the agent and message data. - """ - messages = event.agent.messages - if not messages or messages[-1].get("role") != "user": - return None - content = messages[-1].get("content") - if not content or "text" not in content[0]: - return None - if not self.config.retrieval_config: - # Only retrieve LTM - return None - - user_query = messages[-1]["content"][0]["text"] - - def retrieve_for_namespace(namespace: str, retrieval_config: RetrievalConfig): - """Helper function to retrieve memories for a single namespace.""" - resolved_namespace = namespace.format( - actorId=self.config.actor_id, - sessionId=self.config.session_id, - memoryStrategyId=retrieval_config.strategy_id or "", - ) - - memories = self.memory_client.retrieve_memories( - memory_id=self.config.memory_id, - namespace_path=resolved_namespace, - query=user_query, - top_k=retrieval_config.top_k, - ) - if retrieval_config.relevance_score: - memories = [m for m in memories if m.get("score", 0.0) >= retrieval_config.relevance_score] - context_items = [] - for memory in memories: - if isinstance(memory, dict): - content = memory.get("content", {}) - if isinstance(content, dict): - text = content.get("text", "").strip() - if text: - context_items.append(text) - return context_items - - try: - # Retrieve customer context from all namespaces in parallel - all_context = [] - - with ThreadPoolExecutor() as executor: - future_to_namespace = { - executor.submit(retrieve_for_namespace, namespace, retrieval_config): namespace - for namespace, retrieval_config in self.config.retrieval_config.items() - } - for future in as_completed(future_to_namespace): - try: - context_items = future.result() - all_context.extend(context_items) - except Exception as e: - # Continue processing other futures event if one fails rather than failing the entire operation - namespace = future_to_namespace[future] - logger.error("Failed to retrieve memories for namespace %s: %s", namespace, e) - - # Inject retrieved memory as a content block in the last user message. - # Prepended so the user's query text remains last (avoids assistant-prefill - # errors on Claude 4.6+ and keeps the user request in the position models - # attend to most). - if all_context: - context_text = "\n".join(all_context) - event.agent.messages[-1]["content"].insert( - 0, {"text": f"<{self.config.context_tag}>{context_text}"} - ) - logger.info("Retrieved %s customer context items", len(all_context)) - - except Exception as e: - logger.error("Failed to retrieve customer context: %s", e) - - @override - def register_hooks(self, registry: HookRegistry, **kwargs) -> None: - """Register additional hooks. - - In sync mode (the default), delegates to the base class and adds the - retrieve_customer_context + batching callbacks synchronously, preserving - existing behavior exactly. - - In async mode, registers async callbacks that wrap every per-turn - boto3-backed operation (append_message, sync_agent, buffer flushes, - customer-context retrieval) with asyncio.to_thread, so the asyncio - event loop stays free while boto3 is blocking on the network. - - Note: AgentInitializedEvent cannot be async per Strands' HookRegistry, - so agent restoration (read_session / read_agent / list_messages) still - blocks the calling thread in async mode — see AgentCoreMemoryConfig - docstring for mitigations. - - Args: - registry (HookRegistry): The hook registry to register callbacks with. - **kwargs: Additional keyword arguments. - """ - if not self.config.async_mode: - RepositorySessionManager.register_hooks(self, registry, **kwargs) - registry.add_callback(MessageAddedEvent, lambda event: self.retrieve_customer_context(event)) - - # Only register AfterInvocationEvent hook when batching is enabled - if self.config.batch_size > 1: - registry.add_callback(AfterInvocationEvent, lambda event: self._flush_messages()) - return - - # Async mode: register async callbacks that offload the existing sync - # methods to a worker thread via asyncio.to_thread. AgentInitializedEvent - # and BidiAgentInitializedEvent must stay sync (Strands disallows async - # callbacks for AgentInitializedEvent — see strands/hooks/registry.py:227). - logger.warning( - "AgentCoreMemorySessionManager async_mode=True: the agent must be invoked " - "via the async path (e.g. agent.stream_async(...) or agent.invoke_async(...)). " - "Sync invocation will raise RuntimeError from Strands' hook registry." - ) - - def _offload(method, *event_args): - """Build an async callback that offloads `method(*[a(event) for a in event_args])` to a thread. - - Each entry in `event_args` is a callable that extracts an argument from the event; - pass none for a zero-arg method. - """ - - async def _callback(event): - await asyncio.to_thread(method, *(extract(event) for extract in event_args)) - - return _callback - - registry.add_callback(AgentInitializedEvent, lambda event: self.initialize(event.agent)) - - async def _on_message_added_persist(event: MessageAddedEvent) -> None: - await asyncio.to_thread(self.append_message, event.message, event.agent) - await asyncio.to_thread(self.sync_agent, event.agent) - - registry.add_callback(MessageAddedEvent, _on_message_added_persist) - registry.add_callback(AfterInvocationEvent, _offload(self.sync_agent, lambda e: e.agent)) - registry.add_callback(MessageAddedEvent, _offload(self.retrieve_customer_context, lambda e: e)) - - if self.config.batch_size > 1: - registry.add_callback(AfterInvocationEvent, _offload(self._flush_messages)) - - # Register multi-agent callbacks so async-mode parity matches sync-mode - registry.add_callback(MultiAgentInitializedEvent, _offload(self.initialize_multi_agent, lambda e: e.source)) - registry.add_callback(AfterNodeCallEvent, _offload(self.sync_multi_agent, lambda e: e.source)) - registry.add_callback(AfterMultiAgentInvocationEvent, _offload(self.sync_multi_agent, lambda e: e.source)) - - # Register BidiAgent callbacks so async-mode parity matches sync-mode. - # BidiAgentInitializedEvent dispatches through invoke_callbacks (sync), - # so its callback must stay sync; the other two dispatch through - # invoke_callbacks_async, so async wrappers are safe. - registry.add_callback(BidiAgentInitializedEvent, lambda event: self.initialize_bidi_agent(event.agent)) - - async def _on_bidi_message_added(event: BidiMessageAddedEvent) -> None: - await asyncio.to_thread(self.append_bidi_message, event.message, event.agent) - await asyncio.to_thread(self.sync_bidi_agent, event.agent) - - registry.add_callback(BidiMessageAddedEvent, _on_bidi_message_added) - registry.add_callback(BidiAfterInvocationEvent, _offload(self.sync_bidi_agent, lambda e: e.agent)) - - @override - def initialize(self, agent: "Agent", **kwargs: Any) -> None: - if self.has_existing_agent: - logger.warning( - "An Agent already exists in session %s. We currently support one agent per session.", self.session_id - ) - else: - self.has_existing_agent = True - RepositorySessionManager.initialize(self, agent, **kwargs) - - # endregion RepositorySessionManager overrides - - # region Batching support - - def _update_buffered_message(self, session_message: SessionMessage) -> bool: - """Attempt to update a message that is still in the send buffer. - - When batch_size > 1, messages may not yet be persisted to AgentCore Memory. - This method finds the most recent buffered message matching the session_message's - content role and replaces it with the updated content. - - Args: - session_message (SessionMessage): The message with updated content. - - Returns: - bool: True if a buffered message was found and updated, False otherwise. - """ - updated_messages = self.converter.message_to_payload(session_message) - if not updated_messages: - return False - - is_blob = self.converter.exceeds_conversational_limit(updated_messages[0]) - - with self._message_lock: - # Search from the end (most recent) to find the message to update - for i in range(len(self._message_buffer) - 1, -1, -1): - buf = self._message_buffer[i] - if buf.session_id == self.config.session_id and buf.messages: - # Match by role - the most recent message with the same role - existing_role = buf.messages[0][1] if not buf.is_blob else None - new_role = updated_messages[0][1] if not is_blob else None - if existing_role == new_role: - self._message_buffer[i] = BufferedMessage( - session_id=buf.session_id, - messages=updated_messages, - is_blob=is_blob, - timestamp=buf.timestamp, - metadata=buf.metadata, - ) - return True - return False - - def _flush_messages_only(self) -> list[dict[str, Any]]: - """Flush only buffered messages to AgentCore Memory. - - Call this method to send any remaining buffered messages when batch_size > 1. - This is called when the message buffer reaches batch_size. - Messages are batched by session_id - all conversational messages for the same - session are combined into a single create_event() call to reduce API calls. - Messages that exceed the conversational payload limit are sent as blob events individually - as they require a different API path. - - Returns: - list[dict[str, Any]]: List of created event responses from AgentCore Memory. - - Raises: - SessionException: If message creation fails. On failure, messages remain in the buffer. - """ - if self.persistence_mode is PersistenceMode.NONE: - return [] - - with self._message_lock: - messages_to_send = list(self._message_buffer) - self._message_buffer.clear() - - if not messages_to_send: - return [] - - # Group all messages by session_id, combining conversational and blob messages - # Structure: {session_id: {"payload": [...], "timestamp": latest_timestamp, "metadata": {...}}} - session_groups: dict[str, dict[str, Any]] = {} - - for buffered_msg in messages_to_send: - sid = buffered_msg.session_id - if sid not in session_groups: - session_groups[sid] = {"payload": [], "timestamp": buffered_msg.timestamp, "metadata": {}} - - if buffered_msg.is_blob: - for msg in buffered_msg.messages: - session_groups[sid]["payload"].append({"blob": json.dumps(msg)}) - else: - for text, role in buffered_msg.messages: - session_groups[sid]["payload"].append( - {"conversational": {"content": {"text": text}, "role": role.upper()}} - ) - - # Use the latest timestamp for the combined event - if buffered_msg.timestamp > session_groups[sid]["timestamp"]: - session_groups[sid]["timestamp"] = buffered_msg.timestamp - - # Merge metadata (later entries override earlier for same key) - if buffered_msg.metadata: - session_groups[sid]["metadata"].update(buffered_msg.metadata) - - results = [] - try: - # Send one create_event per session_id with all messages (conversational + blob) - for session_id, group in session_groups.items(): - create_event_kwargs: dict[str, Any] = { - "memoryId": self.config.memory_id, - "actorId": self.config.actor_id, - "sessionId": session_id, - "payload": group["payload"], - "eventTimestamp": group["timestamp"], - } - if group["metadata"]: - create_event_kwargs["metadata"] = group["metadata"] - event = self.memory_client.gmdp_client.create_event(**create_event_kwargs) - results.append(event) - logger.debug( - "Flushed batched event for session %s with %d messages: %s", - session_id, - len(group["payload"]), - event.get("eventId"), - ) - - except Exception as e: - # Restore messages to buffer so they aren't lost - with self._message_lock: - self._message_buffer.extend(messages_to_send) - logger.error("Failed to flush messages to AgentCore Memory: %s", e) - raise SessionException(f"Failed to flush messages: {e}") from e - - logger.info("Flushed %d message events to AgentCore Memory", len(results)) - return results - - def _flush_agent_states_only(self) -> list[dict[str, Any]]: - """Flush only buffered agent states to AgentCore Memory. - - Call this method to send any remaining agent state when batch_size > 1. - Agent states are grouped by agent_id and sent as separate events so that - each event carries the correct agentId metadata for read_agent() lookups. - - Returns: - list[dict[str, Any]]: List of created event responses from AgentCore Memory. - - Raises: - SessionException: If agent state creation fails. On failure, agent states remain in the buffer. - """ - if self.persistence_mode is PersistenceMode.NONE: - return [] - - with self._agent_state_lock: - agent_states_to_send = list(self._agent_state_buffer) - self._agent_state_buffer.clear() - - if not agent_states_to_send: - return [] - - results = [] - try: - # Group agent states by agent_id - agent_groups: dict[str, list[dict]] = {} - for _session_id, session_agent in agent_states_to_send: - agent_id = session_agent.agent_id - if agent_id not in agent_groups: - agent_groups[agent_id] = [] - agent_groups[agent_id].append({"blob": json.dumps(session_agent.to_dict())}) - - # Send one event per agent_id with correct metadata - for agent_id, payloads in agent_groups.items(): - event = self.memory_client.gmdp_client.create_event( - memoryId=self.config.memory_id, - actorId=self.config.actor_id, - sessionId=self.config.session_id, - payload=payloads, - eventTimestamp=self._get_monotonic_timestamp(), - metadata={ - STATE_TYPE_KEY: {"stringValue": StateType.AGENT.value}, - AGENT_ID_KEY: {"stringValue": agent_id}, - }, - ) - results.append(event) - logger.debug("Flushed %d agent states for agent %s: %s", len(payloads), agent_id, event.get("eventId")) - - except Exception as e: - # Restore agent states to buffer so they aren't lost - with self._agent_state_lock: - self._agent_state_buffer.extend(agent_states_to_send) - logger.error("Failed to flush agent states to AgentCore Memory: %s", e) - raise SessionException(f"Failed to flush agent states: {e}") from e - - logger.info("Flushed %d agent state events to AgentCore Memory", len(results)) - return results - - def _flush_messages(self) -> list[dict[str, Any]]: - """Flush all buffered messages and agent state to AgentCore Memory. - - Call this method to send any remaining buffered messages and agent state messages. - This is automatically called when the session is complete (via close() or context manager). - - Returns: - list[dict[str, Any]]: List of created event responses from AgentCore Memory. - - Raises: - SessionException: If any message or agent state creation fails. - """ - results = [] - results.extend(self._flush_messages_only()) - results.extend(self._flush_agent_states_only()) - return results - - def pending_message_count(self) -> int: - """Return the number of messages pending in the buffer. - - Returns: - int: Number of buffered messages waiting to be sent. - """ - with self._message_lock: - return len(self._message_buffer) - - def pending_agent_state_count(self) -> int: - """Return the number of agent states pending in the buffer. - - Returns: - int: Number of buffered agent states waiting to be sent. - """ - with self._agent_state_lock: - return len(self._agent_state_buffer) - - def close(self) -> None: - """Explicitly flush pending messages and close the session manager. - - Call this method when the session is complete to ensure all buffered - messages are sent to AgentCore Memory. Alternatively, use the context - manager protocol (with statement) for automatic cleanup. - """ - self._stop_flush_timer() - self._flush_messages() - - def __enter__(self) -> "AgentCoreMemorySessionManager": - """Enter the context manager. - - Returns: - AgentCoreMemorySessionManager: This session manager instance. - """ - return self - - def __exit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None: - """Exit the context manager and flush any pending messages. - - Args: - exc_type: Exception type if an exception occurred. - exc_val: Exception value if an exception occurred. - exc_tb: Exception traceback if an exception occurred. - """ - try: - self._stop_flush_timer() - self._flush_messages() - except Exception as e: - if exc_type is not None: - logger.error("Failed to flush messages during exception handling: %s", e) - else: - raise - - # endregion Batching support - - # region Interval-based flushing support - - def _start_flush_timer(self) -> None: - """Start the interval-based flush timer. - - This method schedules a recurring timer that flushes the message buffer - at regular intervals if flush_interval_seconds is configured. - """ - with self._timer_lock: - if self._shutdown: - return - - # Cancel existing timer if any - if self._flush_timer is not None: - self._flush_timer.cancel() - - # Schedule next flush - self._flush_timer = threading.Timer( - self.config.flush_interval_seconds, - self._interval_flush_callback, - ) - self._flush_timer.daemon = True - self._flush_timer.start() - logger.debug( - "Scheduled interval flush in %.1f seconds", - self.config.flush_interval_seconds, - ) - - def _interval_flush_callback(self) -> None: - """Callback executed by the flush timer. - - Flushes the buffer if it contains messages or agent states, then reschedules the timer. - """ - try: - # Only flush if there are messages or agent states in the buffer - pending_messages = self.pending_message_count() - pending_agent_states = self.pending_agent_state_count() - if pending_messages > 0 or pending_agent_states > 0: - logger.debug( - "Interval flush triggered: %d message(s) and %d agent state(s) pending", - pending_messages, - pending_agent_states, - ) - self._flush_messages() - else: - logger.debug("Interval flush skipped: buffers are empty") - - # Reschedule the timer (unless shutdown) - if not self._shutdown and self.config.flush_interval_seconds: - self._start_flush_timer() - - except Exception as e: - logger.error("Error during interval flush: %s", e) - # Attempt to reschedule even after error - if not self._shutdown and self.config.flush_interval_seconds: - self._start_flush_timer() - - def _stop_flush_timer(self) -> None: - """Stop the interval-based flush timer. - - This method cancels the timer and prevents it from rescheduling. - Should be called during cleanup (close() or __exit__). - """ - with self._timer_lock: - self._shutdown = True - if self._flush_timer is not None: - self._flush_timer.cancel() - self._flush_timer = None - logger.debug("Stopped interval flush timer") - - # endregion Interval-based flushing support + from .memorysessionmanager.session_manager import ( + AGENT_ID_KEY as AGENT_ID_KEY, + ) + from .memorysessionmanager.session_manager import ( + LEGACY_AGENT_PREFIX as LEGACY_AGENT_PREFIX, + ) + from .memorysessionmanager.session_manager import ( + LEGACY_SESSION_PREFIX as LEGACY_SESSION_PREFIX, + ) + from .memorysessionmanager.session_manager import ( + MAX_FETCH_ALL_RESULTS as MAX_FETCH_ALL_RESULTS, + ) + from .memorysessionmanager.session_manager import ( + MAX_METADATA_KEYS as MAX_METADATA_KEYS, + ) + from .memorysessionmanager.session_manager import ( + RESERVED_METADATA_KEYS as RESERVED_METADATA_KEYS, + ) + from .memorysessionmanager.session_manager import ( + STATE_TYPE_KEY as STATE_TYPE_KEY, + ) + from .memorysessionmanager.session_manager import ( + AgentCoreMemorySessionManager as AgentCoreMemorySessionManager, + ) + from .memorysessionmanager.session_manager import BufferedMessage as BufferedMessage + from .memorysessionmanager.session_manager import StateType as StateType + +sys.modules[__name__] = _canonical_module diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/__init__.py b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_config.py b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_agentcore_memory_config.py similarity index 97% rename from tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_config.py rename to tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_agentcore_memory_config.py index 97f0bdae..0e8f50f7 100644 --- a/tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_config.py +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_agentcore_memory_config.py @@ -3,7 +3,10 @@ import pytest from pydantic import ValidationError -from bedrock_agentcore.memory.integrations.strands.config import AgentCoreMemoryConfig, RetrievalConfig +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.config import ( + AgentCoreMemoryConfig, + RetrievalConfig, +) class TestRetrievalConfig: diff --git a/tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_session_manager.py b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_agentcore_memory_session_manager.py similarity index 97% rename from tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_session_manager.py rename to tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_agentcore_memory_session_manager.py index fa0c4787..0450d490 100644 --- a/tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_session_manager.py +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_agentcore_memory_session_manager.py @@ -26,12 +26,16 @@ from strands.types.exceptions import SessionException from strands.types.session import Session, SessionAgent, SessionMessage, SessionType -from bedrock_agentcore.memory.integrations.strands.bedrock_converter import ( +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.bedrock_converter import ( CONVERSATIONAL_MAX_SIZE, AgentCoreMemoryConverter, ) -from bedrock_agentcore.memory.integrations.strands.config import AgentCoreMemoryConfig, PersistenceMode, RetrievalConfig -from bedrock_agentcore.memory.integrations.strands.session_manager import ( +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.config import ( + AgentCoreMemoryConfig, + PersistenceMode, + RetrievalConfig, +) +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager import ( AgentCoreMemorySessionManager, BufferedMessage, ) @@ -74,7 +78,7 @@ def _create_session_manager(config, mock_memory_client): """Helper to create a session manager with mocked dependencies.""" with ( patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ), patch("boto3.Session") as mock_boto_session, @@ -125,7 +129,9 @@ class TestAgentCoreMemorySessionManager: def test_init_basic(self, agentcore_config): """Test basic initialization.""" - with patch("bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient") as mock_client_class: + with patch( + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient" + ) as mock_client_class: mock_client = Mock() mock_client_class.return_value = mock_client @@ -717,7 +723,7 @@ def test_load_long_term_memories_with_validation_failure(self, mock_memory_clien ) with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -823,7 +829,7 @@ def test_retrieve_contextual_memories_all_namespaces(self, agentcore_config_with ] with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -861,7 +867,7 @@ def test_retrieve_contextual_memories_specific_namespaces( mock_memory_client.retrieve_memories.return_value = [{"content": "User preference memory", "score": 0.9}] with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -902,7 +908,7 @@ def test_retrieve_contextual_memories_no_config(self, session_manager): def test_retrieve_contextual_memories_invalid_namespace(self, agentcore_config_with_retrieval, mock_memory_client): """Test contextual memory retrieval with invalid namespace.""" with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -929,7 +935,7 @@ def test_load_long_term_memories_with_config(self, agentcore_config_with_retriev ] with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -955,7 +961,7 @@ def test_load_long_term_memories_exception_handling( mock_memory_client.retrieve_memories.side_effect = Exception("API Error") with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -1053,7 +1059,7 @@ def mock_retrieve_side_effect(*args, **kwargs): mock_memory_client.retrieve_memories.side_effect = mock_retrieve_side_effect with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -1099,7 +1105,7 @@ def test_initialize_with_ltm_integration(self, agentcore_config_with_retrieval, mock_memory_client.retrieve_memories.return_value = [{"content": "User prefers morning meetings", "score": 0.8}] with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -1129,7 +1135,7 @@ def test_init_with_boto_config(self, agentcore_config, mock_memory_client): boto_config = BotocoreConfig(user_agent_extra="custom-agent") with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -1147,7 +1153,7 @@ def test_init_with_boto_config(self, agentcore_config, mock_memory_client): def test_retrieve_customer_context_no_messages(self, agentcore_config_with_retrieval, mock_memory_client): """Test retrieve_customer_context with no messages.""" with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -1172,7 +1178,7 @@ def test_retrieve_customer_context_no_messages(self, agentcore_config_with_retri def test_retrieve_customer_context_empty_content(self, agentcore_config_with_retrieval, mock_memory_client): """Empty content list on the last message must not raise IndexError.""" with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -1197,7 +1203,7 @@ def test_retrieve_customer_context_empty_content(self, agentcore_config_with_ret def test_retrieve_customer_context_no_config(self, agentcore_config, mock_memory_client): """Test retrieve_customer_context with no retrieval config.""" with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -1226,7 +1232,7 @@ def test_retrieve_customer_context_with_memories(self, agentcore_config_with_ret ] with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -1254,7 +1260,7 @@ def test_retrieve_customer_context_exception(self, agentcore_config_with_retriev mock_memory_client.retrieve_memories.side_effect = Exception("Memory error") with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -1295,7 +1301,7 @@ def test_retrieve_customer_context_filters_by_relevance_score(self, mock_memory_ ) with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -2520,7 +2526,7 @@ def test_retrieve_customer_context_does_not_append_assistant_message( ] with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -2561,7 +2567,7 @@ def test_retrieve_customer_context_no_assistant_message_multi_turn( ] with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -2612,7 +2618,7 @@ def test_retrieve_customer_context_custom_context_tag(self, mock_memory_client): ] with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -2652,7 +2658,7 @@ def test_retrieve_customer_context_default_context_tag(self, mock_memory_client) ] with patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ): with patch("boto3.Session") as mock_boto_session: @@ -2764,7 +2770,7 @@ def test_interval_flush_timer_starts_when_configured(self): with ( patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_client, ), patch("boto3.Session") as mock_boto_session, @@ -2800,7 +2806,7 @@ def test_interval_flush_timer_not_started_when_disabled(self): with ( patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_client, ), patch("boto3.Session") as mock_boto_session, @@ -2832,7 +2838,7 @@ def test_interval_flush_timer_stops_on_close(self): with ( patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_client, ), patch("boto3.Session") as mock_boto_session, @@ -2876,7 +2882,7 @@ def test_interval_flush_timer_stops_on_context_exit(self): with ( patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_client, ), patch("boto3.Session") as mock_boto_session, @@ -2911,7 +2917,7 @@ def test_interval_flush_callback_flushes_when_buffer_has_messages(self): with ( patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_client, ), patch("boto3.Session") as mock_boto_session, @@ -2958,7 +2964,7 @@ def test_interval_flush_callback_skips_when_buffer_empty(self): with ( patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_client, ), patch("boto3.Session") as mock_boto_session, @@ -3009,7 +3015,7 @@ def test_interval_flush_callback_flushes_when_agent_state_pending(self): with ( patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_client, ), patch("boto3.Session") as mock_boto_session, @@ -3065,7 +3071,7 @@ def test_interval_flush_callback_flushes_when_both_buffers_have_data(self): with ( patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_client, ), patch("boto3.Session") as mock_boto_session, @@ -3209,7 +3215,9 @@ def test_metadata_merging_precedence(self, session_manager_with_metadata, mock_m def test_metadata_reserved_keys_rejected(self, session_manager): """ValueError raised when user metadata contains reserved keys.""" - from bedrock_agentcore.memory.integrations.strands.session_manager import RESERVED_METADATA_KEYS + from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager import ( + RESERVED_METADATA_KEYS, + ) session_message = SessionMessage.from_message({"role": "user", "content": [{"text": "hello"}]}, 0) @@ -3224,7 +3232,7 @@ def test_metadata_reserved_keys_rejected(self, session_manager): def test_metadata_max_keys_exceeded(self, session_manager): """ValueError raised when combined metadata exceeds MAX_METADATA_KEYS.""" - from bedrock_agentcore.memory.integrations.strands.session_manager import MAX_METADATA_KEYS + from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager import MAX_METADATA_KEYS session_message = SessionMessage.from_message({"role": "user", "content": [{"text": "hello"}]}, 0) too_many = {f"key_{i}": {"stringValue": f"val_{i}"} for i in range(MAX_METADATA_KEYS + 1)} @@ -3283,7 +3291,9 @@ def test_batched_messages_include_metadata(self, mock_memory_client): def test_blob_message_with_metadata(self, session_manager_with_metadata, mock_memory_client): """Blob messages also receive metadata.""" - from bedrock_agentcore.memory.integrations.strands.bedrock_converter import CONVERSATIONAL_MAX_SIZE + from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.bedrock_converter import ( + CONVERSATIONAL_MAX_SIZE, + ) mock_memory_client.gmdp_client.create_event.return_value = {"event": {"eventId": "blob_1"}} big_text = "x" * (CONVERSATIONAL_MAX_SIZE + 100) @@ -3722,9 +3732,14 @@ def test_async_mode_registers_multi_agent_callbacks(self, mock_memory_client): registry = HookRegistry() manager.register_hooks(registry) - for event_type in (MultiAgentInitializedEvent, AfterNodeCallEvent, AfterMultiAgentInvocationEvent): - callbacks = registry._registered_callbacks.get(event_type, []) - assert callbacks, f"No callbacks registered for {event_type.__name__}" + events = ( + MultiAgentInitializedEvent(source=Mock()), + AfterNodeCallEvent(source=Mock(), node_id="node"), + AfterMultiAgentInvocationEvent(source=Mock()), + ) + for event in events: + callbacks = list(registry.get_callbacks_for(event)) + assert callbacks, f"No callbacks registered for {type(event).__name__}" assert all(asyncio.iscoroutinefunction(cb) for cb in callbacks) def test_async_mode_logs_sync_invocation_warning(self, mock_memory_client, caplog): @@ -3733,7 +3748,9 @@ def test_async_mode_logs_sync_invocation_warning(self, mock_memory_client, caplo manager = _create_session_manager(config, mock_memory_client) registry = HookRegistry() - with caplog.at_level(logging.WARNING, logger="bedrock_agentcore.memory.integrations.strands.session_manager"): + with caplog.at_level( + logging.WARNING, logger="bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager" + ): manager.register_hooks(registry) assert any("async_mode=True" in rec.message and "stream_async" in rec.message for rec in caplog.records) @@ -3746,15 +3763,20 @@ def test_async_mode_registers_bidi_agent_callbacks(self, mock_memory_client): manager.register_hooks(registry) # BidiAgentInitializedEvent dispatches via the sync hook path, so its callback must NOT be a coroutine. - init_callbacks = registry._registered_callbacks.get(BidiAgentInitializedEvent, []) + init_event = BidiAgentInitializedEvent(agent=Mock()) + init_callbacks = list(registry.get_callbacks_for(init_event)) assert init_callbacks, "No callbacks registered for BidiAgentInitializedEvent" assert not any(asyncio.iscoroutinefunction(cb) for cb in init_callbacks) # BidiMessageAddedEvent and BidiAfterInvocationEvent dispatch via invoke_callbacks_async, # so their callbacks should be async to keep the event loop unblocked. - for event_type in (BidiMessageAddedEvent, BidiAfterInvocationEvent): - callbacks = registry._registered_callbacks.get(event_type, []) - assert callbacks, f"No callbacks registered for {event_type.__name__}" + events = ( + BidiMessageAddedEvent(agent=Mock(), message={"role": "user", "content": [{"text": "hello"}]}), + BidiAfterInvocationEvent(agent=Mock()), + ) + for event in events: + callbacks = list(registry.get_callbacks_for(event)) + assert callbacks, f"No callbacks registered for {type(event).__name__}" assert all(asyncio.iscoroutinefunction(cb) for cb in callbacks) diff --git a/tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_session_manager_openai_converter.py b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_agentcore_memory_session_manager_openai_converter.py similarity index 88% rename from tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_session_manager_openai_converter.py rename to tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_agentcore_memory_session_manager_openai_converter.py index e16fd25a..04a71887 100644 --- a/tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_session_manager_openai_converter.py +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_agentcore_memory_session_manager_openai_converter.py @@ -4,9 +4,11 @@ from strands.types.session import Session, SessionMessage, SessionType -from bedrock_agentcore.memory.integrations.strands.config import AgentCoreMemoryConfig -from bedrock_agentcore.memory.integrations.strands.converters import OpenAIConverseConverter -from bedrock_agentcore.memory.integrations.strands.session_manager import AgentCoreMemorySessionManager +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.config import AgentCoreMemoryConfig +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.converters import OpenAIConverseConverter +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager import ( + AgentCoreMemorySessionManager, +) def test_create_message_uses_tool_role_with_openai_converter(): @@ -25,7 +27,7 @@ def test_create_message_uses_tool_role_with_openai_converter(): with ( patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ), patch("boto3.Session") as mock_boto_session, @@ -79,7 +81,7 @@ def test_list_messages_filters_restored_tool_context(): with ( patch( - "bedrock_agentcore.memory.integrations.strands.session_manager.MemoryClient", + "bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager.MemoryClient", return_value=mock_memory_client, ), patch("boto3.Session") as mock_boto_session, diff --git a/tests/bedrock_agentcore/memory/integrations/strands/test_bedrock_converter.py b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_bedrock_converter.py similarity index 97% rename from tests/bedrock_agentcore/memory/integrations/strands/test_bedrock_converter.py rename to tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_bedrock_converter.py index a1f66bd8..45c649a5 100644 --- a/tests/bedrock_agentcore/memory/integrations/strands/test_bedrock_converter.py +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_bedrock_converter.py @@ -5,7 +5,9 @@ from strands.types.session import SessionMessage -from bedrock_agentcore.memory.integrations.strands.bedrock_converter import AgentCoreMemoryConverter +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.bedrock_converter import ( + AgentCoreMemoryConverter, +) def _make_conversational_event(session_messages): @@ -72,7 +74,7 @@ def test_events_to_messages_blob_valid(self): assert len(result) == 1 assert result[0].message["role"] == "user" - @patch("bedrock_agentcore.memory.integrations.strands.bedrock_converter.logger") + @patch("bedrock_agentcore.memory.integrations.strands.memorysessionmanager.bedrock_converter.logger") def test_events_to_messages_blob_invalid_json(self, mock_logger): """Test handling invalid JSON in blob events.""" events = [{"payload": [{"blob": "invalid json"}]}] @@ -82,7 +84,7 @@ def test_events_to_messages_blob_invalid_json(self, mock_logger): assert len(result) == 0 mock_logger.error.assert_called() - @patch("bedrock_agentcore.memory.integrations.strands.bedrock_converter.logger") + @patch("bedrock_agentcore.memory.integrations.strands.memorysessionmanager.bedrock_converter.logger") def test_events_to_messages_blob_invalid_session_message(self, mock_logger): """Test handling invalid SessionMessage in blob events.""" blob_data = ["invalid", "user"] @@ -343,7 +345,7 @@ def test_events_to_messages_mixed_blob_and_conversational_ordering(self): assert result[0].message["content"][0]["text"] == "First" assert result[1].message["content"][0]["text"] == "Second" - @patch("bedrock_agentcore.memory.integrations.strands.bedrock_converter.logger") + @patch("bedrock_agentcore.memory.integrations.strands.memorysessionmanager.bedrock_converter.logger") def test_events_to_messages_malformed_payload_does_not_break_batch(self, mock_logger): """Test a malformed blob payload between two valid conversational payloads in a single event.""" msg1 = SessionMessage( diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_compatibility_imports.py b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_compatibility_imports.py new file mode 100644 index 00000000..097f2497 --- /dev/null +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_compatibility_imports.py @@ -0,0 +1,116 @@ +"""Compatibility tests for pre-existing SessionManager import paths.""" + +import importlib +import subprocess +import sys +from unittest.mock import Mock, patch + +from bedrock_agentcore.memory.integrations.strands import MemoryConverter, OpenAIConverseConverter +from bedrock_agentcore.memory.integrations.strands.bedrock_converter import AgentCoreMemoryConverter +from bedrock_agentcore.memory.integrations.strands.config import AgentCoreMemoryConfig, PersistenceMode, RetrievalConfig +from bedrock_agentcore.memory.integrations.strands.converters import ( + MemoryConverter as CompatibilityMemoryConverter, +) +from bedrock_agentcore.memory.integrations.strands.converters import ( + OpenAIConverseConverter as CompatibilityOpenAIConverseConverter, +) +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager import ( + AgentCoreMemoryConfig as CanonicalAgentCoreMemoryConfig, +) +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager import ( + AgentCoreMemoryConverter as CanonicalAgentCoreMemoryConverter, +) +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager import ( + AgentCoreMemorySessionManager as CanonicalAgentCoreMemorySessionManager, +) +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager import ( + MemoryConverter as CanonicalMemoryConverter, +) +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager import ( + OpenAIConverseConverter as CanonicalOpenAIConverseConverter, +) +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager import ( + PersistenceMode as CanonicalPersistenceMode, +) +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager import ( + RetrievalConfig as CanonicalRetrievalConfig, +) +from bedrock_agentcore.memory.integrations.strands.session_manager import AgentCoreMemorySessionManager + +OLD_MODULE_ROOT = "bedrock_agentcore.memory.integrations.strands" +CANONICAL_MODULE_ROOT = f"{OLD_MODULE_ROOT}.memorysessionmanager" + + +def test_preexisting_import_paths_are_explicit_aliases() -> None: + """Keep existing imports working without warning or duplicate implementations.""" + assert AgentCoreMemorySessionManager is CanonicalAgentCoreMemorySessionManager + assert AgentCoreMemoryConfig is CanonicalAgentCoreMemoryConfig + assert AgentCoreMemoryConverter is CanonicalAgentCoreMemoryConverter + assert PersistenceMode is CanonicalPersistenceMode + assert RetrievalConfig is CanonicalRetrievalConfig + assert MemoryConverter is CompatibilityMemoryConverter is CanonicalMemoryConverter + assert OpenAIConverseConverter is CompatibilityOpenAIConverseConverter is CanonicalOpenAIConverseConverter + + +MODULE_SUFFIXES = ( + "session_manager", + "bedrock_converter", + "config", + "converters", + "converters.protocol", + "converters.openai", +) + + +def test_preexisting_module_paths_are_canonical_module_objects() -> None: + """Compatibility paths must share globals with their canonical modules.""" + for suffix in MODULE_SUFFIXES: + old_module = importlib.import_module(f"{OLD_MODULE_ROOT}.{suffix}") + canonical_module = importlib.import_module(f"{CANONICAL_MODULE_ROOT}.{suffix}") + assert old_module is canonical_module + + +def test_preexisting_module_paths_are_canonical_when_imported_first() -> None: + """Old paths must alias canonical modules even before an explicit canonical import.""" + script = f""" +import importlib + +old_root = {OLD_MODULE_ROOT!r} +canonical_root = {CANONICAL_MODULE_ROOT!r} +suffixes = {MODULE_SUFFIXES!r} +for suffix in suffixes: + old_module = importlib.import_module(f"{{old_root}}.{{suffix}}") + canonical_module = importlib.import_module(f"{{canonical_root}}.{{suffix}}") + assert old_module is canonical_module, suffix +""" + + subprocess.run([sys.executable, "-c", script], check=True) + + +def test_patching_preexisting_memory_client_path_affects_session_manager_runtime() -> None: + """Patching MemoryClient through the old path must affect canonical construction.""" + config = AgentCoreMemoryConfig(memory_id="memory-id", session_id="session-id", actor_id="actor-id") + memory_client = Mock() + boto_session = Mock(region_name="us-west-2") + boto_session.client.return_value = Mock() + + with ( + patch(f"{OLD_MODULE_ROOT}.session_manager.MemoryClient", return_value=memory_client) as memory_client_class, + patch("boto3.Session", return_value=boto_session), + patch("strands.session.repository_session_manager.RepositorySessionManager.__init__", return_value=None), + ): + manager = CanonicalAgentCoreMemorySessionManager(config) + + assert manager.memory_client is memory_client + memory_client_class.assert_called_once_with(region_name=None) + + +def test_patching_preexisting_bedrock_converter_logger_affects_canonical_global() -> None: + """Patching logger through the old path must affect canonical converter behavior.""" + canonical_module = importlib.import_module(f"{CANONICAL_MODULE_ROOT}.bedrock_converter") + + with patch(f"{OLD_MODULE_ROOT}.bedrock_converter.logger") as old_path_logger: + assert canonical_module.logger is old_path_logger + CanonicalAgentCoreMemoryConverter.events_to_messages([{"payload": [{"blob": "invalid json"}]}]) + + old_path_logger.error.assert_called_once_with("Failed to parse blob content: %s", {"blob": "invalid json"}) diff --git a/tests/bedrock_agentcore/memory/integrations/strands/test_openai_converter.py b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_openai_converter.py similarity index 97% rename from tests/bedrock_agentcore/memory/integrations/strands/test_openai_converter.py rename to tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_openai_converter.py index 7b94db92..27a727c0 100644 --- a/tests/bedrock_agentcore/memory/integrations/strands/test_openai_converter.py +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorysessionmanager/test_openai_converter.py @@ -4,7 +4,7 @@ from strands.types.session import SessionMessage -from bedrock_agentcore.memory.integrations.strands.converters import OpenAIConverseConverter +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.converters import OpenAIConverseConverter def test_tool_result_message_serializes_as_openai_tool_role(): diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorystore/__init__.py b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_factory.py b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_factory.py new file mode 100644 index 00000000..32246909 --- /dev/null +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_factory.py @@ -0,0 +1,279 @@ +"""Tests for AgentCore multi-namespace store construction.""" + +from typing import Any +from unittest.mock import Mock + +import pytest +from strands.memory import ExtractionTrigger, ExtractionTriggerContext, MemoryMessageFilter + +from bedrock_agentcore.memory.integrations.strands.memorystore.factory import ( + assert_writable_topology, + create_agentcore_memory_stores, +) +from bedrock_agentcore.memory.integrations.strands.memorystore.store import AgentCoreMemoryStore + + +class FakeTrigger(ExtractionTrigger): + """Minimal custom Strands extraction cadence.""" + + name = "fake" + + def attach(self, context: ExtractionTriggerContext) -> None: + """Accept an extraction context without registering hooks.""" + + +def base_input(**overrides: Any) -> dict[str, Any]: + """Build a two-namespace writable factory input.""" + result: dict[str, Any] = { + "memory_id": "mem-1", + "actor_id": "actor-1", + "session_id": "sess-1", + "namespaces": [ + {"namespace": "/strategy/s/actor/{actorId}/facts"}, + {"namespace": "/strategy/s/actor/{actorId}/preferences"}, + ], + "extraction": {"cadence": FakeTrigger()}, + "client": Mock(), + } + result.update(overrides) + return result + + +def test_returns_one_store_per_namespace_with_one_default_writer() -> None: + """The factory creates one store per read namespace and one write sink.""" + stores = create_agentcore_memory_stores(**base_input()) + assert len(stores) == 2 + writers = [store for store in stores if store.writable] + assert len(writers) == 1 + assert writers[0].name == "strategy-s-actor-facts" + assert writers[0].extraction is not None + assert next(store for store in stores if not store.writable).extraction is None + + +def test_custom_cadence_and_filter_reach_writer() -> None: + """Translate factory extraction options to Strands ``ExtractionConfig`` keys.""" + trigger = FakeTrigger() + message_filter = MemoryMessageFilter(exclude=["toolUse", "toolResult", "image"]) + stores = create_agentcore_memory_stores(**base_input(extraction={"cadence": trigger, "filter": message_filter})) + extraction = next(store.extraction for store in stores if store.writable) + assert extraction == {"trigger": trigger, "filter": message_filter} + + +def test_explicit_writer_flag_selects_non_first_namespace() -> None: + """Honor the namespace explicitly designated as the write sink.""" + stores = create_agentcore_memory_stores( + **base_input( + namespaces=[ + {"namespace": "/strategy/s/actor/{actorId}/facts"}, + { + "namespace": "/strategy/s/actor/{actorId}/preferences", + "writable": True, + }, + ] + ) + ) + writers = [store for store in stores if store.writable] + assert [store.name for store in writers] == ["strategy-s-actor-preferences"] + + +def test_explicit_opt_out_skips_first_default_writer_candidate() -> None: + """Do not override a namespace's ``writable=False`` opt-out.""" + stores = create_agentcore_memory_stores( + **base_input( + namespaces=[ + { + "namespace": "/strategy/s/actor/{actorId}/facts", + "writable": False, + }, + {"namespace": "/strategy/s/actor/{actorId}/preferences"}, + ], + extraction=True, + ) + ) + assert [store.name for store in stores if store.writable] == ["strategy-s-actor-preferences"] + + +def test_all_explicit_opt_outs_reject_enabled_extraction() -> None: + """Enabled extraction requires an eligible write sink.""" + with pytest.raises(ValueError, match="every namespace is marked writable: false"): + create_agentcore_memory_stores( + **base_input( + namespaces=[ + {"namespace": "/a/{actorId}", "writable": False}, + {"namespace": "/b/{actorId}", "writable": False}, + ], + extraction=True, + ) + ) + + +def test_multiple_explicit_writers_are_rejected() -> None: + """Namespace-free ``create_event`` would otherwise duplicate writes.""" + with pytest.raises(ValueError, match="at most one store may be writable"): + create_agentcore_memory_stores( + **base_input( + namespaces=[ + {"namespace": "/a/{actorId}", "writable": True}, + {"namespace": "/b/{actorId}", "writable": True}, + ] + ) + ) + + +@pytest.mark.parametrize("extraction", [None, False]) +def test_recall_only_has_no_writer(extraction: object) -> None: + """Omitted and false extraction both construct read-only stores.""" + input_data = base_input(extraction=extraction) + if extraction is None: + input_data.pop("extraction") + stores = create_agentcore_memory_stores(**input_data) + assert all(not store.writable for store in stores) + assert all(store.extraction is None for store in stores) + + +def test_extraction_true_passes_framework_default_shorthand() -> None: + """Let MemoryManager choose its standard cadence.""" + stores = create_agentcore_memory_stores(**base_input(extraction=True)) + assert next(store.extraction for store in stores if store.writable) is True + + +def test_names_are_derived_or_respected() -> None: + """Use explicit names and a fallback for placeholder-only namespaces.""" + stores = create_agentcore_memory_stores( + **base_input( + namespaces=[ + {"namespace": "/a/{actorId}", "name": "alpha"}, + {"namespace": "/b/{actorId}", "name": "beta"}, + ] + ) + ) + assert [store.name for store in stores] == ["alpha", "beta"] + fallback = create_agentcore_memory_stores(**base_input(namespaces=[{"namespace": "{actorId}"}])) + assert fallback[0].name == "agentcore-memory" + + +def test_factory_shares_one_client_across_stores() -> None: + """Construct or accept one boto3 client for the complete topology.""" + client = Mock() + stores = create_agentcore_memory_stores(**base_input(client=client)) + assert all(isinstance(store, AgentCoreMemoryStore) and store._client is client for store in stores) + + +@pytest.mark.parametrize( + "namespaces", + [[], [{"namespace": " "}], [{}], [None]], +) +def test_rejects_missing_or_invalid_namespaces(namespaces: list[object]) -> None: + """Require at least one non-empty namespace string.""" + expected = "at least one namespace" if not namespaces else r"namespaces\[0\]\.namespace" + with pytest.raises(ValueError, match=expected): + create_agentcore_memory_stores(**base_input(namespaces=namespaces)) + + +def test_namespace_validation_uses_python_strip_semantics() -> None: + """Reject Python whitespace-only namespaces and retain BOM content.""" + stores = create_agentcore_memory_stores(**base_input(namespaces=[{"namespace": "\ufeff"}])) + assert len(stores) == 1 + with pytest.raises(ValueError, match=r"namespaces\[0\]\.namespace must be a non-empty"): + create_agentcore_memory_stores(**base_input(namespaces=[{"namespace": "\u0085"}])) + + +@pytest.mark.parametrize( + "override", + [{"actor_id": ""}, {"session_id": " "}, {"memory_id": ""}], +) +def test_identity_validation_propagates_from_store(override: dict[str, str]) -> None: + """Keep flat identity validation consistent with direct construction.""" + with pytest.raises(ValueError, match="must be a non-empty string"): + create_agentcore_memory_stores(**base_input(**override)) + + +def test_unresolved_placeholder_validation_propagates() -> None: + """Reject unsupported namespace placeholders in the factory path too.""" + with pytest.raises(ValueError, match=r"\{memoryStrategyId\}"): + create_agentcore_memory_stores( + **base_input(namespaces=[{"namespace": "/strategies/{memoryStrategyId}/actors/{actorId}"}]) + ) + + +@pytest.mark.parametrize("value", [0, -1, 2.5, True]) +def test_factory_validates_event_cap_even_for_recall_only(value: object) -> None: + """Validate tuning even when no sender is built.""" + with pytest.raises(ValueError, match="positive integer"): + create_agentcore_memory_stores(**base_input(extraction=False, max_turns_per_event=value)) + + +def hand_built_store(*, name: str = "facts", writable: bool = False) -> AgentCoreMemoryStore: + """Build one store for topology assertions.""" + return AgentCoreMemoryStore( + memory_id="mem-1", + actor_id="actor-1", + session_id="sess-1", + namespace="/users/{actorId}/facts", + name=name, + writable=writable, + client=Mock(), + ) + + +def test_assert_writable_topology_accepts_zero_or_one_writer() -> None: + """Recall-only and exactly-one-writer topologies are valid.""" + assert_writable_topology([hand_built_store(), hand_built_store(name="prefs")]) + assert_writable_topology([hand_built_store(writable=True), hand_built_store(name="prefs")]) + + +def test_assert_writable_topology_rejects_multiple_writers() -> None: + """Hand-built store sets can use the same exported guard.""" + with pytest.raises(ValueError, match="at most one store may be writable"): + assert_writable_topology([hand_built_store(name="a", writable=True), hand_built_store(name="b", writable=True)]) + + +def test_assert_writable_topology_can_require_writer() -> None: + """Expected extraction turns zero writers into an error.""" + with pytest.raises(ValueError, match="no store is writable"): + assert_writable_topology([hand_built_store()], True) + assert_writable_topology([hand_built_store()], False) + + +@pytest.mark.parametrize("value", [101, 1000]) +def test_factory_accepts_event_caps_above_python_service_assumption(value: int) -> None: + """Match source validation, which only requires a positive integer.""" + stores = create_agentcore_memory_stores(**base_input(extraction=True, max_turns_per_event=value)) + writer = next(store for store in stores if store.writable) + assert writer._sender is not None + assert writer._sender._max_turns_per_event == value + + +def test_factory_binds_actor_and_session_into_distinct_namespaces() -> None: + """Each factory call resolves its own actor/session identity.""" + namespace = "/users/{actorId}/sessions/{sessionId}/facts" + first = create_agentcore_memory_stores( + **base_input(actor_id="actor-a", session_id="session-a", namespaces=[{"namespace": namespace}]) + ) + second = create_agentcore_memory_stores( + **base_input(actor_id="actor-b", session_id="session-b", namespaces=[{"namespace": namespace}]) + ) + assert first[0]._resolved_namespace == "/users/actor-a/sessions/session-a/facts" + assert second[0]._resolved_namespace == "/users/actor-b/sessions/session-b/facts" + + +def test_factory_preserves_explicit_falsey_client(monkeypatch: pytest.MonkeyPatch) -> None: + """Use nullish client selection rather than truthiness.""" + from bedrock_agentcore.memory.integrations.strands.memorystore import factory as factory_module + + class FalseyClient: + def __bool__(self) -> bool: + return False + + def create_event(self, **_kwargs: Any) -> dict[str, Any]: + return {} + + def retrieve_memory_records(self, **_kwargs: Any) -> dict[str, Any]: + return {"memoryRecordSummaries": []} + + client = FalseyClient() + create = Mock(side_effect=AssertionError("must not construct a replacement client")) + monkeypatch.setattr(factory_module, "_create_data_plane_client", create) + stores = create_agentcore_memory_stores(**base_input(client=client)) + assert all(store._client is client for store in stores) + create.assert_not_called() diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_format.py b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_format.py new file mode 100644 index 00000000..327fa914 --- /dev/null +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_format.py @@ -0,0 +1,70 @@ +"""Tests for AgentCore event message formatting.""" + +import pytest +from strands.types.content import Message + +from bedrock_agentcore.memory.integrations.strands.memorystore._format import ( + extract_text, + is_user_or_assistant_with_text, + map_role, +) + + +def message(role: str, content: list[dict[str, object]]) -> Message: + """Build a minimally typed Strands message.""" + return {"role": role, "content": content} # type: ignore[typeddict-item] + + +@pytest.mark.parametrize(("role", "expected"), [("user", "USER"), ("assistant", "ASSISTANT")]) +def test_map_role(role: str, expected: str) -> None: + """Map the two Strands conversation roles.""" + assert map_role(message(role, [])) == expected + + +def test_extract_text_concatenates_blocks_and_ignores_non_text() -> None: + """Trim and join only non-empty text blocks.""" + actual = extract_text( + message( + "user", + [ + {"text": " hello "}, + {"toolUse": {"toolUseId": "t1", "name": "noop", "input": {}}}, + {"text": " "}, + {"text": "world"}, + ], + ) + ) + assert actual == "hello\nworld" + + +def test_extract_text_uses_python_strip_semantics() -> None: + """Remove Python whitespace and retain BOM content.""" + actual = extract_text( + message( + "user", + [ + {"text": " \u0085 "}, + {"text": "\ufeff"}, + ], + ) + ) + assert actual == "\ufeff" + + +def test_extract_text_returns_empty_for_tool_only_message() -> None: + """Tool-only messages have no AgentCore conversational text.""" + assert extract_text(message("assistant", [{"toolUse": {}}])) == "" + + +@pytest.mark.parametrize( + ("role", "content", "expected"), + [ + ("user", [{"text": "hi"}], True), + ("assistant", [{"text": "hi"}], True), + ("user", [{"toolUse": {}}], False), + ("assistant", [{"text": " "}], False), + ], +) +def test_is_user_or_assistant_with_text(role: str, content: list[dict[str, object]], expected: bool) -> None: + """Accept only supported roles with non-blank text.""" + assert is_user_or_assistant_with_text(message(role, content)) is expected diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_package_exports.py b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_package_exports.py new file mode 100644 index 00000000..dee9aff0 --- /dev/null +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_package_exports.py @@ -0,0 +1,35 @@ +"""Package export tests for the contained MemoryStore integration.""" + +import bedrock_agentcore.memory.integrations.strands as strands_integration +import bedrock_agentcore.memory.integrations.strands.memorystore as memorystore + + +def test_memorystore_package_explicitly_exports_public_surface() -> None: + """Expose MemoryStore APIs from their contained canonical package.""" + assert set(memorystore.__all__) == { + "RESERVED_METADATA_PREFIX", + "AgentCoreEventSender", + "AgentCoreEventSenderConfig", + "AgentCoreExtractionConfig", + "AgentCoreExactNamespaceStoreConfig", + "AgentCoreMemoryStore", + "AgentCoreMemoryStoreConfig", + "AgentCoreNamespaceConfig", + "AgentCoreSubtreeStoreConfig", + "CreateAgentCoreMemoryStoresInput", + "ExtractionMode", + "MetadataProvider", + "MetadataValue", + "assert_writable_topology", + "create_agentcore_memory_stores", + "resolve_namespace", + "slugify_namespace", + } + + +def test_strands_root_preserves_converter_exports_only() -> None: + """Do not add MemoryStore APIs to the existing Strands package root.""" + assert strands_integration.__all__ == ["MemoryConverter", "OpenAIConverseConverter"] + assert strands_integration.MemoryConverter is not None + assert strands_integration.OpenAIConverseConverter is not None + assert not hasattr(strands_integration, "AgentCoreMemoryStore") diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_sender.py b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_sender.py new file mode 100644 index 00000000..f6b7d3f3 --- /dev/null +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_sender.py @@ -0,0 +1,471 @@ +"""Tests for the AgentCore event sender.""" + +import asyncio +import re +import threading +from collections.abc import Callable +from typing import Any +from unittest.mock import Mock + +import pytest +from strands.memory import AggregateMemoryError +from strands.types.content import Message + +from bedrock_agentcore.memory.integrations.strands.memorystore.sender import AgentCoreEventSender +from bedrock_agentcore.memory.integrations.strands.memorystore.types import MetadataProvider, MetadataValue + + +def user_message(text: str) -> Message: + """Build a user text message.""" + return {"role": "user", "content": [{"text": text}]} + + +def assistant_message(text: str) -> Message: + """Build an assistant text message.""" + return {"role": "assistant", "content": [{"text": text}]} + + +TOOL_ONLY: Message = { + "role": "user", + "content": [{"toolUse": {"toolUseId": "t1", "name": "noop", "input": {}}}], +} + + +def make_sender( + client: Mock, + *, + max_turns_per_event: int = 50, + run_id: str | None = "run-1", + metadata_provider: MetadataProvider | None = None, + extraction_mode: str | None = None, +) -> AgentCoreEventSender: + """Build a sender with deterministic identity.""" + return AgentCoreEventSender( + client=client, + memory_id="mem-1", + actor_id="actor-1", + session_id="sess-1", + run_id=run_id, + max_turns_per_event=max_turns_per_event, + metadata_provider=metadata_provider, + extraction_mode=extraction_mode, # type: ignore[arg-type] + ) + + +def turns(call: Any) -> list[dict[str, str]]: + """Extract role/text pairs from one mock call.""" + return [ + { + "role": item["conversational"]["role"], + "text": item["conversational"]["content"]["text"], + } + for item in call.kwargs["payload"] + ] + + +async def test_packs_batch_into_one_role_tagged_event() -> None: + """A whole flush becomes one event when under the cap.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client).send_batch([user_message("hello"), assistant_message("hi there"), user_message("again")]) + client.create_event.assert_called_once() + assert client.create_event.call_args.kwargs["memoryId"] == "mem-1" + assert turns(client.create_event.call_args) == [ + {"role": "USER", "text": "hello"}, + {"role": "ASSISTANT", "text": "hi there"}, + {"role": "USER", "text": "again"}, + ] + + +async def test_chunks_batch_at_max_turns() -> None: + """Split a batch into ceil(n / cap) events.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client, max_turns_per_event=2).send_batch([user_message(value) for value in "abcde"]) + assert client.create_event.call_count == 3 + assert [[turn["text"] for turn in turns(call)] for call in client.create_event.call_args_list] == [ + ["a", "b"], + ["c", "d"], + ["e"], + ] + + +async def test_skips_tool_only_empty_and_all_unsendable_batches() -> None: + """Omit messages without extractable user/assistant text.""" + client = Mock() + client.create_event.return_value = {} + sender = make_sender(client) + await sender.send_batch([TOOL_ONLY, user_message("real"), assistant_message(" ")]) + assert turns(client.create_event.call_args) == [{"role": "USER", "text": "real"}] + client.reset_mock() + await sender.send_batch([TOOL_ONLY]) + client.create_event.assert_not_called() + + +async def test_message_text_uses_python_strip_semantics_on_wire() -> None: + """Drop Python whitespace-only text and retain a BOM as content.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client).send_batch([user_message(" \u0085 "), user_message("\ufeff")]) + client.create_event.assert_called_once() + assert turns(client.create_event.call_args) == [{"role": "USER", "text": "\ufeff"}] + + +async def test_omits_client_token_without_complete_sequence_numbers() -> None: + """A token is unsafe unless every covered message has a sequence number.""" + client = Mock() + client.create_event.return_value = {} + sender = make_sender(client) + await sender.send_batch([user_message("x"), user_message("y")]) + assert "clientToken" not in client.create_event.call_args.kwargs + client.reset_mock() + await sender.send_batch([user_message("x"), user_message("y")], [7]) + assert "clientToken" not in client.create_event.call_args.kwargs + + +async def test_sequence_range_token_is_stable_and_chunk_specific() -> None: + """Re-fires reuse a run-scoped deterministic range token.""" + client = Mock() + client.create_event.return_value = {} + sender = make_sender(client, max_turns_per_event=2) + batch = [user_message("a"), user_message("b"), user_message("c")] + await sender.send_batch(batch, [1, 2, 3]) + assert [call.kwargs["clientToken"] for call in client.create_event.call_args_list] == [ + "mem-1-actor-1-run-1-1-2", + "mem-1-actor-1-run-1-3-3", + ] + client.reset_mock() + await sender.send_batch(batch, [1, 2, 3]) + assert client.create_event.call_args_list[0].kwargs["clientToken"] == "mem-1-actor-1-run-1-1-2" + + +async def test_explicit_empty_run_id_is_preserved() -> None: + """Use nullish rather than truthy defaulting, matching the source runtime.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client, run_id="").send_batch([user_message("x")], [0]) + assert client.create_event.call_args.kwargs["clientToken"] == "mem-1-actor-1--0-0" + + +async def test_run_id_distinguishes_sequence_resets_and_defaults_to_uuid() -> None: + """Two runs cannot collide when sequence numbers restart at zero.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client, run_id="run-A").send_batch([user_message("x")], [0]) + await make_sender(client, run_id="run-B").send_batch([user_message("x")], [0]) + assert [call.kwargs["clientToken"] for call in client.create_event.call_args_list] == [ + "mem-1-actor-1-run-A-0-0", + "mem-1-actor-1-run-B-0-0", + ] + client.reset_mock() + await make_sender(client, run_id=None).send_batch([user_message("x")], [0]) + token = client.create_event.call_args.kwargs["clientToken"] + assert re.fullmatch(r"mem-1-actor-1-[0-9a-f-]{36}-0-0", token) + assert "sess-1" not in token + + +async def test_metadata_changes_split_only_consecutive_runs() -> None: + """Metadata is per-event, so A,A,B,C,B forms four events.""" + client = Mock() + client.create_event.return_value = {} + + def provider(message: Message) -> dict[str, MetadataValue]: + return {"topic": message["content"][0]["text"]} + + await make_sender(client, metadata_provider=provider).send_batch( + [user_message(value) for value in ["A", "A", "B", "C", "B"]] + ) + assert client.create_event.call_count == 4 + assert [[turn["text"] for turn in turns(call)] for call in client.create_event.call_args_list] == [ + ["A", "A"], + ["B"], + ["C"], + ["B"], + ] + assert client.create_event.call_args_list[0].kwargs["metadata"] == {"topic": {"stringValue": "A"}} + + +async def test_constant_metadata_is_mapped_and_empty_bag_omitted() -> None: + """Map scalar metadata to the boto3 wire shape.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client, metadata_provider=lambda _message: {"source": "support", "priority": 3}).send_batch( + [user_message("x"), user_message("y")] + ) + assert client.create_event.call_count == 1 + assert client.create_event.call_args.kwargs["metadata"] == { + "source": {"stringValue": "support"}, + "priority": {"stringValue": "3"}, + } + client.reset_mock() + await make_sender(client, metadata_provider=lambda _message: {}).send_batch([user_message("x")]) + assert "metadata" not in client.create_event.call_args.kwargs + + +@pytest.mark.parametrize( + "metadata", + [ + {"note": "billing,refund"}, + {"q": "why?"}, + ], +) +async def test_rejects_disallowed_metadata_before_network( + metadata: dict[str, Any], +) -> None: + """Surface AgentCore's metadata charset restriction locally.""" + client = Mock() + with pytest.raises(ValueError, match="characters AgentCore rejects"): + await make_sender(client, metadata_provider=lambda _message: metadata).send_batch([user_message("x")]) + client.create_event.assert_not_called() + + +@pytest.mark.parametrize( + "value", + [ + "tab\tline\nvertical\vform\ffeed\rspace ", + "\u00a0\u1680\u2000\u2001\u2002\u2003\u2004\u2005\u2006\u2007\u2008\u2009\u200a", + "\u2028\u2029\u202f\u205f\u3000", + ], +) +async def test_accepts_service_whitespace_metadata(value: str) -> None: + """Accept whitespace represented by Python regular-expression semantics.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client, metadata_provider=lambda _message: {"space": value}).send_batch([user_message("x")]) + assert client.create_event.call_args.kwargs["metadata"] == {"space": {"stringValue": value}} + + +async def test_accepts_python_next_line_whitespace_metadata() -> None: + r"""Python ``\s`` includes U+0085.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client, metadata_provider=lambda _message: {"space": "\u0085"}).send_batch([user_message("x")]) + assert client.create_event.call_args.kwargs["metadata"] == {"space": {"stringValue": "\u0085"}} + + +@pytest.mark.parametrize("value", [None, float("nan"), float("inf")]) +async def test_rejects_nullish_or_non_finite_metadata_before_network(value: object) -> None: + """Reject values that have no safe scalar representation.""" + client = Mock() + + def provider(_message: Message) -> dict[str, MetadataValue]: + return {"bad": value} # type: ignore[dict-item] + + with pytest.raises(ValueError, match="no valid string representation"): + await make_sender(client, metadata_provider=provider).send_batch([user_message("x")]) + client.create_event.assert_not_called() + + +async def test_aggregates_failures_after_attempting_every_event() -> None: + """Use all-settled behavior and preserve every failed event reason.""" + attempted: list[str] = [] + + def create_event(**kwargs: Any) -> dict[str, Any]: + text = kwargs["payload"][0]["conversational"]["content"]["text"] + attempted.append(text) + if text.startswith("bad"): + raise RuntimeError(f"nope: {text}") + return {} + + client = Mock() + client.create_event.side_effect = create_event + with pytest.raises(AggregateMemoryError, match="2 of 3.*first error: nope:") as raised: + await make_sender(client, max_turns_per_event=1).send_batch( + [user_message("good-1"), user_message("bad-1"), user_message("bad-2")] + ) + assert len(raised.value.errors) == 2 + assert sorted(attempted) == ["bad-1", "bad-2", "good-1"] + assert client.create_event.call_count == 3 + + +async def test_sender_has_no_retry_layer() -> None: + """One event failure produces one network attempt.""" + client = Mock() + client.create_event.side_effect = RuntimeError("throttled by AgentCore") + with pytest.raises(AggregateMemoryError, match="first error: throttled by AgentCore"): + await make_sender(client).send_batch([user_message("x")]) + client.create_event.assert_called_once() + + +@pytest.mark.parametrize("value", [0, -1, 2.5, True]) +def test_rejects_invalid_max_turns(value: object) -> None: + """The event cap must be a positive integer.""" + with pytest.raises(ValueError, match="positive integer"): + AgentCoreEventSender( + client=Mock(), + memory_id="m", + actor_id="a", + session_id="s", + max_turns_per_event=value, # type: ignore[arg-type] + ) + + +async def test_extraction_mode_is_optional_wire_passthrough() -> None: + """Send SKIP exactly when configured.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client, extraction_mode="SKIP").send_batch([user_message("sensitive")]) + assert client.create_event.call_args.kwargs["extractionMode"] == "SKIP" + client.reset_mock() + await make_sender(client).send_batch([user_message("normal")]) + assert "extractionMode" not in client.create_event.call_args.kwargs + + +async def test_blocking_boto_call_runs_off_event_loop() -> None: + """The synchronous boto3 call is delegated through ``asyncio.to_thread``.""" + client = Mock() + client.create_event.return_value = {} + loop_thread_seen: list[bool] = [] + original = asyncio.to_thread + + async def tracked(function: Callable[..., object], /, *args: object, **kwargs: object) -> object: + loop_thread_seen.append(True) + return await original(function, *args, **kwargs) + + with pytest.MonkeyPatch.context() as patch: + patch.setattr(asyncio, "to_thread", tracked) + await make_sender(client).send_batch([user_message("x")]) + assert loop_thread_seen == [True] + + +@pytest.mark.parametrize("value", [101, 1000]) +def test_accepts_event_caps_above_python_service_assumption(value: int) -> None: + """Match the source, which only requires a positive integer.""" + assert make_sender(Mock(), max_turns_per_event=value)._max_turns_per_event == value + + +async def test_oversized_turn_is_forwarded_without_python_only_validation() -> None: + """Leave service payload validation to AgentCore, matching the source.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client).send_batch([user_message("x" * 100_001)]) + assert len(client.create_event.call_args.kwargs["payload"][0]["conversational"]["content"]["text"]) == 100_001 + + +async def test_metadata_does_not_enforce_python_only_key_count_or_length_limits() -> None: + """Only metadata values receive the source's local validation.""" + client = Mock() + client.create_event.return_value = {} + metadata = {f"key-{index}": "v" for index in range(16)} + metadata["k" * 129] = "v" * 257 + await make_sender(client, metadata_provider=lambda _message: metadata).send_batch([user_message("x")]) + wire = client.create_event.call_args.kwargs["metadata"] + assert len(wire) == 17 + assert wire["k" * 129] == {"stringValue": "v" * 257} + + +@pytest.mark.parametrize("value", [None, ["a"], {"nested": "value"}]) +async def test_rejects_dynamic_non_scalar_metadata(value: object) -> None: + """Dynamically supplied null, arrays, and objects fail with a scalar-only message.""" + client = Mock() + + def provider(_message: Message) -> dict[str, MetadataValue]: + return {"bad": value} # type: ignore[dict-item] + + with pytest.raises(ValueError, match="scalar|valid string representation"): + await make_sender(client, metadata_provider=provider).send_batch([user_message("x")]) + client.create_event.assert_not_called() + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + (3, "3"), + (3.0, "3.0"), + (-0.0, "-0.0"), + (True, "true"), + (1e-7, "1e-07"), + (1e20, "1e+20"), + ], +) +async def test_scalar_metadata_uses_python_json_semantics(value: MetadataValue, expected: str) -> None: + """Pass strings through and JSON-encode other finite scalars.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client, metadata_provider=lambda _message: {"value": value}).send_batch([user_message("x")]) + assert client.create_event.call_args.kwargs["metadata"] == {"value": {"stringValue": expected}} + + +async def test_raw_string_and_number_metadata_create_separate_groups() -> None: + """Raw JSON signatures distinguish string and numeric values before wire mapping.""" + client = Mock() + client.create_event.return_value = {} + + def provider(message: Message) -> dict[str, MetadataValue]: + text = message["content"][0]["text"] + return {"v": "3" if text == "string" else 3} + + await make_sender(client, metadata_provider=provider).send_batch([user_message("string"), user_message("number")]) + assert client.create_event.call_count == 2 + assert [call.kwargs["metadata"] for call in client.create_event.call_args_list] == [ + {"v": {"stringValue": "3"}}, + {"v": {"stringValue": "3"}}, + ] + + +async def test_metadata_signature_sorts_keys() -> None: + """Equivalent metadata bags share an event regardless of insertion order.""" + client = Mock() + client.create_event.return_value = {} + + def provider(message: Message) -> dict[str, MetadataValue]: + if message["content"][0]["text"] == "first": + return {"z": "last", "a": "first"} + return {"a": "first", "z": "last"} + + await make_sender(client, metadata_provider=provider).send_batch([user_message("first"), user_message("second")]) + client.create_event.assert_called_once() + assert [turn["text"] for turn in turns(client.create_event.call_args)] == ["first", "second"] + + +async def test_empty_metadata_bag_has_signature_but_no_wire_metadata() -> None: + """Empty provider results remain groupable while wire metadata stays omitted.""" + client = Mock() + client.create_event.return_value = {} + await make_sender(client, metadata_provider=lambda _message: {}).send_batch([user_message("x"), user_message("y")]) + client.create_event.assert_called_once() + assert "metadata" not in client.create_event.call_args.kwargs + + +async def test_cancellation_waits_for_delayed_failure_and_raises_aggregate_error() -> None: + """A cancelled caller cannot detach a failed boto3 write from coordinator rollback.""" + started = threading.Event() + release = threading.Event() + + def create_event(**_kwargs: Any) -> dict[str, Any]: + started.set() + assert release.wait(timeout=2) + raise RuntimeError("delayed failure") + + client = Mock() + client.create_event.side_effect = create_event + task = asyncio.create_task(make_sender(client).send_batch([user_message("x")])) + await asyncio.to_thread(started.wait, 2) + task.cancel() + await asyncio.sleep(0) + assert not task.done() + release.set() + with pytest.raises(AggregateMemoryError, match="delayed failure"): + await task + + +async def test_cancellation_propagates_after_successful_write_settles() -> None: + """Preserve cancellation when all shielded writes eventually succeed.""" + started = threading.Event() + release = threading.Event() + + def create_event(**_kwargs: Any) -> dict[str, Any]: + started.set() + assert release.wait(timeout=2) + return {} + + client = Mock() + client.create_event.side_effect = create_event + task = asyncio.create_task(make_sender(client).send_batch([user_message("x")])) + await asyncio.to_thread(started.wait, 2) + task.cancel() + await asyncio.sleep(0) + assert not task.done() + release.set() + with pytest.raises(asyncio.CancelledError): + await task diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_static_typing.py b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_static_typing.py new file mode 100644 index 00000000..339cedd4 --- /dev/null +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_static_typing.py @@ -0,0 +1,31 @@ +"""Mypy regression tests for the public native-memory integration types.""" + +import subprocess +import sys +from pathlib import Path + +FIXTURES = Path(__file__).with_name("typing_fixtures") + + +def _run_mypy(fixture: str) -> subprocess.CompletedProcess[str]: + return subprocess.run( + [sys.executable, "-m", "mypy", "--no-incremental", "--show-error-codes", str(FIXTURES / fixture)], + check=False, + capture_output=True, + text=True, + ) + + +def test_memorystore_package_types_compose_with_memory_manager() -> None: + """Documented canonical imports retain precise types and factory lists remain composable.""" + result = _run_mypy("valid.py") + assert result.returncode == 0, result.stdout + result.stderr + assert "Success: no issues found" in result.stdout + + +def test_absent_optional_methods_are_not_statically_callable() -> None: + """Protocol conformance must not advertise unsupported runtime capabilities.""" + result = _run_mypy("reject_absent_methods.py") + assert result.returncode != 0 + output = result.stdout + result.stderr + assert output.count('"Never" not callable [misc]') == 3, output diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_store.py b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_store.py new file mode 100644 index 00000000..ca685482 --- /dev/null +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_store.py @@ -0,0 +1,500 @@ +"""Tests for the native Strands AgentCore memory store.""" + +from datetime import datetime, timezone +from typing import Any +from unittest.mock import Mock + +import pytest +from strands.memory import AddMessagesContext + +from bedrock_agentcore.memory.integrations.strands.memorystore.store import AgentCoreMemoryStore +from bedrock_agentcore.memory.integrations.strands.memorystore.types import RESERVED_METADATA_PREFIX + +from .test_sender import assistant_message, user_message + + +def record( + record_id: str, + text: str, + score: float | None = None, + namespaces: list[str] | None = None, +) -> dict[str, Any]: + """Build a data-plane memory record summary.""" + result: dict[str, Any] = { + "memoryRecordId": record_id, + "content": {"text": text}, + "namespaces": namespaces or ["/ns/a"], + "memoryStrategyId": "strategy", + "createdAt": datetime.now(timezone.utc), + } + if score is not None: + result["score"] = score + return result + + +def client_returning(records: list[dict[str, Any]] | None) -> Mock: + """Build a mock boto3 client returning record summaries.""" + client = Mock() + client.retrieve_memory_records.return_value = {"memoryRecordSummaries": records} + client.create_event.return_value = {} + return client + + +def make_store(client: Mock, **overrides: Any) -> AgentCoreMemoryStore: + """Build an exact-mode store with test identity.""" + config: dict[str, Any] = { + "memory_id": "mem-1", + "actor_id": "actor-1", + "session_id": "sess-1", + "namespace": "/strategy/s/actor/{actorId}/preferences", + "name": "prefs", + "writable": False, + "client": client, + } + config.update(overrides) + return AgentCoreMemoryStore(**config) + + +async def test_exact_namespace_resolves_actor_and_uses_namespace() -> None: + """Exact mode emits the exact-prefix wire field.""" + client = client_returning([record("1", "a")]) + await make_store(client).search("q") + kwargs = client.retrieve_memory_records.call_args.kwargs + assert kwargs["namespace"] == "/strategy/s/actor/actor-1/preferences" + assert "namespacePath" not in kwargs + + +async def test_subtree_mode_uses_namespace_path() -> None: + """Subtree mode emits ``namespacePath`` instead of ``namespace``.""" + client = client_returning([record("1", "a")]) + store = make_store( + client, + namespace=None, + namespace_path="/strategy/s/actor/{actorId}", + ) + await store.search("q") + kwargs = client.retrieve_memory_records.call_args.kwargs + assert kwargs["namespacePath"] == "/strategy/s/actor/actor-1" + assert "namespace" not in kwargs + + +async def test_maps_memory_record_summary_to_entry() -> None: + """Map content and reserved metadata using Python datetimes.""" + created_at = datetime(2026, 1, 2, 3, 4, 5, tzinfo=timezone.utc) + client = client_returning( + [ + { + "memoryRecordId": "rec-9", + "content": {"text": "dark mode"}, + "score": 0.8, + "namespaces": ["/ns/x"], + "createdAt": created_at, + } + ] + ) + results = await make_store(client).search("q") + assert len(results) == 1 + assert results[0].content == "dark mode" + assert results[0].metadata == { + "_id": "rec-9", + "_score": 0.8, + "_namespaces": ["/ns/x"], + "_createdAt": "2026-01-02T03:04:05.000Z", + } + assert all(key.startswith(RESERVED_METADATA_PREFIX) for key in results[0].metadata or {}) + + +async def test_top_k_equals_want_without_score_floor() -> None: + """Do not overfetch when no client-side filter can remove results.""" + client = client_returning([]) + await make_store(client, max_search_results=3, over_fetch_factor=10).search("q") + assert client.retrieve_memory_records.call_args.kwargs["searchCriteria"]["topK"] == 3 + + +async def test_score_floor_overfetches_filters_and_trims() -> None: + """Overfetch before applying the client-side relevance floor.""" + client = client_returning( + [ + record("1", "a", 0.9), + record("2", "b", 0.1), + record("3", "c", 0.7), + record("4", "d", 0.2), + record("5", "e", 0.6), + ] + ) + results = await make_store(client, max_search_results=2, min_score=0.5).search("q") + assert client.retrieve_memory_records.call_args.kwargs["searchCriteria"]["topK"] == 8 + assert [result.content for result in results] == ["a", "c"] + + +@pytest.mark.parametrize( + ("want", "factor", "expected"), + [(3, 10, 30), (5, 1.5, 8), (50, 3, 100), (5, 1e308, 100)], +) +async def test_custom_overfetch_is_ceiled_and_capped(want: int, factor: float, expected: int) -> None: + """Keep AgentCore ``topK`` integral and no larger than 100.""" + client = client_returning([]) + await make_store( + client, + max_search_results=want, + min_score=0.5, + over_fetch_factor=factor, + ).search("q") + assert client.retrieve_memory_records.call_args.kwargs["searchCriteria"]["topK"] == expected + + +@pytest.mark.parametrize("value", [0, -1, 2.5, float("nan")]) +async def test_invalid_call_time_result_cap_fails_before_network(value: Any) -> None: + """Validate the effective search option, not only constructor defaults.""" + client = client_returning([]) + with pytest.raises(ValueError, match="max_search_results must be a positive integer"): + await make_store(client).search("q", {"max_search_results": value}) + client.retrieve_memory_records.assert_not_called() + + +async def test_unscored_records_are_zero_under_positive_floor() -> None: + """Treat absent scores as zero for filtering.""" + client = client_returning([record("1", "scored", 0.9), record("2", "unscored")]) + results = await make_store(client, min_score=0.5).search("q") + assert [result.content for result in results] == ["scored"] + + +async def test_non_text_content_maps_to_empty_string() -> None: + """Unknown MemoryContent union members do not become text.""" + item = record("x", "ignored", 0.9) + item["content"] = {"unknown": ["blob", {}]} + results = await make_store(client_returning([item])).search("q") + assert results[0].content == "" + + +async def test_empty_response_returns_empty_list() -> None: + """An absent summary list is an empty search result.""" + assert await make_store(client_returning(None)).search("q") == [] + + +async def test_retrieve_errors_propagate() -> None: + """Let MemoryManager isolate and report store failures.""" + client = Mock() + client.retrieve_memory_records.side_effect = RuntimeError("throttled") + with pytest.raises(RuntimeError, match="throttled"): + await make_store(client).search("q") + + +async def test_writable_store_sends_messages_and_sequence_numbers() -> None: + """Use the sender's batched deterministic-token path.""" + client = client_returning([]) + store = make_store(client, writable=True) + await store.add_messages( + [user_message("first"), assistant_message("second")], + AddMessagesContext(sequence_numbers=[41, 42]), + ) + client.create_event.assert_called_once() + kwargs = client.create_event.call_args.kwargs + assert len(kwargs["payload"]) == 2 + assert kwargs["payload"][0]["conversational"]["role"] == "USER" + assert kwargs["clientToken"].endswith("-41-42") + + +def test_unsupported_optional_methods_are_absent_at_runtime() -> None: + """Keep Strands capability detection aligned with the methods actually supported.""" + store = make_store(client_returning([])) + assert not hasattr(store, "add") + assert not hasattr(store, "initialize") + assert not hasattr(store, "get_tools") + + +async def test_non_writable_store_rejects_add_messages() -> None: + """Guard direct misuse even though MemoryManager will not call this sink.""" + with pytest.raises(ValueError, match="not writable"): + await make_store(client_returning([])).add_messages([user_message("x")]) + + +def test_only_writable_store_carries_extraction(caplog: pytest.LogCaptureFixture) -> None: + """Drop and warn about extraction configuration on a recall-only store.""" + trigger = Mock() + config = {"trigger": trigger} + writable = make_store(client_returning([]), writable=True, extraction=config) + with caplog.at_level("WARNING"): + readonly = make_store(client_returning([]), extraction=config) + assert writable.extraction == config + assert readonly.extraction is None + assert "writable is false" in caplog.text + + +@pytest.mark.parametrize("extraction", [None, False]) +def test_recall_only_store_does_not_warn(extraction: object, caplog: pytest.LogCaptureFixture) -> None: + """No extraction and explicit opt-out are valid recall-only configurations.""" + with caplog.at_level("WARNING"): + make_store(client_returning([]), extraction=extraction) + assert "writable is false" not in caplog.text + + +@pytest.mark.parametrize("name", [None, "", " "]) +def test_store_self_names_from_namespace_when_name_absent(name: str | None) -> None: + """Use a non-degenerate namespace slug.""" + store = make_store( + client_returning([]), + name=name, + namespace="/users/{actorId}/facts", + ) + assert store.name == "users-facts" + + +@pytest.mark.parametrize("field", ["memory_id", "actor_id", "session_id"]) +def test_identity_uses_python_strip_semantics(field: str) -> None: + """Reject Python whitespace-only identity and retain BOM content.""" + client = client_returning([]) + store = make_store(client, **{field: "\ufeff"}) + assert getattr(store, f"_{field}") == "\ufeff" + with pytest.raises(ValueError, match=rf"{field} must be a non-empty"): + make_store(client, **{field: "\u0085"}) + + +def test_explicit_name_uses_python_strip_semantics() -> None: + """Treat Python whitespace as absent and retain a BOM name.""" + client = client_returning([]) + assert make_store(client, name=" \u0085 ").name == "strategy-s-actor-preferences" + assert make_store(client, name="\ufeff").name == "\ufeff" + + +def test_read_target_uses_python_strip_semantics() -> None: + """Reject Python whitespace-only targets and retain BOM content.""" + client = client_returning([]) + assert make_store(client, namespace="\ufeff")._resolved_namespace == "\ufeff" + with pytest.raises(ValueError, match="namespace must be a non-empty"): + make_store(client, namespace="\u0085") + + +def test_writable_defaults_false() -> None: + """A bare store is recall-safe.""" + store = AgentCoreMemoryStore( + memory_id="mem-1", + actor_id="actor-1", + session_id="sess-1", + namespace="/users/{actorId}/facts", + client=client_returning([]), + ) + assert store.writable is False + + +async def test_direct_store_stands_alone_without_factory() -> None: + """Flat identity plus namespace is sufficient for read/write construction.""" + client = client_returning([]) + store = AgentCoreMemoryStore( + memory_id="mem-1", + actor_id="actor-1", + session_id="sess-1", + namespace="/users/{actorId}/facts", + writable=True, + extraction=True, + client=client, + ) + assert store.name == "users-facts" + assert store.extraction is True + await store.add_messages([user_message("hi")]) + assert client.create_event.call_args.kwargs["actorId"] == "actor-1" + + +@pytest.mark.parametrize( + ("field", "value"), + [("memory_id", ""), ("actor_id", " "), ("session_id", "")], +) +def test_rejects_empty_identity(field: str, value: str) -> None: + """Validate each flat identity field.""" + kwargs: dict[str, Any] = { + "memory_id": "mem-1", + "actor_id": "actor-1", + "session_id": "sess-1", + field: value, + } + with pytest.raises(ValueError, match=rf"{field} must be a non-empty"): + AgentCoreMemoryStore( + **kwargs, + namespace="/users/{actorId}/facts", + client=client_returning([]), + ) + + +def test_rejects_empty_or_ambiguous_read_target() -> None: + """Require exactly one non-empty read target.""" + client = client_returning([]) + with pytest.raises(ValueError, match="namespace must be a non-empty"): + make_store(client, namespace=" ") + with pytest.raises(ValueError, match="exactly one"): + make_store(client, namespace=None) + with pytest.raises(ValueError, match="exactly one"): + make_store(client, namespace="/a", namespace_path="/b") + + +@pytest.mark.parametrize( + "target", + [ + {"namespace": "/strategies/{memoryStrategyId}/actors/{actorId}/facts"}, + {"namespace": "/users/{actorId}/we{ird"}, + {"namespace": "/users/{actorId}/weird}"}, + {"namespace": "/a/{strategy/b"}, + {"namespace": None, "namespace_path": "/strategies/{memoryStrategyId}/actors/{actorId}"}, + ], +) +def test_rejects_unresolved_or_malformed_placeholders(target: dict[str, object]) -> None: + """Fail before AgentCore rejects braces at first retrieval.""" + with pytest.raises(ValueError, match="still contains"): + make_store(client_returning([]), **target) + + +def test_actor_dollar_sequences_are_inserted_verbatim() -> None: + """Python substitution does not interpret JavaScript-style replacement syntax.""" + store = make_store( + client_returning([]), + actor_id="a$$b", + namespace="/p/{actorId}/x", + ) + assert store.name == "prefs" + + +@pytest.mark.parametrize("value", [float("nan"), -0.1, 1.5, float("inf")]) +def test_rejects_invalid_min_score(value: float) -> None: + """The relevance floor must be finite and normalized.""" + with pytest.raises(ValueError, match="finite number between 0 and 1"): + make_store(client_returning([]), min_score=value) + + +@pytest.mark.parametrize("value", [0, -1, 2.5]) +def test_rejects_invalid_constructor_result_cap(value: Any) -> None: + """The store-level result cap must be a positive integer.""" + with pytest.raises(ValueError, match="positive integer"): + make_store(client_returning([]), max_search_results=value) + + +@pytest.mark.parametrize("value", [0, 0.5, float("nan"), float("inf")]) +def test_rejects_invalid_overfetch_factor(value: float) -> None: + """Overfetch factors must be finite and at least one.""" + with pytest.raises(ValueError, match="number >= 1"): + make_store(client_returning([]), over_fetch_factor=value) + + +def test_direct_store_rejects_invalid_event_cap() -> None: + """Writable direct construction delegates cap validation to its sender.""" + with pytest.raises(ValueError, match="positive integer"): + make_store(client_returning([]), writable=True, max_turns_per_event=0) + + +async def test_ordered_substitution_rescans_actor_value_during_session_pass() -> None: + """Match source runtime chaining when actor substitution introduces a session token.""" + client = client_returning([]) + store = make_store( + client, + actor_id="literal-{sessionId}", + session_id="runtime-parity", + namespace="/p/{actorId}/x", + ) + await store.search("q") + assert client.retrieve_memory_records.call_args.kwargs["namespace"] == "/p/literal-runtime-parity/x" + + +async def test_search_sends_only_source_equivalent_top_k() -> None: + """Do not add the target-only top-level ``maxResults`` request field.""" + client = client_returning([]) + await make_store(client, max_search_results=10, min_score=0.5).search("q") + kwargs = client.retrieve_memory_records.call_args.kwargs + assert kwargs["searchCriteria"]["topK"] == 40 + assert "maxResults" not in kwargs + + +async def test_result_cap_above_100_is_allowed_without_score_floor() -> None: + """Match source positive-integer validation; only overfetch topK is capped.""" + client = client_returning([]) + await make_store(client, max_search_results=101).search("q") + assert client.retrieve_memory_records.call_args.kwargs["searchCriteria"]["topK"] == 101 + + +async def test_result_cap_above_100_caps_only_overfetch_top_k() -> None: + """A score-floor request keeps the source's MAX_TOPK overfetch cap.""" + client = client_returning([]) + await make_store(client, max_search_results=101, min_score=0.5).search("q") + assert client.retrieve_memory_records.call_args.kwargs["searchCriteria"]["topK"] == 100 + + +def test_client_region_prefers_explicit_session_without_loading_invalid_ambient_profile( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """An explicit session bypasses an invalid ambient profile and supplies its region.""" + import boto3 + + from bedrock_agentcore.memory.integrations.strands.memorystore import store as store_module + + supplied_session = boto3.Session( + aws_access_key_id="test", + aws_secret_access_key="test", + region_name="session-region", + ) + create_client = Mock(return_value=Mock()) + monkeypatch.setattr(supplied_session, "client", create_client) + monkeypatch.setenv("AWS_PROFILE", "profile-that-does-not-exist") + monkeypatch.setenv("AWS_REGION", "environment-region") + + store_module._create_data_plane_client(boto3_session=supplied_session) + + create_client.assert_called_once() + assert create_client.call_args.kwargs["region_name"] == "session-region" + + +def test_explicit_region_overrides_explicit_session_region(monkeypatch: pytest.MonkeyPatch) -> None: + """Use the caller's explicit region before the selected session region.""" + from bedrock_agentcore.memory.integrations.strands.memorystore import store as store_module + + supplied_session = Mock(region_name="session-region") + supplied_session.client.return_value = Mock() + monkeypatch.setenv("AWS_REGION", "environment-region") + store_module._create_data_plane_client(region_name="explicit-region", boto3_session=supplied_session) + assert supplied_session.client.call_args.kwargs["region_name"] == "explicit-region" + + +def test_client_region_falls_back_through_environment_default_and_us_west( + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Resolve the selected default session, environment, then SDK fallback region.""" + from bedrock_agentcore.memory.integrations.strands.memorystore import store as store_module + + default_session = Mock(region_name="default-region") + default_session.client.return_value = Mock() + session_factory = Mock(return_value=default_session) + monkeypatch.setattr(store_module.boto3, "Session", session_factory) + monkeypatch.setenv("AWS_REGION", "environment-region") + store_module._create_data_plane_client() + assert default_session.client.call_args.kwargs["region_name"] == "default-region" + + default_session.client.reset_mock() + default_session.region_name = None + store_module._create_data_plane_client() + assert default_session.client.call_args.kwargs["region_name"] == "environment-region" + + default_session.client.reset_mock() + monkeypatch.delenv("AWS_REGION") + store_module._create_data_plane_client() + assert default_session.client.call_args.kwargs["region_name"] == "us-west-2" + + +class _FalseyClient: + """Client whose truth value is false but whose methods remain usable.""" + + def __bool__(self) -> bool: + return False + + def retrieve_memory_records(self, **_kwargs: Any) -> dict[str, Any]: + return {"memoryRecordSummaries": []} + + def create_event(self, **_kwargs: Any) -> dict[str, Any]: + return {} + + +def test_store_preserves_explicit_falsey_client(monkeypatch: pytest.MonkeyPatch) -> None: + """Use nullish client selection rather than truthiness.""" + from bedrock_agentcore.memory.integrations.strands.memorystore import store as store_module + + falsey = _FalseyClient() + create = Mock(side_effect=AssertionError("must not construct a replacement client")) + monkeypatch.setattr(store_module, "_create_data_plane_client", create) + store = make_store(falsey) # type: ignore[arg-type] + assert store._client is falsey + create.assert_not_called() diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_types.py b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_types.py new file mode 100644 index 00000000..14fe889a --- /dev/null +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_types.py @@ -0,0 +1,71 @@ +"""Runtime type metadata tests for the public Strands memory-store types.""" + +from typing import get_args, get_type_hints + +import boto3 + +from bedrock_agentcore.memory.integrations.strands.memorystore.types import ( + AgentCoreEventSenderConfig, + AgentCoreExactNamespaceStoreConfig, + AgentCoreExtractionConfig, + AgentCoreNamespaceConfig, + AgentCoreSubtreeStoreConfig, + CreateAgentCoreMemoryStoresInput, + _AgentCoreMemoryConnectionConfig, +) + + +def test_connection_config_runtime_typed_dict_keys_and_hints() -> None: + """Required/optional keys and boto3 hints remain introspectable on Python 3.10+.""" + assert _AgentCoreMemoryConnectionConfig.__required_keys__ == frozenset({"memory_id", "actor_id", "session_id"}) + assert _AgentCoreMemoryConnectionConfig.__optional_keys__ == frozenset( + { + "metadata_provider", + "max_turns_per_event", + "extraction_mode", + "region_name", + "boto3_session", + "client", + } + ) + assert get_args(get_type_hints(_AgentCoreMemoryConnectionConfig)["boto3_session"]) == (boto3.Session,) + + +def test_public_factory_typed_dict_keys_are_correct_at_runtime() -> None: + """Public factory configuration exposes accurate required/optional metadata.""" + assert CreateAgentCoreMemoryStoresInput.__required_keys__ == frozenset( + {"memory_id", "actor_id", "session_id", "namespaces"} + ) + assert CreateAgentCoreMemoryStoresInput.__optional_keys__ == frozenset( + {"extraction", "metadata_provider", "max_turns_per_event", "region_name", "boto3_session", "client"} + ) + hints = get_type_hints(CreateAgentCoreMemoryStoresInput) + assert get_args(hints["boto3_session"]) == (boto3.Session,) + assert "extraction_mode" not in hints + + +def test_other_public_typed_dict_runtime_metadata() -> None: + """NotRequired fields are optional without postponed annotations.""" + assert AgentCoreEventSenderConfig.__required_keys__ == frozenset({"client", "memory_id", "actor_id", "session_id"}) + assert AgentCoreEventSenderConfig.__optional_keys__ == frozenset( + {"metadata_provider", "run_id", "max_turns_per_event", "extraction_mode"} + ) + assert AgentCoreNamespaceConfig.__required_keys__ == frozenset({"namespace"}) + assert AgentCoreNamespaceConfig.__optional_keys__ == frozenset( + {"name", "description", "max_search_results", "min_score", "over_fetch_factor", "writable"} + ) + assert AgentCoreExtractionConfig.__required_keys__ == frozenset() + assert AgentCoreExtractionConfig.__optional_keys__ == frozenset({"cadence", "filter"}) + exact_hints = get_type_hints(AgentCoreExactNamespaceStoreConfig) + subtree_hints = get_type_hints(AgentCoreSubtreeStoreConfig) + assert exact_hints["boto3_session"] == get_type_hints(_AgentCoreMemoryConnectionConfig)["boto3_session"] + assert subtree_hints["boto3_session"] == get_type_hints(_AgentCoreMemoryConnectionConfig)["boto3_session"] + assert AgentCoreExactNamespaceStoreConfig.__required_keys__ == frozenset( + {"memory_id", "actor_id", "session_id", "namespace"} + ) + assert AgentCoreSubtreeStoreConfig.__required_keys__ == frozenset( + {"memory_id", "actor_id", "session_id", "namespace_path"} + ) + get_type_hints(AgentCoreEventSenderConfig) + get_type_hints(AgentCoreNamespaceConfig) + get_type_hints(AgentCoreExtractionConfig) diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorystore/typing_fixtures/reject_absent_methods.py b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/typing_fixtures/reject_absent_methods.py new file mode 100644 index 00000000..20f91ec8 --- /dev/null +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/typing_fixtures/reject_absent_methods.py @@ -0,0 +1,18 @@ +"""Expected-failure mypy fixture for intentionally absent optional methods.""" + +from typing import cast + +from bedrock_agentcore.memory.integrations.strands.memorystore import AgentCoreMemoryStore +from bedrock_agentcore.memory.integrations.strands.memorystore.types import AgentCoreDataPlaneClient + +client = cast(AgentCoreDataPlaneClient, object()) +store = AgentCoreMemoryStore( + memory_id="memory", + actor_id="actor", + session_id="session", + namespace="/facts/{actorId}", + client=client, +) +store.add("content") +store.initialize() +store.get_tools() diff --git a/tests/bedrock_agentcore/memory/integrations/strands/memorystore/typing_fixtures/valid.py b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/typing_fixtures/valid.py new file mode 100644 index 00000000..ce305937 --- /dev/null +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/typing_fixtures/valid.py @@ -0,0 +1,52 @@ +"""Expected-success mypy consumer fixture for package-root exports.""" + +from typing import cast + +from strands.memory import MemoryManager, MemoryStore +from typing_extensions import assert_type + +from bedrock_agentcore.memory.integrations.strands.memorystore import ( + AgentCoreEventSender, + AgentCoreEventSenderConfig, + AgentCoreMemoryStore, + AgentCoreMemoryStoreConfig, + create_agentcore_memory_stores, +) +from bedrock_agentcore.memory.integrations.strands.memorystore.types import AgentCoreDataPlaneClient + +client = cast(AgentCoreDataPlaneClient, object()) +store = AgentCoreMemoryStore( + memory_id="memory", + actor_id="actor", + session_id="session", + namespace="/facts/{actorId}", + client=client, +) +protocol_store: MemoryStore = store +MemoryManager(stores=[store]) + +stores = create_agentcore_memory_stores( + memory_id="memory", + actor_id="actor", + session_id="session", + namespaces=[{"namespace": "/facts/{actorId}"}], + client=client, +) +assert_type(stores, list[MemoryStore]) +MemoryManager(stores=stores) +assert_type(AgentCoreEventSender, type[AgentCoreEventSender]) + +sender_config: AgentCoreEventSenderConfig = { + "client": client, + "memory_id": "memory", + "actor_id": "actor", + "session_id": "session", +} +store_config: AgentCoreMemoryStoreConfig = { + "memory_id": "memory", + "actor_id": "actor", + "session_id": "session", + "namespace": "/facts/{actorId}", +} +assert protocol_store.name == store.name +assert sender_config["memory_id"] == store_config["memory_id"] diff --git a/tests_integ/memory/integrations/strands/__init__.py b/tests_integ/memory/integrations/strands/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests_integ/memory/integrations/strands/memorysessionmanager/__init__.py b/tests_integ/memory/integrations/strands/memorysessionmanager/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests_integ/memory/integrations/test_session_manager.py b/tests_integ/memory/integrations/strands/memorysessionmanager/test_session_manager.py similarity index 97% rename from tests_integ/memory/integrations/test_session_manager.py rename to tests_integ/memory/integrations/strands/memorysessionmanager/test_session_manager.py index f420d0bd..65afa891 100644 --- a/tests_integ/memory/integrations/test_session_manager.py +++ b/tests_integ/memory/integrations/strands/memorysessionmanager/test_session_manager.py @@ -1,7 +1,7 @@ """ Integration tests for AgentCore Memory Session Manager. -Run with: python -m pytest tests_integ/memory/integrations/test_session_manager.py -v +Run with: python -m pytest tests_integ/memory/integrations/strands/memorysessionmanager/test_session_manager.py -v """ import json @@ -16,9 +16,16 @@ from strands.types.session import Session, SessionAgent, SessionType from bedrock_agentcore.memory import MemoryClient -from bedrock_agentcore.memory.integrations.strands.bedrock_converter import AgentCoreMemoryConverter -from bedrock_agentcore.memory.integrations.strands.config import AgentCoreMemoryConfig, RetrievalConfig -from bedrock_agentcore.memory.integrations.strands.session_manager import AgentCoreMemorySessionManager +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.bedrock_converter import ( + AgentCoreMemoryConverter, +) +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.config import ( + AgentCoreMemoryConfig, + RetrievalConfig, +) +from bedrock_agentcore.memory.integrations.strands.memorysessionmanager.session_manager import ( + AgentCoreMemorySessionManager, +) from bedrock_agentcore.memory.models.filters import ( EventMetadataFilter, LeftExpression, diff --git a/tests_integ/memory/integrations/strands/memorystore/__init__.py b/tests_integ/memory/integrations/strands/memorystore/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests_integ/memory/integrations/strands/memorystore/test_memory_store.py b/tests_integ/memory/integrations/strands/memorystore/test_memory_store.py new file mode 100644 index 00000000..42dfd8aa --- /dev/null +++ b/tests_integ/memory/integrations/strands/memorystore/test_memory_store.py @@ -0,0 +1,280 @@ +"""Live AgentCore Memory tests for the native Strands ``MemoryStore`` integration. + +Run with:: + + AWS_PROFILE= BEDROCK_TEST_REGION=us-west-2 \ + MEMORY_PREPOPULATED_ID= \ + uv run pytest tests_integ/memory/integrations/strands/memorystore/test_memory_store.py -xvs +""" + +import asyncio +import os +import time +import uuid +from collections.abc import Awaitable, Callable +from typing import Any, cast + +import boto3 +import pytest +from strands import Agent +from strands.memory import AddMessagesContext, MemoryEntry, MemoryManager +from strands.models import BedrockModel +from strands.types.content import Message + +from bedrock_agentcore.memory.integrations.strands.memorystore import ( + AgentCoreMemoryStore, + create_agentcore_memory_stores, +) + +REGION = os.environ.get("BEDROCK_TEST_REGION", "us-west-2") +MODEL_ID = os.environ.get("STRANDS_TEST_MODEL_ID", "global.anthropic.claude-sonnet-4-6") +FACTS_NAMESPACE = "/facts/{actorId}/" +SUMMARY_NAMESPACE = "/summaries/{actorId}/{sessionId}/" + + +async def poll_for_records( + search: Callable[[], Awaitable[list[MemoryEntry]]], + timeout_seconds: int = 240, +) -> list[MemoryEntry]: + """Poll eventual long-term extraction until records appear or time expires.""" + deadline = time.monotonic() + timeout_seconds + while True: + results = await search() + if results or time.monotonic() > deadline: + return results + await asyncio.sleep(10) + + +@pytest.fixture(scope="module") +def data_plane_client() -> Any: + """Create a live boto3 AgentCore data-plane client.""" + return boto3.client("bedrock-agentcore", region_name=REGION) + + +@pytest.fixture(scope="module") +def semantic_memory() -> dict[str, str]: + """Use a pre-provisioned memory; this test must never create an untagged resource.""" + memory_id = os.environ.get("MEMORY_PREPOPULATED_ID") + if not memory_id: + pytest.skip("MEMORY_PREPOPULATED_ID is required for native Strands memory-store tests") + return {"id": memory_id} + + +@pytest.mark.integration +class TestAgentCoreMemoryStore: + """Store-level tests against the live AgentCore data plane.""" + + async def test_write_idempotency_batching_and_recall(self, semantic_memory: Any, data_plane_client: Any) -> None: + """Write one batched event, re-fire it idempotently, and recall extracted facts.""" + actor_id = f"batch-actor-{uuid.uuid4().hex}" + session_id = f"batch-session-{uuid.uuid4().hex}" + create_event_calls = 0 + + class CountingClient: + """Count create-event calls while delegating to the live client.""" + + def create_event(self, **kwargs: Any) -> dict[str, Any]: + nonlocal create_event_calls + create_event_calls += 1 + return cast(dict[str, Any], data_plane_client.create_event(**kwargs)) + + def retrieve_memory_records(self, **kwargs: Any) -> dict[str, Any]: + return cast(dict[str, Any], data_plane_client.retrieve_memory_records(**kwargs)) + + store = AgentCoreMemoryStore( + memory_id=semantic_memory["id"], + actor_id=actor_id, + session_id=session_id, + namespace=FACTS_NAMESPACE, + writable=True, + extraction=True, + client=CountingClient(), + ) + messages: list[Message] = [ + {"role": "user", "content": [{"text": "I am a pilot based in Denver and fly Cessnas."}]}, + {"role": "assistant", "content": [{"text": "Flying Cessnas out of Denver — nice."}]}, + {"role": "user", "content": [{"text": "I also play cello in my spare time."}]}, + {"role": "assistant", "content": [{"text": "A pilot and a cellist!"}]}, + ] + context = AddMessagesContext(sequence_numbers=[0, 1, 2, 3]) + await store.add_messages(messages, context) + assert create_event_calls == 1 + await store.add_messages(messages, context) + assert create_event_calls == 2 # same token; the service accepts/deduplicates the re-fire + + results = await poll_for_records(lambda: store.search("what does the user do and where")) + assert results, "No records surfaced before the AgentCore extraction timeout" + joined = " ".join(result.content.lower() for result in results) + assert any(term in joined for term in ("pilot", "denver", "cello", "cessna")) + expected_namespace = FACTS_NAMESPACE.replace("{actorId}", actor_id) + assert any(expected_namespace in (result.metadata or {}).get("_namespaces", []) for result in results), ( + f"Expected recalled _namespaces to contain {expected_namespace}" + ) + + async def test_extraction_mode_wire_passthrough_and_recall_only_guard( + self, semantic_memory: Any, data_plane_client: Any + ) -> None: + """Prove live ``SKIP`` acceptance and direct recall-only write rejection.""" + captured: dict[str, Any] = {} + + class CapturingClient: + """Capture create-event parameters while delegating live calls.""" + + def create_event(self, **kwargs: Any) -> dict[str, Any]: + captured.update(kwargs) + return cast(dict[str, Any], data_plane_client.create_event(**kwargs)) + + def retrieve_memory_records(self, **kwargs: Any) -> dict[str, Any]: + return cast(dict[str, Any], data_plane_client.retrieve_memory_records(**kwargs)) + + identity = { + "memory_id": semantic_memory["id"], + "actor_id": f"skip-actor-{uuid.uuid4().hex}", + "session_id": f"skip-session-{uuid.uuid4().hex}", + "namespace": FACTS_NAMESPACE, + "client": CapturingClient(), + } + writer = AgentCoreMemoryStore( + **identity, + writable=True, + extraction=True, + extraction_mode="SKIP", + ) + await writer.add_messages( + [{"role": "user", "content": [{"text": "Short-term only temporary data."}]}], + AddMessagesContext(sequence_numbers=[0]), + ) + assert captured["extractionMode"] == "SKIP" + + captured.clear() + default_writer = AgentCoreMemoryStore(**identity, writable=True, extraction=True) + await default_writer.add_messages( + [{"role": "user", "content": [{"text": "Normal extraction event."}]}], + AddMessagesContext(sequence_numbers=[1]), + ) + assert "extractionMode" not in captured + + readonly = AgentCoreMemoryStore(**identity) + assert readonly.writable is False + with pytest.raises(ValueError, match="not writable"): + await readonly.add_messages([{"role": "user", "content": [{"text": "x"}]}]) + + async def test_exact_and_subtree_retrieval_fields_are_accepted_live( + self, semantic_memory: Any, data_plane_client: Any + ) -> None: + """Exercise both AgentCore retrieval target arms against the service.""" + actor_id = f"read-actor-{uuid.uuid4().hex}" + identity = { + "memory_id": semantic_memory["id"], + "actor_id": actor_id, + "session_id": f"read-session-{uuid.uuid4().hex}", + "client": data_plane_client, + } + exact = AgentCoreMemoryStore(**identity, namespace=FACTS_NAMESPACE) + subtree = AgentCoreMemoryStore(**identity, namespace_path=f"/facts/{actor_id}") + assert isinstance(await exact.search("anything"), list) + assert isinstance(await subtree.search("anything"), list) + + async def test_direct_store_and_factory_work_with_memory_manager( + self, semantic_memory: Any, data_plane_client: Any + ) -> None: + """Validate the direct primitive and factory output against real MemoryManager.""" + actor_id = f"manager-actor-{uuid.uuid4().hex}" + session_id = f"manager-session-{uuid.uuid4().hex}" + stores = create_agentcore_memory_stores( + memory_id=semantic_memory["id"], + actor_id=actor_id, + session_id=session_id, + namespaces=[{"namespace": FACTS_NAMESPACE}], + extraction=True, + client=data_plane_client, + ) + manager = MemoryManager(stores=stores) + assert len(stores) == 1 and stores[0].writable + assert isinstance(await manager.search("anything"), list) + + direct = AgentCoreMemoryStore( + memory_id=semantic_memory["id"], + actor_id=f"direct-{uuid.uuid4().hex}", + session_id=f"direct-session-{uuid.uuid4().hex}", + namespace=FACTS_NAMESPACE, + writable=True, + extraction=True, + client=data_plane_client, + ) + await direct.add_messages( + [{"role": "user", "content": [{"text": "I collect vinyl records."}]}], + AddMessagesContext(sequence_numbers=[0]), + ) + assert isinstance(await direct.search("what does the user collect"), list) + + +@pytest.mark.integration +async def test_session_scoped_namespace_drift(semantic_memory: Any, data_plane_client: Any) -> None: + """A ``{sessionId}`` namespace does not leak records across sessions.""" + actor_id = f"drift-actor-{uuid.uuid4().hex}" + session_a = f"drift-session-a-{uuid.uuid4().hex}" + session_b = f"drift-session-b-{uuid.uuid4().hex}" + writer = AgentCoreMemoryStore( + memory_id=semantic_memory["id"], + actor_id=actor_id, + session_id=session_a, + namespace=SUMMARY_NAMESPACE, + writable=True, + extraction=True, + client=data_plane_client, + ) + await writer.add_messages( + [ + {"role": "user", "content": [{"text": "We are planning a spring trip to Japan."}]}, + {"role": "assistant", "content": [{"text": "Cherry blossom season is lovely."}]}, + {"role": "user", "content": [{"text": "Book a Kyoto ryokan for two nights."}]}, + ], + AddMessagesContext(sequence_numbers=[0, 1, 2]), + ) + store_a = AgentCoreMemoryStore( + memory_id=semantic_memory["id"], + actor_id=actor_id, + session_id=session_a, + namespace=SUMMARY_NAMESPACE, + client=data_plane_client, + ) + store_b = AgentCoreMemoryStore( + memory_id=semantic_memory["id"], + actor_id=actor_id, + session_id=session_b, + namespace=SUMMARY_NAMESPACE, + client=data_plane_client, + ) + from_a = await poll_for_records(lambda: store_a.search("What trip is planned?")) + assert from_a, "No session-A summary surfaced before the AgentCore extraction timeout" + assert await store_b.search("What trip is planned?") == [] + + +@pytest.mark.integration +async def test_real_agent_memory_manager_round_trip(semantic_memory: Any, data_plane_client: Any) -> None: + """Drive extraction through a real Strands agent and poll manager recall.""" + actor_id = f"e2e-actor-{uuid.uuid4().hex}" + stores = create_agentcore_memory_stores( + memory_id=semantic_memory["id"], + actor_id=actor_id, + session_id=f"e2e-session-{uuid.uuid4().hex}", + namespaces=[{"namespace": FACTS_NAMESPACE}], + extraction=True, + client=data_plane_client, + ) + manager = MemoryManager(stores=stores) + agent = Agent( + model=BedrockModel(region_name=REGION, model_id=MODEL_ID), + system_prompt="Use long-term memory to personalize answers.", + memory_manager=manager, + ) + await agent.invoke_async("Remember this: my dog Pixel is a corgi.") + await manager.flush() + results = await poll_for_records(lambda: manager.search("What is the user's dog named?")) + assert results, "No E2E records surfaced before the AgentCore extraction timeout" + assert "pixel" in " ".join(result.content.lower() for result in results) + expected_namespace = FACTS_NAMESPACE.replace("{actorId}", actor_id) + assert any(expected_namespace in (result.metadata or {}).get("_namespaces", []) for result in results), ( + f"Expected recalled _namespaces to contain {expected_namespace}" + ) diff --git a/uv.lock b/uv.lock index 1bd84534..7825afc3 100644 --- a/uv.lock +++ b/uv.lock @@ -443,8 +443,8 @@ dev = [ requires-dist = [ { name = "a2a-sdk", extras = ["http-server"], marker = "extra == 'a2a'", specifier = ">=0.3,<1.0" }, { name = "ag-ui-protocol", marker = "extra == 'ag-ui'", specifier = ">=0.1.10" }, - { name = "boto3", specifier = ">=1.43.31" }, - { name = "botocore", specifier = ">=1.43.31" }, + { name = "boto3", specifier = ">=1.43.35" }, + { name = "botocore", specifier = ">=1.43.35" }, { name = "httpx", marker = "extra == 'langgraph'", specifier = ">=0.27.0" }, { name = "jinja2", marker = "extra == 'simulation'", specifier = ">=3.1.0" }, { name = "langchain", marker = "extra == 'langgraph'", specifier = ">=1.0.0" }, @@ -454,7 +454,7 @@ requires-dist = [ { name = "pydantic", specifier = ">=2.0.0,<2.41.3" }, { name = "requests", marker = "extra == 'datasets'", specifier = ">=2.31.0" }, { name = "starlette", specifier = ">=0.46.2" }, - { name = "strands-agents", marker = "extra == 'strands-agents'", specifier = ">=1.20.0" }, + { name = "strands-agents", marker = "extra == 'strands-agents'", specifier = ">=1.46.0" }, { name = "strands-agents-evals", marker = "extra == 'simulation'", specifier = ">=0.1.0" }, { name = "strands-agents-evals", marker = "extra == 'strands-agents-evals'", specifier = ">=0.1.0" }, { name = "typing-extensions", specifier = ">=4.13.2,<5.0.0" }, @@ -482,7 +482,7 @@ dev = [ { name = "pytest-order", specifier = ">=1.3.0" }, { name = "pytest-rerunfailures", specifier = ">=15.0" }, { name = "ruff", specifier = ">=0.12.0" }, - { name = "strands-agents", specifier = ">=1.20.0" }, + { name = "strands-agents", specifier = ">=1.46.0" }, { name = "strands-agents-evals", specifier = ">=0.1.0" }, { name = "websockets", specifier = ">=14.1" }, { name = "wheel", specifier = ">=0.45.1" }, @@ -490,30 +490,30 @@ dev = [ [[package]] name = "boto3" -version = "1.43.36" +version = "1.43.35" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "botocore" }, { name = "jmespath" }, { name = "s3transfer" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/ff/9f/897287e955db0f50b12fd69ef45956e4fd2c7ddb48c736872f7ea2314443/boto3-1.43.36.tar.gz", hash = "sha256:587d7ee92a12e440ad12b0e7f11f3358f0c4d65b19f64726efc94aaf194aff28", size = 112690, upload-time = "2026-06-23T02:47:14.561Z" } +sdist = { url = "https://files.pythonhosted.org/packages/b7/f7/9ba34c655158a9827d0fdd8c6211a3836d9fbf29d48b82783995927b3907/boto3-1.43.35.tar.gz", hash = "sha256:392ba41e82629a77ad70de38c43f8fc0f5f453c965e381e4e87d6fcedc80d6c4", size = 112638, upload-time = "2026-06-22T20:26:39.039Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/9f/f1/274303f52483ecf199eae6f8d9b6f5951670397ee4d72c06cfd4eb644612/boto3-1.43.36-py3-none-any.whl", hash = "sha256:42942dde254673abcbc9e6e60017c88341a4f49d99d24e1f2e290fb38138c26f", size = 140031, upload-time = "2026-06-23T02:47:13.178Z" }, + { url = "https://files.pythonhosted.org/packages/6a/72/a92face6b5284d1283f2cccae224d6a10ebc5735c5d6fa93837109ebb6a1/boto3-1.43.35-py3-none-any.whl", hash = "sha256:461d5fb1ee1422d4a0a439666cebb29e0d5b4bf4308d3c2564449c77a1f4ae52", size = 140032, upload-time = "2026-06-22T20:26:36.638Z" }, ] [[package]] name = "botocore" -version = "1.43.36" +version = "1.43.35" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "jmespath" }, { name = "python-dateutil" }, { name = "urllib3" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/7c/37/da9e7f6ca73ac73afd7f0bb7f238aa5daba35c081e98d7f48a7c399599c0/botocore-1.43.36.tar.gz", hash = "sha256:4cae47d1b2d426316b85a0087d9e69e048f13bc003b5177d74639fe9dfd28205", size = 15625488, upload-time = "2026-06-23T02:47:03.192Z" } +sdist = { url = "https://files.pythonhosted.org/packages/17/d8/28177fef65f6ef8c7fdd5b44024254816961578ce8925bc336229e236cdc/botocore-1.43.35.tar.gz", hash = "sha256:ea16aaad7db67b5af67719e9da3302474af97a7fdf27161c2ecb30adb3570d50", size = 15625607, upload-time = "2026-06-22T20:26:26.637Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/5c/19/934f81592527a3f7f9b943c893e334c721a4644948642bc33885d584e9ec/botocore-1.43.36-py3-none-any.whl", hash = "sha256:3c65fdc39ed01d8dfde1e961b34038aed03c459f8ddf80717a12ac006475e49d", size = 15313630, upload-time = "2026-06-23T02:46:59.327Z" }, + { url = "https://files.pythonhosted.org/packages/82/6c/5807aca5a873714144beb69ed5adc04bcdb8e9496be6bf7b72a76a6c9773/botocore-1.43.35-py3-none-any.whl", hash = "sha256:51c5893b5404b74012edde03bbd9c199c4bd7ffa35ddad173d0835efa7de281b", size = 15313659, upload-time = "2026-06-22T20:26:22.017Z" }, ] [package.optional-dependencies] @@ -3223,7 +3223,7 @@ wheels = [ [[package]] name = "strands-agents" -version = "1.38.0" +version = "1.46.0" source = { registry = "https://pypi.org/simple" } dependencies = [ { name = "boto3" }, @@ -3239,9 +3239,9 @@ dependencies = [ { name = "typing-extensions" }, { name = "watchdog" }, ] -sdist = { url = "https://files.pythonhosted.org/packages/11/89/3e722f4b5bd913531bc32a23bf88aaa77a434774f294bba5bfa88690ec46/strands_agents-1.38.0.tar.gz", hash = "sha256:02a68ec321ad457f9137dfd6a99cf72cf0e86081fee35de85fbe29b9ac0af2b2", size = 858950, upload-time = "2026-04-30T16:57:43.244Z" } +sdist = { url = "https://files.pythonhosted.org/packages/2e/64/f96e656ca18df422e3783cc961a1e23223e5ccf4439f337c29d154ed036f/strands_agents-1.46.0.tar.gz", hash = "sha256:80fdc142904ce1c0e0c85645f5deef50f63d76207375705150732344c9776e8d", size = 1141936, upload-time = "2026-07-08T14:11:19.259Z" } wheels = [ - { url = "https://files.pythonhosted.org/packages/cc/06/de8d8ab14a2e92dcb0fa82db0a4cb102418a1eda139412bbe5b5725e28df/strands_agents-1.38.0-py3-none-any.whl", hash = "sha256:9dc3de17e25d70e367d37f9151f2a4c7b3ac8fc9f6237e9e1f34d00bfbfd001b", size = 422354, upload-time = "2026-04-30T16:57:41.094Z" }, + { url = "https://files.pythonhosted.org/packages/43/eb/888a34e5f0c408cdf2fc2971ca1d23085333adfb4b4dfdb34ac0a330dc7c/strands_agents-1.46.0-py3-none-any.whl", hash = "sha256:69796aaac0937182d1f4b9d2b5a1e47589ac4a3e0021102e246989d46bf27b63", size = 596775, upload-time = "2026-07-08T14:11:17.318Z" }, ] [[package]]