Skip to content
Open
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
8 changes: 5 additions & 3 deletions packages/toolbox-core/src/toolbox_core/client.py
Original file line number Diff line number Diff line change
Expand Up @@ -493,6 +493,11 @@ async def load_toolset(
)
tools.append(tool)

# Track usage for both modes. Skipping this under strict=True left
# overall-used sets empty and made the final toolset check always fail.
overall_used_auth_keys.update(used_auth_keys)
overall_used_bound_params.update(used_bound_keys)

if strict:
validate_unused_requirements(
provided_auth_keys,
Expand All @@ -502,9 +507,6 @@ async def load_toolset(
tool_name,
is_toolset=False,
)
else:
overall_used_auth_keys.update(used_auth_keys)
overall_used_bound_params.update(used_bound_keys)

validate_unused_requirements(
provided_auth_keys,
Expand Down
53 changes: 53 additions & 0 deletions packages/toolbox-core/tests/test_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -515,6 +515,32 @@ async def test_load_tool_with_unused_bound_param_fail(
TOOL_NAME, bound_params={"unused_param": "some_value"}
)

@pytest.mark.asyncio
async def test_load_toolset_strict_with_fully_used_bound_param_success(
self, mock_transport, tool_schema_with_param_P
):
"""Tests that load_toolset succeeds in strict mode when every tool uses all bound params."""
TOOL_P1 = "tool_with_p1"
TOOL_P2 = "tool_with_p2"
manifest = ManifestSchema(
serverVersion="0.0.0",
tools={
TOOL_P1: tool_schema_with_param_P,
TOOL_P2: tool_schema_with_param_P,
},
)
mock_transport.tools_list_mock.return_value = manifest

async with ToolboxClient(TEST_BASE_URL) as client:
client._ToolboxClient__transport = mock_transport
tools = await client.load_toolset(
bound_params={"param_P": "some_value"}, strict=True
)

assert len(tools) == 2
assert {t.__name__ for t in tools} == {TOOL_P1, TOOL_P2}
assert all("param_P" not in t.__signature__.parameters for t in tools)

@pytest.mark.asyncio
async def test_load_toolset_strict_with_partially_used_bound_param_fail(
self, mock_transport, tool_schema_with_param_P, tool_schema_minimal
Expand Down Expand Up @@ -559,6 +585,33 @@ async def test_load_toolset_non_strict_with_unused_bound_param_fail(
TOOLSET_NAME, bound_params={"param_Z": "some_value"}
)

@pytest.mark.asyncio
async def test_load_toolset_strict_with_fully_used_auth_success(
self, mock_transport, tool_schema_requires_auth_X
):
"""Tests that load_toolset succeeds in strict mode when every tool uses all auth tokens."""
TOOL_AUTH1 = "tool_with_auth1"
TOOL_AUTH2 = "tool_with_auth2"
manifest = ManifestSchema(
serverVersion="0.0.0",
tools={
TOOL_AUTH1: tool_schema_requires_auth_X,
TOOL_AUTH2: tool_schema_requires_auth_X,
},
)
mock_transport.tools_list_mock.return_value = manifest

async with ToolboxClient(TEST_BASE_URL) as client:
client._ToolboxClient__transport = mock_transport
tools = await client.load_toolset(
auth_token_getters={"auth_service_X": lambda: "token"}, strict=True
)

assert len(tools) == 2
assert {t.__name__ for t in tools} == {TOOL_AUTH1, TOOL_AUTH2}
assert all(t._required_authn_params == {} for t in tools)
assert all(t._required_authz_tokens == () for t in tools)

@pytest.mark.asyncio
async def test_load_toolset_strict_with_partially_used_auth_fail(
self, mock_transport, tool_schema_requires_auth_X, tool_schema_minimal
Expand Down
104 changes: 104 additions & 0 deletions packages/toolbox-core/tests/test_sync_client.py

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Can we also add tests for auth params as well just like in async client?

Original file line number Diff line number Diff line change
Expand Up @@ -232,6 +232,110 @@ def test_sync_load_toolset_success(
assert result1 == f"{TOOL1_NAME}_ok"


def test_sync_load_toolset_strict_with_fully_used_bound_param_success(
sync_client, mock_transport
):
"""Sync path: strict=True succeeds when every tool uses all bound params."""
TOOL_P1 = "tool_with_p1"
TOOL_P2 = "tool_with_p2"
schema = ToolSchema(
description="Tool with Parameter P",
parameters=[
ParameterSchema(name="param_P", type="string", description="Parameter P"),
],
)
manifest = ManifestSchema(
serverVersion="0.0.0", tools={TOOL_P1: schema, TOOL_P2: schema}
)
mock_transport.tools_list_mock.return_value = manifest
sync_client._ToolboxSyncClient__async_client._ToolboxClient__transport = (
mock_transport
)

tools = sync_client.load_toolset(
bound_params={"param_P": "some_value"}, strict=True
)

assert len(tools) == 2
assert {t.__name__ for t in tools} == {TOOL_P1, TOOL_P2}
assert all("param_P" not in t.__signature__.parameters for t in tools)


def test_sync_load_toolset_strict_with_fully_used_auth_success(
sync_client, mock_transport
):
"""Sync path: strict=True succeeds when every tool uses all auth tokens."""
TOOL_AUTH1 = "tool_with_auth1"
TOOL_AUTH2 = "tool_with_auth2"
schema = ToolSchema(
description="Tool Requiring Auth X",
parameters=[
ParameterSchema(
name="auth_param_X",
type="string",
description="Auth X Token",
authSources=["auth_service_X"],
),
ParameterSchema(name="data", type="string", description="Some data"),
],
)
manifest = ManifestSchema(
serverVersion="0.0.0", tools={TOOL_AUTH1: schema, TOOL_AUTH2: schema}
)
mock_transport.tools_list_mock.return_value = manifest
sync_client._ToolboxSyncClient__async_client._ToolboxClient__transport = (
mock_transport
)

tools = sync_client.load_toolset(
auth_token_getters={"auth_service_X": lambda: "token"}, strict=True
)

assert len(tools) == 2
assert {t.__name__ for t in tools} == {TOOL_AUTH1, TOOL_AUTH2}
assert all(t._required_authn_params == {} for t in tools)
assert all(t._required_authz_tokens == () for t in tools)


def test_sync_load_toolset_strict_with_partially_used_auth_fail(
sync_client, mock_transport
):
"""Sync path: strict=True fails if an auth token is only used by some tools."""
TOOL_AUTH = "tool_with_auth"
TOOL_MIN = "minimal_tool"
schema_auth = ToolSchema(
description="Tool Requiring Auth X",
parameters=[
ParameterSchema(
name="auth_param_X",
type="string",
description="Auth X Token",
authSources=["auth_service_X"],
),
],
)
schema_min = ToolSchema(
description="Minimal Test Tool",
parameters=[],
)
manifest = ManifestSchema(
serverVersion="0.0.0",
tools={TOOL_AUTH: schema_auth, TOOL_MIN: schema_min},
)
mock_transport.tools_list_mock.return_value = manifest
sync_client._ToolboxSyncClient__async_client._ToolboxClient__transport = (
mock_transport
)

with pytest.raises(
ValueError,
match=f"Validation failed for tool '{TOOL_MIN}': unused auth tokens: auth_service_X.",
):
sync_client.load_toolset(
auth_token_getters={"auth_service_X": lambda: "token"}, strict=True
)


def test_sync_invoke_tool_server_error(
test_tool_str_schema, sync_client, mock_transport
):
Expand Down
Loading