refactor(model_tools,toolsets): dispatcher phase helpers, shared registry accessor, table-built distributions

This commit is contained in:
Teknium
2026-09-02 18:17:45 -07:00
parent 28c9b74a43
commit 47ac76a03a
3 changed files with 220 additions and 545 deletions
+138 -179
View File
@@ -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'</?(?:tool_call|function_call|result|response|output|input|system|assistant|user)>',
re.IGNORECASE,
_TOOL_ERROR_STRIP_RES = (
re.compile(
r'</?(?:tool_call|function_call|result|response|output|input|system|assistant|user)>',
re.IGNORECASE,
),
re.compile(r'^\s*```(?:json|xml|html|markdown)?\s*', re.MULTILINE),
re.compile(r'\s*```\s*$', re.MULTILINE),
re.compile(r'<!\[CDATA\[.*?\]\]>', 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'<!\[CDATA\[.*?\]\]>', 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)
+56 -304
View File
@@ -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")
+26 -62
View File
@@ -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.<method>() 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]: