diff --git a/.github/workflows/integration-testing.yml b/.github/workflows/integration-testing.yml index e3194e2d..05bcbf37 100644 --- a/.github/workflows/integration-testing.yml +++ b/.github/workflows/integration-testing.yml @@ -113,8 +113,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/test_session_manager.py tests_integ/memory/integrations/test_memory_store.py + timeout: 30 extra-deps: "" ignore: "" - group: evaluation @@ -186,7 +186,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 1a700b76..65ad6829 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>=1.0.3,<2.0.0", "langchain>=1.0.0", "langgraph>=1.0.0", @@ -164,7 +164,7 @@ a2a = ["a2a-sdk[http-server]>=0.3,<0.4"] a2a-v1 = ["a2a-sdk[http-server]>=1.0.1,<2.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/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..751ee654 --- /dev/null +++ b/src/bedrock_agentcore/memory/integrations/strands/memorystore/store.py @@ -0,0 +1,288 @@ +"""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. AgentCore returns at most + 100 records per search, so larger values are clamped to that limit. + 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 + self._warned_result_cap_clamped = False + # 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) + # AgentCore rejects topK above MAX_TOPK, so clamp on every path rather than only the + # score-filter one: a plausible max_search_results of 101 otherwise reached the wire + # verbatim and failed as an opaque service ValidationException, while the same value + # with min_score set succeeded. Warn once so the reduced cap is not silent. + if top_k > MAX_TOPK: + top_k = MAX_TOPK + if want > MAX_TOPK and not self._warned_result_cap_clamped: + self._warned_result_cap_clamped = True + logger.warning( + "AgentCoreMemoryStore %s: max_search_results=%d exceeds AgentCore's topK limit of %d; " + "searches return at most %d records per request.", + self.name, + want, + MAX_TOPK, + MAX_TOPK, + ) + 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/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..205cb4ee --- /dev/null +++ b/tests/bedrock_agentcore/memory/integrations/strands/memorystore/test_store.py @@ -0,0 +1,529 @@ +"""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 + + +@pytest.mark.parametrize(("want", "expected_top_k"), [(100, 100), (101, 100), (200, 100)]) +async def test_result_cap_never_exceeds_service_top_k_limit(want: int, expected_top_k: int) -> None: + """Never send a topK above AgentCore's limit, with or without a score floor. + + Without the clamp on this path, 101 and 200 reached the service verbatim and failed with a + raw ``ValidationException`` while the same values succeeded once ``min_score`` was set. + """ + client = client_returning([]) + await make_store(client, max_search_results=want).search("q") + assert client.retrieve_memory_records.call_args.kwargs["searchCriteria"]["topK"] == expected_top_k + + floor_client = client_returning([]) + await make_store(floor_client, max_search_results=want, min_score=0.5).search("q") + assert floor_client.retrieve_memory_records.call_args.kwargs["searchCriteria"]["topK"] == expected_top_k + + +@pytest.mark.parametrize("want", [101, 200]) +async def test_call_time_result_cap_above_limit_is_clamped(want: int) -> None: + """Clamp a per-call cap too: ``search_memory`` options reach the same wire field.""" + client = client_returning([]) + await make_store(client).search("q", {"max_search_results": want}) + assert client.retrieve_memory_records.call_args.kwargs["searchCriteria"]["topK"] == 100 + + +async def test_result_cap_above_limit_warns_once(caplog: pytest.LogCaptureFixture) -> None: + """Surface the reduced cap instead of silently under-delivering on every call.""" + client = client_returning([]) + store = make_store(client, max_search_results=200) + with caplog.at_level("WARNING"): + await store.search("q") + await store.search("q") + warnings = [record for record in caplog.records if "exceeds AgentCore's topK limit" in record.message] + assert len(warnings) == 1 + + +async def test_result_cap_at_limit_does_not_warn(caplog: pytest.LogCaptureFixture) -> None: + """A cap at the service limit is a normal request, not a misconfiguration.""" + client = client_returning([]) + with caplog.at_level("WARNING"): + await make_store(client, max_search_results=100).search("q") + assert not [record for record in caplog.records if "topK limit" in record.message] + + +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/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_session_manager.py b/tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_session_manager.py index fa0c4787..cb45866a 100644 --- a/tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_session_manager.py +++ b/tests/bedrock_agentcore/memory/integrations/strands/test_agentcore_memory_session_manager.py @@ -3722,9 +3722,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): @@ -3746,15 +3751,19 @@ 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_callbacks = list(registry.get_callbacks_for(BidiAgentInitializedEvent(agent=Mock()))) 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_integ/memory/integrations/test_memory_store.py b/tests_integ/memory/integrations/test_memory_store.py new file mode 100644 index 00000000..852a0c63 --- /dev/null +++ b/tests_integ/memory/integrations/test_memory_store.py @@ -0,0 +1,279 @@ +"""Live AgentCore Memory tests for the native Strands ``MemoryStore`` integration. + +Requires ``MEMORY_PREPOPULATED_ID`` (a pre-provisioned memory with semantic and summary +strategies), ``BEDROCK_TEST_REGION``, and credentials allowed to invoke a Bedrock model. +Optionally set ``STRANDS_TEST_MODEL_ID`` to override the default test model. Long-term +extraction is eventually consistent, so these tests poll and may take several minutes. +""" + +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 cb467767..9e3aca74 100644 --- a/uv.lock +++ b/uv.lock @@ -540,8 +540,8 @@ requires-dist = [ { name = "a2a-sdk", extras = ["http-server"], marker = "extra == 'a2a-v1'", specifier = ">=1.0.1,<2.0" }, { name = "ag-ui-protocol", marker = "extra == 'ag-ui'", specifier = ">=0.1.10" }, { name = "autoevals", marker = "extra == 'autoevals'", specifier = ">=0.0.50" }, - { 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 = "deepeval", marker = "extra == 'deepeval'", specifier = ">=2.0.0" }, { name = "httpx", marker = "extra == 'langgraph'", specifier = ">=0.27.0" }, { name = "jinja2", marker = "extra == 'simulation'", specifier = ">=3.1.0" }, @@ -552,7 +552,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 == 'autoevals'", specifier = ">=1.0.3,<2.0.0" }, { name = "strands-agents-evals", marker = "extra == 'deepeval'", specifier = ">=1.0.3,<2.0.0" }, { name = "strands-agents-evals", marker = "extra == 'simulation'", specifier = ">=1.0.3,<2.0.0" }, @@ -582,7 +582,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 = ">=1.0.3,<2.0.0" }, { name = "websockets", specifier = ">=14.1" }, { name = "wheel", specifier = ">=0.45.1" },