refactor(model_tools,mcp_serve): final compaction pass (signatures, comments, baseline helper)

This commit is contained in:
Teknium
2026-09-02 18:26:58 -07:00
parent 7fd8ccff62
commit 24574b8534
2 changed files with 49 additions and 99 deletions
+18 -23
View File
@@ -77,7 +77,7 @@ def _close_quietly(db, what: str) -> None:
def _get_session_db(): def _get_session_db():
"""Get a SessionDB instance for reading message transcripts.""" """SessionDB instance for reading message transcripts, or None."""
try: try:
from hermes_state import get_shared_session_db from hermes_state import get_shared_session_db
return get_shared_session_db() return get_shared_session_db()
@@ -112,6 +112,13 @@ def _load_sessions_index() -> dict:
return _load_sessions_index_from_db() or _load_sessions_index_from_json() return _load_sessions_index_from_db() or _load_sessions_index_from_json()
def _iso(ts) -> str:
try:
return datetime.fromtimestamp(float(ts)).isoformat() if ts else ""
except (TypeError, ValueError, OSError):
return ""
def _row_to_index_entry(row: dict) -> dict: def _row_to_index_entry(row: dict) -> dict:
"""Convert a state.db gateway session row to the sessions.json entry shape.""" """Convert a state.db gateway session row to the sessions.json entry shape."""
origin = {} origin = {}
@@ -126,12 +133,6 @@ def _row_to_index_entry(row: dict) -> dict:
origin = {k: row.get(k) for k in ("chat_id", "chat_type", "thread_id", "user_id")} origin = {k: row.get(k) for k in ("chat_id", "chat_type", "thread_id", "user_id")}
origin["platform"] = row.get("source", "") origin["platform"] = row.get("source", "")
def _iso(ts) -> str:
try:
return datetime.fromtimestamp(float(ts)).isoformat() if ts else ""
except (TypeError, ValueError, OSError):
return ""
input_tokens = int(row.get("input_tokens") or 0) input_tokens = int(row.get("input_tokens") or 0)
output_tokens = int(row.get("output_tokens") or 0) output_tokens = int(row.get("output_tokens") or 0)
return { return {
@@ -156,10 +157,7 @@ def _load_sessions_index_from_db() -> dict:
lister = getattr(db, "list_gateway_sessions", None) lister = getattr(db, "list_gateway_sessions", None)
if not callable(lister): if not callable(lister):
return {} return {}
return { return {row["session_key"]: _row_to_index_entry(row) for row in lister(active_only=True) if row.get("session_key")}
row["session_key"]: _row_to_index_entry(row)
for row in lister(active_only=True) if row.get("session_key")
}
except Exception as e: except Exception as e:
logger.debug("Failed to load gateway sessions from state.db: %s", e) logger.debug("Failed to load gateway sessions from state.db: %s", e)
return {} return {}
@@ -302,9 +300,9 @@ class EventBridge:
logger.debug("EventBridge started") logger.debug("EventBridge started")
def stop(self): def stop(self):
"""Stop the background polling thread.""" """Stop the background polling thread and wake any waiters."""
self._running = False self._running = False
self._new_event.set() # Wake any waiters self._new_event.set()
if self._thread: if self._thread:
self._thread.join(timeout=5) self._thread.join(timeout=5)
logger.debug("EventBridge stopped") logger.debug("EventBridge stopped")
@@ -364,12 +362,11 @@ class EventBridge:
def _establish_baseline(self) -> None: def _establish_baseline(self) -> None:
db = _get_session_db() db = _get_session_db()
if not db: if db:
return try:
try: self._establish_baseline_with_db(db)
self._establish_baseline_with_db(db) finally:
finally: _close_quietly(db, "baseline")
_close_quietly(db, "baseline")
def _establish_baseline_with_db(self, db) -> None: def _establish_baseline_with_db(self, db) -> None:
"""Record per-session latest timestamps and the state.db mtime WITHOUT """Record per-session latest timestamps and the state.db mtime WITHOUT
@@ -385,10 +382,9 @@ class EventBridge:
if not session_id: if not session_id:
continue continue
try: try:
messages = db.get_messages(session_id) latest = _latest_ts(db.get_messages(session_id))
except Exception: except Exception:
continue continue
latest = _latest_ts(messages)
if latest > 0.0: if latest > 0.0:
self._last_poll_timestamps[session_key] = latest self._last_poll_timestamps[session_key] = latest
@@ -418,8 +414,7 @@ class EventBridge:
""" """
db_mtime = _read_state_db_mtime() db_mtime = _read_state_db_mtime()
if db_mtime == self._state_db_mtime: if db_mtime == self._state_db_mtime:
return # Nothing changed since last poll — skip entirely return
self._state_db_mtime = db_mtime self._state_db_mtime = db_mtime
# Refresh the index on every change tick: one indexed query, never lags messages. # Refresh the index on every change tick: one indexed query, never lags messages.
self._cached_sessions_index = _load_sessions_index() self._cached_sessions_index = _load_sessions_index()
+31 -76
View File
@@ -110,13 +110,11 @@ def _run_async(coro):
loop = asyncio.get_running_loop() loop = asyncio.get_running_loop()
except RuntimeError: except RuntimeError:
loop = None loop = None
if loop and loop.is_running(): if loop and loop.is_running():
# Inside a running loop: run in a fresh thread whose loop we keep a # 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 # reference to, so on timeout we can cancel the task inside it
# (ThreadPoolExecutor.cancel() is a no-op on a running worker). # (ThreadPoolExecutor.cancel() is a no-op on a running worker).
import concurrent.futures import concurrent.futures
worker_loop: Optional[asyncio.AbstractEventLoop] = None worker_loop: Optional[asyncio.AbstractEventLoop] = None
loop_ready = threading.Event() loop_ready = threading.Event()
@@ -191,10 +189,8 @@ _LEGACY_TOOLSET_MAP = {
"image_tools": ["image_generate"], "image_tools": ["image_generate"],
"skills_tools": ["skills_list", "skill_view", "skill_manage"], "skills_tools": ["skills_list", "skill_view", "skill_manage"],
"browser_tools": [ "browser_tools": [
"browser_navigate", "browser_snapshot", "browser_click", "browser_navigate", "browser_snapshot", "browser_click", "browser_type", "browser_scroll",
"browser_type", "browser_scroll", "browser_back", "browser_back", "browser_press", "browser_get_images", "browser_vision", "browser_console",
"browser_press", "browser_get_images",
"browser_vision", "browser_console"
], ],
"cronjob_tools": ["cronjob_manage"], "cronjob_tools": ["cronjob_manage"],
"file_tools": ["read_file", "write_file", "patch", "search_files"], "file_tools": ["read_file", "write_file", "patch", "search_files"],
@@ -248,8 +244,7 @@ def get_tool_definitions(
if cache_key is None: if cache_key is None:
return list(result) return list(result)
with _tool_defs_cache_lock: with _tool_defs_cache_lock:
# Another thread may have filled this key meanwhile; reuse it. cached = _tool_defs_cache.get(cache_key) # another thread may have filled it meanwhile
cached = _tool_defs_cache.get(cache_key)
if cached is None: if cached is None:
if len(_tool_defs_cache) >= _TOOL_DEFS_CACHE_MAX: if len(_tool_defs_cache) >= _TOOL_DEFS_CACHE_MAX:
_tool_defs_cache.pop(next(iter(_tool_defs_cache))) _tool_defs_cache.pop(next(iter(_tool_defs_cache)))
@@ -257,16 +252,13 @@ def get_tool_definitions(
else: else:
global _last_resolved_tool_names global _last_resolved_tool_names
_last_resolved_tool_names = [t["function"]["name"] for t in cached] _last_resolved_tool_names = [t["function"]["name"] for t in cached]
# Always a shallow copy: run_agent appends memory/LCM schemas to its list, and # Always a shallow copy: run_agent appends memory/LCM schemas to its list; a
# a shared list would accumulate duplicate tool names across agent inits # shared list would accumulate duplicate names (HTTP 400 from DeepSeek/Kimi/MiMo).
# (rejected with HTTP 400 by DeepSeek/Kimi/MiMo).
return list(cached) return list(cached)
def _tool_defs_cache_key( def _tool_defs_cache_key(
enabled_toolsets: Optional[List[str]], enabled_toolsets: Optional[List[str]], disabled_toolsets: Optional[List[str]], skip_tool_search_assembly: bool,
disabled_toolsets: Optional[List[str]],
skip_tool_search_assembly: bool,
) -> Optional[tuple]: ) -> Optional[tuple]:
"""Memo key for get_tool_definitions, or None when caching must be bypassed. """Memo key for get_tool_definitions, or None when caching must be bypassed.
@@ -301,20 +293,16 @@ def _apply_toolset_selection(tools: set, names: List[str], quiet_mode: bool, *,
if validate_toolset(name): if validate_toolset(name):
label = f"{verb} toolset" label = f"{verb} toolset"
if disable and (name.startswith("hermes-") or (get_toolset(name) or {}).get("posture")): if disable and (name.startswith("hermes-") or (get_toolset(name) or {}).get("posture")):
# Platform bundles and posture toolsets re-list the core tools # Bundles/postures re-list the core tools without owning them;
# without owning them; subtracting the whole set would empty # subtracting the whole set would empty the list — remove only the non-core delta.
# the tool list. Remove only the non-core delta.
resolved = sorted(bundle_non_core_tools(name)) resolved = sorted(bundle_non_core_tools(name))
if not quiet_mode and name.startswith("hermes-") and name not in _WARNED_DISABLED_BUNDLES: if not quiet_mode and name.startswith("hermes-") and name not in _WARNED_DISABLED_BUNDLES:
_WARNED_DISABLED_BUNDLES.add(name) _WARNED_DISABLED_BUNDLES.add(name)
logger.info( logger.info(
"agent.disabled_toolsets contains platform-bundle " "agent.disabled_toolsets contains platform-bundle name '%s'; core tools are "
"name '%s'; core tools are preserved and only its " "preserved and only its platform-specific tools (%s) are removed. Bundle names "
"platform-specific tools (%s) are removed. Bundle " "usually belong in `toolsets:`, not `disabled_toolsets` (#33924).",
"names usually belong in `toolsets:`, not " name, ", ".join(resolved) if resolved else "none",
"`disabled_toolsets` (#33924).",
name,
", ".join(resolved) if resolved else "none",
) )
else: else:
resolved = resolve_toolset(name) resolved = resolve_toolset(name)
@@ -576,18 +564,14 @@ def _resolve_active_context_length() -> int:
# handle_function_call (the main dispatcher) # handle_function_call (the main dispatcher)
# ============================================================================= # =============================================================================
# Tools the agent loop (run_agent.py) intercepts because they need agent-level # Intercepted by the agent loop (need agent-level state); dispatch returns a stub error.
# state. The registry still holds their schemas; dispatch returns a stub error.
_AGENT_LOOP_TOOLS = {"todo_list", "memory", "session_search", "delegate_task"} _AGENT_LOOP_TOOLS = {"todo_list", "memory", "session_search", "delegate_task"}
# Legacy tool-name aliases (2026-08 renames), accepted at every dispatch seam so # Legacy tool-name aliases accepted at every dispatch seam (old sessions/saved
# old sessions and saved prompts keep working; schemas advertise only new names. # prompts keep working); schemas advertise only new names.
_LEGACY_TOOL_ALIASES = { _LEGACY_TOOL_ALIASES = {
"todo": "todo_list", "todo": "todo_list", "cronjob": "cronjob_manage", "process": "process_manage",
"cronjob": "cronjob_manage", "tour": "gui_tour", "tip": "show_tip",
"process": "process_manage",
"tour": "gui_tour",
"tip": "show_tip",
} }
_READ_SEARCH_TOOLS = {"read_file", "search_files"} _READ_SEARCH_TOOLS = {"read_file", "search_files"}
@@ -597,10 +581,7 @@ _READ_SEARCH_TOOLS = {"read_file", "search_files"}
# model will read, and cap length (cap shared with tools/registry.py so text never # model will read, and cap length (cap shared with tools/registry.py so text never
# passes two different caps with two different markers). # passes two different caps with two different markers).
_TOOL_ERROR_STRIP_RES = ( _TOOL_ERROR_STRIP_RES = (
re.compile( re.compile(r'</?(?:tool_call|function_call|result|response|output|input|system|assistant|user)>', re.IGNORECASE),
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*```(?:json|xml|html|markdown)?\s*', re.MULTILINE),
re.compile(r'\s*```\s*$', re.MULTILINE), re.compile(r'\s*```\s*$', re.MULTILINE),
re.compile(r'<!\[CDATA\[.*?\]\]>', re.DOTALL), re.compile(r'<!\[CDATA\[.*?\]\]>', re.DOTALL),
@@ -653,19 +634,11 @@ def _tool_result_observer_fields(tool_name: str, result: Any) -> tuple[str, Opti
def _emit_post_tool_call_hook( def _emit_post_tool_call_hook(
*, *, function_name: str, function_args: Dict[str, Any], result: Any,
function_name: str, task_id: Optional[str] = None, session_id: Optional[str] = None, tool_call_id: Optional[str] = None,
function_args: Dict[str, Any], turn_id: Optional[str] = None, api_request_id: Optional[str] = None,
result: Any, duration_ms: int = 0, status: Optional[str] = None,
task_id: Optional[str] = None, error_type: Optional[str] = None, error_message: Optional[str] = None,
session_id: Optional[str] = None,
tool_call_id: Optional[str] = None,
turn_id: Optional[str] = None,
api_request_id: Optional[str] = None,
duration_ms: int = 0,
status: Optional[str] = None,
error_type: Optional[str] = None,
error_message: Optional[str] = None,
middleware_trace: Optional[List[Dict[str, Any]]] = None, middleware_trace: Optional[List[Dict[str, Any]]] = None,
) -> None: ) -> None:
"""Emit the ``post_tool_call`` observer hook; gated on has_hook, and ok/error """Emit the ``post_tool_call`` observer hook; gated on has_hook, and ok/error
@@ -689,10 +662,8 @@ def _emit_post_tool_call_hook(
def _dispatch_bridge_tool( def _dispatch_bridge_tool(
function_name: str, function_name: str, function_args: Dict[str, Any],
function_args: Dict[str, Any], enabled_toolsets: Optional[List[str]], disabled_toolsets: Optional[List[str]],
enabled_toolsets: Optional[List[str]],
disabled_toolsets: Optional[List[str]],
): ):
"""Handle a Tool Search bridge call (tool_search / tool_describe / tool_call). """Handle a Tool Search bridge call (tool_search / tool_describe / tool_call).
@@ -739,10 +710,7 @@ def _dispatch_bridge_tool(
def _apply_request_middleware( def _apply_request_middleware(
function_name: str, function_name: str, function_args: Dict[str, Any], ids: _CallIds, trace: List[Dict[str, Any]],
function_args: Dict[str, Any],
ids: _CallIds,
trace: List[Dict[str, Any]],
) -> Tuple[Dict[str, Any], Dict[str, Any], 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.""" """tool_request middleware: returns (args, original_args, trace); fail-open."""
try: try:
@@ -755,11 +723,8 @@ def _apply_request_middleware(
def _pre_dispatch_guards( def _pre_dispatch_guards(
function_name: str, function_name: str, function_args: Dict[str, Any], skip_pre_tool_call_hook: bool,
function_args: Dict[str, Any], ids: _CallIds, middleware_trace: List[Dict[str, Any]],
skip_pre_tool_call_hook: bool,
ids: _CallIds,
middleware_trace: List[Dict[str, Any]],
) -> Tuple[Dict[str, Any], Optional[Tuple[Any, str, Optional[str]]]]: ) -> Tuple[Dict[str, Any], Optional[Tuple[Any, str, Optional[str]]]]:
"""Plugin pre_tool_call hook, then ACP edit approval. """Plugin pre_tool_call hook, then ACP edit approval.
@@ -817,14 +782,8 @@ def _approval_observability(ids: _CallIds):
def _execute_tool( def _execute_tool(
function_name: str, function_name: str, function_args: Dict[str, Any], original_args: Dict[str, Any], ids: _CallIds,
function_args: Dict[str, Any], *, user_task: Optional[str], enabled_tools: Optional[List[str]], skip_tool_execution_middleware: bool,
original_args: Dict[str, Any],
ids: _CallIds,
*,
user_task: Optional[str],
enabled_tools: Optional[List[str]],
skip_tool_execution_middleware: bool,
) -> Any: ) -> Any:
"""Run the registry handler (through tool-execution middleware unless skipped) """Run the registry handler (through tool-execution middleware unless skipped)
with the approval observability context bound for the duration.""" with the approval observability context bound for the duration."""
@@ -849,11 +808,7 @@ def _execute_tool(
def _apply_transform_tool_result_hook( def _apply_transform_tool_result_hook(
function_name: str, function_name: str, function_args: Dict[str, Any], result: Any, duration_ms: int, ids: _CallIds,
function_args: Dict[str, Any],
result: Any,
duration_ms: int,
ids: _CallIds,
) -> Any: ) -> Any:
"""transform_tool_result: plugins may replace the final result string. """transform_tool_result: plugins may replace the final result string.