-
Notifications
You must be signed in to change notification settings - Fork 153
fix(memory): use current Strands bidi session hooks #664
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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, | ||
|
|
@@ -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__) | ||
|
|
||
|
|
@@ -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: | ||
|
|
@@ -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 | ||
|
|
@@ -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, | ||
|
|
@@ -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(...)). " | ||
|
|
@@ -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: | ||
|
|
@@ -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) | ||
| 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)) | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I might be missing something here, but with batching enabled, doesn’t since Bidi doesn’t fire
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 | ||
|
|
||
Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.
There was a problem hiding this comment.
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, willretrieve_customer_contextrun for bidi too? It looks like it updatesagent.messages, while the model still receives the original input event. could we skip retrieval for bidi or inject it into the outgoing event?Uh oh!
There was an error while loading. Please reload this page.
There was a problem hiding this comment.
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).
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Makes sense!