Files
Alberto Schiabel 53c024a9c8 fix(py): keep local results when remote multi-execute transport fails (#4386)
## Summary

Follow-up to #4310. The TypeScript SDK (since #3800) catches a thrown
backend call for the remote half of a mixed
`COMPOSIO_MULTI_EXECUTE_TOOL` batch and turns it into one failure entry
per remote slug, so completed local results are not lost. The Python SDK
still let the exception escape `_route_multi_execute`, discarding every
local result that had already run.

This ports the TypeScript behavior so both SDKs return the same shape on
a remote transport failure.

Fixes #

## Changes

- Catch the remote future's exception in `_route_multi_execute` and keep
`str(error)`, falling back to `Remote tool execution failed` when the
message is empty (same fallback as TS).
- Synthesize `{response: {successful: False, data: {}, error},
tool_slug, error}` for each remote index and merge them in original
request order.
- Recompute `total_count` / `success_count` / `error_count` on transport
failure, as TS does.
- Add two regression tests mirroring the TS cases in
`customToolRouting.test.ts`: local results preserved with per-tool
remote errors, and the empty-message fallback.

## Type of change

- [x] Bug fix
- [ ] New feature
- [ ] Refactor/Chore
- [ ] Documentation
- [ ] Breaking change

## How Has This Been Tested?

- `uv run --locked --group dev pytest tests/test_custom_tools.py -q`: 87
passed.
- `uv run --locked --group dev nox -s chk`: ruff and mypy clean.
- Without the source change, the new
`test_remote_transport_failure_keeps_local_results` raises
`RuntimeError: remote unavailable` out of `_route_multi_execute`.

## Checklist

- [x] I have read the Code of Conduct and this PR adheres to it
- [x] I ran linters/tests locally and they passed
- [ ] I updated documentation as needed
- [x] I added tests or explain why not applicable
- [ ] I added a changeset if this change affects published packages
(Python does not use Changesets)

## Additional context

TypeScript reference:
`ts/packages/core/src/models/ToolRouterSession.ts`, the
`remoteErrorMessage` branch, and the test "should preserve successful
local results when remote transport fails".

https://claude.ai/code/session_01PAXMbiZd3qPoJ8Z9uPvEAb
EOF -R ComposioHQ/composio
2026-09-08 18:11:25 +02:00

1430 lines
53 KiB
Python

"""Tests for custom tools in tool router sessions.
Covers:
- Decorator API (inference, overrides, validation)
- Toolkit builder with @toolkit.tool() decorator
- Slug prefixing and collision detection
- Serialization to API format
- Custom tool execution (validation, defaults, errors)
- Routing map building and lookup
- SessionContextImpl sibling routing
- ToolRouterSession integration (execute routing, custom_tools(), custom_toolkits())
- Multi-execute routing (all-local, all-remote, mixed, parallel)
"""
from dataclasses import replace
from unittest.mock import MagicMock
import pytest
from composio_client import omit
from composio_client.types.tool_router.session_execute_response import (
SessionExecuteResponse,
)
from composio_client.types.tool_router.session_proxy_execute_response import (
BinaryData,
SessionProxyExecuteResponse,
)
from pydantic import BaseModel, Field
from composio.core.models.custom_tool import (
ExperimentalToolkit,
assert_no_custom_tool_slugs_in_preload,
build_custom_tools_map,
build_custom_tools_map_from_response,
serialize_custom_toolkits,
serialize_custom_tools,
)
from composio.core.models.custom_tool_execution import (
assert_unambiguous_custom_tool_slug,
execute_custom_tool,
find_custom_tool,
)
from composio.core.models.experimental import ExperimentalAPI
from composio.core.models.session_context import SessionContextImpl
from composio.core.models.tool_router_session import ToolRouterSession
from composio.exceptions import ValidationError
# ────────────────────────────────────────────────────────────────
# Fixtures
# ────────────────────────────────────────────────────────────────
exp = ExperimentalAPI()
class GrepInput(BaseModel):
pattern: str = Field(description="Pattern to search for")
path: str = Field(default=".", description="File path")
class EmailInput(BaseModel):
to: str = Field(description="Recipient email")
subject: str = Field(description="Subject")
class SetRoleInput(BaseModel):
user_id: str = Field(description="User ID")
role: str = Field(default="viewer", description="New role")
@pytest.fixture
def grep_tool():
@exp.tool()
def grep(input: GrepInput, ctx):
"""Search for patterns in files."""
return {"matches": [input.pattern], "path": input.path}
return grep
@pytest.fixture
def duplicate_slug_tools():
@exp.tool(slug="GREP")
def alpha_grep(input: GrepInput, ctx):
"""Search Alpha."""
return {"toolkit": "alpha", "matches": [input.pattern]}
@exp.tool(slug="GREP")
def beta_grep(input: GrepInput, ctx):
"""Search Beta."""
return {"toolkit": "beta", "matches": [input.pattern]}
return alpha_grep, beta_grep
@pytest.fixture
def email_tool():
@exp.tool(extends_toolkit="gmail")
def get_emails(input: EmailInput, ctx):
"""Get emails from Gmail."""
return {"emails": []}
return get_emails
@pytest.fixture
def set_role_tool():
@exp.tool()
def set_role(input: SetRoleInput, ctx):
"""Set a user's role."""
return {"user_id": input.user_id, "role": input.role, "updated": True}
return set_role
@pytest.fixture
def role_toolkit(set_role_tool):
tk = ExperimentalToolkit(
slug="ROLE_MANAGER", name="Role Manager", description="Manage user roles"
)
tk._tools.append(set_role_tool)
return tk
class MockSessionContext:
user_id = "test-user"
def execute(self, tool_slug, arguments):
return {"data": {}, "error": None, "successful": True}
def proxy_execute(self, **kwargs):
return {"status": 200, "data": {}, "headers": {}}
# ────────────────────────────────────────────────────────────────
# Decorator API tests
# ────────────────────────────────────────────────────────────────
class TestDecoratorTool:
def test_bare_decorator(self):
@exp.tool
def my_tool(input: GrepInput, ctx):
"""My tool description."""
return {}
assert my_tool.slug == "MY_TOOL"
assert my_tool.name == "My Tool"
assert my_tool.description == "My tool description."
assert my_tool.input_schema["type"] == "object"
assert "pattern" in my_tool.input_schema["properties"]
def test_decorator_with_parens(self):
@exp.tool()
def search(input: GrepInput, ctx):
"""Search stuff."""
return {}
assert search.slug == "SEARCH"
def test_decorator_with_extends_toolkit(self):
@exp.tool(extends_toolkit="gmail")
def fetch_mail(input: EmailInput, ctx):
"""Fetch mail."""
return {}
assert fetch_mail.extends_toolkit == "gmail"
def test_decorator_with_preload(self):
@exp.tool(preload=True)
def fetch_context(input: GrepInput, ctx):
"""Fetch context."""
return {}
assert fetch_context.preload is True
def test_decorator_explicit_overrides(self):
@exp.tool(slug="CUSTOM", name="Custom Name", description="Custom desc")
def whatever(input: GrepInput, ctx):
"""Original docstring."""
return {}
assert whatever.slug == "CUSTOM"
assert whatever.name == "Custom Name"
assert whatever.description == "Custom desc"
def test_infers_from_function_name(self):
@exp.tool()
def search_user_by_email(input: GrepInput, ctx):
"""Find a user."""
return {}
assert search_user_by_email.slug == "SEARCH_USER_BY_EMAIL"
assert search_user_by_email.name == "Search User By Email"
def test_infers_description_from_docstring(self):
@exp.tool()
def my_func(input: GrepInput, ctx):
"""Trimmed docstring."""
return {}
assert my_func.description == "Trimmed docstring."
def test_multiline_docstring(self):
@exp.tool()
def multi(input: GrepInput, ctx):
"""Search for users by email.
This tool searches the internal database for users
matching the given pattern. Returns a list of matches.
Supports wildcards.
"""
return {}
assert multi.description.startswith("Search for users by email.")
# inspect.cleandoc strips indentation
assert "\n " not in multi.description
assert "Supports wildcards." in multi.description
def test_indented_docstring_cleaned(self):
@exp.tool()
def indented(input: GrepInput, ctx):
"""
Leading and trailing whitespace removed.
Indentation normalized.
"""
return {}
assert indented.description == (
"Leading and trailing whitespace removed.\nIndentation normalized."
)
def test_missing_docstring_raises(self):
with pytest.raises(ValidationError, match="description is required"):
@exp.tool()
def no_doc(input: GrepInput, ctx):
return {}
def test_missing_basemodel_annotation_raises(self):
with pytest.raises(ValidationError, match="BaseModel subclass"):
@exp.tool()
def bad(x: str, ctx):
"""Bad tool."""
return {}
def test_no_params_rejected(self):
with pytest.raises(ValidationError, match="at least one parameter"):
@exp.tool()
def no_params():
"""No params."""
return {}
def test_too_many_params_rejected(self):
with pytest.raises(ValidationError, match="at most 2"):
@exp.tool()
def three_params(input: GrepInput, ctx, extra):
"""Three params."""
return {}
def test_first_param_must_be_basemodel(self):
with pytest.raises(ValidationError, match="first parameter"):
@exp.tool()
def swapped(ctx, input: GrepInput):
"""Swapped order — ctx first, input second."""
return {}
def test_future_annotations_are_resolved(self):
namespace = {"exp": exp}
exec(
"""
from __future__ import annotations
from pydantic import BaseModel, Field
class FutureInput(BaseModel):
pattern: str = Field(description="Pattern to search for")
@exp.tool()
def future_tool(input: FutureInput, ctx):
'''Resolved from postponed annotations.'''
return {}
""",
namespace,
)
future_tool = namespace["future_tool"]
future_input = namespace["FutureInput"]
assert future_tool.input_params is future_input
def test_local_string_annotations_are_resolved(self):
def build_tool():
class LocalInput(BaseModel):
pattern: str = Field(description="Pattern to search for")
@exp.tool
def nested_search(input: "LocalInput", ctx):
"""Resolved from local string annotations."""
return {}
return nested_search, LocalInput
local_tool, local_input = build_tool()
assert local_tool.input_params is local_input
def test_async_function_rejected(self):
with pytest.raises(ValidationError, match="async"):
@exp.tool()
async def async_tool(input: GrepInput, ctx):
"""Async tool."""
return {}
def test_async_single_param_rejected(self):
"""Async must be caught even when wrapper hides it (single-param path)."""
with pytest.raises(ValidationError, match="async"):
@exp.tool()
async def async_single(input: GrepInput):
"""Async single param."""
return {}
def test_root_model_rejected(self):
from pydantic import RootModel
class ListInput(RootModel[list[int]]):
pass
with pytest.raises(ValidationError, match="RootModel"):
@exp.tool()
def bad(input: ListInput, ctx):
"""Bad tool."""
return {}
def test_single_param_function(self):
"""Function with only input param (no ctx) should work."""
@exp.tool()
def simple(input: GrepInput):
"""Simple tool."""
return {"pattern": input.pattern}
m = build_custom_tools_map([simple])
entry = find_custom_tool(m, "SIMPLE")
result = execute_custom_tool(entry, {"pattern": "test"}, MockSessionContext()) # type: ignore
assert result["successful"] is True
assert result["data"]["pattern"] == "test"
def test_slug_validation(self):
with pytest.raises(ValidationError, match="LOCAL_"):
@exp.tool(slug="LOCAL_BAD")
def bad(input: GrepInput, ctx):
"""Bad tool."""
return {}
def test_creates_frozen_custom_tool(self):
@exp.tool()
def frozen(input: GrepInput, ctx):
"""Frozen tool."""
return {}
with pytest.raises(AttributeError):
frozen.slug = "NEW" # type: ignore
def test_input_schema_includes_defaults(self, grep_tool):
assert "pattern" in grep_tool.input_schema.get("required", [])
assert "path" not in grep_tool.input_schema.get("required", [])
class TestToolkitBuilder:
def test_toolkit_with_decorator(self):
tk = ExperimentalToolkit(slug="MY_TK", name="My TK", description="Desc")
@tk.tool()
def tool_a(input: GrepInput, ctx):
"""Tool A."""
return {}
@tk.tool()
def tool_b(input: GrepInput, ctx):
"""Tool B."""
return {}
assert tk.slug == "MY_TK"
assert len(tk.tools) == 2
assert tk.tools[0].slug == "TOOL_A"
assert tk.tools[1].slug == "TOOL_B"
def test_toolkit_bare_decorator(self):
tk = ExperimentalToolkit(slug="TK2", name="TK2", description="Desc")
@tk.tool
def tool_c(input: GrepInput, ctx):
"""Tool C."""
return {}
assert len(tk.tools) == 1
def test_toolkit_slug_validation(self):
with pytest.raises(ValidationError, match="LOCAL_"):
ExperimentalToolkit(slug="LOCAL_BAD", name="Bad", description="Desc")
def test_toolkit_name_required(self):
with pytest.raises(ValidationError, match="name is required"):
ExperimentalToolkit(slug="TK", name="", description="Desc")
def test_toolkit_via_experimental_api(self):
tk = exp.Toolkit(slug="API_TK", name="API TK", description="Via API")
assert isinstance(tk, ExperimentalToolkit)
assert tk.slug == "API_TK"
# ────────────────────────────────────────────────────────────────
# Serialization tests
# ────────────────────────────────────────────────────────────────
class TestSerialization:
def test_serialize_standalone_tool(self, grep_tool):
result = serialize_custom_tools([grep_tool])
assert len(result) == 1
assert result[0]["slug"] == "GREP"
assert result[0]["input_schema"]["type"] == "object"
assert "extends_toolkit" not in result[0]
def test_serialize_extension_tool(self, email_tool):
result = serialize_custom_tools([email_tool])
assert result[0]["extends_toolkit"] == "gmail"
def test_serialize_custom_tool_preload_hint(self, grep_tool):
preloaded = replace(grep_tool, preload=True)
result = serialize_custom_tools([preloaded])
assert result[0]["preload"] is True
def test_serialize_custom_tool_omits_redundant_preload_false(self, grep_tool):
search_only = replace(grep_tool, preload=False)
result = serialize_custom_tools([search_only])
assert "preload" not in result[0]
def test_serialize_toolkit(self, role_toolkit):
result = serialize_custom_toolkits([role_toolkit])
assert len(result) == 1
assert result[0]["slug"] == "ROLE_MANAGER"
assert len(result[0]["tools"]) == 1
def test_serialize_toolkit_preload_hint(self, role_toolkit):
role_toolkit.preload = True
result = serialize_custom_toolkits([role_toolkit])
assert result[0]["preload"] is True
assert result[0]["tools"][0]["preload"] is True
# ────────────────────────────────────────────────────────────────
# Routing map tests
# ────────────────────────────────────────────────────────────────
class TestCustomToolsMap:
def test_build_map_standalone(self, grep_tool):
m = build_custom_tools_map([grep_tool])
assert "LOCAL_GREP" in m.by_final_slug
def test_build_map_extension(self, email_tool):
m = build_custom_tools_map([email_tool])
assert "LOCAL_GMAIL_GET_EMAILS" in m.by_final_slug
def test_build_map_toolkit(self, role_toolkit):
m = build_custom_tools_map([], [role_toolkit])
assert "LOCAL_ROLE_MANAGER_SET_ROLE" in m.by_final_slug
def test_allows_same_slug_in_different_toolkits(self, duplicate_slug_tools):
alpha_tool, beta_tool = duplicate_slug_tools
alpha = ExperimentalToolkit(
slug="ALPHA", name="Alpha", description="Alpha tools"
)
alpha._tools.append(alpha_tool)
beta = ExperimentalToolkit(slug="BETA", name="Beta", description="Beta tools")
beta._tools.append(beta_tool)
m = build_custom_tools_map([], [alpha, beta])
assert set(m.by_final_slug) == {"LOCAL_ALPHA_GREP", "LOCAL_BETA_GREP"}
assert "GREP" not in m.by_original_slug
assert m.ambiguous_original_slugs == {"GREP"}
def test_ambiguous_original_slug_requires_final_slug(self, duplicate_slug_tools):
alpha_tool, beta_tool = duplicate_slug_tools
alpha = ExperimentalToolkit(
slug="ALPHA", name="Alpha", description="Alpha tools"
)
alpha._tools.append(alpha_tool)
beta = ExperimentalToolkit(slug="BETA", name="Beta", description="Beta tools")
beta._tools.append(beta_tool)
m = build_custom_tools_map([], [alpha, beta])
assert find_custom_tool(m, "LOCAL_ALPHA_GREP").toolkit == "ALPHA"
assert find_custom_tool(m, "LOCAL_BETA_GREP").toolkit == "BETA"
assert find_custom_tool(m, "GREP") is None
with pytest.raises(ValidationError, match="Ambiguous custom tool slug"):
assert_unambiguous_custom_tool_slug(m, "GREP")
def test_preload_rejects_ambiguous_original_slug(self, duplicate_slug_tools):
alpha_tool, beta_tool = duplicate_slug_tools
alpha = ExperimentalToolkit(
slug="ALPHA", name="Alpha", description="Alpha tools"
)
alpha._tools.append(alpha_tool)
beta = ExperimentalToolkit(slug="BETA", name="Beta", description="Beta tools")
beta._tools.append(beta_tool)
m = build_custom_tools_map([], [alpha, beta])
with pytest.raises(ValidationError, match="not supported in preload.tools"):
assert_no_custom_tool_slugs_in_preload(["grep"], m)
def test_collision_detection(self, grep_tool):
with pytest.raises(ValidationError, match="collision"):
build_custom_tools_map([grep_tool, grep_tool])
def test_rejects_standalone_and_toolkit_same_slug(self, duplicate_slug_tools):
alpha_tool, beta_tool = duplicate_slug_tools
beta = ExperimentalToolkit(slug="BETA", name="Beta", description="Beta tools")
beta._tools.append(beta_tool)
with pytest.raises(ValidationError, match="collision"):
build_custom_tools_map([alpha_tool], [beta])
class TestFindCustomTool:
def test_find_by_final_slug(self, grep_tool):
m = build_custom_tools_map([grep_tool])
assert find_custom_tool(m, "LOCAL_GREP") is not None
def test_find_case_insensitive(self, grep_tool):
m = build_custom_tools_map([grep_tool])
assert find_custom_tool(m, "grep") is not None
def test_find_nonexistent(self, grep_tool):
m = build_custom_tools_map([grep_tool])
assert find_custom_tool(m, "NONEXISTENT") is None
def test_find_none_map(self):
assert find_custom_tool(None, "GREP") is None
class TestBuildMapFromResponse:
def test_builds_from_response(self, grep_tool, email_tool, role_toolkit):
mock_exp = MagicMock()
mock_ct1 = MagicMock(
slug="LOCAL_GREP", original_slug="GREP", extends_toolkit=None
)
mock_ct2 = MagicMock(
slug="LOCAL_GMAIL_GET_EMAILS",
original_slug="GET_EMAILS",
extends_toolkit="gmail",
)
mock_exp.custom_tools = [mock_ct1, mock_ct2]
mock_ctk = MagicMock(slug="ROLE_MANAGER")
mock_ctk.tools = [
MagicMock(slug="LOCAL_ROLE_MANAGER_SET_ROLE", original_slug="SET_ROLE")
]
mock_exp.custom_toolkits = [mock_ctk]
m = build_custom_tools_map_from_response(
[grep_tool, email_tool], [role_toolkit], mock_exp
)
assert "LOCAL_GREP" in m.by_final_slug
assert "LOCAL_GMAIL_GET_EMAILS" in m.by_final_slug
assert "LOCAL_ROLE_MANAGER_SET_ROLE" in m.by_final_slug
def test_allows_same_slug_in_different_toolkits(self, duplicate_slug_tools):
alpha_tool, beta_tool = duplicate_slug_tools
alpha = ExperimentalToolkit(
slug="ALPHA", name="Alpha", description="Alpha tools"
)
alpha._tools.append(alpha_tool)
beta = ExperimentalToolkit(slug="BETA", name="Beta", description="Beta tools")
beta._tools.append(beta_tool)
mock_exp = MagicMock()
mock_exp.custom_tools = []
mock_exp.custom_toolkits = [
MagicMock(
slug="ALPHA",
tools=[MagicMock(slug="LOCAL_ALPHA_GREP", original_slug="GREP")],
),
MagicMock(
slug="BETA",
tools=[MagicMock(slug="LOCAL_BETA_GREP", original_slug="GREP")],
),
]
m = build_custom_tools_map_from_response([], [alpha, beta], mock_exp)
assert set(m.by_final_slug) == {"LOCAL_ALPHA_GREP", "LOCAL_BETA_GREP"}
assert m.by_final_slug["LOCAL_ALPHA_GREP"].toolkit == "ALPHA"
assert m.by_final_slug["LOCAL_BETA_GREP"].toolkit == "BETA"
assert m.ambiguous_original_slugs == {"GREP"}
def _duplicate_toolkits(self, duplicate_slug_tools):
alpha_tool, beta_tool = duplicate_slug_tools
alpha = ExperimentalToolkit(
slug="ALPHA", name="Alpha", description="Alpha tools"
)
alpha._tools.append(alpha_tool)
beta = ExperimentalToolkit(slug="BETA", name="Beta", description="Beta tools")
beta._tools.append(beta_tool)
return alpha, beta
def test_rejects_toolkit_child_without_exact_match(self, duplicate_slug_tools):
alpha, beta = self._duplicate_toolkits(duplicate_slug_tools)
mock_exp = MagicMock()
mock_exp.custom_tools = []
mock_exp.custom_toolkits = [
MagicMock(
slug="GAMMA",
tools=[MagicMock(slug="LOCAL_GAMMA_GREP", original_slug="GREP")],
),
]
with pytest.raises(ValidationError, match="no exact local match"):
build_custom_tools_map_from_response([], [alpha, beta], mock_exp)
def test_never_binds_toolkit_child_to_another_toolkit(self, duplicate_slug_tools):
alpha_tool, _ = duplicate_slug_tools
alpha = ExperimentalToolkit(
slug="ALPHA", name="Alpha", description="Alpha tools"
)
alpha._tools.append(alpha_tool)
mock_exp = MagicMock()
mock_exp.custom_tools = []
mock_exp.custom_toolkits = [
MagicMock(
slug="BETA",
tools=[MagicMock(slug="LOCAL_BETA_GREP", original_slug="GREP")],
),
]
with pytest.raises(ValidationError, match='for toolkit "BETA"'):
build_custom_tools_map_from_response([], [alpha], mock_exp)
def test_standalone_falls_back_to_unique_bare_match(self, email_tool):
mock_exp = MagicMock()
mock_exp.custom_tools = [
MagicMock(
slug="LOCAL_GMAIL_GET_EMAILS",
original_slug="GET_EMAILS",
extends_toolkit=None,
)
]
mock_exp.custom_toolkits = []
m = build_custom_tools_map_from_response([email_tool], None, mock_exp)
assert m.by_final_slug["LOCAL_GMAIL_GET_EMAILS"].toolkit == "gmail"
def test_skips_response_tools_not_defined_locally(self, grep_tool):
mock_exp = MagicMock()
mock_exp.custom_tools = [
MagicMock(slug="LOCAL_GREP", original_slug="GREP", extends_toolkit=None),
MagicMock(slug="LOCAL_STALE", original_slug="STALE", extends_toolkit=None),
]
mock_exp.custom_toolkits = []
m = build_custom_tools_map_from_response([grep_tool], None, mock_exp)
assert set(m.by_final_slug) == {"LOCAL_GREP"}
def test_ambiguity_follows_local_definitions(self, duplicate_slug_tools):
alpha, beta = self._duplicate_toolkits(duplicate_slug_tools)
mock_exp = MagicMock()
mock_exp.custom_tools = []
mock_exp.custom_toolkits = [
MagicMock(
slug="ALPHA",
tools=[MagicMock(slug="LOCAL_ALPHA_GREP", original_slug="GREP")],
),
]
m = build_custom_tools_map_from_response([], [alpha, beta], mock_exp)
assert set(m.by_final_slug) == {"LOCAL_ALPHA_GREP"}
assert "GREP" not in m.by_original_slug
assert m.ambiguous_original_slugs == {"GREP"}
assert find_custom_tool(m, "GREP") is None
def test_rejects_duplicate_qualified_response_entries(self, duplicate_slug_tools):
alpha, beta = self._duplicate_toolkits(duplicate_slug_tools)
mock_exp = MagicMock()
mock_exp.custom_tools = []
mock_exp.custom_toolkits = [
MagicMock(
slug="ALPHA",
tools=[
MagicMock(slug="LOCAL_ALPHA_GREP", original_slug="GREP"),
MagicMock(slug="LOCAL_ALPHA_GREP_2", original_slug="GREP"),
],
),
]
with pytest.raises(ValidationError, match="already registered for toolkit"):
build_custom_tools_map_from_response([], [alpha, beta], mock_exp)
# ────────────────────────────────────────────────────────────────
# Execution tests
# ────────────────────────────────────────────────────────────────
class TestExecuteCustomTool:
def test_successful_execution(self, grep_tool):
m = build_custom_tools_map([grep_tool])
entry = find_custom_tool(m, "GREP")
result = execute_custom_tool(entry, {"pattern": "hello"}, MockSessionContext()) # type: ignore
assert result["successful"] is True
assert result["data"]["matches"] == ["hello"]
def test_validation_failure(self, grep_tool):
m = build_custom_tools_map([grep_tool])
entry = find_custom_tool(m, "GREP")
result = execute_custom_tool(entry, {}, MockSessionContext()) # type: ignore
assert result["successful"] is False
def test_execute_error(self):
@exp.tool()
def bad(input: GrepInput, ctx):
"""Bad."""
raise RuntimeError("boom")
m = build_custom_tools_map([bad])
entry = find_custom_tool(m, "BAD")
result = execute_custom_tool(entry, {"pattern": "x"}, MockSessionContext()) # type: ignore
assert result["successful"] is False
assert "boom" in result["error"]
# ────────────────────────────────────────────────────────────────
# SessionContextImpl tests
# ────────────────────────────────────────────────────────────────
class TestSessionContextImpl:
def test_sibling_routing(self, grep_tool):
m = build_custom_tools_map([grep_tool])
ctx = SessionContextImpl(
client=MagicMock(), user_id="u", session_id="s", custom_tools_map=m
)
result = ctx.execute("GREP", {"pattern": "test"})
assert isinstance(result, SessionExecuteResponse)
assert result.error is None
assert result.log_id == ""
assert result.data["matches"] == ["test"]
def test_sibling_routing_rejects_ambiguous_slug(self, duplicate_slug_tools):
alpha_tool, beta_tool = duplicate_slug_tools
alpha = ExperimentalToolkit(
slug="ALPHA", name="Alpha", description="Alpha tools"
)
alpha._tools.append(alpha_tool)
beta = ExperimentalToolkit(slug="BETA", name="Beta", description="Beta tools")
beta._tools.append(beta_tool)
m = build_custom_tools_map([], [alpha, beta])
mock_client = MagicMock()
ctx = SessionContextImpl(
client=mock_client, user_id="u", session_id="s", custom_tools_map=m
)
with pytest.raises(ValidationError, match="Ambiguous custom tool slug"):
ctx.execute("GREP", {"pattern": "x"})
mock_client.tool_router.session.execute.assert_not_called()
result = ctx.execute("LOCAL_BETA_GREP", {"pattern": "x"})
assert result.data["toolkit"] == "beta"
def test_remote_fallback(self, grep_tool):
m = build_custom_tools_map([grep_tool])
mock_client = MagicMock()
mock_client.tool_router.session.execute.return_value = SessionExecuteResponse(
data={"remote": True}, error=None, log_id="log_123"
)
ctx = SessionContextImpl(
client=mock_client, user_id="u", session_id="s", custom_tools_map=m
)
ctx.execute("NONEXISTENT", {"arg": "val"})
mock_client.tool_router.session.execute.assert_called_once_with(
session_id="s",
tool_slug="NONEXISTENT",
arguments={"arg": "val"},
experimental=omit,
)
def test_remote_fallback_passes_inline_custom_tools(self):
mock_client = MagicMock()
mock_client.tool_router.session.execute.return_value = SessionExecuteResponse(
data={"remote": True}, error=None, log_id="log_123"
)
inline_payload = {
"custom_tools": [
{
"slug": "GREP",
"name": "Grep",
"description": "Search local text",
"input_schema": {"type": "object", "properties": {}},
}
]
}
ctx = SessionContextImpl(
client=mock_client,
user_id="u",
session_id="s",
inline_custom_tools_payload=inline_payload,
)
ctx.execute("GMAIL_SEND_EMAIL", {"to": "a@b.com"})
mock_client.tool_router.session.execute.assert_called_once_with(
session_id="s",
tool_slug="GMAIL_SEND_EMAIL",
arguments={"to": "a@b.com"},
experimental=inline_payload,
)
def test_proxy_execute(self):
mock_client = MagicMock()
mock_client.tool_router.session.proxy_execute.return_value = (
SessionProxyExecuteResponse(
status=200, data={"ok": True}, headers={}, binary_data=None
)
)
ctx = SessionContextImpl(client=mock_client, user_id="u", session_id="s")
result = ctx.proxy_execute(
toolkit="gmail", endpoint="https://example.com", method="GET"
)
assert result == {"status": 200, "data": {"ok": True}, "headers": {}}
assert "binary_data" not in result
def test_proxy_execute_narrows_status_to_int(self):
"""The generated model types ``status`` as ``float`` and pydantic coerces.
Equality alone cannot catch the leak, because ``200 == 200.0``. Only the
type assertion distinguishes ``200`` from the ``200.0`` a raw read returns.
"""
mock_client = MagicMock()
mock_client.tool_router.session.proxy_execute.return_value = (
SessionProxyExecuteResponse(
status=200, data=None, headers=None, binary_data=None
)
)
ctx = SessionContextImpl(client=mock_client, user_id="u", session_id="s")
result = ctx.proxy_execute(
toolkit="gmail", endpoint="https://example.com", method="GET"
)
assert isinstance(result["status"], int)
assert result == {"status": 200, "data": None, "headers": None}
def test_proxy_execute_projects_binary_data(self):
mock_client = MagicMock()
mock_client.tool_router.session.proxy_execute.return_value = (
SessionProxyExecuteResponse(
status=200,
data={"ok": True},
headers={"content-type": "application/pdf"},
binary_data=BinaryData(
content_type="application/pdf",
size=123,
url="https://example.com/file.pdf",
expires_at="2026-08-18T00:00:00Z",
),
)
)
ctx = SessionContextImpl(client=mock_client, user_id="u", session_id="s")
result = ctx.proxy_execute(
toolkit="gmail", endpoint="https://example.com", method="GET"
)
assert result == {
"status": 200,
"data": {"ok": True},
"headers": {"content-type": "application/pdf"},
"binary_data": {
"content_type": "application/pdf",
"size": 123,
"url": "https://example.com/file.pdf",
"expires_at": "2026-08-18T00:00:00Z",
},
}
assert isinstance(result["binary_data"]["size"], int)
def test_proxy_execute_binary_data_without_expiry(self):
"""``expires_at`` is optional on the generated model; the key stays present."""
mock_client = MagicMock()
mock_client.tool_router.session.proxy_execute.return_value = (
SessionProxyExecuteResponse(
status=200,
data=None,
headers=None,
binary_data=BinaryData(
content_type="image/png",
size=7,
url="https://example.com/file.png",
),
)
)
ctx = SessionContextImpl(client=mock_client, user_id="u", session_id="s")
result = ctx.proxy_execute(
toolkit="gmail", endpoint="https://example.com", method="GET"
)
assert result["binary_data"] == {
"content_type": "image/png",
"size": 7,
"url": "https://example.com/file.png",
"expires_at": None,
}
# ────────────────────────────────────────────────────────────────
# ToolRouterSession integration
# ────────────────────────────────────────────────────────────────
@pytest.fixture
def mock_session_deps(grep_tool, email_tool, role_toolkit):
return {
"client": MagicMock(),
"provider": MagicMock(),
"experimental": MagicMock(),
"tools_map": build_custom_tools_map([grep_tool, email_tool], [role_toolkit]),
}
def _session(deps, **overrides):
kwargs = dict(
client=deps["client"],
provider=deps["provider"],
dangerously_allow_auto_upload_download_files=True,
session_id="s",
mcp=MagicMock(),
experimental=deps["experimental"],
custom_tools_map=deps["tools_map"],
user_id="u",
)
kwargs.update(overrides)
return ToolRouterSession(**kwargs)
class TestToolRouterSessionCustomTools:
def test_execute_local(self, mock_session_deps):
s = _session(mock_session_deps)
result = s.execute("GREP", arguments={"pattern": "x"})
# Returns SessionExecuteResponse — same type as remote
assert isinstance(result, SessionExecuteResponse)
assert result.error is None
assert result.log_id == ""
assert result.data["matches"] == ["x"]
mock_session_deps["client"].tool_router.session.execute.assert_not_called()
def test_execute_remote(self, mock_session_deps):
mock_response = SessionExecuteResponse(
data={"sent": True}, error=None, log_id="log_123"
)
mock_session_deps[
"client"
].tool_router.session.execute.return_value = mock_response
s = _session(mock_session_deps)
result = s.execute("GMAIL_SEND_EMAIL", arguments={"to": "a@b.com"})
mock_session_deps["client"].tool_router.session.execute.assert_called_once()
call_args = mock_session_deps["client"].tool_router.session.execute.call_args
assert call_args.kwargs["session_id"] == "s"
assert call_args.kwargs["tool_slug"] == "GMAIL_SEND_EMAIL"
assert call_args.kwargs["arguments"] == {"to": "a@b.com"}
assert "extra_body" not in call_args.kwargs
assert "enable_auto_workbench_offload" not in call_args.kwargs
# Remote returns client model as-is (backward compat, supports attribute access)
assert isinstance(result, SessionExecuteResponse)
assert result.data == {"sent": True}
assert result.log_id == "log_123"
def test_execute_remote_passes_inline_custom_tools(self, mock_session_deps):
mock_response = SessionExecuteResponse(
data={"sent": True}, error=None, log_id="log_123"
)
mock_session_deps[
"client"
].tool_router.session.execute.return_value = mock_response
inline_payload = {
"custom_tools": [
{
"slug": "GREP",
"name": "Grep",
"description": "Search local text",
"input_schema": {"type": "object", "properties": {}},
}
]
}
s = _session(mock_session_deps, inline_custom_tools_payload=inline_payload)
s.execute("GMAIL_SEND_EMAIL", arguments={"to": "a@b.com"})
call_args = mock_session_deps["client"].tool_router.session.execute.call_args
assert call_args.kwargs["experimental"] == inline_payload
def test_proxy_execute(self, mock_session_deps):
mock_session_deps[
"client"
].tool_router.session.proxy_execute.return_value = SessionProxyExecuteResponse(
status=200, data={"ok": True}, headers={}, binary_data=None
)
s = _session(mock_session_deps)
result = s.proxy_execute(
toolkit="gmail", endpoint="https://example.com", method="GET"
)
assert result == {"status": 200, "data": {"ok": True}, "headers": {}}
assert isinstance(result["status"], int)
assert "binary_data" not in result
def test_proxy_execute_projects_binary_data(self, mock_session_deps):
mock_session_deps[
"client"
].tool_router.session.proxy_execute.return_value = SessionProxyExecuteResponse(
status=200,
data={"ok": True},
headers={"content-type": "application/pdf"},
binary_data=BinaryData(
content_type="application/pdf",
size=123,
url="https://example.com/file.pdf",
expires_at="2026-08-18T00:00:00Z",
),
)
s = _session(mock_session_deps)
result = s.proxy_execute(
toolkit="gmail", endpoint="https://example.com", method="GET"
)
assert result == {
"status": 200,
"data": {"ok": True},
"headers": {"content-type": "application/pdf"},
"binary_data": {
"content_type": "application/pdf",
"size": 123,
"url": "https://example.com/file.pdf",
"expires_at": "2026-08-18T00:00:00Z",
},
}
assert isinstance(result["binary_data"]["size"], int)
def test_custom_tools_list(self, mock_session_deps):
s = _session(mock_session_deps)
assert len(s.custom_tools()) == 3
def test_custom_tools_filter(self, mock_session_deps):
s = _session(mock_session_deps)
assert len(s.custom_tools(toolkit="gmail")) == 1
def test_custom_toolkits_list(self, mock_session_deps):
s = _session(mock_session_deps)
tks = s.custom_toolkits()
assert len(tks) == 1
assert tks[0].slug == "ROLE_MANAGER"
def test_custom_toolkits_list_uses_qualified_slugs(self, duplicate_slug_tools):
alpha_tool, beta_tool = duplicate_slug_tools
alpha = ExperimentalToolkit(
slug="ALPHA", name="Alpha", description="Alpha tools"
)
alpha._tools.append(alpha_tool)
beta = ExperimentalToolkit(slug="BETA", name="Beta", description="Beta tools")
beta._tools.append(beta_tool)
custom_tools_map = build_custom_tools_map([], [alpha, beta])
s = ToolRouterSession(
client=MagicMock(),
provider=MagicMock(),
dangerously_allow_auto_upload_download_files=True,
session_id="s",
mcp=MagicMock(),
experimental=MagicMock(),
custom_tools_map=custom_tools_map,
user_id="u",
)
toolkits = s.custom_toolkits()
assert [tk.tools[0].slug for tk in toolkits] == [
"LOCAL_ALPHA_GREP",
"LOCAL_BETA_GREP",
]
def test_custom_toolkits_list_never_borrows_other_toolkit_slug(
self, duplicate_slug_tools
):
alpha_tool, beta_tool = duplicate_slug_tools
alpha = ExperimentalToolkit(
slug="ALPHA", name="Alpha", description="Alpha tools"
)
alpha._tools.append(alpha_tool)
beta = ExperimentalToolkit(slug="BETA", name="Beta", description="Beta tools")
beta._tools.append(beta_tool)
mock_exp = MagicMock()
mock_exp.custom_tools = []
mock_exp.custom_toolkits = [
MagicMock(
slug="ALPHA",
tools=[MagicMock(slug="LOCAL_ALPHA_GREP", original_slug="GREP")],
),
]
custom_tools_map = build_custom_tools_map_from_response(
[], [alpha, beta], mock_exp
)
s = ToolRouterSession(
client=MagicMock(),
provider=MagicMock(),
dangerously_allow_auto_upload_download_files=True,
session_id="s",
mcp=MagicMock(),
experimental=MagicMock(),
custom_tools_map=custom_tools_map,
user_id="u",
)
assert [tk.tools[0].slug for tk in s.custom_toolkits()] == [
"LOCAL_ALPHA_GREP",
"GREP",
]
def test_execute_rejects_ambiguous_original_slug(self, duplicate_slug_tools):
alpha_tool, beta_tool = duplicate_slug_tools
alpha = ExperimentalToolkit(
slug="ALPHA", name="Alpha", description="Alpha tools"
)
alpha._tools.append(alpha_tool)
beta = ExperimentalToolkit(slug="BETA", name="Beta", description="Beta tools")
beta._tools.append(beta_tool)
custom_tools_map = build_custom_tools_map([], [alpha, beta])
client = MagicMock()
s = ToolRouterSession(
client=client,
provider=MagicMock(),
dangerously_allow_auto_upload_download_files=True,
session_id="s",
mcp=MagicMock(),
experimental=MagicMock(),
custom_tools_map=custom_tools_map,
user_id="u",
)
with pytest.raises(ValidationError, match="Ambiguous custom tool slug"):
s.execute("GREP", arguments={"pattern": "x"})
client.tool_router.session.execute.assert_not_called()
alpha_result = s.execute("LOCAL_ALPHA_GREP", arguments={"pattern": "x"})
beta_result = s.execute("LOCAL_BETA_GREP", arguments={"pattern": "x"})
assert alpha_result.data["toolkit"] == "alpha"
assert beta_result.data["toolkit"] == "beta"
def test_empty_when_no_map(self):
s = ToolRouterSession(
client=MagicMock(),
provider=MagicMock(),
dangerously_allow_auto_upload_download_files=True,
session_id="s",
mcp=MagicMock(),
experimental=MagicMock(),
)
assert s.custom_tools() == []
# ────────────────────────────────────────────────────────────────
# Multi-execute routing
# ────────────────────────────────────────────────────────────────
class TestMultiExecuteRouting:
def _make_session(self, *tools, inline_custom_tools_payload=None):
m = build_custom_tools_map(list(tools))
return ToolRouterSession(
client=MagicMock(),
provider=MagicMock(),
dangerously_allow_auto_upload_download_files=True,
session_id="s",
mcp=MagicMock(),
experimental=MagicMock(),
custom_tools_map=m,
user_id="u",
inline_custom_tools_payload=inline_custom_tools_payload,
)
def test_single_local(self, grep_tool):
s = self._make_session(grep_tool)
result = s._route_multi_execute(
{"tools": [{"tool_slug": "GREP", "arguments": {"pattern": "x"}}]},
MagicMock(),
)
assert result["successful"] is True
assert result["data"]["matches"] == ["x"]
def test_rejects_ambiguous_original_slug(self, duplicate_slug_tools):
alpha_tool, beta_tool = duplicate_slug_tools
alpha = ExperimentalToolkit(
slug="ALPHA", name="Alpha", description="Alpha tools"
)
alpha._tools.append(alpha_tool)
beta = ExperimentalToolkit(slug="BETA", name="Beta", description="Beta tools")
beta._tools.append(beta_tool)
s = ToolRouterSession(
client=MagicMock(),
provider=MagicMock(),
dangerously_allow_auto_upload_download_files=True,
session_id="s",
mcp=MagicMock(),
experimental=MagicMock(),
custom_tools_map=build_custom_tools_map([], [alpha, beta]),
user_id="u",
)
tm = MagicMock()
backend = MagicMock()
tm._wrap_execute_tool_for_tool_router.return_value = backend
with pytest.raises(ValidationError, match="Ambiguous custom tool slug"):
s._route_multi_execute(
{
"tools": [
{
"tool_slug": "LOCAL_ALPHA_GREP",
"arguments": {"pattern": "x"},
},
{"tool_slug": "GREP", "arguments": {"pattern": "x"}},
]
},
tm,
)
backend.assert_not_called()
def test_all_remote(self, grep_tool):
s = self._make_session(grep_tool)
tm = MagicMock()
remote = {"data": {"results": []}, "error": None, "successful": True}
tm._wrap_execute_tool_for_tool_router.return_value = lambda slug, args: remote
result = s._route_multi_execute(
{"tools": [{"tool_slug": "REMOTE", "arguments": {}}]}, tm
)
assert result == remote
def test_all_remote_passes_inline_custom_tools(self, grep_tool):
inline_payload = {
"custom_tools": [
{
"slug": "GREP",
"name": "Grep",
"description": "Search local text",
"input_schema": {"type": "object", "properties": {}},
}
]
}
s = self._make_session(grep_tool, inline_custom_tools_payload=inline_payload)
tm = MagicMock()
remote = {"data": {"results": []}, "error": None, "successful": True}
tm._wrap_execute_tool_for_tool_router.return_value = lambda slug, args: remote
s._route_multi_execute(
{"tools": [{"tool_slug": "REMOTE", "arguments": {}}]}, tm
)
tm._wrap_execute_tool_for_tool_router.assert_called_once_with(
session_id="s",
modifiers=None,
inline_custom_tools_payload=inline_payload,
)
def test_mixed(self, grep_tool):
s = self._make_session(grep_tool)
tm = MagicMock()
remote = {
"data": {
"results": [
{"tool_slug": "R", "response": {"successful": True, "data": {}}},
],
"total_count": 1,
"success_count": 1,
"error_count": 0,
},
"error": None,
"successful": True,
}
tm._wrap_execute_tool_for_tool_router.return_value = lambda slug, args: remote
result = s._route_multi_execute(
{
"tools": [
{"tool_slug": "GREP", "arguments": {"pattern": "x"}},
{"tool_slug": "REMOTE", "arguments": {}},
]
},
tm,
)
results = result["data"]["results"]
assert len(results) == 2
# Results preserve the original request order.
assert results[0]["tool_slug"] == "GREP"
assert results[1]["tool_slug"] == "R"
assert results[0]["index"] == 0
assert results[1]["index"] == 1
assert result["data"]["total_count"] == 2
assert result["data"]["success_count"] == 2
assert result["data"]["error_count"] == 0
def test_failure_propagated(self):
@exp.tool()
def ok_tool(input: GrepInput, ctx):
"""OK."""
return {"ok": True}
@exp.tool()
def bad_tool(input: GrepInput, ctx):
"""Bad."""
raise RuntimeError("boom")
s = self._make_session(ok_tool, bad_tool)
result = s._route_multi_execute(
{
"tools": [
{"tool_slug": "OK_TOOL", "arguments": {"pattern": "x"}},
{"tool_slug": "BAD_TOOL", "arguments": {"pattern": "y"}},
]
},
MagicMock(),
)
assert result["successful"] is False
assert "1 out of 2" in result["error"]
def test_remote_null_data_does_not_crash(self, grep_tool):
s = self._make_session(grep_tool)
tm = MagicMock()
remote = {"data": None, "error": None, "successful": True}
tm._wrap_execute_tool_for_tool_router.return_value = lambda slug, args: remote
result = s._route_multi_execute(
{
"tools": [
{"tool_slug": "GREP", "arguments": {"pattern": "x"}},
{"tool_slug": "REMOTE", "arguments": {}},
]
},
tm,
)
assert result["successful"] is True
assert result["data"]["results"][0]["tool_slug"] == "GREP"
def test_remote_transport_failure_keeps_local_results(self, grep_tool):
"""A raised backend error becomes per-tool failures (matches TS)."""
s = self._make_session(grep_tool)
tm = MagicMock()
def failing_backend(slug, args):
raise RuntimeError("remote unavailable")
tm._wrap_execute_tool_for_tool_router.return_value = failing_backend
result = s._route_multi_execute(
{
"tools": [
{"tool_slug": "GREP", "arguments": {"pattern": "x"}},
{"tool_slug": "REMOTE", "arguments": {}},
{"tool_slug": "OTHER_REMOTE", "arguments": {}},
]
},
tm,
)
assert result["successful"] is False
assert result["error"] == "2 out of 3 tools failed"
results = result["data"]["results"]
assert [r["tool_slug"] for r in results] == ["GREP", "REMOTE", "OTHER_REMOTE"]
assert [r["index"] for r in results] == [0, 1, 2]
assert results[0]["response"] == {
"successful": True,
"data": {"matches": ["x"], "path": "."},
}
assert "error" not in results[0]
for entry in results[1:]:
assert entry["error"] == "remote unavailable"
assert entry["response"] == {
"successful": False,
"data": {},
"error": "remote unavailable",
}
assert result["data"]["total_count"] == 3
assert result["data"]["success_count"] == 1
assert result["data"]["error_count"] == 2
def test_remote_transport_failure_uses_fallback_message(self, grep_tool):
s = self._make_session(grep_tool)
tm = MagicMock()
def failing_backend(slug, args):
raise RuntimeError()
tm._wrap_execute_tool_for_tool_router.return_value = failing_backend
result = s._route_multi_execute(
{
"tools": [
{"tool_slug": "GREP", "arguments": {"pattern": "x"}},
{"tool_slug": "REMOTE", "arguments": {}},
]
},
tm,
)
assert result["successful"] is False
assert result["data"]["results"][1]["error"] == "Remote tool execution failed"
def test_remote_batch_error_without_item_errors_uses_batch_message(self, grep_tool):
s = self._make_session(grep_tool)
tm = MagicMock()
remote = {
"data": None,
"error": "Remote batch failed before per-tool results were produced",
"successful": False,
}
tm._wrap_execute_tool_for_tool_router.return_value = lambda slug, args: remote
result = s._route_multi_execute(
{
"tools": [
{"tool_slug": "GREP", "arguments": {"pattern": "x"}},
{"tool_slug": "REMOTE", "arguments": {}},
]
},
tm,
)
assert result["successful"] is False
assert (
result["error"]
== "Remote batch failed before per-tool results were produced"
)