Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
5 changes: 5 additions & 0 deletions temporalio/contrib/strands/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down
22 changes: 20 additions & 2 deletions temporalio/contrib/strands/_temporal_activity_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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"],
Expand All @@ -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
Expand Down
24 changes: 15 additions & 9 deletions temporalio/contrib/strands/_temporal_mcp_tool.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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)
172 changes: 172 additions & 0 deletions tests/contrib/strands/test_tool_errors.py
Original file line number Diff line number Diff line change
@@ -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()
)
Loading