From 3e5e89e265aaa33580d0cc5dbe6891f5d36a7c5f Mon Sep 17 00:00:00 2001 From: Patrick Gray Date: Tue, 29 Sep 2026 12:06:47 -0400 Subject: [PATCH] feat(bidi): support cancellation from custom tools (#4664) --- .../sdk/bidirectional-streaming/agent.mdx | 4 +- .../sdk/bidirectional-streaming/events.mdx | 3 +- .../sdk/bidirectional-streaming/io.mdx | 16 +- .../models/bedrock.mdx | 3 +- .../bidirectional-streaming/models/google.mdx | 4 +- .../bidirectional-streaming/models/openai.mdx | 4 +- .../bidirectional-streaming/quickstart.mdx | 34 ++-- .../src/strands/experimental/bidi/__init__.py | 4 +- .../strands/experimental/bidi/agent/agent.py | 19 +- .../strands/experimental/bidi/agent/loop.py | 26 +-- .../experimental/bidi/tools/__init__.py | 18 -- .../bidi/tools/stop_conversation.py | 32 ---- .../src/strands/tools/executors/_executor.py | 2 +- strands-py/src/strands/types/agent.py | 4 + .../strands/agent/test_agent_cancellation.py | 4 +- .../experimental/bidi/agent/test_agent.py | 38 +++- .../experimental/bidi/agent/test_loop.py | 170 +++++++----------- strands-py/tests_typing/test_local_agent.py | 3 +- 18 files changed, 160 insertions(+), 228 deletions(-) delete mode 100644 strands-py/src/strands/experimental/bidi/tools/__init__.py delete mode 100644 strands-py/src/strands/experimental/bidi/tools/stop_conversation.py diff --git a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/agent.mdx b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/agent.mdx index f23bf7752..b1821de42 100644 --- a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/agent.mdx +++ b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/agent.mdx @@ -177,7 +177,9 @@ Each grouped tool-use message stays adjacent to its matching result message, eve Tool-result messages use `metadata["custom"]["bidi"]["kind"]` to distinguish `tool_dispatch` acknowledgements from `tool_result` messages. -After the group finishes, the agent loop checks `request_state["stop_event_loop"]` to trigger graceful shutdown instead of sending tool results back to the model. Any tool can set this flag to stop the conversation. The SDK's experimental `stop` tool uses this mechanism. +To let a tool end the conversation, call `tool_context.agent.cancel()`. Cancellation takes effect only after the tool group completes. Requests from other contexts remain pending until then. + +See [Graceful shutdown](./quickstart.mdx#graceful-shutdown) for a custom tool example. ### Connection Lifecycle diff --git a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/events.mdx b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/events.mdx index 0c65aa29d..d1eb5a4b7 100644 --- a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/events.mdx +++ b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/events.mdx @@ -225,8 +225,7 @@ Emitted when the streaming connection is closed. - `"timeout"`: Connection timed out - `"error"`: Error occurred - `"complete"`: Conversation completed normally - - `"user_request"`: User requested closure (via the SDK's experimental `stop` - tool or any tool that sets `request_state["stop_event_loop"]`) + - `"user_request"`: Cancellation requested through `agent.cancel()` ### Response Lifecycle Events diff --git a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/io.mdx b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/io.mdx index 4c741f513..524b9651e 100644 --- a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/io.mdx +++ b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/io.mdx @@ -88,12 +88,10 @@ Pass your I/O streams to the agent's `run()` method to connect them to the agent import asyncio from strands.experimental.bidi.agent import BidiAgent -from strands.experimental.tools import stop async def main(): - # stop tool allows user to verbally stop agent execution. - agent = BidiAgent(tools=[stop]) + agent = BidiAgent() await agent.run(inputs=[MyInputStream()], outputs=[MyOutputStream()]) @@ -104,6 +102,8 @@ The `run()` method handles startup, execution, and shutdown for the agent and it streams. Inputs and outputs run concurrently, so you can mix and match implementations. If an I/O task fails, `run()` cancels the remaining tasks, stops the streams, and re-raises the exception. +For a tool that lets users end the conversation, see +[Graceful shutdown](./quickstart.mdx#graceful-shutdown). ## Audio I/O @@ -126,12 +126,10 @@ import asyncio from strands.experimental.bidi.agent import BidiAgent from strands.experimental.bidi.io import AudioIO -from strands.experimental.tools import stop async def main(): - # stop tool allows user to verbally stop agent execution. - agent = BidiAgent(tools=[stop]) + agent = BidiAgent() audio_io = AudioIO(input_device_index=1) await agent.run( @@ -232,7 +230,6 @@ import asyncio from strands.experimental.bidi.agent import BidiAgent from strands.experimental.bidi.io import ConsoleIO from strands.experimental.bidi.models import OpenAIRealtimeModel -from strands.experimental.tools import stop async def main(): @@ -241,7 +238,7 @@ async def main(): transcription_model_id=None, params={"output_modalities": ["text"]}, ) - agent = BidiAgent(model=model, tools=[stop]) + agent = BidiAgent(model=model) console_io = ConsoleIO() await agent.run( @@ -280,11 +277,10 @@ import asyncio from strands.experimental.bidi.agent import BidiAgent from strands.experimental.bidi.io import AudioIO, ConsoleIO -from strands.experimental.tools import stop async def main(): - agent = BidiAgent(tools=[stop]) + agent = BidiAgent() console_io = ConsoleIO(placeholder="Type or speak…") audio_io = AudioIO(console=console_io) diff --git a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/bedrock.mdx b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/bedrock.mdx index 6bcc76eec..c22870aaf 100644 --- a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/bedrock.mdx +++ b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/bedrock.mdx @@ -49,7 +49,6 @@ import asyncio from strands.experimental.bidi.agent import BidiAgent from strands.experimental.bidi.io import AudioIO from strands.experimental.bidi.models import BedrockNovaSonicModel -from strands.experimental.tools import stop from strands.vended_tools import notebook @@ -59,7 +58,7 @@ async def main() -> None: region="us-east-1", voice="tiffany", ) - agent = BidiAgent(model=model, tools=[notebook, stop]) + agent = BidiAgent(model=model, tools=[notebook]) audio_io = AudioIO() await agent.run(inputs=[audio_io.input()], outputs=[audio_io.output()]) diff --git a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/google.mdx b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/google.mdx index 87314d351..09d460af3 100644 --- a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/google.mdx +++ b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/google.mdx @@ -45,7 +45,6 @@ import asyncio from strands.experimental.bidi.agent import BidiAgent from strands.experimental.bidi.io import AudioIO from strands.experimental.bidi.models import GoogleGeminiLiveModel -from strands.experimental.tools import stop from strands.vended_tools import notebook @@ -55,8 +54,7 @@ async def main() -> None: voice="Kore", client_args={"api_key": ""}, ) - # stop tool allows user to verbally stop agent execution. - agent = BidiAgent(model=model, tools=[notebook, stop]) + agent = BidiAgent(model=model, tools=[notebook]) audio_io = AudioIO() await agent.run(inputs=[audio_io.input()], outputs=[audio_io.output()]) diff --git a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/openai.mdx b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/openai.mdx index 46c3096e6..bd900544d 100644 --- a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/openai.mdx +++ b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/models/openai.mdx @@ -44,7 +44,6 @@ import asyncio from strands.experimental.bidi.agent import BidiAgent from strands.experimental.bidi.io import AudioIO from strands.experimental.bidi.models import OpenAIRealtimeModel -from strands.experimental.tools import stop from strands.vended_tools import notebook @@ -55,8 +54,7 @@ async def main() -> None: voice="coral", api_key="", ) - # stop tool allows user to verbally stop agent execution. - agent = BidiAgent(model=model, tools=[notebook, stop]) + agent = BidiAgent(model=model, tools=[notebook]) audio_io = AudioIO() await agent.run(inputs=[audio_io.input()], outputs=[audio_io.output()]) diff --git a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/quickstart.mdx b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/quickstart.mdx index 0ac985195..5e4ebfb02 100644 --- a/site/src/content/docs/user-guide/sdk/bidirectional-streaming/quickstart.mdx +++ b/site/src/content/docs/user-guide/sdk/bidirectional-streaming/quickstart.mdx @@ -406,22 +406,28 @@ See [Controlling Conversation Lifecycle](#controlling-conversation-lifecycle) fo ## Graceful Shutdown -Use the SDK's experimental `stop` tool to allow users to end conversations -naturally. It sets `request_state["stop_event_loop"]`, which the agent loop checks -to trigger a graceful shutdown: +To let users end a conversation by voice, define a tool that calls `agent.cancel()`: ```python import asyncio +from strands import LocalAgent, ToolContext, tool from strands.experimental.bidi.agent import BidiAgent from strands.experimental.bidi.io import AudioIO from strands.experimental.bidi.models import BedrockNovaSonicModel -from strands.experimental.tools import stop + + +@tool(context=True) +def end_conversation(tool_context: ToolContext[LocalAgent]) -> str: + """End the conversation when the user asks to stop.""" + tool_context.agent.cancel() + return "Ending conversation" + model = BedrockNovaSonicModel(model_id="amazon.nova-2-sonic-v1:0") agent = BidiAgent( model=model, - tools=[stop], - system_prompt="You are a helpful assistant. When the user says 'stop conversation', use the stop tool." + tools=[end_conversation], + system_prompt="You are a helpful assistant.", ) audio_io = AudioIO() @@ -431,24 +437,12 @@ async def main(): inputs=[audio_io.input()], outputs=[audio_io.output()] ) - # Conversation ends when user says "stop conversation" + # run() returns after the agent calls end_conversation. asyncio.run(main()) ``` -You can also create custom stop tools using the `request_state["stop_event_loop"]` flag: - -```python -from strands import tool - -@tool -def end_session(request_state: dict) -> str: - request_state["stop_event_loop"] = True - return "Goodbye!" -``` - -The agent will gracefully close the connection when any tool sets `request_state["stop_event_loop"] = True`. - +Bidi checks for cancellation after a tool group completes and its results are recorded. Requests from other contexts remain pending until that checkpoint. ## Debug Logs diff --git a/strands-py/src/strands/experimental/bidi/__init__.py b/strands-py/src/strands/experimental/bidi/__init__.py index 06c814e19..62c6134a4 100644 --- a/strands-py/src/strands/experimental/bidi/__init__.py +++ b/strands-py/src/strands/experimental/bidi/__init__.py @@ -1,7 +1,7 @@ """Experimental bidirectional streaming APIs.""" -from . import agent, hooks, io, models, tools, types +from . import agent, hooks, io, models, types from .agent import BidiAgent as BidiAgent # Compatibility for AgentCore's root import; remove after AgentCore imports from bidi.agent. -__all__ = ["agent", "hooks", "io", "models", "tools", "types"] +__all__ = ["agent", "hooks", "io", "models", "types"] diff --git a/strands-py/src/strands/experimental/bidi/agent/agent.py b/strands-py/src/strands/experimental/bidi/agent/agent.py index 5d6961504..2b173496f 100644 --- a/strands-py/src/strands/experimental/bidi/agent/agent.py +++ b/strands-py/src/strands/experimental/bidi/agent/agent.py @@ -153,7 +153,6 @@ class BidiAgent(LocalAgent): self.messages = messages if messages is not None else [] self._storage: Storage | None = storage self._sandbox: Sandbox = NotASandboxLocalEnvironment() - # Never set yet: bidirectional agents do not act on a cancellation signal. self._cancel_signal = threading.Event() # Agent identification @@ -289,9 +288,20 @@ class BidiAgent(LocalAgent): def event_loop_metrics(self, value: "EventLoopMetrics") -> None: raise NotImplementedError("event_loop_metrics is not supported by bidirectional agents yet") + def cancel(self) -> None: + """Request cancellation of the current conversation. + + This method is thread-safe and idempotent. Cancellation takes effect + only after a tool group completes. + """ + self._cancel_signal.set() + @property def cancel_signal(self) -> threading.Event: - """The cancellation signal; never set yet, because bidirectional agents do not act on it.""" + """The cancellation signal for the current conversation. + + Treat as read-only; call ``cancel()`` to request cancellation. + """ return self._cancel_signal def add_hook( @@ -445,7 +455,10 @@ class BidiAgent(LocalAgent): closes the connection to the model provider. """ self._started = False - await self._loop.stop() + try: + await self._loop.stop() + finally: + self._cancel_signal.clear() def take_snapshot( self, diff --git a/strands-py/src/strands/experimental/bidi/agent/loop.py b/strands-py/src/strands/experimental/bidi/agent/loop.py index b22c0e4c7..21dde3603 100644 --- a/strands-py/src/strands/experimental/bidi/agent/loop.py +++ b/strands-py/src/strands/experimental/bidi/agent/loop.py @@ -6,7 +6,6 @@ The agent loop handles the events received from the model and executes tools whe import asyncio import logging import time -import warnings from collections.abc import AsyncGenerator from dataclasses import dataclass from typing import TYPE_CHECKING, Any, Literal, cast @@ -198,6 +197,8 @@ class _AgentLoop: self._model_task = self._task_pool.create(self._run_model(self._generation)) self._invocation_state = invocation_state if invocation_state is not None else {} + # Retained for compatibility with shared tools that expect request_state. + self._invocation_state.setdefault("request_state", {}) self._send_gate.set() self._started = True @@ -741,7 +742,6 @@ class _AgentLoop: async def _run_tools(self, tool_uses: list[ToolUse]) -> None: """Execute a provider's tool group concurrently and send its results together.""" - invocation_state = self._invocation_state try: async with _TaskGroup() as task_group: tasks = [task_group.create_task(self._run_tool(tool_use)) for tool_use in tool_uses] @@ -759,18 +759,8 @@ class _AgentLoop: await self._agent._append_messages(tool_use_message, tool_result_message) await self._event_queue.put(ToolResultMessageEvent(tool_result_message)) - should_stop = invocation_state.get("request_state", {}).get("stop_event_loop", False) - if not should_stop and any(tool_use["name"] == "stop_conversation" for tool_use in tool_uses): - warnings.warn( - "Stopping the event loop by tool name 'stop_conversation' is deprecated. " - "Use request_state['stop_event_loop'] = True instead.", - DeprecationWarning, - stacklevel=2, - ) - should_stop = True - - if should_stop: - logger.info("stop_event_loop= | stopping conversation") + if self._agent.cancel_signal.is_set(): + logger.info("cancellation requested | stopping conversation") connection_id = getattr(self._agent.model, "_connection_id", "unknown") await self._event_queue.put(BidiConnectionStopEvent(connection_id=connection_id, reason="user_request")) return @@ -803,11 +793,6 @@ class _AgentLoop: tool_results: list[ToolResult] = [] - # Ensure request_state exists for tools like strands_tools.stop - invocation_state = self._invocation_state - if "request_state" not in invocation_state: - invocation_state["request_state"] = {} - tool_call_span = self._tracer.start_tool_call_span(tool_use, parent_span=self._session_span) tool_result: ToolResult | None = None tool_error: Exception | None = None @@ -817,7 +802,7 @@ class _AgentLoop: self._agent, tool_use, tool_results, - invocation_state, + self._invocation_state, structured_output_context=None, ) @@ -829,7 +814,6 @@ class _AgentLoop: await self._event_queue.put(tool_event) - # Normal flow for all tools (including stop_conversation) tool_result_event = cast(ToolResultEvent, tool_event) tool_result = tool_result_event.tool_result diff --git a/strands-py/src/strands/experimental/bidi/tools/__init__.py b/strands-py/src/strands/experimental/bidi/tools/__init__.py deleted file mode 100644 index de67040de..000000000 --- a/strands-py/src/strands/experimental/bidi/tools/__init__.py +++ /dev/null @@ -1,18 +0,0 @@ -"""Built-in tools for bidirectional agents. - -.. deprecated:: - The built-in ``stop_conversation`` tool is deprecated. Use ``strands_tools.stop`` or set - ``request_state["stop_event_loop"] = True`` in any custom tool instead. - -To stop a bidirectional conversation, use the standard ``stop`` tool from strands_tools:: - - from strands_tools import stop - agent = BidiAgent(tools=[stop, ...]) - -The stop tool sets ``request_state["stop_event_loop"] = True``, which signals the -BidiAgent to gracefully close the connection. -""" - -from .stop_conversation import stop_conversation - -__all__ = ["stop_conversation"] diff --git a/strands-py/src/strands/experimental/bidi/tools/stop_conversation.py b/strands-py/src/strands/experimental/bidi/tools/stop_conversation.py deleted file mode 100644 index 21b530552..000000000 --- a/strands-py/src/strands/experimental/bidi/tools/stop_conversation.py +++ /dev/null @@ -1,32 +0,0 @@ -"""Tool to gracefully stop a bidirectional connection. - -.. deprecated:: - The ``stop_conversation`` tool is deprecated and will be removed in a future version. - Use ``strands_tools.stop`` or set ``request_state["stop_event_loop"] = True`` in any custom tool instead. -""" - -import warnings - -from ....tools.decorator import tool - - -@tool -def stop_conversation() -> str: - """Stop the bidirectional conversation gracefully. - - .. deprecated:: - Use ``strands_tools.stop`` or set ``request_state["stop_event_loop"] = True`` in a custom tool instead. - - Use ONLY when user says "stop conversation" exactly. - Do NOT use for: "stop", "goodbye", "bye", "exit", "quit", "end" or other farewells or phrases. - - Returns: - Success message confirming the conversation will end. - """ - warnings.warn( - "stop_conversation is deprecated and will be removed in a future version. " - "Use strands_tools.stop or set request_state['stop_event_loop'] = True in any custom tool instead.", - DeprecationWarning, - stacklevel=2, - ) - return "Ending conversation" diff --git a/strands-py/src/strands/tools/executors/_executor.py b/strands-py/src/strands/tools/executors/_executor.py index 0c376433f..16008258d 100644 --- a/strands-py/src/strands/tools/executors/_executor.py +++ b/strands-py/src/strands/tools/executors/_executor.py @@ -58,7 +58,7 @@ class ToolExecutor(abc.ABC): if cast(dict[str, Any], after_event.result).get("cancelled") is True: return False if not ToolExecutor._is_agent(agent): - return True + return not agent.cancel_signal.is_set() return not cast("Agent", agent)._observe_cancellation() async def _execute_background( diff --git a/strands-py/src/strands/types/agent.py b/strands-py/src/strands/types/agent.py index ebbbc1f1b..5868ca8f0 100644 --- a/strands-py/src/strands/types/agent.py +++ b/strands-py/src/strands/types/agent.py @@ -104,6 +104,10 @@ class LocalAgent(Protocol): """The cancellation signal for the current invocation.""" ... + def cancel(self) -> None: + """Request cancellation at the agent's next supported checkpoint.""" + ... + def add_hook( self, callback: HookCallback[_TEvent], diff --git a/strands-py/tests/strands/agent/test_agent_cancellation.py b/strands-py/tests/strands/agent/test_agent_cancellation.py index dcf3a760a..4d63355df 100644 --- a/strands-py/tests/strands/agent/test_agent_cancellation.py +++ b/strands-py/tests/strands/agent/test_agent_cancellation.py @@ -7,7 +7,7 @@ from unittest.mock import ANY import pytest -from strands import Agent, tool +from strands import Agent, LocalAgent, ToolContext, tool from strands.hooks import AfterModelCallEvent, BeforeModelCallEvent, BeforeToolCallEvent, BeforeToolsEvent from tests.fixtures.mocked_model_provider import MockedModelProvider @@ -575,7 +575,7 @@ async def test_hook_cancelled_tool_batch_does_not_replay_the_stored_tool_use(): def _approver_agent(ran, cancel_on_first_run=False): @tool(context=True) - def approver(tool_context) -> str: + def approver(tool_context: ToolContext[LocalAgent]) -> str: """Require approval before doing the work.""" tool_context.interrupt("approve", reason="proceed?") ran.append("executed") diff --git a/strands-py/tests/strands/experimental/bidi/agent/test_agent.py b/strands-py/tests/strands/experimental/bidi/agent/test_agent.py index 2bdfff21d..3a4d2d559 100644 --- a/strands-py/tests/strands/experimental/bidi/agent/test_agent.py +++ b/strands-py/tests/strands/experimental/bidi/agent/test_agent.py @@ -2,7 +2,6 @@ import asyncio import sys -import threading import unittest.mock from contextlib import nullcontext from uuid import uuid4 @@ -18,6 +17,7 @@ from strands.experimental.bidi.types import ( BidiConnectionStartEvent, BidiConnectionStopEvent, BidiMessage, + BidiToolUseBlocksEvent, BidiTranscriptDeltaEvent, BidiTranscriptStartEvent, InputStream, @@ -313,12 +313,40 @@ def test_bidi_agent_sandbox_defaults_to_host_environment(mock_model): assert agent.sandbox is agent.sandbox -def test_bidi_agent_cancel_signal_is_never_set(mock_model): +def test_cancel_sets_signal(mock_model): agent = BidiAgent(model=mock_model) + signal = agent.cancel_signal - assert isinstance(agent.cancel_signal, threading.Event) - assert not agent.cancel_signal.is_set() - assert agent.cancel_signal is agent.cancel_signal + assert not signal.is_set() + + agent.cancel() + agent.cancel() + + assert signal.is_set() + + +@pytest.mark.asyncio +async def test_run_cancel_cleans_up_and_allows_reuse(mock_model): + @tool(context=True) + def end_conversation(tool_context: ToolContext[LocalAgent]) -> str: + """End the conversation.""" + tool_context.agent.cancel() + return "Ending conversation" + + mock_model.set_events( + [BidiToolUseBlocksEvent([{"toolUseId": "end", "name": end_conversation.tool_name, "input": {}}])] + ) + agent = BidiAgent(model=mock_model, tools=[end_conversation]) + + for _ in range(2): + input_ = unittest.mock.AsyncMock(spec=InputStream, side_effect=asyncio.Queue().get) + output = unittest.mock.AsyncMock(spec=OutputStream) + await asyncio.wait_for(agent.run(inputs=[input_], outputs=[output]), 2) + + input_.stop.assert_awaited_once() + output.stop.assert_awaited_once() + assert not agent.cancel_signal.is_set() + assert not mock_model._started def test_bidi_agent_tool_context_receives_cancel_signal(mock_model): diff --git a/strands-py/tests/strands/experimental/bidi/agent/test_loop.py b/strands-py/tests/strands/experimental/bidi/agent/test_loop.py index 525930ee7..3acef72c0 100644 --- a/strands-py/tests/strands/experimental/bidi/agent/test_loop.py +++ b/strands-py/tests/strands/experimental/bidi/agent/test_loop.py @@ -1,11 +1,10 @@ import asyncio import unittest.mock -import warnings import pytest import pytest_asyncio -from strands import ToolContext, tool +from strands import LocalAgent, ToolContext, tool from strands.experimental.bidi.agent import BidiAgent from strands.experimental.bidi.agent.loop import _ReaderError from strands.experimental.bidi.hooks import BidiAgentStopEvent, BidiBeforeConnectionRestartEvent @@ -1869,120 +1868,87 @@ async def test_tool_exchanges_remain_paired_when_results_finish_out_of_order(str @pytest.mark.asyncio -async def test_bidi_agent_loop_request_state_initialized_for_tools(loop, agent, agenerator): - """Test that request_state is initialized in invocation_state before tool execution. - - This ensures request_state exists for tools that may need it via invocation_state, - even when invocation_state is not provided by the user. - """ - tool_use = {"toolUseId": "t2", "name": "time_tool", "input": {}} - tool_use_event = BidiToolUseBlocksEvent([tool_use]) - - agent.model.receive = unittest.mock.Mock(return_value=agenerator([tool_use_event])) - - # Start without providing invocation_state - await loop.start() - - tru_events = [] - async for event in loop.receive(): - tru_events.append(event) - if len(tru_events) >= 3: - break - - # Verify tool executed successfully - tool_result_event = tru_events[1] - assert isinstance(tool_result_event, ToolResultEvent) - assert tool_result_event.tool_result["status"] == "success" - - # Verify request_state was initialized in invocation_state - assert "request_state" in loop._invocation_state - assert isinstance(loop._invocation_state["request_state"], dict) - - -@pytest.mark.asyncio +@pytest.mark.parametrize("retry", [False, True]) +@pytest.mark.parametrize("cancel_source", ["tool", "hook"]) @pytest.mark.parametrize("tool_count", [1, 2]) -async def test_bidi_agent_loop_stop_event_loop_flag(agent, agenerator, alist, tool_count): - """Complete the tool group before honoring the stop flag.""" - loop = agent._loop - tool_uses = [{"toolUseId": f"call-{index}", "name": "time_tool", "input": {}} for index in range(tool_count)] - tool_use_event = BidiToolUseBlocksEvent(tool_uses) - agent.model.receive = unittest.mock.Mock(return_value=agenerator([tool_use_event])) - await loop.start(invocation_state={"request_state": {"stop_event_loop": True}}) +async def test_receive_cancel_after_tool(agent, agenerator, alist, retry, cancel_source, tool_count): + @tool(context=True) + def end_conversation(tool_context: ToolContext[LocalAgent]) -> str: + """End the conversation.""" + if cancel_source == "tool": + tool_context.agent.cancel() + return "Ending conversation" - results = [ - {"toolUseId": call["toolUseId"], "status": "success", "content": [{"text": "12:00"}]} for call in tool_uses + def after_tool(event: AfterToolCallEvent[LocalAgent]) -> None: + if cancel_source == "hook": + event.agent.cancel() + event.retry = retry + + agent.tool_registry.register_tool(end_conversation) + agent.hooks.add_callback(AfterToolCallEvent, after_tool) + + tool_uses = [ + {"toolUseId": f"call-{index}", "name": end_conversation.tool_name, "input": {}} for index in range(tool_count) ] - tru_events = await alist(loop.receive()) + tool_results = [ + {"toolUseId": call["toolUseId"], "status": "success", "content": [{"text": "Ending conversation"}]} + for call in tool_uses + ] + tool_use_event = BidiToolUseBlocksEvent(tool_uses) + + agent.model.receive = unittest.mock.Mock(return_value=agenerator([tool_use_event])) + + async with agent: + tru_events = await asyncio.wait_for(alist(agent.receive()), 2) + + exp_result_message = { + "role": "user", + "content": [{"toolResult": tool_result} for tool_result in tool_results], + "metadata": {"custom": {"bidi": {"kind": "tool_result"}}}, + "tracking_id": unittest.mock.ANY, + } exp_events = [ tool_use_event, - *[ToolResultEvent(result) for result in results], - ToolResultMessageEvent( - { - "role": "user", - "content": [{"toolResult": result} for result in results], - "metadata": {"custom": {"bidi": {"kind": "tool_result"}}}, - "tracking_id": unittest.mock.ANY, - } - ), - BidiConnectionStopEvent(connection_id=unittest.mock.ANY, reason="user_request"), + *[ToolResultEvent(tool_result) for tool_result in tool_results], + ToolResultMessageEvent(exp_result_message), + BidiConnectionStopEvent(connection_id="unknown", reason="user_request"), ] assert tru_events == exp_events - agent.model.send.assert_not_awaited() - - -@pytest.mark.asyncio -async def test_bidi_agent_loop_stop_conversation_deprecated_but_works(loop, agent, agenerator): - """Test that stop_conversation tool still works but emits a deprecation warning. - - The stop_conversation tool is deprecated in favor of request_state["stop_event_loop"], - but should continue to work for backward compatibility via the name-based check. - """ - from strands.experimental.bidi.tools import stop_conversation - - agent.tool_registry.register_tool(stop_conversation) - - tool_use = {"toolUseId": "t5", "name": "stop_conversation", "input": {}} - tool_use_event = BidiToolUseBlocksEvent([tool_use]) - - agent.model.receive = unittest.mock.Mock(return_value=agenerator([tool_use_event])) - - await loop.start() - - tru_events = [] - with warnings.catch_warnings(record=True) as caught_warnings: - warnings.simplefilter("always") - async for event in loop.receive(): - tru_events.append(event) - - # Should receive: tool_use_event, tool_result_event, tool_result_message, connection_stop - assert len(tru_events) == 4 - - # Verify tool executed successfully - tool_result_event = tru_events[1] - assert isinstance(tool_result_event, ToolResultEvent) - assert tool_result_event.tool_result["status"] == "success" - assert "Ending conversation" in tool_result_event.tool_result["content"][0]["text"] - - # Verify connection stop event was emitted - connection_stop_event = tru_events[3] - assert isinstance(connection_stop_event, BidiConnectionStopEvent) - assert connection_stop_event["reason"] == "user_request" - - # Verify model.send was NOT called (tool result not sent to model) + assert agent.messages[-1] == exp_result_message agent.model.send.assert_not_called() - # Verify deprecation warnings were emitted (from both the tool itself and the loop name check) - deprecation_warnings = [w for w in caught_warnings if issubclass(w.category, DeprecationWarning)] - assert len(deprecation_warnings) >= 1 - assert any("stop_conversation" in str(w.message).lower() for w in deprecation_warnings) + +@pytest.mark.asyncio +async def test_receive_cancel_pending_until_tool_completes(streaming_agent, alist): + agent = streaming_agent + request = BidiToolUseBlocksEvent([{"toolUseId": "time", "name": "time_tool", "input": {}}]) + audio = BidiAudioDeltaEvent("audio", "pcm", 24000, 1, content_id="audio") + reader = agent.receive() + try: + agent.cancel() + await agent.model.emit(audio) + assert await asyncio.wait_for(anext(reader), 2) == audio + + await agent.model.emit(request) + tru_events = await asyncio.wait_for(alist(reader), 2) + finally: + await reader.aclose() + + assert [type(event) for event in tru_events] == [ + BidiToolUseBlocksEvent, + ToolResultEvent, + ToolResultMessageEvent, + BidiConnectionStopEvent, + ] @pytest.mark.asyncio -@pytest.mark.parametrize("invocation_state", [{}, {"custom_data": "preserved"}]) +@pytest.mark.parametrize( + "invocation_state", [{}, {"custom_data": "preserved"}, {"request_state": {"custom_data": "preserved"}}] +) async def test_tools_share_invocation_state(agent, agenerator, invocation_state): """Tools, hooks, and the caller share state throughout the invocation.""" - invocation_state = dict(invocation_state) - exp_state = {**invocation_state, "call_count": 2, "request_state": {}} + exp_state = {"request_state": {}, **invocation_state, "call_count": 2} tool_states = [] @tool(context=True) diff --git a/strands-py/tests_typing/test_local_agent.py b/strands-py/tests_typing/test_local_agent.py index 7ac68aea0..5651eea32 100644 --- a/strands-py/tests_typing/test_local_agent.py +++ b/strands-py/tests_typing/test_local_agent.py @@ -32,6 +32,7 @@ def agent_tool(tool_context: ToolContext[Agent]) -> str: @tool(context=True) def local_agent_tool(tool_context: ToolContext[LocalAgent]) -> str: assert_type(tool_context.agent, LocalAgent) + tool_context.agent.cancel() return tool_context.agent.name @@ -41,6 +42,7 @@ def before_tool_call(event: BeforeToolCallEvent) -> None: def before_local_tool_call(event: BeforeToolCallEvent[LocalAgent]) -> None: assert_type(event.agent, LocalAgent) + event.agent.cancel() async def after_local_tool_call(event: AfterToolCallEvent[LocalAgent]) -> None: @@ -100,7 +102,6 @@ def register_hooks(agent: Agent, bidi_agent: BidiAgent, local_agent: LocalAgent) def local_agent_excludes_agent_only_members(local_agent: LocalAgent) -> None: local_agent.cleanup() # type: ignore[attr-defined] - local_agent.cancel() # type: ignore[attr-defined] local_agent.conversation_manager # type: ignore[attr-defined] # noqa: B018 local_agent.tool_executor # type: ignore[attr-defined] # noqa: B018