diff --git a/agent/conversation_loop.py b/agent/conversation_loop.py index 21c3c493e0..086e476d1b 100644 --- a/agent/conversation_loop.py +++ b/agent/conversation_loop.py @@ -656,6 +656,20 @@ def _persist_system_prompt(agent, failure_message: str, *, persist_tools: bool = logger.warning(failure_message, agent.session_id, exc) +def _restore_pinned_tools(agent, session_row) -> list: + """Pin ``agent.tools`` to the session's persisted array (tools freeze); returns the names + this surface built BEFORE the pin merged a previous surface's tools back in.""" + from tools.mcp_tool_agent import agent_tool_names, restore_agent_tool_prefix + built_for_this_surface = agent_tool_names(agent) + try: + saved_tools = session_row.get("tool_names") if session_row else None + if saved_tools: + restore_agent_tool_prefix(agent, json.loads(saved_tools)) + except Exception: + logger.debug("tool prefix restore skipped", exc_info=True) + return built_for_this_surface + + def _restore_or_build_system_prompt(agent, system_message, conversation_history): """Restore the cached system prompt from the session DB or build it fresh. @@ -719,17 +733,9 @@ def _restore_or_build_system_prompt(agent, system_message, conversation_history) # ADDS what the new surface brought (a tui -> desktop switch pays a break no freeze can # avoid), and what it carries FORWARD is named in the note instead, so a tool that can # only answer ``tool_error("desktop only")`` here does not read as a live capability. - try: - saved_tools = session_row.get("tool_names") if session_row else None - if saved_tools: - from tools.mcp_tool_agent import agent_tool_names, restore_agent_tool_prefix - # Captured BEFORE the pin merges the previous surface's tools back in. - built_for_this_surface = agent_tool_names(agent) if announced_switch else [] - restore_agent_tool_prefix(agent, json.loads(saved_tools)) - if announced_switch: - note_inert_pinned_tools(agent, built_for_this_surface) - except Exception: - logger.debug("tool prefix restore skipped", exc_info=True) + built_for_this_surface = _restore_pinned_tools(agent, session_row) + if announced_switch: + note_inert_pinned_tools(agent, built_for_this_surface) # Prompt-section callbacks are new-session-only; recover their frozen bytes # from the persisted prompt so a compression rebuild keeps them. The static # prefix is not persisted either; rebuild it for the early cache breakpoint or @@ -757,7 +763,11 @@ def _restore_or_build_system_prompt(agent, system_message, conversation_history) agent.session_id, stored_state, ) - # First turn of a new session (or recovering from a broken stored prompt). + # First turn of a new session (or recovering from a broken stored prompt). Rebuilding an + # EXISTING session's prompt (cwd drift, model switch) still keeps its pinned tools[]: this + # surface's own build (the -q footprint, its tool_search catalog) would otherwise be + # persisted over the pin below. Pinned first, so the prompt describes the tools sent. + built_for_this_surface = _restore_pinned_tools(agent, session_row) agent._cached_system_prompt = agent._build_system_prompt(system_message) # The rebuilt prompt describes the CURRENT surface, but a surface note left in the @@ -765,6 +775,7 @@ def _restore_or_build_system_prompt(agent, system_message, conversation_history) # unrelated reason (a model switch) would leave the newest interface statement in the # request naming a surface the conversation has left (#104414). stage_surface_switch_note(agent, agent._cached_system_prompt, conversation_history) + note_inert_pinned_tools(agent, built_for_this_surface) # Persistence-disabled forks share their parent's session ID and are not real sessions. if not getattr(agent, "_persist_disabled", False): diff --git a/hermes_state.py b/hermes_state.py index 5d453c6731..debb49dfe8 100644 --- a/hermes_state.py +++ b/hermes_state.py @@ -513,16 +513,18 @@ class SessionDB( def _delete_unreferenced_system_prompts(conn) -> None: conn.execute( "DELETE FROM system_prompts WHERE NOT EXISTS (" - "SELECT 1 FROM sessions WHERE sessions.system_prompt_hash = system_prompts.hash)" + "SELECT 1 FROM sessions WHERE sessions.system_prompt_hash = system_prompts.hash) AND NOT EXISTS (" + "SELECT 1 FROM sessions WHERE sessions.tool_names = system_prompts.hash)" ) @staticmethod def _session_row_dict(row: sqlite3.Row) -> Dict[str, Any]: data = dict(row) - if "_system_prompt_resolved" in data: - resolved = data.pop("_system_prompt_resolved") - if "system_prompt" in data: - data["system_prompt"] = resolved + for column in ("system_prompt", "tool_names"): + if f"_{column}_resolved" in data: + resolved = data.pop(f"_{column}_resolved") + if column in data: + data[column] = resolved return data @staticmethod diff --git a/hermes_state_common.py b/hermes_state_common.py index 03b3467991..3d36055e76 100644 --- a/hermes_state_common.py +++ b/hermes_state_common.py @@ -669,6 +669,8 @@ CREATE INDEX IF NOT EXISTS idx_sessions_handoff_state ON sessions(handoff_state, started_at); CREATE INDEX IF NOT EXISTS idx_sessions_system_prompt_hash ON sessions(system_prompt_hash); +CREATE INDEX IF NOT EXISTS idx_sessions_tool_names + ON sessions(tool_names); -- Recent-session browsing must never derive recency by scanning messages. -- This expression is the durable, indexable approximation used to preselect -- a small candidate set before compression-chain and preview hydration. diff --git a/hermes_state_sessions.py b/hermes_state_sessions.py index 68a6ce85b7..6285a3e0b8 100644 --- a/hermes_state_sessions.py +++ b/hermes_state_sessions.py @@ -661,11 +661,17 @@ class SessionSessionsMixin: self._delete_unreferenced_system_prompts(conn) self._execute_write(_do) - def update_session_tool_names(self, session_id: str, tool_names: Optional[List[str]]) -> None: - """Persist the resolved ``tools[]`` name order so a rebuilt AIAgent can't fork the cached tool - prefix on a flipped check_fn verdict; ``None`` clears.""" - payload = json.dumps(list(tool_names)) if tool_names is not None else None - self._write_sql("UPDATE sessions SET tool_names = ? WHERE id = ?", (payload, session_id)) + def update_session_tool_names(self, session_id: str, tools: Optional[List[Any]]) -> None: + """Persist the session's ``tools[]`` pin so a rebuilt AIAgent sends the same bytes; ``None`` + clears. The array repeats across sessions like a system prompt does, so it is stored in the + same content-addressed ``system_prompts`` table and the column holds its hash (legacy rows: + an inline JSON name list); ``get_session`` resolves either.""" + payload = json.dumps(list(tools)) if tools is not None else None + def _do(conn): + conn.execute("UPDATE sessions SET tool_names = ? WHERE id = ?", + (self._store_system_prompt(conn, payload), session_id)) + self._delete_unreferenced_system_prompts(conn) + self._execute_write(_do) def update_session_model( self, session_id: str, model: str, provider: Optional[str] = None, *, @@ -781,8 +787,10 @@ class SessionSessionsMixin: """Get a session by ID (drains queued token deltas first so cost readers see exact totals).""" self.flush_token_counts() row = self._read_one( - "SELECT s.*, COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved " - "FROM sessions s LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash WHERE s.id = ?", + "SELECT s.*, COALESCE(sp.prompt, s.system_prompt) AS _system_prompt_resolved, " + "COALESCE(tp.prompt, s.tool_names) AS _tool_names_resolved " + "FROM sessions s LEFT JOIN system_prompts sp ON sp.hash = s.system_prompt_hash " + "LEFT JOIN system_prompts tp ON tp.hash = s.tool_names WHERE s.id = ?", (session_id,), ) return self._session_row_dict(row) if row else None diff --git a/tests/agent/test_system_prompt_restore.py b/tests/agent/test_system_prompt_restore.py index d372b27156..833ef72746 100644 --- a/tests/agent/test_system_prompt_restore.py +++ b/tests/agent/test_system_prompt_restore.py @@ -15,7 +15,9 @@ instead of rebuilding). Covers: from __future__ import annotations +import json import logging +from types import SimpleNamespace from unittest.mock import MagicMock import pytest @@ -274,6 +276,37 @@ class TestStoredPromptReuse: assert any("stale runtime identity" in r.getMessage() for r in caplog.records) + def test_rebuilding_an_existing_sessions_prompt_keeps_its_pinned_tools(self, tmp_path): + """A continuing session whose stored prompt goes stale (model switch, cwd drift) is rebuilt + by whichever surface resumes it — ``-q --resume`` builds without skill_manage. tools[] sits + ahead of the prompt and only /new, /reload-mcp and compaction may re-derive it, so the + rebuild keeps the pinned array and never persists its own surface's build over the pin.""" + from unittest.mock import patch as _patch + + from hermes_state import SessionDB + + def _tool(name): + return {"type": "function", "function": {"name": name, "description": f"{name} v1", "parameters": {}}} + + pinned = [_tool("read_file"), _tool("skill_manage"), _tool("terminal")] + with SessionDB(db_path=tmp_path / "state.db") as db: + db.create_session("test-session-id", source="tui") + db.update_system_prompt("test-session-id", "Model: old-model\nProvider: openrouter") + db.update_session_tool_names("test-session-id", pinned) + agent = _make_agent(session_db=db) + agent.side_agent = False + agent._bot_mode_protocol = False + agent.tools = [_tool("read_file"), _tool("terminal")] # the -q footprint pruned skill_manage + registered = [SimpleNamespace(name=t["function"]["name"]) for t in pinned] + with _patch("tools.registry.registry.get_all_entries", return_value=registered): + _restore_or_build_system_prompt(agent, None, [{"role": "user", "content": "hi"}]) + + agent._build_system_prompt.assert_called_once() + assert agent.tools == pinned + assert "skill_manage" in agent.valid_tool_names + assert json.loads(db.get_session("test-session-id")["tool_names"]) == pinned + + # --------------------------------------------------------------------------- # Legitimate fresh-build paths (no history, no DB) # --------------------------------------------------------------------------- diff --git a/tests/tools/test_refresh_agent_mcp_tools.py b/tests/tools/test_refresh_agent_mcp_tools.py index a0e9befb71..676c0abebf 100644 --- a/tests/tools/test_refresh_agent_mcp_tools.py +++ b/tests/tools/test_refresh_agent_mcp_tools.py @@ -336,6 +336,45 @@ def test_eviction_rebuild_restores_the_sessions_saved_tool_order(monkeypatch): assert rebuilt.valid_tool_names == set(saved) +def test_resume_on_another_surface_restores_the_pinned_tool_bytes(monkeypatch, tmp_path): + """One durable session hops gateway -> ``-q --resume``: the new process derives different + bytes for the SAME tools (tool_search's per-surface deferred catalog, a dynamic schema + override, the one-shot footprint pruning skill_manage). tools[] heads every request, so + the pin must hand back exactly what the session already sent, or every hop re-prefills.""" + from hermes_state import SessionDB + from tools import registry as registry_mod + + def _described(name, description): + tool = _tool(name) + tool["function"]["description"] = description + return tool + + sent = _agent([]) + sent.tools = [_tool("read_file"), _described("skill_manage", "lands in /home/u/.hermes/skills"), + _described("tool_search", "Search 6 additional tools.")] + static = {"skill_manage": _described("skill_manage", "lands in the profile's skills dir")["function"]} + monkeypatch.setattr(registry_mod.registry, "get_all_entries", + lambda: [types.SimpleNamespace(name=n) for n in ("read_file", "skill_manage")], raising=False) + monkeypatch.setattr(registry_mod.registry, "get_entry", + lambda name, **kw: types.SimpleNamespace(name=name, schema=static[name]), raising=False) + with SessionDB(db_path=tmp_path / "state.db") as db: + sent._session_db = db + for sid in ("s1", "s2"): + db.create_session(sid, source="tui") + sent.session_id = sid + _mcp_agent.persist_agent_tool_names(sent) + # Stored once, like the system prompt: a ~50KB array per session row would bloat state.db. + stored = db._conn.execute("SELECT COUNT(*) FROM system_prompts").fetchone()[0] + + resumed = _agent([]) + resumed.tools = [_tool("read_file"), _described("tool_search", "Search 5 additional tools.")] + _mcp_agent.restore_agent_tool_prefix(resumed, json.loads(db.get_session("s1")["tool_names"])) + + assert json.dumps(resumed.tools) == json.dumps(sent.tools) + assert resumed.valid_tool_names == {"read_file", "skill_manage", "tool_search"} + assert stored == 1 + + def test_reprobe_tool_availability_drops_cached_check_fn_verdicts(monkeypatch): """/reload-mcp is the explicit hatch: a cached False must be re-probed.""" from tools import registry as registry_mod diff --git a/tools/mcp_tool_agent.py b/tools/mcp_tool_agent.py index 148cc117fb..e73e4d21ac 100644 --- a/tools/mcp_tool_agent.py +++ b/tools/mcp_tool_agent.py @@ -149,47 +149,55 @@ def reprobe_tool_availability() -> None: def persist_agent_tool_names(agent) -> None: - """Best-effort: write ``agent.tools`` names to the session row (freeze pin).""" + """Best-effort: write ``agent.tools`` to the session row (freeze pin). The full definitions, + not just names: another process or surface derives different bytes for the same tool.""" db = getattr(agent, "_session_db", None) session_id = getattr(agent, "session_id", None) if not db or not session_id: return try: - db.update_session_tool_names(session_id, [_def_name(t) for t in _agent_tool_defs(agent)]) + db.update_session_tool_names(session_id, _agent_tool_defs(agent)) except Exception: # noqa: BLE001 logger.debug("tool_names persist skipped", exc_info=True) -def restore_agent_tool_prefix(agent, saved_names: list) -> bool: - """Fold a freshly built agent's ``tools`` onto the session's saved order; True if changed. - After agent-cache eviction the gateway rebuilds a NEW AIAgent with no predecessor to - preserve, so the saved name list stands in (``_merge_preserving_prefix`` rule; a saved - tool still registered but failing its probe is carried forward from the registry schema).""" - if not saved_names: +def restore_agent_tool_prefix(agent, saved: list) -> bool: + """Fold a freshly built agent's ``tools`` onto the session's pinned array; True if changed. + A fresh AIAgent (gateway cache eviction, ``--resume`` in a new process, a surface hop) has no + predecessor to preserve, so the pin stands in. A pinned tool still available here (built + fresh, or registered but failing its probe) keeps its pinned BYTES: this process derives + others for it (tool_search's per-surface catalog, dynamic schema overrides, the ``-q`` + footprint) and tools[] heads every request. Deregistered tools drop, tools new to this + process append at the tail. Legacy name-only pins take the fresh/registry schema once.""" + if not saved: return False from tools.registry import registry fresh_defs = _agent_tool_defs(agent) fresh = {_def_name(t): t for t in fresh_defs} + registered_names = {entry.name for entry in registry.get_all_entries()} - def _saved_def(name): - if name in fresh: - return fresh[name] - entry = registry.get_entry(name) + def _pinned_def(item): + if isinstance(item, dict): + return item + if item in fresh: + return fresh[item] + entry = registry.get_entry(item) return None if entry is None else {"type": "function", "function": {**entry.schema, "name": entry.name}} - saved_defs = [d for d in map(_saved_def, saved_names) if d is not None] - registered_names = {entry.name for entry in registry.get_all_entries()} - merged, merged_names = _merge_preserving_prefix(saved_defs, fresh_defs, registered_names) + merged = [d for d in map(_pinned_def, saved) if d and (_def_name(d) in fresh or _def_name(d) in registered_names)] + pinned_names = {_def_name(d) for d in merged} + merged.extend(t for t in fresh_defs if _def_name(t) not in pinned_names) + merged_names = {_def_name(t) for t in merged} _reinject_authorized_dynamic_tools(agent, merged, merged_names) merged, merged_names = _drop_side_agent_tools(agent, merged, merged_names) - with _agent_tools_lock: - if merged == fresh_defs: - return False - agent.tools = merged - agent.valid_tool_names = merged_names - if [_def_name(t) for t in merged] != list(saved_names): + changed = merged != fresh_defs + if changed: + with _agent_tools_lock: + agent.tools = merged + agent.valid_tool_names = merged_names + if merged != list(saved): persist_agent_tool_names(agent) - return True + return changed def _merge_preserving_prefix(current_defs: list, new_defs: list, registered_names: set) -> tuple[list, set]: