mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
feat(bidi): support cancellation from custom tools (#4664)
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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()])
|
||||
|
||||
|
||||
@@ -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": "<GOOGLE_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()])
|
||||
|
||||
@@ -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="<OPENAI_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()])
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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"]
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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=<True> | 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
|
||||
|
||||
|
||||
@@ -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"]
|
||||
@@ -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"
|
||||
@@ -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(
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user