From 28782f90b3c358c9d5a354f3b56436214e8aa11e Mon Sep 17 00:00:00 2001 From: Henry Su Date: Fri, 24 Jul 2026 16:44:55 -0500 Subject: [PATCH] fix(core): allow load_toolset strict success when all params are used Accumulate used auth tokens and bound params under strict=True so the final toolset unused-requirement check no longer rejects valid loads. --- .../toolbox-core/src/toolbox_core/client.py | 8 +- packages/toolbox-core/tests/test_client.py | 53 +++++++++ .../toolbox-core/tests/test_sync_client.py | 104 ++++++++++++++++++ 3 files changed, 162 insertions(+), 3 deletions(-) diff --git a/packages/toolbox-core/src/toolbox_core/client.py b/packages/toolbox-core/src/toolbox_core/client.py index 8f9d0cd75..cdd5c1539 100644 --- a/packages/toolbox-core/src/toolbox_core/client.py +++ b/packages/toolbox-core/src/toolbox_core/client.py @@ -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, @@ -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, diff --git a/packages/toolbox-core/tests/test_client.py b/packages/toolbox-core/tests/test_client.py index bd67fde9b..2551f9fee 100644 --- a/packages/toolbox-core/tests/test_client.py +++ b/packages/toolbox-core/tests/test_client.py @@ -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 @@ -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 diff --git a/packages/toolbox-core/tests/test_sync_client.py b/packages/toolbox-core/tests/test_sync_client.py index 6a6649b9f..92c8b94d0 100644 --- a/packages/toolbox-core/tests/test_sync_client.py +++ b/packages/toolbox-core/tests/test_sync_client.py @@ -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 ):