feat(bidi): support cancellation from custom tools (#4664)

This commit is contained in:
Patrick Gray
2026-09-29 12:06:47 -04:00
committed by GitHub
parent 43479b8da2
commit 3e5e89e265
18 changed files with 160 additions and 228 deletions
@@ -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(
+4
View File
@@ -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)
+2 -1
View File
@@ -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