From 24574b853498da2eeed2e1e80938ec5a6c73129e Mon Sep 17 00:00:00 2001 From: Teknium <127238744+teknium1@users.noreply.github.com> Date: Wed, 2 Sep 2026 18:26:58 -0700 Subject: [PATCH] refactor(model_tools,mcp_serve): final compaction pass (signatures, comments, baseline helper) --- mcp_serve.py | 41 +++++++++---------- model_tools.py | 107 ++++++++++++++----------------------------------- 2 files changed, 49 insertions(+), 99 deletions(-) diff --git a/mcp_serve.py b/mcp_serve.py index c3e517c252..e1c52f8907 100644 --- a/mcp_serve.py +++ b/mcp_serve.py @@ -77,7 +77,7 @@ def _close_quietly(db, what: str) -> None: def _get_session_db(): - """Get a SessionDB instance for reading message transcripts.""" + """SessionDB instance for reading message transcripts, or None.""" try: from hermes_state import 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() +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: """Convert a state.db gateway session row to the sessions.json entry shape.""" 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["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) output_tokens = int(row.get("output_tokens") or 0) return { @@ -156,10 +157,7 @@ def _load_sessions_index_from_db() -> dict: lister = getattr(db, "list_gateway_sessions", None) if not callable(lister): return {} - return { - row["session_key"]: _row_to_index_entry(row) - for row in lister(active_only=True) if row.get("session_key") - } + return {row["session_key"]: _row_to_index_entry(row) for row in lister(active_only=True) if row.get("session_key")} except Exception as e: logger.debug("Failed to load gateway sessions from state.db: %s", e) return {} @@ -302,9 +300,9 @@ class EventBridge: logger.debug("EventBridge started") def stop(self): - """Stop the background polling thread.""" + """Stop the background polling thread and wake any waiters.""" self._running = False - self._new_event.set() # Wake any waiters + self._new_event.set() if self._thread: self._thread.join(timeout=5) logger.debug("EventBridge stopped") @@ -364,12 +362,11 @@ class EventBridge: def _establish_baseline(self) -> None: db = _get_session_db() - if not db: - return - try: - self._establish_baseline_with_db(db) - finally: - _close_quietly(db, "baseline") + if db: + try: + self._establish_baseline_with_db(db) + finally: + _close_quietly(db, "baseline") def _establish_baseline_with_db(self, db) -> None: """Record per-session latest timestamps and the state.db mtime WITHOUT @@ -385,10 +382,9 @@ class EventBridge: if not session_id: continue try: - messages = db.get_messages(session_id) + latest = _latest_ts(db.get_messages(session_id)) except Exception: continue - latest = _latest_ts(messages) if latest > 0.0: self._last_poll_timestamps[session_key] = latest @@ -418,8 +414,7 @@ class EventBridge: """ db_mtime = _read_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 # Refresh the index on every change tick: one indexed query, never lags messages. self._cached_sessions_index = _load_sessions_index() diff --git a/model_tools.py b/model_tools.py index 86352e6176..345a64305e 100644 --- a/model_tools.py +++ b/model_tools.py @@ -110,13 +110,11 @@ def _run_async(coro): loop = asyncio.get_running_loop() except RuntimeError: loop = None - if loop and loop.is_running(): # 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 loop_ready = threading.Event() @@ -191,10 +189,8 @@ _LEGACY_TOOLSET_MAP = { "image_tools": ["image_generate"], "skills_tools": ["skills_list", "skill_view", "skill_manage"], "browser_tools": [ - "browser_navigate", "browser_snapshot", "browser_click", - "browser_type", "browser_scroll", "browser_back", - "browser_press", "browser_get_images", - "browser_vision", "browser_console" + "browser_navigate", "browser_snapshot", "browser_click", "browser_type", "browser_scroll", + "browser_back", "browser_press", "browser_get_images", "browser_vision", "browser_console", ], "cronjob_tools": ["cronjob_manage"], "file_tools": ["read_file", "write_file", "patch", "search_files"], @@ -248,8 +244,7 @@ def get_tool_definitions( if cache_key is None: return list(result) with _tool_defs_cache_lock: - # Another thread may have filled this key meanwhile; reuse it. - cached = _tool_defs_cache.get(cache_key) + cached = _tool_defs_cache.get(cache_key) # another thread may have filled it meanwhile if cached is None: if len(_tool_defs_cache) >= _TOOL_DEFS_CACHE_MAX: _tool_defs_cache.pop(next(iter(_tool_defs_cache))) @@ -257,16 +252,13 @@ def get_tool_definitions( else: global _last_resolved_tool_names _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 - # a shared list would accumulate duplicate tool names across agent inits - # (rejected with HTTP 400 by DeepSeek/Kimi/MiMo). + # Always a shallow copy: run_agent appends memory/LCM schemas to its list; a + # shared list would accumulate duplicate names (HTTP 400 from DeepSeek/Kimi/MiMo). return list(cached) def _tool_defs_cache_key( - enabled_toolsets: Optional[List[str]], - disabled_toolsets: Optional[List[str]], - skip_tool_search_assembly: bool, + enabled_toolsets: Optional[List[str]], disabled_toolsets: Optional[List[str]], skip_tool_search_assembly: bool, ) -> Optional[tuple]: """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): label = f"{verb} toolset" if disable and (name.startswith("hermes-") or (get_toolset(name) or {}).get("posture")): - # Platform bundles and posture toolsets re-list the core tools - # without owning them; subtracting the whole set would empty - # the tool list. Remove only the non-core delta. + # Bundles/postures re-list the core tools without owning them; + # subtracting the whole set would empty the list — remove only the non-core delta. resolved = sorted(bundle_non_core_tools(name)) if not quiet_mode and name.startswith("hermes-") and name not in _WARNED_DISABLED_BUNDLES: _WARNED_DISABLED_BUNDLES.add(name) logger.info( - "agent.disabled_toolsets contains platform-bundle " - "name '%s'; core tools are preserved and only its " - "platform-specific tools (%s) are removed. Bundle " - "names usually belong in `toolsets:`, not " - "`disabled_toolsets` (#33924).", - name, - ", ".join(resolved) if resolved else "none", + "agent.disabled_toolsets contains platform-bundle name '%s'; core tools are " + "preserved and only its platform-specific tools (%s) are removed. Bundle names " + "usually belong in `toolsets:`, not `disabled_toolsets` (#33924).", + name, ", ".join(resolved) if resolved else "none", ) else: resolved = resolve_toolset(name) @@ -576,18 +564,14 @@ def _resolve_active_context_length() -> int: # handle_function_call (the main dispatcher) # ============================================================================= -# Tools the agent loop (run_agent.py) intercepts because they need agent-level -# state. The registry still holds their schemas; dispatch returns a stub error. +# Intercepted by the agent loop (need agent-level state); dispatch returns a stub error. _AGENT_LOOP_TOOLS = {"todo_list", "memory", "session_search", "delegate_task"} -# Legacy tool-name aliases (2026-08 renames), accepted at every dispatch seam so -# old sessions and saved prompts keep working; schemas advertise only new names. +# Legacy tool-name aliases accepted at every dispatch seam (old sessions/saved +# prompts keep working); schemas advertise only new names. _LEGACY_TOOL_ALIASES = { - "todo": "todo_list", - "cronjob": "cronjob_manage", - "process": "process_manage", - "tour": "gui_tour", - "tip": "show_tip", + "todo": "todo_list", "cronjob": "cronjob_manage", "process": "process_manage", + "tour": "gui_tour", "tip": "show_tip", } _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 # passes two different caps with two different markers). _TOOL_ERROR_STRIP_RES = ( - re.compile( - r'', - re.IGNORECASE, - ), + re.compile(r'', re.IGNORECASE), re.compile(r'^\s*```(?:json|xml|html|markdown)?\s*', re.MULTILINE), re.compile(r'\s*```\s*$', re.MULTILINE), re.compile(r'', re.DOTALL), @@ -653,19 +634,11 @@ def _tool_result_observer_fields(tool_name: str, result: Any) -> tuple[str, Opti def _emit_post_tool_call_hook( - *, - function_name: str, - function_args: Dict[str, Any], - result: Any, - task_id: 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, + *, function_name: str, function_args: Dict[str, Any], result: Any, + task_id: 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, ) -> None: """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( - function_name: str, - function_args: Dict[str, Any], - enabled_toolsets: Optional[List[str]], - disabled_toolsets: Optional[List[str]], + function_name: str, function_args: Dict[str, Any], + enabled_toolsets: Optional[List[str]], disabled_toolsets: Optional[List[str]], ): """Handle a Tool Search bridge call (tool_search / tool_describe / tool_call). @@ -739,10 +710,7 @@ def _dispatch_bridge_tool( def _apply_request_middleware( - function_name: str, - function_args: Dict[str, Any], - ids: _CallIds, - trace: List[Dict[str, Any]], + 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: @@ -755,11 +723,8 @@ def _apply_request_middleware( def _pre_dispatch_guards( - function_name: str, - function_args: Dict[str, Any], - skip_pre_tool_call_hook: bool, - ids: _CallIds, - middleware_trace: List[Dict[str, Any]], + function_name: str, function_args: 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]]]]: """Plugin pre_tool_call hook, then ACP edit approval. @@ -817,14 +782,8 @@ def _approval_observability(ids: _CallIds): def _execute_tool( - function_name: str, - function_args: Dict[str, Any], - original_args: Dict[str, Any], - ids: _CallIds, - *, - user_task: Optional[str], - enabled_tools: Optional[List[str]], - skip_tool_execution_middleware: bool, + function_name: str, function_args: Dict[str, Any], original_args: Dict[str, Any], ids: _CallIds, + *, user_task: Optional[str], enabled_tools: Optional[List[str]], skip_tool_execution_middleware: bool, ) -> Any: """Run the registry handler (through tool-execution middleware unless skipped) with the approval observability context bound for the duration.""" @@ -849,11 +808,7 @@ def _execute_tool( def _apply_transform_tool_result_hook( - function_name: str, - function_args: Dict[str, Any], - result: Any, - duration_ms: int, - ids: _CallIds, + function_name: str, function_args: Dict[str, Any], result: Any, duration_ms: int, ids: _CallIds, ) -> Any: """transform_tool_result: plugins may replace the final result string.