From 73163e3acb93dda97bdbabaef928360bf0c29117 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 12:40:59 -0700 Subject: [PATCH] =?UTF-8?q?refactor(agent):=20model=5Ftools=20=E2=80=94=20?= =?UTF-8?q?extract=20=5Fdispatch=5Fbridge=5Ftool=20from=20handle=5Ffunctio?= =?UTF-8?q?n=5Fcall;=20pin=20bridge=20dispatch=20with=20tests?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- model_tools.py | 139 +++++++++++++++++++++----------------- tests/test_model_tools.py | 38 +++++++++++ 2 files changed, 114 insertions(+), 63 deletions(-) diff --git a/model_tools.py b/model_tools.py index 542f44572f..f214b5df70 100644 --- a/model_tools.py +++ b/model_tools.py @@ -976,6 +976,59 @@ def _emit_post_tool_call_hook( logger.debug("post_tool_call hook error: %s", _hook_err) +def _dispatch_bridge_tool( + function_name: str, + function_args: Dict[str, Any], + enabled_toolsets: Optional[List[str]], + disabled_toolsets: Optional[List[str]], +): + """Handle a Tool Search bridge call (tool_search / tool_describe / tool_call). + + Returns None when *function_name* is not a bridge tool. Otherwise returns + ``(result, None)`` for a finished catalog read or error, or + ``(None, (underlying_name, underlying_args))`` when a validated tool_call + should be re-dispatched as the real tool. + """ + try: + from tools import tool_search as ts + except Exception: + return None + if not ts.is_bridge_tool(function_name): + return None + # Read the un-collapsed catalog, scoped to the session's toolsets so a + # restricted session (subagent, kanban worker) cannot see or invoke the + # whole process registry through the bridge. + try: + current_defs = get_tool_definitions( + enabled_toolsets=enabled_toolsets, + disabled_toolsets=disabled_toolsets, + quiet_mode=True, skip_tool_search_assembly=True, + ) or [] + except Exception: + current_defs = [] + args = function_args or {} + if function_name == ts.TOOL_SEARCH_NAME: + return ts.dispatch_tool_search(args, current_tool_defs=current_defs), None + if function_name == ts.TOOL_DESCRIBE_NAME: + return ts.dispatch_tool_describe(args, current_tool_defs=current_defs), None + underlying_name, underlying_args, err = ts.resolve_underlying_call(args) + if err or not underlying_name: + return tool_error(err or "tool_call could not be resolved"), None + # Defense in depth: resolve_underlying_call only checks the global + # registry; also require membership in the session-scoped catalog. + if underlying_name not in ts.scoped_deferrable_names(current_defs): + return tool_error( + f"'{underlying_name}' is not available in this session. " + "Use tool_search to find tools you can call." + ), None + # Validate against the deferred tool's concrete schema — the generic + # ``arguments: object`` bridge schema can't enforce it. + probe_err = ts.validate_deferred_call_args(underlying_name, underlying_args) + if probe_err is not None: + return probe_err, None + return None, (underlying_name, underlying_args) + + def handle_function_call( function_name: str, function_args: Dict[str, Any], @@ -1032,69 +1085,29 @@ def handle_function_call( # Tool Search bridge: tool_search / tool_describe are catalog reads handled # inline; tool_call is unwrapped so every downstream hook (pre/post, edit # approval, guardrails) sees the real tool name, never the bridge. - try: - from tools import tool_search as _ts_mod - except Exception: - _ts_mod = None - - if _ts_mod is not None and _ts_mod.is_bridge_tool(function_name): - # Read the un-collapsed catalog, scoped to the session's toolsets so a - # restricted session (subagent, kanban worker) cannot see or invoke the - # whole process registry through the bridge. - try: - current_defs = get_tool_definitions( - enabled_toolsets=enabled_toolsets, - disabled_toolsets=disabled_toolsets, - quiet_mode=True, skip_tool_search_assembly=True, - ) or [] - except Exception: - current_defs = [] - - def _elapsed() -> int: - return int((time.monotonic() - _dispatch_start) * 1000) - - if function_name == _ts_mod.TOOL_SEARCH_NAME: - return _emit(_ts_mod.dispatch_tool_search(function_args or {}, current_tool_defs=current_defs), - duration_ms=_elapsed()) - if function_name == _ts_mod.TOOL_DESCRIBE_NAME: - return _emit(_ts_mod.dispatch_tool_describe(function_args or {}, current_tool_defs=current_defs), - duration_ms=_elapsed()) - if function_name == _ts_mod.TOOL_CALL_NAME: - underlying_name, underlying_args, err = _ts_mod.resolve_underlying_call(function_args or {}) - if err or not underlying_name: - return _emit(tool_error(err or "tool_call could not be resolved"), duration_ms=_elapsed()) - # Defense in depth: resolve_underlying_call only checks the global - # registry; also require membership in the session-scoped catalog. - if underlying_name not in _ts_mod.scoped_deferrable_names(current_defs): - return _emit( - tool_error( - f"'{underlying_name}' is not available in this session. " - "Use tool_search to find tools you can call." - ), - duration_ms=_elapsed(), - ) - # Validate against the deferred tool's concrete schema — the generic - # ``arguments: object`` bridge schema can't enforce it. - _probe_err = _ts_mod.validate_deferred_call_args(underlying_name, underlying_args) - if _probe_err is not None: - return _emit(_probe_err, duration_ms=_elapsed()) - return handle_function_call( - function_name=underlying_name, - function_args=underlying_args, - task_id=task_id, - tool_call_id=tool_call_id, - session_id=session_id, - turn_id=turn_id, - api_request_id=api_request_id, - user_task=user_task, - enabled_tools=enabled_tools, - skip_pre_tool_call_hook=skip_pre_tool_call_hook, - skip_tool_request_middleware=skip_tool_request_middleware, - skip_tool_execution_middleware=skip_tool_execution_middleware, - tool_request_middleware_trace=list(_tool_middleware_trace), - enabled_toolsets=enabled_toolsets, - disabled_toolsets=disabled_toolsets, - ) + bridged = _dispatch_bridge_tool(function_name, function_args, enabled_toolsets, disabled_toolsets) + if bridged is not None: + result, underlying = bridged + if underlying is None: + return _emit(result, duration_ms=int((time.monotonic() - _dispatch_start) * 1000)) + underlying_name, underlying_args = underlying + return handle_function_call( + function_name=underlying_name, + function_args=underlying_args, + task_id=task_id, + tool_call_id=tool_call_id, + session_id=session_id, + turn_id=turn_id, + api_request_id=api_request_id, + user_task=user_task, + enabled_tools=enabled_tools, + skip_pre_tool_call_hook=skip_pre_tool_call_hook, + skip_tool_request_middleware=skip_tool_request_middleware, + skip_tool_execution_middleware=skip_tool_execution_middleware, + tool_request_middleware_trace=list(_tool_middleware_trace), + enabled_toolsets=enabled_toolsets, + disabled_toolsets=disabled_toolsets, + ) _tool_original_args = dict(function_args) if not skip_tool_request_middleware: diff --git a/tests/test_model_tools.py b/tests/test_model_tools.py index 9e1fa9886e..751ca3187c 100644 --- a/tests/test_model_tools.py +++ b/tests/test_model_tools.py @@ -509,3 +509,41 @@ class TestDisabledToolsetsPostureToolset: ) } assert "write_file" not in no_file + + +# ========================================================================= +# Tool Search bridge dispatch +# ========================================================================= + +class TestBridgeDispatch: + """handle_function_call routes tool_search/tool_describe inline, unwraps tool_call, + and refuses tool_call targets outside the session-scoped deferrable catalog.""" + + def test_tool_search_and_describe_return_json_strings(self): + with patch("model_tools.get_tool_definitions", return_value=[]): + out = handle_function_call("tool_search", {"queries": ["anything"]}) + assert isinstance(out, str) and json.loads(out) is not None + out = handle_function_call("tool_describe", {"names": ["nope"]}) + assert isinstance(out, str) and json.loads(out) is not None + + def test_tool_call_bad_args_error(self): + with patch("model_tools.get_tool_definitions", return_value=[]): + result = json.loads(handle_function_call("tool_call", {})) + assert "requires a 'name'" in result["error"] + + def test_tool_call_rejects_out_of_scope_and_unwraps_in_scope(self): + import tools.tool_search as ts + with patch("model_tools.get_tool_definitions", return_value=[]), \ + patch.object(ts, "resolve_underlying_call", return_value=("mcp_x", {"a": 1}, None)), \ + patch.object(ts, "scoped_deferrable_names", return_value=frozenset()): + result = json.loads(handle_function_call("tool_call", {"name": "mcp_x"})) + assert "not available in this session" in result["error"] + + with patch("model_tools.get_tool_definitions", return_value=[]), \ + patch.object(ts, "resolve_underlying_call", return_value=("mcp_x", {"a": 1}, None)), \ + patch.object(ts, "scoped_deferrable_names", return_value=frozenset({"mcp_x"})), \ + patch.object(ts, "validate_deferred_call_args", return_value=None), \ + patch("model_tools.registry.dispatch", return_value='{"ok": true}') as disp: + out = handle_function_call("tool_call", {"name": "mcp_x"}, task_id="t") + assert json.loads(out) == {"ok": True} + assert disp.call_args.args[0] == "mcp_x" and disp.call_args.args[1] == {"a": 1}