Skip to content
Closed
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
Original file line number Diff line number Diff line change
Expand Up @@ -11,11 +11,6 @@

import boto3
from botocore.config import Config as BotocoreConfig
from strands.experimental.hooks.events import (
BidiAfterInvocationEvent,
BidiAgentInitializedEvent,
BidiMessageAddedEvent,
)
from strands.experimental.hooks.multiagent.events import (
AfterMultiAgentInvocationEvent,
AfterNodeCallEvent,
Expand Down Expand Up @@ -923,9 +918,8 @@ def retrieve_for_namespace(namespace: str, retrieval_config: RetrievalConfig):
def register_hooks(self, registry: HookRegistry, **kwargs) -> None:
"""Register additional hooks.

In sync mode (the default), delegates to the base class and adds the
retrieve_customer_context + batching callbacks synchronously, preserving
existing behavior exactly.
In sync mode (the default), registers the standard session callbacks and
adds the 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,7 +936,13 @@ 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(AgentInitializedEvent, lambda event: self.initialize(event.agent))
registry.add_callback(MessageAddedEvent, lambda event: self.append_message(event.message, event.agent))
registry.add_callback(MessageAddedEvent, lambda event: self.sync_agent(event.agent))
registry.add_callback(AfterInvocationEvent, lambda event: self.sync_agent(event.agent))
registry.add_callback(MultiAgentInitializedEvent, lambda event: self.initialize_multi_agent(event.source))
registry.add_callback(AfterNodeCallEvent, lambda event: self.sync_multi_agent(event.source))
registry.add_callback(AfterMultiAgentInvocationEvent, lambda event: self.sync_multi_agent(event.source))
registry.add_callback(MessageAddedEvent, lambda event: self.retrieve_customer_context(event))

# Only register AfterInvocationEvent hook when batching is enabled
Expand All @@ -952,8 +952,8 @@ def register_hooks(self, registry: HookRegistry, **kwargs) -> None:

# 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).
# must stay sync because Strands disallows async callbacks for it (see
# strands/hooks/registry.py:227).
logger.warning(
"AgentCoreMemorySessionManager async_mode=True: the agent must be invoked "
"via the async path (e.g. agent.stream_async(...) or agent.invoke_async(...)). "
Expand Down Expand Up @@ -990,19 +990,6 @@ async def _on_message_added_persist(event: MessageAddedEvent) -> None:
registry.add_callback(AfterNodeCallEvent, _offload(self.sync_multi_agent, lambda e: e.source))
registry.add_callback(AfterMultiAgentInvocationEvent, _offload(self.sync_multi_agent, lambda e: e.source))

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

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

registry.add_callback(BidiMessageAddedEvent, _on_bidi_message_added)
registry.add_callback(BidiAfterInvocationEvent, _offload(self.sync_bidi_agent, lambda e: e.agent))

@override
def initialize(self, agent: "Agent", **kwargs: Any) -> None:
if self.has_existing_agent:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -11,11 +11,6 @@
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.hooks.multiagent.events import (
AfterMultiAgentInvocationEvent,
AfterNodeCallEvent,
Expand Down Expand Up @@ -3742,29 +3737,6 @@ def test_async_mode_logs_sync_invocation_warning(self, mock_memory_client, caplo

assert any("async_mode=True" in rec.message and "stream_async" in rec.message for rec in caplog.records)

def test_async_mode_registers_bidi_agent_callbacks(self, mock_memory_client):
"""async_mode=True: BidiAgent events get callbacks; init stays sync, others are async."""
config = AgentCoreMemoryConfig(memory_id="m", session_id="s", actor_id="a", async_mode=True)
manager = _create_session_manager(config, 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"
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()),
):
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)


class TestFlushAgentStatesRaceCondition:
"""Tests for the copy-and-clear-under-one-lock fix in _flush_agent_states_only."""

Expand Down
Loading