diff --git a/CHANGELOG.md b/CHANGELOG.md index 3de56f2cb..7d68251aa 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -40,6 +40,9 @@ to include examples, links to docs, or any other relevant information. ### Fixed +- `temporalio.contrib.strands` activity and MCP tools now give the model the + Activity's failure message and expose its exception to after-tool hooks. + - Encoding a datetime search attribute without a timezone now raises `ValueError("Timezone must be present on all search attribute dates")` on the typed path, matching the deprecated untyped encoder, instead of sending diff --git a/temporalio/contrib/strands/README.md b/temporalio/contrib/strands/README.md index dacd5c1b7..0941e3172 100644 --- a/temporalio/contrib/strands/README.md +++ b/temporalio/contrib/strands/README.md @@ -405,6 +405,11 @@ Worker( ) ``` +If a tool's activity fails, the agent receives a failed tool result containing +the activity's error message, and `AfterToolCallEvent.exception` contains the +underlying exception. This also applies when an MCP tool's call-tool activity +fails. + ## Hooks Strands' [hook system](https://strandsagents.com/) (`strands.hooks`) lets you subscribe callbacks to events in the agent lifecycle — invocation start/end, model call before/after, tool call before/after, message added. Pass `hooks=[MyHookProvider()]` to `TemporalAgent`: every single-agent hook event fires in workflow context, so deterministic callbacks just work. diff --git a/temporalio/contrib/strands/_temporal_activity_tool.py b/temporalio/contrib/strands/_temporal_activity_tool.py index bb838834e..476b1c457 100644 --- a/temporalio/contrib/strands/_temporal_activity_tool.py +++ b/temporalio/contrib/strands/_temporal_activity_tool.py @@ -9,7 +9,7 @@ from strands.types.tools import AgentTool, ToolGenerator, ToolResult, ToolSpec, ToolUse from temporalio import activity, workflow -from temporalio.exceptions import ActivityError, ApplicationError +from temporalio.exceptions import ActivityError, ApplicationError, FailureError from ._failure_converter import STRANDS_INTERRUPT_TYPE @@ -76,7 +76,8 @@ async def stream( ): yield ToolInterruptEvent(tool_use, [Interrupt(**cause.details[0])]) return - raise + yield _activity_error_event(tool_use["toolUseId"], e) + return yield ToolResultEvent( ToolResult( toolUseId=tool_use["toolUseId"], @@ -86,6 +87,23 @@ async def stream( ) +def _activity_error_event(tool_use_id: str, error: ActivityError) -> ToolResultEvent: + # ActivityError reports a generic message; its cause holds the activity's failure. + cause = error.__cause__ + exception = cause if isinstance(cause, Exception) else error + message = ( + exception.message if isinstance(exception, FailureError) else str(exception) + ) + return ToolResultEvent( + ToolResult( + toolUseId=tool_use_id, + status="error", + content=[{"text": message}], + ), + exception=exception, + ) + + def _to_text(result: Any) -> str: if isinstance(result, str): return result diff --git a/temporalio/contrib/strands/_temporal_mcp_tool.py b/temporalio/contrib/strands/_temporal_mcp_tool.py index 885b1a7e2..2e5dfd061 100644 --- a/temporalio/contrib/strands/_temporal_mcp_tool.py +++ b/temporalio/contrib/strands/_temporal_mcp_tool.py @@ -4,7 +4,9 @@ from strands.types.tools import AgentTool, ToolGenerator, ToolResult, ToolSpec, ToolUse from temporalio import workflow +from temporalio.exceptions import ActivityError +from ._temporal_activity_tool import _activity_error_event from ._temporal_mcp_client import _CallToolArgs, _MCPToolInfo @@ -53,13 +55,17 @@ async def stream( **kwargs: Any, ) -> ToolGenerator: """Execute the tool by dispatching to the per-server call-tool activity.""" - result: ToolResult = await workflow.execute_activity( - f"{self._server}-call-tool", - _CallToolArgs( - tool_name=self._info.name, - arguments=tool_use["input"], - tool_use_id=tool_use["toolUseId"], - ), - **self._options, - ) + try: + result: ToolResult = await workflow.execute_activity( + f"{self._server}-call-tool", + _CallToolArgs( + tool_name=self._info.name, + arguments=tool_use["input"], + tool_use_id=tool_use["toolUseId"], + ), + **self._options, + ) + except ActivityError as error: + yield _activity_error_event(tool_use["toolUseId"], error) + return yield ToolResultEvent(result) diff --git a/tests/contrib/strands/test_tool_errors.py b/tests/contrib/strands/test_tool_errors.py new file mode 100644 index 000000000..5c2b91136 --- /dev/null +++ b/tests/contrib/strands/test_tool_errors.py @@ -0,0 +1,172 @@ +from collections.abc import AsyncIterable +from datetime import timedelta +from typing import Any +from uuid import uuid4 + +import pytest +from strands.hooks import HookProvider, HookRegistry +from strands.hooks.events import AfterToolCallEvent +from strands.types.content import Messages, SystemContentBlock +from strands.types.streaming import StreamEvent +from strands.types.tools import ToolChoice, ToolResult, ToolSpec + +from temporalio import activity, workflow +from temporalio.client import Client +from temporalio.contrib.strands import StrandsPlugin, TemporalAgent +from temporalio.contrib.strands._temporal_mcp_client import _CallToolArgs, _MCPToolInfo +from temporalio.contrib.strands._temporal_mcp_tool import TemporalMCPTool +from temporalio.contrib.strands.workflow import activity_as_tool +from temporalio.exceptions import ApplicationError +from temporalio.worker import Replayer, Worker +from tests.contrib.strands.mock_model import MockModel + +_MODEL_TOOL_RESULTS: list[ToolResult] = [] + + +class _RecordingModel(MockModel): + async def stream( + self, + messages: Messages, + tool_specs: list[ToolSpec] | None = None, + system_prompt: str | None = None, + *, + tool_choice: ToolChoice | None = None, + system_prompt_content: list[SystemContentBlock] | None = None, + invocation_state: dict[str, Any] | None = None, + **kwargs: Any, + ) -> AsyncIterable[StreamEvent]: + for message in messages: + for content in message["content"]: + if "toolResult" in content: + _MODEL_TOOL_RESULTS.append(content["toolResult"]) + async for event in super().stream( + messages, + tool_specs, + system_prompt, + tool_choice=tool_choice, + system_prompt_content=system_prompt_content, + invocation_state=invocation_state, + **kwargs, + ): + yield event + + +class _ToolResultHook(HookProvider): + def __init__(self) -> None: + self.failures: list[dict[str, str | None]] = [] + + def register_hooks(self, registry: HookRegistry, **kwargs: object) -> None: + registry.add_callback(AfterToolCallEvent, self._record) + + def _record(self, event: AfterToolCallEvent) -> None: + if event.result["status"] == "error": + error = event.exception + self.failures.append( + { + "text": event.result["content"][0].get("text"), + "exception_type": type(error).__name__ if error else None, + "exception_message": error.message + if isinstance(error, ApplicationError) + else None, + } + ) + + +@activity.defn +async def failing_activity_tool(location: str) -> None: + raise ApplicationError(f"Unknown location: {location}", non_retryable=True) + + +@activity.defn(name="failing-mcp-call-tool") +async def failing_mcp_call_tool(_args: _CallToolArgs) -> None: + raise ApplicationError("MCP server unavailable", non_retryable=True) + + +@workflow.defn +class _FailingToolWorkflow: + @workflow.run + async def run(self, kind: str) -> list[dict[str, str | None]]: + if kind == "activity": + tools = [ + activity_as_tool( + failing_activity_tool, + start_to_close_timeout=timedelta(seconds=15), + ) + ] + else: + tools = [ + TemporalMCPTool( + "failing-mcp", + _MCPToolInfo("list_files", "List files", {"type": "object"}), + {"start_to_close_timeout": timedelta(seconds=15)}, + ) + ] + hook = _ToolResultHook() + agent = TemporalAgent( + model="mock", + start_to_close_timeout=timedelta(seconds=15), + tools=tools, + hooks=[hook], + ) + await agent.invoke_async("Use the tool") + return hook.failures + + +@pytest.mark.parametrize( + ("kind", "tool_name", "tool_input", "message"), + [ + ( + "activity", + "failing_activity_tool", + {"location": "Atlantis"}, + "Unknown location: Atlantis", + ), + ("mcp", "list_files", {"path": "/"}, "MCP server unavailable"), + ], +) +async def test_failed_activity_tool_reports_cause( + client: Client, + kind: str, + tool_name: str, + tool_input: dict[str, Any], + message: str, +) -> None: + _MODEL_TOOL_RESULTS.clear() + task_queue = f"test_failed_tool-{uuid4()}" + plugin = StrandsPlugin( + models={ + "mock": lambda: _RecordingModel( + [{"name": tool_name, "input": tool_input}, "Done!"] + ) + } + ) + + async with Worker( + client, + task_queue=task_queue, + workflows=[_FailingToolWorkflow], + activities=[failing_activity_tool, failing_mcp_call_tool], + plugins=[plugin], + max_cached_workflows=0, + ): + handle = await client.start_workflow( + _FailingToolWorkflow.run, + kind, + id=f"test_failed_tool_{uuid4()}", + task_queue=task_queue, + ) + assert await handle.result() == [ + { + "text": message, + "exception_type": "ApplicationError", + "exception_message": message, + } + ] + + assert len(_MODEL_TOOL_RESULTS) == 1 + assert _MODEL_TOOL_RESULTS[0]["status"] == "error" + assert _MODEL_TOOL_RESULTS[0]["content"] == [{"text": message}] + + await Replayer(workflows=[_FailingToolWorkflow], plugins=[plugin]).replay_workflow( + await handle.fetch_history() + )