Skip to content
Draft
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: 4 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -77,8 +77,12 @@ select = [
"G", # logging format
"I", # isort
"LOG", # logging
"TID251", # flake8-tidy-imports banned-api
]

[tool.ruff.lint.flake8-tidy-imports.banned-api]
"strands.experimental" = { msg = "Use stable Strands APIs instead." }

[tool.ruff.lint.per-file-ignores]
"!src/**/*.py" = ["D"]
"src/bedrock_agentcore/memory/metadata-workflow.ipynb" = ["E501"]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -11,14 +11,13 @@

import boto3
from botocore.config import Config as BotocoreConfig
from strands.experimental.bidi import BidiAgent
from strands.experimental.bidi.hooks import BidiAgentStopEvent
from strands.experimental.hooks.multiagent.events import (
from strands.hooks import (
AfterInvocationEvent,
AfterMultiAgentInvocationEvent,
AfterNodeCallEvent,
MessageAddedEvent,
MultiAgentInitializedEvent,
)
from strands.hooks import AfterInvocationEvent, MessageAddedEvent
from strands.hooks.events import AgentInitializedEvent
from strands.hooks.registry import HookRegistry
from strands.session.repository_session_manager import RepositorySessionManager
Expand Down Expand Up @@ -851,9 +850,6 @@ def retrieve_customer_context(self, event: MessageAddedEvent) -> None:
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 @@ -949,7 +945,6 @@ def register_hooks(self, registry: HookRegistry, **kwargs) -> None:
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))
Expand Down Expand Up @@ -979,7 +974,6 @@ async def _callback(event):
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))

Expand All @@ -996,8 +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))

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

@override
def initialize(self, agent: "LocalAgent", **kwargs: Any) -> None:
if self.has_existing_agent:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -12,14 +12,14 @@
from botocore.config import Config as BotocoreConfig
from botocore.exceptions import ClientError
from strands.agent.agent import Agent
from strands.experimental.bidi import BidiAgent
from strands.experimental.bidi.hooks import BidiAgentStopEvent
from strands.experimental.hooks.multiagent.events import (
from strands.hooks import (
AfterInvocationEvent,
AfterMultiAgentInvocationEvent,
AfterNodeCallEvent,
AgentInitializedEvent,
MessageAddedEvent,
MultiAgentInitializedEvent,
)
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 @@ -2693,33 +2693,6 @@ def test_retrieve_customer_context_default_context_tag(self, mock_memory_client)
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."""
registry = HookRegistry()
Expand Down Expand Up @@ -2784,10 +2757,7 @@ 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
):
async def test_completion_flushes_messages_and_final_state(self, batching_config, mock_memory_client, async_mode):
"""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)
Expand All @@ -2805,7 +2775,7 @@ async def test_completion_flushes_messages_and_final_state(
agent.state.get.return_value = {"status": "stopped"}
registry = HookRegistry()
manager.register_hooks(registry)
await registry.invoke_callbacks_async(event_type(agent=agent))
await registry.invoke_callbacks_async(AfterInvocationEvent(agent=agent))

assert manager.pending_message_count() == 0
assert manager.pending_agent_state_count() == 0
Expand Down Expand Up @@ -3809,8 +3779,8 @@ 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."""
def test_async_mode_registers_agent_callbacks(self, mock_memory_client):
"""async_mode=True: initialization stays sync while message callbacks 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()
Expand All @@ -3820,13 +3790,10 @@ def test_async_mode_registers_bidi_agent_callbacks(self, mock_memory_client):
assert init_callbacks
assert not any(asyncio.iscoroutinefunction(cb) for cb in init_callbacks)

for event in (
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__}"
assert all(asyncio.iscoroutinefunction(cb) for cb in callbacks)
event = MessageAddedEvent(agent=Mock(), message={"role": "user", "content": [{"text": "x"}]})
callbacks = list(registry.get_callbacks_for(event))
assert callbacks
assert all(asyncio.iscoroutinefunction(cb) for cb in callbacks)


class TestFlushAgentStatesRaceCondition:
Expand Down
Loading