Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -149,7 +149,7 @@ dev = [
"ruff>=0.12.0",
"websockets>=14.1",
"wheel>=0.45.1",
"strands-agents>=1.46.0",
"strands-agents>=1.56.0",
"strands-agents-evals>=1.0.3,<2.0.0",
"deepeval>=3.5.0,<5.0.0",
"autoevals>=0.3.0,<1.0.0",
Expand All @@ -169,7 +169,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.46.0",
"strands-agents>=1.56.0",
"mcp>=1.23.0,<2.0.0",
]
langgraph = [
Expand Down
3 changes: 2 additions & 1 deletion src/bedrock_agentcore/memory/integrations/strands/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -47,7 +47,8 @@ class AgentCoreMemoryConfig(BaseModel):
memory_id: Required Bedrock AgentCore Memory ID
session_id: Required unique ID for the session
actor_id: Required unique ID for the agent instance/user
retrieval_config: Optional dictionary mapping namespaces to retrieval configurations
retrieval_config: Optional dictionary mapping namespaces to retrieval configurations.
Automatic context retrieval applies to Agent only.
batch_size: Number of messages to batch before sending to AgentCore Memory.
Default of 1 means immediate sending (no batching). Max 100.
flush_interval_seconds: Optional interval in seconds for automatic buffer flushing.
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -11,11 +11,8 @@

import boto3
from botocore.config import Config as BotocoreConfig
from strands.experimental.hooks.events import (
BidiAfterInvocationEvent,
BidiAgentInitializedEvent,
BidiMessageAddedEvent,
)
from strands.experimental.bidi import BidiAgent
from strands.experimental.bidi.hooks import BidiAgentStopEvent
from strands.experimental.hooks.multiagent.events import (
AfterMultiAgentInvocationEvent,
AfterNodeCallEvent,
Expand Down Expand Up @@ -45,7 +42,7 @@
from .converters import MemoryConverter

if TYPE_CHECKING:
from strands.agent.agent import Agent
from strands.types.agent import LocalAgent

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -829,7 +826,7 @@ def _filter_restored_tool_context(self, messages: list[SessionMessage]) -> list[

# region RepositorySessionManager overrides
@override
def append_message(self, message: Message, agent: "Agent", **kwargs: Any) -> None:
def append_message(self, message: Message, agent: "LocalAgent", **kwargs: Any) -> None:
"""Append a message to the agent's session using AgentCore's eventId as message_id.

Args:
Expand All @@ -844,11 +841,14 @@ def append_message(self, message: Message, agent: "Agent", **kwargs: Any) -> Non
self._latest_agent_message[agent.agent_id] = session_message

def retrieve_customer_context(self, event: MessageAddedEvent) -> None:
"""Retrieve customer LTM context before processing support query.
"""Retrieve customer LTM context for regular Agent invocations.

Args:
event (MessageAddedEvent): The message added event containing the agent and message data.
"""
if isinstance(event.agent, BidiAgent):
return None

messages = event.agent.messages
if not messages or messages[-1].get("role") != "user":
return None
Expand Down Expand Up @@ -924,8 +924,7 @@ def register_hooks(self, registry: HookRegistry, **kwargs) -> None:
"""Register additional hooks.

In sync mode (the default), delegates to the base class and adds the
retrieve_customer_context + batching callbacks synchronously, preserving
existing behavior exactly.
retrieve_customer_context + batching callbacks synchronously.

In async mode, registers async callbacks that wrap every per-turn
boto3-backed operation (append_message, sync_agent, buffer flushes,
Expand All @@ -942,18 +941,18 @@ def register_hooks(self, registry: HookRegistry, **kwargs) -> None:
**kwargs: Additional keyword arguments.
"""
if not self.config.async_mode:
RepositorySessionManager.register_hooks(self, registry, **kwargs)
registry.add_callback(MessageAddedEvent, lambda event: self.retrieve_customer_context(event))

# Only register AfterInvocationEvent hook when batching is enabled
if self.config.batch_size > 1:
# Completion callbacks run in reverse order, so register flushes before state syncs.
registry.add_callback(AfterInvocationEvent, lambda event: self._flush_messages())
registry.add_callback(BidiAgentStopEvent, lambda event: self._flush_messages())

RepositorySessionManager.register_hooks(self, registry, **kwargs)
registry.add_callback(MessageAddedEvent, lambda event: self.retrieve_customer_context(event))
return

# Async mode: register async callbacks that offload the existing sync
# methods to a worker thread via asyncio.to_thread. AgentInitializedEvent
# and BidiAgentInitializedEvent must stay sync (Strands disallows async
# callbacks for AgentInitializedEvent — see strands/hooks/registry.py:227).
# methods to a worker thread via asyncio.to_thread. Initialization
# callbacks must stay synchronous.
logger.warning(
"AgentCoreMemorySessionManager async_mode=True: the agent must be invoked "
"via the async path (e.g. agent.stream_async(...) or agent.invoke_async(...)). "
Expand All @@ -972,6 +971,11 @@ async def _callback(event):

return _callback

if self.config.batch_size > 1:
# Completion callbacks run in reverse order, so register flushes before state syncs.
registry.add_callback(AfterInvocationEvent, _offload(self._flush_messages))
registry.add_callback(BidiAgentStopEvent, _offload(self._flush_messages))

registry.add_callback(AgentInitializedEvent, lambda event: self.initialize(event.agent))

async def _on_message_added_persist(event: MessageAddedEvent) -> None:
Expand All @@ -982,29 +986,15 @@ async def _on_message_added_persist(event: MessageAddedEvent) -> None:
registry.add_callback(AfterInvocationEvent, _offload(self.sync_agent, lambda e: e.agent))
registry.add_callback(MessageAddedEvent, _offload(self.retrieve_customer_context, lambda e: e))

if self.config.batch_size > 1:
registry.add_callback(AfterInvocationEvent, _offload(self._flush_messages))

# Register multi-agent callbacks so async-mode parity matches sync-mode
registry.add_callback(MultiAgentInitializedEvent, _offload(self.initialize_multi_agent, lambda e: e.source))
registry.add_callback(AfterNodeCallEvent, _offload(self.sync_multi_agent, lambda e: e.source))
registry.add_callback(AfterMultiAgentInvocationEvent, _offload(self.sync_multi_agent, lambda e: e.source))

# Register BidiAgent callbacks so async-mode parity matches sync-mode.
# BidiAgentInitializedEvent dispatches through invoke_callbacks (sync),
# so its callback must stay sync; the other two dispatch through
# invoke_callbacks_async, so async wrappers are safe.
registry.add_callback(BidiAgentInitializedEvent, lambda event: self.initialize_bidi_agent(event.agent))

async def _on_bidi_message_added(event: BidiMessageAddedEvent) -> None:
await asyncio.to_thread(self.append_bidi_message, event.message, event.agent)
await asyncio.to_thread(self.sync_bidi_agent, event.agent)

registry.add_callback(BidiMessageAddedEvent, _on_bidi_message_added)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Now that bidi uses MessageAddedEvent, will retrieve_customer_context run for bidi too? It looks like it updates agent.messages, while the model still receives the original input event. could we skip retrieval for bidi or inject it into the outgoing event?

@pgrayy pgrayy Sep 16, 2026 •

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added a skip. This also highlights an important behavioral difference. It may actually make sense to go back to BidiMessageAddedEvent. I'll give this some thought. If I do update, I'll be sure to update here as well.

Another approach could be that Bidi just doesn't emit the event at all. And this is also something to consider for model providers that allow server side conversation management (e.g., OpenAI responses).

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Makes sense!

registry.add_callback(BidiAfterInvocationEvent, _offload(self.sync_bidi_agent, lambda e: e.agent))
registry.add_callback(BidiAgentStopEvent, _offload(self.sync_agent, lambda e: e.agent))

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I might be missing something here, but with batching enabled, doesn’t sync_agent() just add the state to the buffer?

since Bidi doesn’t fire AfterInvocationEvent, what flushes the pending messages and state after agent.stop()?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Added a flush and also adjusted the ordering for AfterInvocationEvent. The inline comment explains under the batch_size > 1 blocks.


@override
def initialize(self, agent: "Agent", **kwargs: Any) -> None:
def initialize(self, agent: "LocalAgent", **kwargs: Any) -> None:
if self.has_existing_agent:
logger.warning(
"An Agent already exists in session %s. We currently support one agent per session.", self.session_id
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import asyncio
import inspect
import json
import logging
import time
from datetime import datetime, timedelta, timezone
Expand All @@ -11,17 +12,14 @@
from botocore.config import Config as BotocoreConfig
from botocore.exceptions import ClientError
from strands.agent.agent import Agent
from strands.experimental.hooks.events import (
BidiAfterInvocationEvent,
BidiAgentInitializedEvent,
BidiMessageAddedEvent,
)
from strands.experimental.bidi import BidiAgent
from strands.experimental.bidi.hooks import BidiAgentStopEvent
from strands.experimental.hooks.multiagent.events import (
AfterMultiAgentInvocationEvent,
AfterNodeCallEvent,
MultiAgentInitializedEvent,
)
from strands.hooks import AfterInvocationEvent, MessageAddedEvent
from strands.hooks import AfterInvocationEvent, AgentInitializedEvent, MessageAddedEvent
from strands.hooks.registry import HookRegistry
from strands.types.exceptions import SessionException
from strands.types.session import Session, SessionAgent, SessionMessage, SessionType
Expand Down Expand Up @@ -2679,8 +2677,35 @@ def test_retrieve_customer_context_default_context_tag(self, mock_memory_client)
assert "</user_context>" in content[0]["text"]


class TestAfterInvocationHook:
"""Test AfterInvocationEvent hook integration."""
class TestSessionHooks:
"""Test session lifecycle hook integration."""

@pytest.mark.parametrize("async_mode", [False, True])
async def test_bidi_message_persists_without_retrieval(
self, agentcore_config_with_retrieval, mock_memory_client, async_mode
):
"""Bidi messages are persisted without retrieving or injecting context."""
agentcore_config_with_retrieval.async_mode = async_mode
manager = _create_session_manager(agentcore_config_with_retrieval, mock_memory_client)
manager.session_repository = Mock()
manager._latest_agent_message = {}
agent = Mock(
spec=BidiAgent,
agent_id="test-agent",
messages=[{"role": "user", "content": [{"text": "Hello"}]}],
state=Mock(),
)
agent.state.get.return_value = {}
mock_memory_client.retrieve_memories.return_value = [{"content": {"text": "User prefers blue"}, "score": 1.0}]
registry = HookRegistry()
manager.register_hooks(registry)

await registry.invoke_callbacks_async(MessageAddedEvent(agent=agent, message=agent.messages[0]))

mock_memory_client.retrieve_memories.assert_not_called()
assert agent.messages == [{"role": "user", "content": [{"text": "Hello"}]}]
mock_memory_client.create_event.assert_called_once()
manager.session_repository.update_agent.assert_called_once()

def test_after_invocation_hook_registered(self, batching_session_manager):
"""Test that AfterInvocationEvent hook is registered when batching is enabled."""
Expand Down Expand Up @@ -2745,6 +2770,35 @@ def spy_add_callback(event_type, callback):
]
assert len(flush_callbacks) == 0

@pytest.mark.parametrize("async_mode", [False, True])
@pytest.mark.parametrize("event_type", [AfterInvocationEvent, BidiAgentStopEvent])
async def test_completion_flushes_messages_and_final_state(
self, batching_config, mock_memory_client, async_mode, event_type
):
"""Completion flushes a partial batch, including the final state update."""
batching_config.async_mode = async_mode
manager = _create_session_manager(batching_config, mock_memory_client)
manager.session_repository = manager
agent = Mock(agent_id="test-agent")
agent.state.get.return_value = {}
manager.create_agent(manager.session_id, SessionAgent.from_agent(agent))
manager.create_message(
manager.session_id,
agent.agent_id,
SessionMessage(message={"role": "user", "content": [{"text": "Hello"}]}, message_id=0),
)
manager.memory_client.gmdp_client.create_event.assert_not_called()

agent.state.get.return_value = {"status": "stopped"}
registry = HookRegistry()
manager.register_hooks(registry)
await registry.invoke_callbacks_async(event_type(agent=agent))

assert manager.pending_message_count() == 0
assert manager.pending_agent_state_count() == 0
state_payloads = manager.memory_client.gmdp_client.create_event.call_args.kwargs["payload"]
assert [json.loads(payload["blob"])["state"] for payload in state_payloads] == [{}, {"status": "stopped"}]


class TestIntervalFlush:
"""Test interval-based flush mechanism for long-running agents."""
Expand Down Expand Up @@ -3749,16 +3803,13 @@ def test_async_mode_registers_bidi_agent_callbacks(self, mock_memory_client):
registry = HookRegistry()
manager.register_hooks(registry)

# BidiAgentInitializedEvent dispatches via the sync hook path, so its callback must NOT be a coroutine.
init_callbacks = list(registry.get_callbacks_for(BidiAgentInitializedEvent(agent=Mock())))
assert init_callbacks, "No callbacks registered for BidiAgentInitializedEvent"
init_callbacks = list(registry.get_callbacks_for(AgentInitializedEvent(agent=Mock())))
assert init_callbacks
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 in (
BidiMessageAddedEvent(agent=Mock(), message={"role": "user", "content": [{"text": "x"}]}),
BidiAfterInvocationEvent(agent=Mock()),
MessageAddedEvent(agent=Mock(), message={"role": "user", "content": [{"text": "x"}]}),
BidiAgentStopEvent(agent=Mock()),
):
callbacks = list(registry.get_callbacks_for(event))
assert callbacks, f"No callbacks registered for {type(event).__name__}"
Expand Down
10 changes: 5 additions & 5 deletions uv.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading