fix: one session sends byte-identical tools[] across TUI, oneshot and gateway hops

The session tools pin (sessions.tool_names) stored names only, so every fresh
process re-materialized the bytes from its own surface and every surface hop
of one durable session was a full prompt-cache miss:

* tool_search's deferred catalog is built per process ("Search 6 additional
  tools" in the TUI gateway vs 5 in -q);
* a pinned tool missing from the fresh build (skill_manage under the -q
  footprint) came back from the static registry schema, without its
  dynamic_schema_overrides;
* a -q --resume that rebuilt the stored prompt (model switch, cwd drift)
  persisted its own pruned array over the pin.

The pin now stores the full definitions and restore replays a pinned tool that
is still available byte-for-byte (deregistered tools drop, new ones append at
the tail, legacy name-only pins still work). A continuing session whose prompt
is rebuilt applies the pin before building it, matching the freeze policy
(tools[] only changes on /new, /reload-mcp, compaction). The array is
content-addressed in the existing system_prompts store like the prompt itself,
so identical arrays across sessions are stored once; get_session resolves it.
This commit is contained in:
teknium1
2026-09-23 06:58:16 -07:00
committed by Teknium
parent 01217c5fc2
commit 7a31c365c3
7 changed files with 149 additions and 46 deletions

View File

@@ -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):

View File

@@ -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

View File

@@ -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.

View File

@@ -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

View File

@@ -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)
# ---------------------------------------------------------------------------

View File

@@ -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

View File

@@ -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]: