From 47ac76a03a00c80464f8c249a5640aae3ef0ad42 Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 18:17:45 -0700 Subject: [PATCH] refactor(model_tools,toolsets): dispatcher phase helpers, shared registry accessor, table-built distributions --- model_tools.py | 317 +++++++++++++++------------------- toolset_distributions.py | 360 ++++++--------------------------------- toolsets.py | 88 +++------- 3 files changed, 220 insertions(+), 545 deletions(-) diff --git a/model_tools.py b/model_tools.py index f0df359105..26d9d97b5f 100644 --- a/model_tools.py +++ b/model_tools.py @@ -12,7 +12,7 @@ import json import re import asyncio from contextlib import contextmanager -from dataclasses import dataclass +from dataclasses import asdict, dataclass from contextvars import ContextVar import logging import threading @@ -78,13 +78,10 @@ def _is_dispatcher_owned_worker() -> bool: return True -# ============================================================================= -# Async Bridging (single source of truth -- used by registry.dispatch too) -# ============================================================================= +# --- Async bridging (single source of truth; registry.dispatch uses it too) --- # Loops are persistent (never asyncio.run per call): cached httpx/AsyncOpenAI -# clients stay bound to a live loop, so their GC cleanup can't hit -# "Event loop is closed". Main thread shares one loop; worker threads -# (parallel tool execution) each own a thread-local loop to avoid contention. +# clients stay bound to a live loop, so their GC cleanup can't hit "Event loop +# is closed". Main thread shares one loop; worker threads own thread-local loops. _tool_loop = None # persistent loop for the main (CLI) thread _tool_loop_lock = threading.Lock() @@ -118,10 +115,9 @@ def _run_async(coro): loop = None if loop and loop.is_running(): - # Already inside an event loop: run in a fresh thread whose loop we - # hold a reference to, so on timeout we can cancel the task inside it - # (ThreadPoolExecutor.cancel() is a no-op on a running worker and - # would leak the thread on every 300 s timeout). + # Inside a running loop: run in a fresh thread whose loop we keep a + # reference to, so on timeout we can cancel the task inside it + # (ThreadPoolExecutor.cancel() is a no-op on a running worker). import concurrent.futures worker_loop: Optional[asyncio.AbstractEventLoop] = None @@ -173,28 +169,20 @@ def _run_async(coro): return _get_tool_loop().run_until_complete(coro) -# ============================================================================= -# Tool Discovery (importing each module triggers its registry.register calls) -# ============================================================================= - +# --- Tool discovery (importing each tools/*.py triggers registry.register) --- discover_builtin_tools() # MCP discovery is deliberately NOT run here: it blocks up to 120 s and the -# gateway lazy-imports this module inside its event loop. Each entry point -# (gateway/run.py, cli.py, tui_gateway, acp_adapter) runs it at its own startup. - -# Plugin tool discovery (user/project/pip plugins) -try: +# gateway lazy-imports this module inside its event loop; each entry point +# (gateway/run.py, cli.py, tui_gateway, acp_adapter) runs it at startup. +try: # plugin tool discovery (user/project/pip plugins) from hermes_cli.plugins import discover_plugins discover_plugins() except Exception as e: logger.debug("Plugin discovery failed: %s", e) -# ============================================================================= -# Backward-compat constants (built once after discovery) -# ============================================================================= - +# Backward-compat constants (built once after discovery) TOOL_TO_TOOLSET_MAP: Dict[str, str] = registry.get_tool_to_toolset_map() TOOLSET_REQUIREMENTS: Dict[str, dict] = registry.get_toolset_requirements() @@ -203,10 +191,7 @@ TOOLSET_REQUIREMENTS: Dict[str, dict] = registry.get_toolset_requirements() _last_resolved_tool_names: List[str] = [] -# ============================================================================= -# Legacy toolset name mapping (old _tools-suffixed names -> tool name lists) -# ============================================================================= - +# Legacy toolset names (old _tools-suffixed names -> tool name lists) _LEGACY_TOOLSET_MAP = { "web_tools": ["web_search", "web_extract"], "terminal_tools": ["terminal"], @@ -225,10 +210,7 @@ _LEGACY_TOOLSET_MAP = { } -# ============================================================================= -# get_tool_definitions (the main schema provider) -# ============================================================================= - +# --- get_tool_definitions (the main schema provider) -------------------------- # Memo for get_tool_definitions(), active only with quiet_mode=True (the # non-quiet path prints). Hot callers (gateway runner, AIAgent.__init__) hit it # every turn; a miss costs ~7 ms of registry walk + check_fn probing. The key @@ -236,9 +218,7 @@ _LEGACY_TOOLSET_MAP = { # invalidation is transparent; check_fn drift is handled by registry.py's 30 s TTL. _tool_defs_cache: Dict[tuple, List[Dict[str, Any]]] = {} _tool_defs_cache_lock = threading.Lock() - -# FIFO cap: a long-lived gateway sees many toolset/config fingerprints; 8 -# covers the warm working set of platform/toolset combos it actually serves. +# FIFO cap: 8 covers a long-lived gateway's warm set of platform/toolset combos. _TOOL_DEFS_CACHE_MAX = 8 @@ -395,19 +375,21 @@ def _select_tool_names( # --- Dynamic schema rewrites ------------------------------------------------- -# Each rewriter receives the tool definition and the set of tool names that -# passed check_fn filtering, and returns the (possibly replaced) definition or -# None to drop the tool. Cross-references must use that set (not the requested -# names) so the model is never told about a tool that isn't in the list. +# Each rewriter gets (tool definition, set of tool names that passed check_fn) +# and returns the (possibly replaced) definition, or None to drop the tool. +# Cross-references must use that set so the model never hears of an absent tool. _BROWSER_NAVIGATE_WEB_HINT = " For simple information retrieval, prefer web_search or web_extract (faster, cheaper)." +def _fn_def(schema: Dict[str, Any]) -> Dict[str, Any]: + return {"type": "function", "function": schema} + + def _rewrite_execute_code(td: Dict[str, Any], available: set) -> Optional[Dict[str, Any]]: """List only sandbox tools that are actually available.""" from tools.code_execution_tool import SANDBOX_ALLOWED_TOOLS, build_execute_code_schema, _get_execution_mode - schema = build_execute_code_schema(SANDBOX_ALLOWED_TOOLS & available, mode=_get_execution_mode()) - return {"type": "function", "function": schema} + return _fn_def(build_execute_code_schema(SANDBOX_ALLOWED_TOOLS & available, mode=_get_execution_mode())) def _discord_rewriter(schema_fn_name: str): @@ -418,7 +400,7 @@ def _discord_rewriter(schema_fn_name: str): dynamic = getattr(_dt, schema_fn_name)() except Exception: dynamic = None - return None if dynamic is None else {"type": "function", "function": dynamic} + return None if dynamic is None else _fn_def(dynamic) return _rewrite @@ -427,14 +409,13 @@ def _rewrite_browser_navigate(td: Dict[str, Any], available: set) -> Optional[Di if {"web_search", "web_extract"} & available: return td desc = td["function"].get("description", "").replace(_BROWSER_NAVIGATE_WEB_HINT, "") - return {"type": "function", "function": {**td["function"], "description": desc}} + return _fn_def({**td["function"], "description": desc}) def _rewrite_browser_exec(td: Dict[str, Any], available: set) -> Optional[Dict[str, Any]]: - """browser_exec runs arbitrary host Python; a session without the terminal - surface must not regain host execution through the browser toolset. This is - a session-level gate rather than a check_fn: check_fns are TTL-cached - process-wide while one gateway serves sessions with different toolsets.""" + """browser_exec runs arbitrary host Python: a session without the terminal surface + must not regain host execution via the browser toolset. Session-level gate rather + than a check_fn because check_fns are TTL-cached process-wide across sessions.""" return td if "terminal" in available else None @@ -451,17 +432,14 @@ def _rewrite_delegate_task(td: Dict[str, Any], available: set) -> Optional[Dict[ full_offvariant = "delegate_task, clarify, memory, or cronjob" full_onvariant = "clarify, memory, or cronjob" if full_offvariant in desc: - full, keep_self = full_offvariant, True + full, names = full_offvariant, ["delegate_task"] + blocked_present elif full_onvariant in desc: - full, keep_self = full_onvariant, False + full, names = full_onvariant, blocked_present else: return td - names = (["delegate_task"] if keep_self else []) + blocked_present if blocked_present: - if len(names) == 1: - replacement = names[0] - elif len(names) == 2: - replacement = f"{names[0]} or {names[1]}" + if len(names) <= 2: + replacement = " or ".join(names) else: replacement = ", ".join(names[:-1]) + ", or " + names[-1] desc = desc.replace(full, replacement) @@ -498,6 +476,15 @@ def _apply_dynamic_schemas(tool_defs: List[Dict[str, Any]]) -> List[Dict[str, An return out +_TOOL_SEARCH_LISTING_FORMS = { + "full": "catalog listing embedded", + "names": "names-only listing embedded", + "mixed": "listing embedded (oversized servers summarized)", + "groups": "server summary embedded (search-only discovery)", + "none": "no listing (search-only)", +} + + def _compute_tool_definitions( enabled_toolsets: Optional[List[str]] = None, disabled_toolsets: Optional[List[str]] = None, @@ -541,15 +528,11 @@ def _compute_tool_definitions( config=ts_cfg, ) if assembly.activated and not quiet_mode: - _forms = {"full": "catalog listing embedded", - "names": "names-only listing embedded", - "mixed": "listing embedded (oversized servers summarized)", - "groups": "server summary embedded (search-only discovery)", - "none": "no listing (search-only)"} print( f"šŸ”Ž Tool Search (tier {assembly.tier}): {assembly.deferred_count} " f"MCP/plugin tools deferred (~{assembly.deferred_tokens} tokens) behind " - f"tool_search/describe/call — {_forms.get(assembly.listing_form, assembly.listing_form)}." + f"tool_search/describe/call — " + f"{_TOOL_SEARCH_LISTING_FORMS.get(assembly.listing_form, assembly.listing_form)}." ) filtered_tools = assembly.tool_defs except Exception as e: # pragma: no cover — never break tool loading @@ -573,11 +556,10 @@ def _active_model_config() -> Tuple[str, Dict[str, Any]]: def _resolve_active_context_length() -> int: """Active model's context length for the tool-search gate (0 if unresolvable). - Order: explicit `model.context_length` in config.yaml; provider-aware - resolution (Codex OAuth enforces a smaller window than the direct API for - the same slug); the on-disk metadata cache (a slightly stale window is fine - for picking a disclosure tier and avoids a ~200 ms /models probe per CLI - startup); then the full live resolver. + Order: explicit `model.context_length`; provider-aware resolution (Codex OAuth + enforces a smaller window than the direct API for the same slug); the on-disk + metadata cache (slightly stale is fine for picking a tier and avoids a ~200 ms + /models probe per CLI startup); then the full live resolver. """ try: model_id, model_cfg = _active_model_config() @@ -611,11 +593,8 @@ def _resolve_active_context_length() -> int: except Exception: pass return int(get_model_context_length( - model_id, - base_url=base_url, - api_key=api_key, - config_context_length=config_ctx, - provider=provider, + model_id, base_url=base_url, api_key=api_key, + config_context_length=config_ctx, provider=provider, ) or 0) except Exception as e: logger.debug("Could not resolve active context length: %s", e) @@ -642,20 +621,20 @@ _LEGACY_TOOL_ALIASES = { _READ_SEARCH_TOOLS = {"read_file", "search_files"} -# ========================================================================= -# Tool error sanitization -# ========================================================================= +# --- Tool error sanitization -------------------------------------------------- # Defense-in-depth: json.dumps already prevents framing escape, but the model # still reads the text, so strip role tags / CDATA / code fences from exception # messages and cap length. The cap is shared with tools/registry.py so text never # passes two different caps with two different markers. -_TOOL_ERROR_ROLE_TAG_RE = re.compile( - r'', - re.IGNORECASE, +_TOOL_ERROR_STRIP_RES = ( + re.compile( + r'', + re.IGNORECASE, + ), + re.compile(r'^\s*```(?:json|xml|html|markdown)?\s*', re.MULTILINE), + re.compile(r'\s*```\s*$', re.MULTILINE), + re.compile(r'', re.DOTALL), ) -_TOOL_ERROR_FENCE_OPEN_RE = re.compile(r'^\s*```(?:json|xml|html|markdown)?\s*', re.MULTILINE) -_TOOL_ERROR_FENCE_CLOSE_RE = re.compile(r'\s*```\s*$', re.MULTILINE) -_TOOL_ERROR_CDATA_RE = re.compile(r'', re.DOTALL) from tools.registry import _MAX_TOOL_ERROR_CHARS as _TOOL_ERROR_MAX_LEN @@ -663,10 +642,9 @@ def _sanitize_tool_error(error_msg: str) -> str: """Strip structural framing tokens from a tool error before the model sees it.""" if not error_msg: return "[TOOL_ERROR] " - sanitized = _TOOL_ERROR_ROLE_TAG_RE.sub("", error_msg) - sanitized = _TOOL_ERROR_FENCE_OPEN_RE.sub("", sanitized) - sanitized = _TOOL_ERROR_FENCE_CLOSE_RE.sub("", sanitized) - sanitized = _TOOL_ERROR_CDATA_RE.sub("", sanitized) + sanitized = error_msg + for pattern in _TOOL_ERROR_STRIP_RES: + sanitized = pattern.sub("", sanitized) if len(sanitized) > _TOOL_ERROR_MAX_LEN: sanitized = sanitized[:_TOOL_ERROR_MAX_LEN - 3] + "..." return f"[TOOL_ERROR] {sanitized}" @@ -683,13 +661,7 @@ class _CallIds: def hook_kwargs(self) -> Dict[str, str]: """The same fields with None normalized to "" (hook/middleware wire contract).""" - return { - "task_id": self.task_id or "", - "session_id": self.session_id or "", - "tool_call_id": self.tool_call_id or "", - "turn_id": self.turn_id or "", - "api_request_id": self.api_request_id or "", - } + return {k: v or "" for k, v in asdict(self).items()} def _tool_result_observer_fields( @@ -730,12 +702,9 @@ def _emit_post_tool_call_hook( error_message: Optional[str] = None, middleware_trace: Optional[List[Dict[str, Any]]] = None, ) -> None: - """Emit the ``post_tool_call`` observer hook. - - Gated on has_hook so the no-listener path costs one dict lookup; when - ``status`` is None the ok/error fields are derived from the result only - after that gate. - """ + """Emit the ``post_tool_call`` observer hook (gated on has_hook so the + no-listener path costs one dict lookup; ok/error fields are derived from + the result only after that gate when ``status`` is None).""" if _post_tool_call_hook_suppressed.get(): return try: @@ -743,10 +712,7 @@ def _emit_post_tool_call_hook( if not has_hook("post_tool_call"): return if status is None: - status, error_type, error_message = _tool_result_observer_fields( - function_name, - result, - ) + status, error_type, error_message = _tool_result_observer_fields(function_name, result) invoke_hook( "post_tool_call", tool_name=function_name, @@ -816,6 +782,23 @@ def _dispatch_bridge_tool( return None, (underlying_name, underlying_args) +def _apply_request_middleware( + function_name: str, + function_args: Dict[str, Any], + ids: _CallIds, + trace: List[Dict[str, Any]], +) -> Tuple[Dict[str, Any], Dict[str, Any], List[Dict[str, Any]]]: + """tool_request middleware: returns (args, original_args, trace); fail-open.""" + try: + from hermes_cli.middleware import apply_tool_request_middleware + + mw = apply_tool_request_middleware(function_name, function_args, **ids.hook_kwargs()) + return mw.payload, mw.original_payload, mw.trace + except Exception as _mw_err: + logger.debug("tool_request middleware error: %s", _mw_err) + return function_args, dict(function_args), trace + + def _pre_dispatch_guards( function_name: str, function_args: Dict[str, Any], @@ -862,6 +845,26 @@ def _pre_dispatch_guards( return function_args, None +@contextmanager +def _approval_observability(ids: _CallIds): + """Bind the approval observability context (turn/tool_call/session ids) for the block.""" + try: + from tools.approval import reset_current_observability_context, set_current_observability_context + tokens = set_current_observability_context( + turn_id=ids.turn_id or "", tool_call_id=ids.tool_call_id or "", session_id=ids.session_id or "", + ) + except Exception: + yield + return + try: + yield + finally: + try: + reset_current_observability_context(tokens) + except Exception: + pass + + def _execute_tool( function_name: str, function_args: Dict[str, Any], @@ -874,34 +877,18 @@ def _execute_tool( ) -> Any: """Run the registry handler (through tool-execution middleware unless skipped) with the approval observability context bound for the duration.""" - approval_tokens = None - reset_obs = None - try: - from tools.approval import ( - reset_current_observability_context as reset_obs, - set_current_observability_context, - ) - approval_tokens = set_current_observability_context( - turn_id=ids.turn_id or "", - tool_call_id=ids.tool_call_id or "", - session_id=ids.session_id or "", - ) - except Exception: - reset_obs = None - try: - dispatch_kwargs: Dict[str, Any] = {"task_id": ids.task_id, "session_id": ids.session_id} - if function_name == "execute_code": - # Prefer the caller's list so subagents can't overwrite the - # parent's tool set via the process-global. - dispatch_kwargs["enabled_tools"] = ( - enabled_tools if enabled_tools is not None else _last_resolved_tool_names - ) - else: - dispatch_kwargs["user_task"] = user_task + dispatch_kwargs: Dict[str, Any] = {"task_id": ids.task_id, "session_id": ids.session_id} + if function_name == "execute_code": + # Prefer the caller's list so subagents can't overwrite the parent's + # tool set via the process-global. + dispatch_kwargs["enabled_tools"] = enabled_tools if enabled_tools is not None else _last_resolved_tool_names + else: + dispatch_kwargs["user_task"] = user_task - def _dispatch(next_args: Dict[str, Any]) -> Any: - return registry.dispatch(function_name, next_args, **dispatch_kwargs) + def _dispatch(next_args: Dict[str, Any]) -> Any: + return registry.dispatch(function_name, next_args, **dispatch_kwargs) + with _approval_observability(ids): if skip_tool_execution_middleware: return _dispatch(function_args) from hermes_cli.middleware import run_tool_execution_middleware @@ -909,12 +896,6 @@ def _execute_tool( return run_tool_execution_middleware( function_name, function_args, _dispatch, original_args=original_args, **ids.hook_kwargs(), ) - finally: - if approval_tokens is not None and reset_obs is not None: - try: - reset_obs(approval_tokens) - except Exception: - pass def _apply_transform_tool_result_hook( @@ -953,6 +934,10 @@ def _apply_transform_tool_result_hook( return result +def _elapsed_ms(start: float) -> int: + return int((time.monotonic() - start) * 1000) + + def handle_function_call( function_name: str, function_args: Dict[str, Any], @@ -986,17 +971,16 @@ def handle_function_call( function_args = coerce_tool_args(function_name, function_args) if not isinstance(function_args, dict): function_args = {} - _tool_middleware_trace = list(tool_request_middleware_trace or []) + trace = list(tool_request_middleware_trace or []) function_name = _LEGACY_TOOL_ALIASES.get(function_name, function_name) ids = _CallIds(task_id, session_id, tool_call_id, turn_id, api_request_id) - _dispatch_start = time.monotonic() + start = time.monotonic() def _emit(result: Any, **extra: Any) -> Any: """Emit post_tool_call with this call's identity fields; returns *result*.""" _emit_post_tool_call_hook( function_name=function_name, function_args=function_args, result=result, - task_id=task_id, session_id=session_id, tool_call_id=tool_call_id, turn_id=turn_id, - api_request_id=api_request_id, middleware_trace=list(_tool_middleware_trace), **extra, + **asdict(ids), middleware_trace=list(trace), **extra, ) return result @@ -1007,45 +991,26 @@ def handle_function_call( 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 _emit(result, duration_ms=_elapsed_ms(start)) 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, + *underlying, 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_request_middleware_trace=list(trace), + enabled_toolsets=enabled_toolsets, disabled_toolsets=disabled_toolsets, ) - _tool_original_args = dict(function_args) + original_args = dict(function_args) if not skip_tool_request_middleware: - try: - from hermes_cli.middleware import apply_tool_request_middleware - - _tool_request_mw = apply_tool_request_middleware(function_name, function_args, **ids.hook_kwargs()) - function_args = _tool_request_mw.payload - _tool_original_args = _tool_request_mw.original_payload - _tool_middleware_trace = _tool_request_mw.trace - except Exception as _mw_err: - logger.debug("tool_request middleware error: %s", _mw_err) + function_args, original_args, trace = _apply_request_middleware(function_name, function_args, ids, trace) try: if function_name in _AGENT_LOOP_TOOLS: return tool_error(f"{function_name} must be handled by the agent loop") - function_args, blocked = _pre_dispatch_guards( - function_name, function_args, skip_pre_tool_call_hook, ids, _tool_middleware_trace, - ) + function_args, blocked = _pre_dispatch_guards(function_name, function_args, skip_pre_tool_call_hook, ids, trace) if blocked is not None: result, error_type, error_message = blocked return _emit(result, status="blocked", error_type=error_type, error_message=error_message) @@ -1059,16 +1024,14 @@ def handle_function_call( pass # file_tools may not be loaded yet # duration_ms (monotonic) is exposed to post_tool_call / transform_tool_result. - _dispatch_start = time.monotonic() + start = time.monotonic() result = _execute_tool( - function_name, function_args, _tool_original_args, ids, + function_name, function_args, original_args, ids, user_task=user_task, enabled_tools=enabled_tools, skip_tool_execution_middleware=skip_tool_execution_middleware, ) - duration_ms = int((time.monotonic() - _dispatch_start) * 1000) - + duration_ms = _elapsed_ms(start) _emit(result, duration_ms=duration_ms) - return _apply_transform_tool_result_hook(function_name, function_args, result, duration_ms, ids) except Exception as e: @@ -1076,37 +1039,33 @@ def handle_function_call( logger.exception(error_msg) return _emit( tool_error(_sanitize_tool_error(error_msg)), - duration_ms=int((time.monotonic() - _dispatch_start) * 1000), - status="error", - error_type=type(e).__name__, - error_message=str(e), + duration_ms=_elapsed_ms(start), status="error", + error_type=type(e).__name__, error_message=str(e), ) # ============================================================================= -# Backward-compat wrapper functions +# Backward-compat wrapper functions (registry pass-throughs) # ============================================================================= def get_all_tool_names() -> List[str]: - """Return all registered tool names.""" return registry.get_all_tool_names() def get_toolset_for_tool(tool_name: str) -> Optional[str]: - """Return the toolset a tool belongs to.""" return registry.get_toolset_for_tool(tool_name) def get_available_toolsets() -> Dict[str, dict]: - """Return toolset availability info for UI display.""" + """Toolset availability info for UI display.""" return registry.get_available_toolsets() def check_toolset_requirements() -> Dict[str, bool]: - """Return {toolset: available_bool} for every registered toolset.""" + """{toolset: available_bool} for every registered toolset.""" return registry.check_toolset_requirements() def check_tool_availability(quiet: bool = False) -> Tuple[List[str], List[dict]]: - """Return (available_toolsets, unavailable_info).""" + """(available_toolsets, unavailable_info).""" return registry.check_tool_availability(quiet=quiet) diff --git a/toolset_distributions.py b/toolset_distributions.py index fa643d312e..c20372524d 100644 --- a/toolset_distributions.py +++ b/toolset_distributions.py @@ -1,22 +1,8 @@ #!/usr/bin/env python3 -""" -Toolset Distributions Module +"""Toolset distributions for batch data-generation runs. -This module defines distributions of toolsets for data generation runs. -Each distribution specifies which toolsets should be used and their probability -of being selected for any given prompt during the batch processing. - -A distribution is a dictionary mapping toolset names to their selection probability (%). -Probabilities should sum to 100, but the system will normalize if they don't. - -Usage: - from toolset_distributions import get_distribution, list_distributions - - # Get a specific distribution - dist = get_distribution("image_gen") - - # List all available distributions - all_dists = list_distributions() +A distribution maps toolset names to the % chance each is enabled for a prompt +(sampled independently, so several toolsets can be active at once). """ from typing import Dict, List, Optional @@ -24,335 +10,101 @@ import random from toolsets import validate_toolset -# Distribution definitions -# Each key is a distribution name, and the value is a dict of toolset_name: probability_percentage +def _dist(description: str, **toolsets: int) -> Dict[str, object]: + return {"description": description, "toolsets": toolsets} + + DISTRIBUTIONS = { - # Default: All tools available 100% of the time - "default": { - "description": "All available tools, all the time", - "toolsets": { - "web": 100, - "vision": 100, - "image_gen": 100, - "terminal": 100, - "file": 100, - "browser": 100 - } - }, - - # Image generation focused distribution - "image_gen": { - "description": "Heavy focus on image generation with vision and web support", - "toolsets": { - "image_gen": 90, # 80% chance of image generation tools - "vision": 90, # 60% chance of vision tools - "web": 55, # 40% chance of web tools - "terminal": 45 - } - }, - - # Research-focused distribution - "research": { - "description": "Web research with vision analysis and reasoning", - "toolsets": { - "web": 90, # 90% chance of web tools - "browser": 70, # 70% chance of browser tools for deep research - "vision": 50, # 50% chance of vision tools - "terminal": 10 # 10% chance of terminal tools - } - }, - - # Scientific problem solving focused distribution - "science": { - "description": "Scientific research with web, terminal, file, and browser capabilities", - "toolsets": { - "web": 94, # 94% chance of web tools - "terminal": 94, # 94% chance of terminal tools - "file": 94, # 94% chance of file tools - "vision": 65, # 65% chance of vision tools - "browser": 50, # 50% chance of browser for accessing papers/databases - "image_gen": 15 # 15% chance of image generation tools - } - }, - - # Development-focused distribution - "development": { - "description": "Terminal, file tools, and reasoning with occasional web lookup", - "toolsets": { - "terminal": 80, # 80% chance of terminal tools - "file": 80, # 80% chance of file tools (read, write, patch, search) - "web": 30, # 30% chance of web tools - "vision": 10 # 10% chance of vision tools - } - }, - - # Safe mode (no terminal) - "safe": { - "description": "All tools except terminal for safety", - "toolsets": { - "web": 80, - "browser": 70, # Browser is safe (no local filesystem access) - "vision": 60, - "image_gen": 60 - } - }, - - # Balanced distribution - "balanced": { - "description": "Equal probability of all toolsets", - "toolsets": { - "web": 50, - "vision": 50, - "image_gen": 50, - "terminal": 50, - "file": 50, - "browser": 50 - } - }, - - # Minimal (web only) - "minimal": { - "description": "Only web tools for basic research", - "toolsets": { - "web": 100 - } - }, - - # Terminal only - "terminal_only": { - "description": "Terminal and file tools for code execution tasks", - "toolsets": { - "terminal": 100, - "file": 100 - } - }, - - # Terminal + web (common for coding tasks that need docs) - "terminal_web": { - "description": "Terminal and file tools with web search for documentation lookup", - "toolsets": { - "terminal": 100, - "file": 100, - "web": 100 - } - }, - - # Creative (vision + image generation) - "creative": { - "description": "Image generation and vision analysis focus", - "toolsets": { - "image_gen": 90, - "vision": 90, - "web": 30 - } - }, - - # Reasoning heavy - "reasoning": { - "description": "Heavy research/reasoning distribution with minimal other tools", - "toolsets": { - "web": 90, - "file": 60, - "terminal": 20 - } - }, - - - # Browser-based web interaction - "browser_use": { - "description": "Full browser-based web interaction with search, vision, and page control", - "toolsets": { - "browser": 100, # All browser tools always available - "web": 80, # Web search for finding URLs and quick lookups - "vision": 70 # Vision analysis for images found on pages - } - }, - - # Browser only (no other tools) - "browser_only": { - "description": "Only browser automation tools for pure web interaction tasks", - "toolsets": { - "browser": 100 - } - }, - - # Browser-focused tasks distribution (for browser-use-tasks.jsonl) - "browser_tasks": { - "description": "Browser-focused distribution (browser toolset includes web_search for finding URLs since Google blocks direct browser searches)", - "toolsets": { - "browser": 97, # 97% - browser tools (includes web_search) almost always available - "vision": 12, # 12% - vision analysis occasionally - "terminal": 15 # 15% - terminal occasionally for local operations - } - }, - - # Terminal-focused tasks distribution (for nous-terminal-tasks.jsonl) - "terminal_tasks": { - "description": "Terminal-focused distribution with high terminal/file availability, occasional other tools", - "toolsets": { - "terminal": 97, # 97% - terminal almost always available - "file": 97, # 97% - file tools almost always available - "web": 97, # 15% - web search/scrape for documentation - "browser": 75, # 10% - browser occasionally for web interaction - "vision": 50, # 8% - vision analysis rarely - "image_gen": 10 # 3% - image generation very rarely - } - }, - - # Mixed browser+terminal tasks distribution (for mixed-browser-terminal-tasks.jsonl) - "mixed_tasks": { - "description": "Mixed distribution with high browser, terminal, and file availability for complex tasks", - "toolsets": { - "browser": 92, # 92% - browser tools highly available - "terminal": 92, # 92% - terminal highly available - "file": 92, # 92% - file tools highly available - "web": 35, # 35% - web search/scrape fairly common - "vision": 15, # 15% - vision analysis occasionally - "image_gen": 15 # 15% - image generation occasionally - } - } + "default": _dist("All available tools, all the time", + web=100, vision=100, image_gen=100, terminal=100, file=100, browser=100), + "image_gen": _dist("Heavy focus on image generation with vision and web support", + image_gen=90, vision=90, web=55, terminal=45), + "research": _dist("Web research with vision analysis and reasoning", + web=90, browser=70, vision=50, terminal=10), + "science": _dist("Scientific research with web, terminal, file, and browser capabilities", + web=94, terminal=94, file=94, vision=65, browser=50, image_gen=15), + "development": _dist("Terminal, file tools, and reasoning with occasional web lookup", + terminal=80, file=80, web=30, vision=10), + "safe": _dist("All tools except terminal for safety", + web=80, browser=70, vision=60, image_gen=60), + "balanced": _dist("Equal probability of all toolsets", + web=50, vision=50, image_gen=50, terminal=50, file=50, browser=50), + "minimal": _dist("Only web tools for basic research", web=100), + "terminal_only": _dist("Terminal and file tools for code execution tasks", terminal=100, file=100), + "terminal_web": _dist("Terminal and file tools with web search for documentation lookup", + terminal=100, file=100, web=100), + "creative": _dist("Image generation and vision analysis focus", image_gen=90, vision=90, web=30), + "reasoning": _dist("Heavy research/reasoning distribution with minimal other tools", + web=90, file=60, terminal=20), + "browser_use": _dist("Full browser-based web interaction with search, vision, and page control", + browser=100, web=80, vision=70), + "browser_only": _dist("Only browser automation tools for pure web interaction tasks", browser=100), + # browser-use-tasks.jsonl: the browser toolset includes web_search since Google blocks direct browser searches + "browser_tasks": _dist( + "Browser-focused distribution (browser toolset includes web_search for finding URLs since Google blocks direct browser searches)", + browser=97, vision=12, terminal=15, + ), + # nous-terminal-tasks.jsonl + "terminal_tasks": _dist( + "Terminal-focused distribution with high terminal/file availability, occasional other tools", + terminal=97, file=97, web=97, browser=75, vision=50, image_gen=10, + ), + # mixed-browser-terminal-tasks.jsonl + "mixed_tasks": _dist( + "Mixed distribution with high browser, terminal, and file availability for complex tasks", + browser=92, terminal=92, file=92, web=35, vision=15, image_gen=15, + ), } def get_distribution(name: str) -> Optional[Dict[str, any]]: - """ - Get a toolset distribution by name. - - Args: - name (str): Name of the distribution - - Returns: - Dict: Distribution definition with description and toolsets - None: If distribution not found - """ + """Distribution definition (description + toolsets), or None if unknown.""" return DISTRIBUTIONS.get(name) def list_distributions() -> Dict[str, Dict]: - """ - List all available distributions. - - Returns: - Dict: All distribution definitions - """ return DISTRIBUTIONS.copy() def sample_toolsets_from_distribution(distribution_name: str) -> List[str]: - """ - Sample toolsets based on a distribution's probabilities. - - Each toolset in the distribution has a % chance of being included. - This allows multiple toolsets to be active simultaneously. - - Args: - distribution_name (str): Name of the distribution to sample from - - Returns: - List[str]: List of sampled toolset names - - Raises: - ValueError: If distribution name is not found + """Sample toolset names, each included independently with its % probability. + + Falls back to the highest-probability toolset when nothing was rolled. + Raises ValueError for an unknown distribution. """ dist = get_distribution(distribution_name) if not dist: raise ValueError(f"Unknown distribution: {distribution_name}") - - # Sample each toolset independently based on its probability + selected_toolsets = [] - for toolset_name, probability in dist["toolsets"].items(): - # Validate toolset exists if not validate_toolset(toolset_name): print(f"āš ļø Warning: Toolset '{toolset_name}' in distribution '{distribution_name}' is not valid") continue - - # Roll the dice - if random value is less than probability, include this toolset if random.random() * 100 < probability: selected_toolsets.append(toolset_name) - - # If no toolsets were selected (can happen with low probabilities), - # ensure at least one toolset is selected by picking the highest probability one + if not selected_toolsets and dist["toolsets"]: - # Find toolset with highest probability highest_prob_toolset = max(dist["toolsets"].items(), key=lambda x: x[1])[0] if validate_toolset(highest_prob_toolset): selected_toolsets.append(highest_prob_toolset) - + return selected_toolsets def validate_distribution(distribution_name: str) -> bool: - """ - Check if a distribution name is valid. - - Args: - distribution_name (str): Distribution name to validate - - Returns: - bool: True if valid, False otherwise - """ return distribution_name in DISTRIBUTIONS def print_distribution_info(distribution_name: str) -> None: - """ - Print detailed information about a distribution. - - Args: - distribution_name (str): Distribution name - """ + """Print a distribution's description and toolset probabilities (highest first).""" dist = get_distribution(distribution_name) if not dist: print(f"āŒ Unknown distribution: {distribution_name}") return - + print(f"\nšŸ“Š Distribution: {distribution_name}") print(f" Description: {dist['description']}") print(" Toolsets:") for toolset, prob in sorted(dist["toolsets"].items(), key=lambda x: x[1], reverse=True): print(f" • {toolset:15} : {prob:3}% chance") - - -if __name__ == "__main__": - """ - Demo and testing of the distributions system - """ - print("šŸ“Š Toolset Distributions Demo") - print("=" * 60) - - # List all distributions - print("\nšŸ“‹ Available Distributions:") - print("-" * 40) - for name, dist in list_distributions().items(): - print(f"\n {name}:") - print(f" {dist['description']}") - toolset_list = ", ".join([f"{ts}({p}%)" for ts, p in dist["toolsets"].items()]) - print(f" Toolsets: {toolset_list}") - - # Demo sampling - print("\n\nšŸŽ² Sampling Examples:") - print("-" * 40) - - test_distributions = ["image_gen", "research", "balanced", "default"] - - for dist_name in test_distributions: - print(f"\n{dist_name}:") - # Sample 5 times to show variability - samples = [] - for _ in range(5): - sampled = sample_toolsets_from_distribution(dist_name) - samples.append(sorted(sampled)) - - print(f" Sample 1: {samples[0]}") - print(f" Sample 2: {samples[1]}") - print(f" Sample 3: {samples[2]}") - print(f" Sample 4: {samples[3]}") - print(f" Sample 5: {samples[4]}") - - # Show detailed info - print("\n\nšŸ“Š Detailed Distribution Info:") - print("-" * 40) - print_distribution_info("image_gen") - print_distribution_info("research") - diff --git a/toolsets.py b/toolsets.py index eda4f8f5fc..a41d025631 100644 --- a/toolsets.py +++ b/toolsets.py @@ -3,45 +3,29 @@ from typing import Dict, List, Any, Set, Optional, Tuple -# Shared tool list for CLI and all messaging platform toolsets. -# Edit this once to update all platforms simultaneously. +# Shared tool list for CLI and all messaging platform toolsets (edit once, all +# platforms follow). Desktop GUI affordances are deliberately NOT here: they live +# in `desktop_ui`/`project`, enabled per desktop-sourced session by the GUI gateway +# (tui_gateway/server.py::_load_enabled_toolsets). HA, kanban and computer_use +# entries are further gated by their tools' check_fns. _HERMES_CORE_TOOLS = [ - # Web "web_search", "web_extract", - # Terminal + process management "terminal", "process_manage", - # Desktop GUI affordances (read_terminal, open_preview, project_*) are NOT - # here: they live in `desktop_ui`/`project`, enabled only by the GUI gateway - # per desktop-sourced session (tui_gateway/server.py::_load_enabled_toolsets). - # File manipulation "read_file", "write_file", "patch", "search_files", - # Vision + image generation "vision_analyze", "image_generate", - # Skills "skills_list", "skill_view", "skill_manage", - # Browser automation "browser_navigate", "browser_snapshot", "browser_click", "browser_type", "browser_scroll", "browser_back", "browser_press", "browser_get_images", "browser_vision", "browser_console", "browser_cdp", "browser_dialog", - # replaces other tools when browser.backend is "browser-use" - "browser_exec", - # Text-to-speech + "browser_exec", # replaces the other browser tools when browser.backend is "browser-use" "text_to_speech", - # Planning & memory "todo_list", "memory", - # Session history search "session_search", - # Clarifying questions "clarify", - # Code execution + delegation "execute_code", "delegate_task", - # Cronjob management "cronjob_manage", - # Home Assistant smart home control (gated on HASS_TOKEN via check_fn) "ha_list_entities", "ha_get_state", "ha_list_services", "ha_call_service", - # Kanban coordination — check_fn in tools/kanban_tools.py admits these only - # for kanban workers (HERMES_KANBAN_TASK) or profiles enabling `kanban`. "kanban_show", "kanban_list", "kanban_complete", "kanban_block", "kanban_request_review", "kanban_request_changes", @@ -49,17 +33,11 @@ _HERMES_CORE_TOOLS = [ "kanban_comment", "kanban_create", "kanban_link", "kanban_unblock", "kanban_attach", "kanban_attach_url", "kanban_attachments", - # Computer use (macOS, gated on cua-driver being installed via check_fn) "computer_use", ] # Webhook payloads are untrusted third-party content: no file/system execution. -_HERMES_WEBHOOK_SAFE_TOOLS = [ - "web_search", - "web_extract", - "vision_analyze", - "clarify", -] +_HERMES_WEBHOOK_SAFE_TOOLS = ["web_search", "web_extract", "vision_analyze", "clarify"] def _ts(description, tools=(), includes=(), **extra): @@ -303,8 +281,7 @@ TOOLSETS = { "hermes-yuanbao": { "description": "Yuanbao Bot å…ƒå®ę¶ˆęÆå¹³å°å·„å…·é›† - ē¾¤äæ”ęÆć€ęˆå‘˜ęŸ„čÆ¢ć€ē§čŠć€č““ēŗøč”Øęƒ…", "tools": _HERMES_CORE_TOOLS + [ - "yb_query_group_info", "yb_query_group_members", "yb_send_dm", - "yb_search_sticker", "yb_send_sticker", + "yb_query_group_info", "yb_query_group_members", "yb_send_dm", "yb_search_sticker", "yb_send_sticker", ], "module": "tools.yuanbao_tools", "includes": [], @@ -337,19 +314,22 @@ def _registry(): return None +def _registry_call(method: str, default): + """registry.() or *default* when the registry is unavailable or the call fails.""" + registry = _registry() + if registry is None: + return default + try: + return getattr(registry, method)() + except Exception: + return default + + def _registry_generation() -> Tuple[int, int]: reg = _registry() return (id(reg), getattr(reg, "_generation", 0)) if reg is not None else (0, 0) -def _static_copy(toolset: Dict[str, Any]) -> Dict[str, Any]: - return { - **toolset, - "tools": list(toolset.get("tools", [])), - "includes": list(toolset.get("includes", [])), - } - - def get_toolset(name: str, *, include_registry: bool = True) -> Optional[Dict[str, Any]]: """Return a toolset definition, or None if unknown. @@ -360,7 +340,7 @@ def get_toolset(name: str, *, include_registry: bool = True) -> Optional[Dict[st """ toolset = TOOLSETS.get(name) if not include_registry: - return _static_copy(toolset) if toolset else None + return {**toolset, "tools": list(toolset["tools"]), "includes": list(toolset["includes"])} if toolset else None registry = _registry() if registry is None: @@ -399,12 +379,12 @@ def bundle_non_core_tools(toolset_name: str) -> Set[str]: ts_def = get_toolset(toolset_name) if not (ts_def and "tools" in ts_def): return set(resolve_toolset(toolset_name)) - core - to_remove = set(ts_def["tools"]) - core + to_remove = set(ts_def["tools"]) for inc in ts_def.get("includes", []): inc_def = get_toolset(inc) if inc_def and "tools" in inc_def: - to_remove.update(set(inc_def["tools"]) - core) - return to_remove + to_remove.update(inc_def["tools"]) + return to_remove - core # Memo keyed on (name, include_registry, id(registry), registry generation). @@ -479,23 +459,11 @@ def resolve_toolset(name: str, visited: Set[str] = None, *, include_registry: bo def _get_plugin_toolset_names() -> Set[str]: """Registry toolset names absent from the static TOOLSETS dict.""" - registry = _registry() - if registry is None: - return set() - try: - return {n for n in registry.get_registered_toolset_names() if n not in TOOLSETS} - except Exception: - return set() + return {n for n in _registry_call("get_registered_toolset_names", ()) if n not in TOOLSETS} def _get_registry_toolset_aliases() -> Dict[str, str]: - registry = _registry() - if registry is None: - return {} - try: - return registry.get_registered_toolset_aliases() - except Exception: - return {} + return _registry_call("get_registered_toolset_aliases", {}) def _display_alias(ts_name: str, aliases: Dict[str, str]) -> Optional[str]: @@ -539,11 +507,7 @@ def create_custom_toolset( includes: List[str] = None ) -> None: """Register a runtime toolset in TOOLSETS.""" - TOOLSETS[name] = { - "description": description, - "tools": tools or [], - "includes": includes or [] - } + TOOLSETS[name] = _ts(description, tools or [], includes or []) def get_toolset_info(name: str) -> Dict[str, Any]: