refactor(model_tools,toolsets,mcp_serve): compact docstrings/comments, derive toolset lists from core, collapse defensive layers

This commit is contained in:
Teknium
2026-09-02 18:24:17 -07:00
parent 47ac76a03a
commit 7fd8ccff62
3 changed files with 180 additions and 386 deletions
+73 -162
View File
@@ -35,9 +35,7 @@ except ImportError:
MCPServer = None # type: ignore[assignment,misc]
# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------
# --- Helpers -----------------------------------------------------------------
def _hermes_home() -> Path:
try:
@@ -117,23 +115,16 @@ def _load_sessions_index() -> dict:
def _row_to_index_entry(row: dict) -> dict:
"""Convert a state.db gateway session row to the sessions.json entry shape."""
origin = {}
origin_json = row.get("origin_json")
if origin_json:
if row.get("origin_json"):
try:
parsed = json.loads(origin_json)
parsed = json.loads(row["origin_json"])
if isinstance(parsed, dict):
origin = parsed
except (TypeError, ValueError):
pass
if not origin:
# Pre-origin_json rows: synthesize the minimal origin from columns.
origin = {
"platform": row.get("source", ""),
"chat_id": row.get("chat_id"),
"chat_type": row.get("chat_type"),
"thread_id": row.get("thread_id"),
"user_id": row.get("user_id"),
}
if not origin: # pre-origin_json rows: synthesize the minimal origin from columns
origin = {k: row.get(k) for k in ("chat_id", "chat_type", "thread_id", "user_id")}
origin["platform"] = row.get("source", "")
def _iso(ts) -> str:
try:
@@ -144,16 +135,14 @@ def _row_to_index_entry(row: dict) -> dict:
input_tokens = int(row.get("input_tokens") or 0)
output_tokens = int(row.get("output_tokens") or 0)
return {
"session_id": str(row.get("id", "")),
"session_key": row.get("session_key", ""),
"session_id": str(row.get("id", "")), "session_key": row.get("session_key", ""),
"platform": row.get("source", ""),
"chat_type": row.get("chat_type") or origin.get("chat_type", ""),
"display_name": row.get("display_name") or origin.get("chat_name") or "",
"origin": origin,
"created_at": _iso(row.get("started_at")),
"updated_at": _iso(row.get("last_active") or row.get("started_at")),
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"input_tokens": input_tokens, "output_tokens": output_tokens,
"total_tokens": input_tokens + output_tokens,
}
@@ -221,21 +210,20 @@ def _extract_attachments(msg: dict) -> List[dict]:
attachments = []
content = msg.get("content", "")
if isinstance(content, list):
for part in content:
if not isinstance(part, dict):
continue
ptype = part.get("type", "")
if ptype == "image_url":
url = part.get("image_url", {}).get("url", "") if isinstance(part.get("image_url"), dict) else ""
if url:
attachments.append({"type": "image", "url": url})
elif ptype == "image":
url = part.get("url", part.get("source", {}).get("url", ""))
if url:
attachments.append({"type": "image", "url": url})
elif ptype != "text":
for part in content if isinstance(content, list) else ():
if not isinstance(part, dict):
continue
ptype = part.get("type", "")
if ptype == "image_url":
url = part.get("image_url", {}).get("url", "") if isinstance(part.get("image_url"), dict) else ""
elif ptype == "image":
url = part.get("url", part.get("source", {}).get("url", ""))
else:
if ptype != "text":
attachments.append({"type": ptype, "data": part})
continue
if url:
attachments.append({"type": "image", "url": url})
text = _extract_message_content(msg)
if text:
@@ -245,9 +233,7 @@ def _extract_attachments(msg: dict) -> List[dict]:
return attachments
# ---------------------------------------------------------------------------
# Event Bridge — polls SessionDB for new messages, maintains event queue
# ---------------------------------------------------------------------------
# --- Event Bridge — polls SessionDB for new messages, maintains event queue ---
QUEUE_LIMIT = 1000
POLL_INTERVAL = 0.2 # seconds between DB polls (200ms)
@@ -269,15 +255,15 @@ def _ts_float(ts) -> float:
"""Normalize a message timestamp (epoch int/float or ISO string) to float."""
if isinstance(ts, (int, float)):
return float(ts)
if isinstance(ts, str) and ts:
if not (isinstance(ts, str) and ts):
return 0.0
try:
return float(ts)
except ValueError:
try:
return float(ts)
except ValueError:
try:
return datetime.fromisoformat(ts).timestamp()
except Exception:
return 0.0
return 0.0
return datetime.fromisoformat(ts).timestamp()
except Exception:
return 0.0
def _latest_ts(messages) -> float:
@@ -299,8 +285,7 @@ class EventBridge:
self._thread: Optional[threading.Thread] = None
self._last_poll_timestamps: Dict[str, float] = {} # session_key -> unix timestamp
self._pending_approvals: Dict[str, dict] = {} # populated from events
# mtime cache — skip expensive work when state.db hasn't changed
self._state_db_mtime: float = 0.0
self._state_db_mtime: float = 0.0 # skip polling work when state.db is unchanged
self._cached_sessions_index: dict = {}
def start(self):
@@ -336,12 +321,7 @@ class EventBridge:
events = self._matching(after_cursor, session_key, limit)
return {"events": events, "next_cursor": events[-1]["cursor"] if events else after_cursor}
def wait_for_event(
self,
after_cursor: int = 0,
session_key: Optional[str] = None,
timeout_ms: int = 30000,
) -> Optional[dict]:
def wait_for_event(self, after_cursor: int = 0, session_key: Optional[str] = None, timeout_ms: int = 30000) -> Optional[dict]:
"""Block until a matching event arrives or timeout expires."""
deadline = time.monotonic() + (timeout_ms / 1000.0)
while time.monotonic() < deadline:
@@ -366,11 +346,9 @@ class EventBridge:
approval = self._pending_approvals.pop(approval_id, None)
if not approval:
return {"error": f"Approval not found: {approval_id}"}
self._enqueue(QueueEvent(
cursor=0, # set by _enqueue
type="approval_resolved",
session_key=approval.get("session_key", ""),
data={"approval_id": approval_id, "decision": decision},
self._enqueue(QueueEvent( # cursor is assigned by _enqueue
0, "approval_resolved", approval.get("session_key", ""),
{"approval_id": approval_id, "decision": decision},
))
return {"resolved": True, "approval_id": approval_id, "decision": decision}
@@ -466,26 +444,17 @@ class EventBridge:
content = _extract_message_content(msg)
if not content:
continue
self._enqueue(QueueEvent(
cursor=0,
type="message",
session_key=session_key,
data={
"role": msg.get("role", ""),
"content": content[:500],
"timestamp": str(msg.get("timestamp", "")),
"message_id": str(msg.get("id", "")),
},
))
self._enqueue(QueueEvent(0, "message", session_key, {
"role": msg.get("role", ""), "content": content[:500],
"timestamp": str(msg.get("timestamp", "")), "message_id": str(msg.get("id", "")),
}))
latest = _latest_ts(messages)
if latest > last_seen:
self._last_poll_timestamps[session_key] = latest
# ---------------------------------------------------------------------------
# MCP Server
# ---------------------------------------------------------------------------
# --- MCP Server ---------------------------------------------------------------
def _conversation_messages(session_key: str):
"""(messages, error_json) for a conversation; exactly one is None."""
@@ -515,12 +484,7 @@ class _ToolHandlers:
def __init__(self, bridge: EventBridge):
self.bridge = bridge
def conversations_list(
self,
platform: Optional[str] = None,
limit: int = 50,
search: Optional[str] = None,
) -> str:
def conversations_list(self, platform: Optional[str] = None, limit: int = 50, search: Optional[str] = None) -> str:
"""List active messaging conversations across connected platforms.
Returns conversations with their session keys (needed for messages_read),
@@ -540,21 +504,13 @@ class _ToolHandlers:
continue
display_name = entry.get("display_name", "")
chat_name = origin.get("chat_name", "")
if search:
search_lower = search.lower()
if (search_lower not in display_name.lower()
and search_lower not in chat_name.lower()
and search_lower not in key.lower()):
continue
if search and not any(search.lower() in s.lower() for s in (display_name, chat_name, key)):
continue
conversations.append({
"session_key": key,
"session_id": entry.get("session_id", ""),
"platform": entry_platform,
"session_key": key, "session_id": entry.get("session_id", ""), "platform": entry_platform,
"chat_type": entry.get("chat_type", origin.get("chat_type", "")),
"display_name": display_name,
"chat_name": chat_name,
"user_name": origin.get("user_name", ""),
"updated_at": entry.get("updated_at", ""),
"display_name": display_name, "chat_name": chat_name,
"user_name": origin.get("user_name", ""), "updated_at": entry.get("updated_at", ""),
})
conversations.sort(key=lambda c: c.get("updated_at", ""), reverse=True)
@@ -577,22 +533,14 @@ class _ToolHandlers:
"platform": entry.get("platform") or origin.get("platform", ""),
"chat_type": entry.get("chat_type", origin.get("chat_type", "")),
"display_name": entry.get("display_name", ""),
"user_name": origin.get("user_name", ""),
"chat_name": origin.get("chat_name", ""),
"chat_id": origin.get("chat_id", ""),
"thread_id": origin.get("thread_id"),
"updated_at": entry.get("updated_at", ""),
"created_at": entry.get("created_at", ""),
"input_tokens": entry.get("input_tokens", 0),
"output_tokens": entry.get("output_tokens", 0),
"user_name": origin.get("user_name", ""), "chat_name": origin.get("chat_name", ""),
"chat_id": origin.get("chat_id", ""), "thread_id": origin.get("thread_id"),
"updated_at": entry.get("updated_at", ""), "created_at": entry.get("created_at", ""),
"input_tokens": entry.get("input_tokens", 0), "output_tokens": entry.get("output_tokens", 0),
"total_tokens": entry.get("total_tokens", 0),
}, indent=2)
def messages_read(
self,
session_key: str,
limit: int = 50,
) -> str:
def messages_read(self, session_key: str, limit: int = 50) -> str:
"""Read recent messages from a conversation.
Returns the message history in chronological order with role, content,
@@ -609,28 +557,19 @@ class _ToolHandlers:
filtered = []
for msg in all_messages:
role = msg.get("role", "")
if role in {"user", "assistant"}:
content = _extract_message_content(msg)
if content:
filtered.append({
"id": str(msg.get("id", "")),
"role": role,
"content": content[:2000],
"timestamp": msg.get("timestamp", ""),
})
content = _extract_message_content(msg) if role in {"user", "assistant"} else ""
if content:
filtered.append({
"id": str(msg.get("id", "")), "role": role,
"content": content[:2000], "timestamp": msg.get("timestamp", ""),
})
messages = filtered[-limit:]
return json.dumps({
"session_key": session_key,
"count": len(messages),
"total_in_session": len(filtered),
"messages": messages,
"session_key": session_key, "count": len(messages),
"total_in_session": len(filtered), "messages": messages,
}, indent=2)
def attachments_fetch(
self,
session_key: str,
message_id: str,
) -> str:
def attachments_fetch(self, session_key: str, message_id: str) -> str:
"""List non-text attachments for a message in a conversation.
Extracts images, media files, and other non-text content blocks
@@ -647,18 +586,9 @@ class _ToolHandlers:
if not target_msg:
return json.dumps({"error": f"Message not found: {message_id}"})
attachments = _extract_attachments(target_msg)
return json.dumps({
"message_id": message_id,
"count": len(attachments),
"attachments": attachments,
}, indent=2)
return json.dumps({"message_id": message_id, "count": len(attachments), "attachments": attachments}, indent=2)
def events_poll(
self,
after_cursor: int = 0,
session_key: Optional[str] = None,
limit: int = 20,
) -> str:
def events_poll(self, after_cursor: int = 0, session_key: Optional[str] = None, limit: int = 20) -> str:
"""Poll for new conversation events since a cursor position.
Returns events that have occurred since the given cursor. Use the
@@ -676,12 +606,7 @@ class _ToolHandlers:
result = self.bridge.poll_events(after_cursor=after_cursor, session_key=session_key, limit=limit)
return json.dumps(result, indent=2)
def events_wait(
self,
after_cursor: int = 0,
session_key: Optional[str] = None,
timeout_ms: int = 30000,
) -> str:
def events_wait(self, after_cursor: int = 0, session_key: Optional[str] = None, timeout_ms: int = 30000) -> str:
"""Wait for the next conversation event (long-poll).
Blocks until a matching event arrives or the timeout expires.
@@ -699,11 +624,7 @@ class _ToolHandlers:
return json.dumps({"event": event}, indent=2)
return json.dumps({"event": None, "reason": "timeout"}, indent=2)
def messages_send(
self,
target: str,
message: str,
) -> str:
def messages_send(self, target: str, message: str) -> str:
"""Send a message to a platform conversation.
The target format is "platform:chat_id" — same format used by the
@@ -754,8 +675,7 @@ class _ToolHandlers:
continue
seen.add(target_str)
targets.append({
"target": target_str,
"platform": p,
"target": target_str, "platform": p,
"name": entry.get("display_name") or origin.get("chat_name", ""),
"chat_type": entry.get("chat_type", origin.get("chat_type", "")),
})
@@ -769,10 +689,8 @@ class _ToolHandlers:
if isinstance(ch, dict):
chat_id = ch.get("id", ch.get("chat_id", ""))
channels.append({
"target": f"{plat}:{chat_id}" if chat_id else plat,
"platform": plat,
"name": ch.get("name", ch.get("display_name", "")),
"chat_type": ch.get("type", ""),
"target": f"{plat}:{chat_id}" if chat_id else plat, "platform": plat,
"name": ch.get("name", ch.get("display_name", "")), "chat_type": ch.get("type", ""),
})
return json.dumps({"count": len(channels), "channels": channels}, indent=2)
@@ -786,11 +704,7 @@ class _ToolHandlers:
approvals = self.bridge.list_pending_approvals()
return json.dumps({"count": len(approvals), "approvals": approvals}, indent=2)
def permissions_respond(
self,
id: str,
decision: str,
) -> str:
def permissions_respond(self, id: str, decision: str) -> str:
"""Respond to a pending approval request.
Args:
@@ -820,14 +734,11 @@ def create_mcp_server(event_bridge: Optional[EventBridge] = None) -> "MCPServer"
"MCP server requires the 'mcp' package. "
f"Install with: {sys.executable} -m pip install 'mcp'"
)
mcp = MCPServer(
"hermes",
instructions=(
"Hermes Agent messaging bridge. Use these tools to interact with "
"conversations across Telegram, Discord, Slack, WhatsApp, Signal, "
"Matrix, and other connected platforms."
),
)
mcp = MCPServer("hermes", instructions=(
"Hermes Agent messaging bridge. Use these tools to interact with "
"conversations across Telegram, Discord, Slack, WhatsApp, Signal, "
"Matrix, and other connected platforms."
))
handlers = _ToolHandlers(event_bridge or EventBridge())
for name in _TOOL_NAMES:
mcp.tool()(getattr(handlers, name))
+69 -136
View File
@@ -1,10 +1,9 @@
"""
Model Tools Module
"""Thin orchestration layer over the tool registry.
Thin orchestration layer over the tool registry: importing this module runs tool
discovery (each tools/*.py self-registers via tools.registry.register()), then
exposes get_tool_definitions() (schemas sent to the model, toolset-filtered) and
handle_function_call() (dispatch with hooks/middleware) plus registry wrappers.
Importing runs tool discovery (each tools/*.py self-registers via
tools.registry.register()); exposes get_tool_definitions() (toolset-filtered
schemas sent to the model) and handle_function_call() (dispatch with
hooks/middleware) plus registry pass-throughs.
"""
import os
@@ -61,7 +60,6 @@ _WARNED_DISABLED_BUNDLES: set = set()
def _is_delegated_child_context() -> bool:
try:
from agent.delegation_context import is_delegated_child_context
return is_delegated_child_context()
except Exception:
return False
@@ -72,7 +70,6 @@ def _is_dispatcher_owned_worker() -> bool:
(delegate_task child, or a cron job fired in-process from a worker)."""
try:
from agent.delegation_context import is_dispatcher_owned_worker_context
return is_dispatcher_owned_worker_context()
except Exception:
return True
@@ -131,15 +128,12 @@ def _run_async(coro):
asyncio.set_event_loop(worker_loop)
return worker_loop.run_until_complete(coro)
finally:
try:
# Drain tasks still pending after an external cancel.
try: # drain tasks still pending after an external cancel
pending = asyncio.all_tasks(worker_loop)
for t in pending:
t.cancel()
if pending:
worker_loop.run_until_complete(
asyncio.gather(*pending, return_exceptions=True)
)
worker_loop.run_until_complete(asyncio.gather(*pending, return_exceptions=True))
except Exception:
pass
worker_loop.close()
@@ -147,7 +141,6 @@ def _run_async(coro):
pool = concurrent.futures.ThreadPoolExecutor(max_workers=1)
# Carry profile + approval/sudo context so get_hermes_home() resolves correctly.
from tools.thread_context import propagate_context_to_thread
future = pool.submit(propagate_context_to_thread(_run_in_worker))
try:
return future.result(timeout=300)
@@ -161,8 +154,7 @@ def _run_async(coro):
pass # loop already closed
raise
finally:
# wait=False: never block the caller on a stuck coroutine.
pool.shutdown(wait=False)
pool.shutdown(wait=False) # never block the caller on a stuck coroutine
if threading.current_thread() is not threading.main_thread():
return _get_worker_loop().run_until_complete(coro)
@@ -236,24 +228,23 @@ def get_tool_definitions(
) -> List[Dict[str, Any]]:
"""Tool definitions for model API calls, filtered by toolset.
Args:
enabled_toolsets: Only include tools from these toolsets (None = all).
disabled_toolsets: Toolsets subtracted after enabling.
quiet_mode: Suppress status prints (and enable memoization).
skip_tool_search_assembly: Return the pre-assembly list (raw schemas for
every enabled tool). Only the tool_search bridge should use this so
it reads the real catalog rather than the collapsed one.
enabled_toolsets None = all; disabled_toolsets are subtracted after enabling.
quiet_mode suppresses status prints and enables memoization.
skip_tool_search_assembly returns raw schemas for every enabled tool — only
the tool_search bridge should use it (it reads the real, uncollapsed catalog).
"""
if not quiet_mode:
def compute():
return _compute_tool_definitions(enabled_toolsets, disabled_toolsets, quiet_mode,
skip_tool_search_assembly=skip_tool_search_assembly)
if not quiet_mode:
return compute()
cache_key = _tool_defs_cache_key(enabled_toolsets, disabled_toolsets, skip_tool_search_assembly)
with _tool_defs_cache_lock:
cached = _tool_defs_cache.get(cache_key) if cache_key is not None else None
if cached is None:
result = _compute_tool_definitions(enabled_toolsets, disabled_toolsets, quiet_mode,
skip_tool_search_assembly=skip_tool_search_assembly)
result = compute()
if cache_key is None:
return list(result)
with _tool_defs_cache_lock:
@@ -279,10 +270,9 @@ def _tool_defs_cache_key(
) -> Optional[tuple]:
"""Memo key for get_tool_definitions, or None when caching must be bypassed.
Covers every argument plus everything that changes the result without an
argument changing: registry generation, config.yaml mtime/size (dynamic
schemas: execute_code mode, discord allowlist), kanban context, profile
scope. check_fn results are TTL-cached inside registry.get_definitions.
Covers every argument plus everything that changes the result without one:
registry generation, config.yaml mtime/size (dynamic schemas), kanban
context, profile scope. check_fn results are TTL-cached in the registry.
"""
profile_scope = check_fn_cache_scope()
if profile_scope == CHECK_FN_CACHE_BYPASS:
@@ -297,13 +287,9 @@ def _tool_defs_cache_key(
registry.current_scope_key(),
frozenset(enabled_toolsets) if enabled_toolsets is not None else None,
frozenset(disabled_toolsets) if disabled_toolsets else None,
registry._generation,
cfg_fp,
bool(os.environ.get("HERMES_KANBAN_TASK")),
bool(skip_tool_search_assembly),
_is_delegated_child_context(),
_is_dispatcher_owned_worker(),
profile_scope,
registry._generation, cfg_fp,
bool(os.environ.get("HERMES_KANBAN_TASK")), bool(skip_tool_search_assembly),
_is_delegated_child_context(), _is_dispatcher_owned_worker(), profile_scope,
)
@@ -429,26 +415,20 @@ def _rewrite_delegate_task(td: Dict[str, Any], available: set) -> Optional[Dict[
return td
fn = td.get("function", {})
desc = fn.get("description", "")
full_offvariant = "delegate_task, clarify, memory, or cronjob"
full_onvariant = "clarify, memory, or cronjob"
if full_offvariant in desc:
full, names = full_offvariant, ["delegate_task"] + blocked_present
elif full_onvariant in desc:
full, names = full_onvariant, blocked_present
for full, self_named in (("delegate_task, clarify, memory, or cronjob", True), ("clarify, memory, or cronjob", False)):
if full in desc:
break
else:
return td
if blocked_present:
if len(names) <= 2:
replacement = " or ".join(names)
else:
replacement = ", ".join(names[:-1]) + ", or " + names[-1]
names = (["delegate_task"] if self_named else []) + blocked_present
replacement = " or ".join(names) if len(names) <= 2 else ", ".join(names[:-1]) + ", or " + names[-1]
desc = desc.replace(full, replacement)
else:
# Both variants end at the following newline.
start = desc.find("- Children cannot call " + full)
if start != -1:
end = desc.index("\n", start) + 1
desc = desc[:start] + desc[end:]
desc = desc[:start] + desc[desc.index("\n", start) + 1:]
return {**td, "function": {**fn, "description": desc}}
@@ -495,17 +475,15 @@ def _compute_tool_definitions(
tools_to_include = _select_tool_names(enabled_toolsets, disabled_toolsets, quiet_mode)
# Registry returns only tools whose check_fn passes.
filtered_tools = _apply_dynamic_schemas(registry.get_definitions(tools_to_include, quiet=quiet_mode))
global _last_resolved_tool_names
_last_resolved_tool_names = [t["function"]["name"] for t in filtered_tools]
if not quiet_mode:
if filtered_tools:
tool_names = [t["function"]["name"] for t in filtered_tools]
print(f"🛠️ Final tool selection ({len(filtered_tools)} tools): {', '.join(tool_names)}")
print(f"🛠️ Final tool selection ({len(filtered_tools)} tools): {', '.join(_last_resolved_tool_names)}")
else:
print("🛠️ No tools selected (all filtered out or unavailable)")
global _last_resolved_tool_names
_last_resolved_tool_names = [t["function"]["name"] for t in filtered_tools]
# Normalize schema shapes llama.cpp's grammar converter rejects (bare
# "type": "object", string-valued nodes from malformed MCP servers).
try:
@@ -522,11 +500,7 @@ def _compute_tool_definitions(
from tools.tool_search import assemble_tool_defs, load_config as _load_ts_config
ts_cfg = _load_ts_config()
if not skip_tool_search_assembly and ts_cfg.enabled != "off":
assembly = assemble_tool_defs(
filtered_tools,
context_length=_resolve_active_context_length(),
config=ts_cfg,
)
assembly = assemble_tool_defs(filtered_tools, context_length=_resolve_active_context_length(), config=ts_cfg)
if assembly.activated and not quiet_mode:
print(
f"🔎 Tool Search (tier {assembly.tier}): {assembly.deferred_count} "
@@ -580,11 +554,8 @@ def _resolve_active_context_length() -> int:
base_url = str(rt.get("base_url") or base_url or "").strip()
api_key = str(rt.get("api_key") or "").strip()
except Exception as rt_exc:
logger.debug(
"Runtime credential resolution failed for tool-search "
"context gate (provider=%s): %s — using config values only",
provider, rt_exc,
)
logger.debug("Runtime credential resolution failed for tool-search "
"context gate (provider=%s): %s — using config values only", provider, rt_exc)
if config_ctx is None and base_url:
try:
cached_ctx = get_cached_context_length(model_id, base_url)
@@ -622,10 +593,9 @@ _READ_SEARCH_TOOLS = {"read_file", "search_files"}
# --- 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.
# Defense-in-depth: strip role tags / CDATA / code fences from exception text the
# model will read, and cap length (cap shared with tools/registry.py so text never
# passes two different caps with two different markers).
_TOOL_ERROR_STRIP_RES = (
re.compile(
r'</?(?:tool_call|function_call|result|response|output|input|system|assistant|user)>',
@@ -660,14 +630,11 @@ class _CallIds:
api_request_id: Optional[str] = None
def hook_kwargs(self) -> Dict[str, str]:
"""The same fields with None normalized to "" (hook/middleware wire contract)."""
"""Same fields with None -> "" (hook/middleware wire contract)."""
return {k: v or "" for k, v in asdict(self).items()}
def _tool_result_observer_fields(
tool_name: str,
result: Any,
) -> tuple[str, Optional[str], Optional[str]]:
def _tool_result_observer_fields(tool_name: str, result: Any) -> tuple[str, Optional[str], Optional[str]]:
"""Derive (status, error_type, error_message) from a tool result for observer hooks."""
try:
parsed_result = json.loads(result) if isinstance(result, str) else result
@@ -677,7 +644,6 @@ def _tool_result_observer_fields(
pass
try:
from agent.display import _detect_tool_failure
failed, suffix = _detect_tool_failure(tool_name, result)
if failed:
return "error", "tool_error", suffix.strip().strip("[]") or None
@@ -702,9 +668,8 @@ 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; ok/error fields are derived from
the result only after that gate when ``status`` is None)."""
"""Emit the ``post_tool_call`` observer hook; gated on has_hook, and ok/error
fields are derived from the result only past that gate when status is None."""
if _post_tool_call_hook_suppressed.get():
return
try:
@@ -714,15 +679,9 @@ def _emit_post_tool_call_hook(
if status is None:
status, error_type, error_message = _tool_result_observer_fields(function_name, result)
invoke_hook(
"post_tool_call",
tool_name=function_name,
args=function_args,
result=result,
"post_tool_call", tool_name=function_name, args=function_args, result=result,
**_CallIds(task_id, session_id, tool_call_id, turn_id, api_request_id).hook_kwargs(),
duration_ms=duration_ms,
status=status,
error_type=error_type,
error_message=error_message,
duration_ms=duration_ms, status=status, error_type=error_type, error_message=error_message,
middleware_trace=list(middleware_trace or []),
)
except Exception as _hook_err:
@@ -737,10 +696,9 @@ def _dispatch_bridge_tool(
):
"""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.
None when *function_name* is not a bridge tool; ``(result, None)`` for a
finished catalog read or error; ``(None, (name, args))`` when a validated
tool_call should be re-dispatched as the real tool.
"""
try:
from tools import tool_search as ts
@@ -748,13 +706,11 @@ def _dispatch_bridge_tool(
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.
# Un-collapsed catalog scoped to the session's toolsets, so a restricted
# session (subagent, kanban worker) can't reach the whole registry via the bridge.
try:
current_defs = get_tool_definitions(
enabled_toolsets=enabled_toolsets,
disabled_toolsets=disabled_toolsets,
enabled_toolsets=enabled_toolsets, disabled_toolsets=disabled_toolsets,
quiet_mode=True, skip_tool_search_assembly=True,
) or []
except Exception:
@@ -791,7 +747,6 @@ def _apply_request_middleware(
"""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:
@@ -808,12 +763,11 @@ def _pre_dispatch_guards(
) -> Tuple[Dict[str, Any], Optional[Tuple[Any, str, Optional[str]]]]:
"""Plugin pre_tool_call hook, then ACP edit approval.
Returns ``(args, None)`` to proceed (args possibly modified by a plugin), or
``(args, (result, error_type, error_message))`` when the call is blocked.
``(args, None)`` to proceed (args possibly plugin-modified), or
``(args, (result, error_type, error_message))`` when blocked.
"""
# pre_tool_call fires exactly once per execution: _dispatch_pre_tool_call_hooks
# returns the block message (block/approve) and modified args (modify) from a
# single invoke_hook pass. skip=True means the caller already fired it.
# pre_tool_call fires exactly once per execution: one invoke_hook pass yields
# both the block message and modified args. skip=True: caller already fired it.
if not skip_pre_tool_call_hook:
block_message: Optional[str] = None
try:
@@ -832,16 +786,13 @@ def _pre_dispatch_guards(
# via ContextVar only for ACP sessions, so CLI/gateway paths are unaffected.
try:
from acp_adapter.edit_approval import maybe_require_edit_approval
edit_block_message = maybe_require_edit_approval(function_name, function_args)
if edit_block_message is not None:
return function_args, (edit_block_message, "edit_approval_denied", None)
except Exception as _edit_approval_err:
logger.debug("ACP edit approval guard error: %s", _edit_approval_err)
if function_name in {"write_file", "patch"}:
return function_args, (
tool_error("Edit approval denied: approval guard failed"), "edit_approval_error", None,
)
return function_args, (tool_error("Edit approval denied: approval guard failed"), "edit_approval_error", None)
return function_args, None
@@ -892,7 +843,6 @@ def _execute_tool(
if skip_tool_execution_middleware:
return _dispatch(function_args)
from hermes_cli.middleware import run_tool_execution_middleware
return run_tool_execution_middleware(
function_name, function_args, _dispatch, original_args=original_args, **ids.hook_kwargs(),
)
@@ -907,24 +857,17 @@ def _apply_transform_tool_result_hook(
) -> Any:
"""transform_tool_result: plugins may replace the final result string.
Runs after post_tool_call (observational) and before the result enters
context. Fail-open; first valid string return wins; non-strings ignored.
Gated on has_hook so the no-listener path skips result-field derivation.
Runs after post_tool_call and before the result enters context. Fail-open;
first string return wins. Gated on has_hook so the no-listener path is cheap.
"""
try:
from hermes_cli.lifecycle import has_hook, invoke_hook
if has_hook("transform_tool_result"):
status, error_type, error_message = _tool_result_observer_fields(function_name, result)
hook_results = invoke_hook(
"transform_tool_result",
tool_name=function_name,
args=function_args,
result=result,
**ids.hook_kwargs(),
duration_ms=duration_ms,
status=status,
error_type=error_type,
error_message=error_message,
"transform_tool_result", tool_name=function_name, args=function_args, result=result,
**ids.hook_kwargs(), duration_ms=duration_ms,
status=status, error_type=error_type, error_message=error_message,
)
for hook_result in hook_results:
if isinstance(hook_result, str):
@@ -957,16 +900,11 @@ def handle_function_call(
) -> str:
"""Route a tool call through hooks/middleware to the registry; returns a JSON string.
Args:
task_id: Terminal/browser session isolation key.
user_task: The user's original task (browser_snapshot context).
enabled_tools: Session tool names; execute_code uses them to pick sandbox
tools (falls back to the process-global ``_last_resolved_tool_names``).
skip_pre_tool_call_hook: Caller already fired pre_tool_call (single-fire contract).
enabled_toolsets / disabled_toolsets: The session's toolset selection,
used to scope the Tool Search bridge catalog so tool_search /
tool_describe / tool_call only see tools this session was granted.
None = no restriction, matching get_tool_definitions semantics.
task_id isolates terminal/browser sessions; user_task feeds browser_snapshot.
enabled_tools picks execute_code's sandbox tools (default: the process-global
``_last_resolved_tool_names``). skip_pre_tool_call_hook: caller already fired
it (single-fire contract). enabled/disabled_toolsets scope the Tool Search
bridge catalog to this session's grant (None = unrestricted).
"""
function_args = coerce_tool_args(function_name, function_args)
if not isinstance(function_args, dict):
@@ -978,10 +916,8 @@ def handle_function_call(
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,
**asdict(ids), middleware_trace=list(trace), **extra,
)
_emit_post_tool_call_hook(function_name=function_name, function_args=function_args, result=result,
**asdict(ids), middleware_trace=list(trace), **extra)
return result
# Tool Search bridge: tool_search / tool_describe are catalog reads handled
@@ -1037,11 +973,8 @@ def handle_function_call(
except Exception as e:
error_msg = f"Error executing {function_name}: {str(e)}"
logger.exception(error_msg)
return _emit(
tool_error(_sanitize_tool_error(error_msg)),
duration_ms=_elapsed_ms(start), status="error",
error_type=type(e).__name__, error_message=str(e),
)
return _emit(tool_error(_sanitize_tool_error(error_msg)), duration_ms=_elapsed_ms(start),
status="error", error_type=type(e).__name__, error_message=str(e))
# =============================================================================
+38 -88
View File
@@ -38,6 +38,12 @@ _HERMES_CORE_TOOLS = [
# Webhook payloads are untrusted third-party content: no file/system execution.
_HERMES_WEBHOOK_SAFE_TOOLS = ["web_search", "web_extract", "vision_analyze", "clarify"]
_HA_TOOLS = ["ha_list_entities", "ha_get_state", "ha_list_services", "ha_call_service"]
_FEISHU_TOOLS = [
"feishu_doc_read", "feishu_drive_list_comments", "feishu_drive_list_comment_replies",
"feishu_drive_reply_comment", "feishu_drive_add_comment",
]
_YUANBAO_TOOLS = ["yb_query_group_info", "yb_query_group_members", "yb_send_dm", "yb_search_sticker", "yb_send_sticker"]
def _ts(description, tools=(), includes=(), **extra):
@@ -57,11 +63,7 @@ def _core_without(*excluded, kanban=True):
# Coding posture: everything you reach for while pairing on code; drops messaging,
# tts, image_gen, home-assistant, cron, kanban and computer-use.
_CODING_TOOLS = _core_without(
"image_generate", "text_to_speech", "cronjob_manage", "computer_use",
"ha_list_entities", "ha_get_state", "ha_list_services", "ha_call_service",
kanban=False,
)
_CODING_TOOLS = _core_without("image_generate", "text_to_speech", "cronjob_manage", "computer_use", *_HA_TOOLS, kanban=False)
# Core toolset definitions: individual tools or references to other toolsets.
TOOLSETS = {
@@ -107,12 +109,7 @@ TOOLSETS = {
"browser": _ts(
"Browser automation for web interaction (navigate, click, type, scroll, "
"iframes, hold-click) with web search for finding URLs",
[
"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",
"browser_exec", "web_search",
],
[t for t in _HERMES_CORE_TOOLS if t.startswith("browser_")] + ["web_search"],
),
"cronjob": _ts(
"Cronjob management tool - create, list, update, pause, resume, remove, and "
@@ -156,10 +153,7 @@ TOOLSETS = {
["execute_code"],
),
"delegation": _ts("Spawn subagents with isolated context for complex subtasks", ["delegate_task"]),
"homeassistant": _ts(
"Home Assistant smart home control and monitoring",
["ha_list_entities", "ha_get_state", "ha_list_services", "ha_call_service"],
),
"homeassistant": _ts("Home Assistant smart home control and monitoring", _HA_TOOLS),
"kanban": _ts(
"Kanban multi-agent coordination — only active when the agent is spawned by "
"the kanban dispatcher (HERMES_KANBAN_TASK env set). The dispatcher runs "
@@ -168,12 +162,7 @@ TOOLSETS = {
"first-class review (request_review — not a block), return review changes, "
"block for human input, heartbeat during long ops, comment on threads, attach "
"files, and (for orchestrators) list, unblock, and fan out tasks.",
[
"kanban_show", "kanban_list", "kanban_complete", "kanban_block",
"kanban_request_review", "kanban_request_changes", "kanban_heartbeat",
"kanban_comment", "kanban_create", "kanban_link", "kanban_unblock",
"kanban_attach", "kanban_attach_url", "kanban_attachments",
],
[t for t in _HERMES_CORE_TOOLS if t.startswith("kanban_")],
),
"discord": _ts(
"Discord read and participate tools (fetch messages, search members, create threads)",
@@ -183,21 +172,9 @@ TOOLSETS = {
"Discord server management (list channels/roles, pin messages, assign roles)",
["discord_admin"],
),
"yuanbao": _ts(
"Yuanbao platform tools - group info, member queries, DM, stickers",
[
"yb_query_group_info", "yb_query_group_members", "yb_send_dm", "yb_search_sticker",
"yb_send_sticker",
],
),
"yuanbao": _ts("Yuanbao platform tools - group info, member queries, DM, stickers", _YUANBAO_TOOLS),
"feishu_doc": _ts("Read Feishu/Lark document content", ["feishu_doc_read"]),
"feishu_drive": _ts(
"Feishu/Lark document comment operations (list, reply, add)",
[
"feishu_drive_list_comments", "feishu_drive_list_comment_replies",
"feishu_drive_reply_comment", "feishu_drive_add_comment",
],
),
"feishu_drive": _ts("Feishu/Lark document comment operations (list, reply, add)", _FEISHU_TOOLS[1:]),
"spotify": _ts(
"Native Spotify playback, search, playlist, album, and library tools",
[
@@ -264,14 +241,7 @@ TOOLSETS = {
"hermes-mattermost": _bundle("Mattermost bot toolset - self-hosted team messaging (full access)"),
"hermes-matrix": _bundle("Matrix bot toolset - decentralized encrypted messaging (full access)"),
"hermes-dingtalk": _bundle("DingTalk bot toolset - enterprise messaging platform (full access)"),
"hermes-feishu": _bundle(
"Feishu/Lark bot toolset - enterprise messaging via Feishu/Lark (full access)",
[
"feishu_doc_read", "feishu_drive_list_comments",
"feishu_drive_list_comment_replies", "feishu_drive_reply_comment",
"feishu_drive_add_comment",
],
),
"hermes-feishu": _bundle("Feishu/Lark bot toolset - enterprise messaging via Feishu/Lark (full access)", _FEISHU_TOOLS),
"hermes-weixin": _bundle("Weixin bot toolset - personal WeChat messaging via iLink (full access)"),
"hermes-qqbot": _bundle("QQBot toolset - QQ messaging via Official Bot API v2 (full access)"),
"hermes-wecom": _bundle("WeCom bot toolset - enterprise WeChat messaging (full access)"),
@@ -280,9 +250,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",
],
"tools": _HERMES_CORE_TOOLS + _YUANBAO_TOOLS,
"module": "tools.yuanbao_tools",
"includes": [],
},
@@ -331,12 +299,12 @@ def _registry_generation() -> Tuple[int, int]:
def get_toolset(name: str, *, include_registry: bool = True) -> Optional[Dict[str, Any]]:
"""Return a toolset definition, or None if unknown.
"""Toolset definition, or None if unknown.
include_registry=True merges tools plugins/overlays registered into this
toolset and resolves registry-only (plugin/MCP) toolsets and aliases.
include_registry=False returns only the static TOOLSETS entry (copied), so
platform reverse-mapping (#49622) is unaffected by registry additions.
include_registry=True merges plugin/overlay tools registered into this toolset
and resolves registry-only (plugin/MCP) toolsets and aliases; False returns a
copy of the static TOOLSETS entry only, so platform reverse-mapping is
unaffected by registry additions.
"""
toolset = TOOLSETS.get(name)
if not include_registry:
@@ -371,9 +339,9 @@ def get_toolset(name: str, *, include_registry: bool = True) -> Optional[Dict[st
def bundle_non_core_tools(toolset_name: str) -> Set[str]:
"""A bundle's tools minus _HERMES_CORE_TOOLS (one level of includes).
Bundles are `_HERMES_CORE_TOOLS + extras`; disabling one must not strip the
core tools every other toolset shares. One `includes` pass suffices because
only hermes-gateway nests bundles. Unknown names: full resolution minus core.
Disabling a `core + extras` bundle must not strip the core tools every other
toolset shares. One `includes` pass suffices (only hermes-gateway nests
bundles). Unknown names: full resolution minus core.
"""
core = set(_HERMES_CORE_TOOLS)
ts_def = get_toolset(toolset_name)
@@ -387,8 +355,8 @@ def bundle_non_core_tools(toolset_name: str) -> Set[str]:
return to_remove - core
# Memo keyed on (name, include_registry, id(registry), registry generation).
# Engages only at the public entry (visited is None); recursion is untouched.
# Memo keyed on (name, include_registry, id(registry), registry generation);
# engages only at the public entry (visited is None).
_resolve_toolset_memo: Dict[Tuple[str, bool, int, int], List[str]] = {}
@@ -402,23 +370,16 @@ def _plugin_platform_bundle(name: str) -> List[str]:
from gateway.platform_registry import platform_registry
if not platform_registry.is_registered(platform_name):
return []
tools = set(_HERMES_CORE_TOOLS)
registry = _registry()
if registry is not None:
try:
tools.update(e.name for e in registry.get_all_entries() if e.toolset == platform_name)
except Exception:
pass
return list(tools)
except Exception:
return []
tools = set(_HERMES_CORE_TOOLS)
tools.update(e.name for e in _registry_call("get_all_entries", ()) if e.toolset == platform_name)
return list(tools)
def resolve_toolset(name: str, visited: Set[str] = None, *, include_registry: bool = True) -> List[str]:
"""Recursively resolve a toolset (and its includes) to a sorted tool-name list.
include_registry=False resolves the static TOOLSETS view only (#49622).
"""
include_registry=False resolves the static TOOLSETS view only."""
external_call = visited is None
if external_call:
memo_key = (name, include_registry, *_registry_generation())
@@ -434,8 +395,7 @@ def resolve_toolset(name: str, visited: Set[str] = None, *, include_registry: bo
all_tools.update(resolve_toolset(toolset_name, visited.copy(), include_registry=include_registry))
return sorted(all_tools)
# Diamond include or cycle: return [] silently — the tools were (or will
# be) collected via another path, so this is not an error.
# Diamond include or cycle: [] silently — the tools are collected via another path.
if name in visited:
return []
visited.add(name)
@@ -450,10 +410,9 @@ def resolve_toolset(name: str, visited: Set[str] = None, *, include_registry: bo
result = sorted(tools)
if external_call:
# Stale-generation entries are never hit again; bound the memo.
if len(_resolve_toolset_memo) >= 256:
if len(_resolve_toolset_memo) >= 256: # stale-generation entries are never hit again
_resolve_toolset_memo.clear()
_resolve_toolset_memo[(name, include_registry, *_registry_generation())] = list(result)
_resolve_toolset_memo[memo_key] = list(result)
return result
@@ -495,17 +454,11 @@ def get_toolset_names() -> List[str]:
def validate_toolset(name: str) -> bool:
if name in {"all", "*"} or name in TOOLSETS:
return True
return name in _get_plugin_toolset_names() or name in _get_registry_toolset_aliases()
return (name in {"all", "*"} or name in TOOLSETS
or name in _get_plugin_toolset_names() or name in _get_registry_toolset_aliases())
def create_custom_toolset(
name: str,
description: str,
tools: List[str] = None,
includes: List[str] = None
) -> None:
def create_custom_toolset(name: str, description: str, tools: List[str] = None, includes: List[str] = None) -> None:
"""Register a runtime toolset in TOOLSETS."""
TOOLSETS[name] = _ts(description, tools or [], includes or [])
@@ -517,11 +470,8 @@ def get_toolset_info(name: str) -> Dict[str, Any]:
return None
resolved_tools = resolve_toolset(name)
return {
"name": name,
"description": toolset["description"],
"direct_tools": toolset["tools"],
"includes": toolset["includes"],
"resolved_tools": resolved_tools,
"tool_count": len(resolved_tools),
"is_composite": bool(toolset["includes"])
"name": name, "description": toolset["description"],
"direct_tools": toolset["tools"], "includes": toolset["includes"],
"resolved_tools": resolved_tools, "tool_count": len(resolved_tools),
"is_composite": bool(toolset["includes"]),
}