Files
hermes-agent/tests/test_project_recall_provenance.py
T

192 lines
9.5 KiB
Python

"""Durable source dependencies travel with the exact API sidecar, never source files."""
import json
import sqlite3
import threading
from contextlib import nullcontext
from types import SimpleNamespace
import pytest
from agent.session_persistence import _db_flush_row
from hermes_state import SessionDB
@pytest.fixture
def db_factory(tmp_path):
handles = []
def open_db():
db = SessionDB(db_path=tmp_path / "provenance.db")
handles.append(db)
return db
yield open_db
for db in handles:
db.close()
def test_provenance_survives_durable_replay_replace_and_clone(db_factory):
db = db_factory()
db.create_session("origin", source="cli")
sources = {"version": 1, "sources": [{"session_id": "source", "row_id": 7}],
"future_extension": {"unknown": [False, None, "opaque"]}}
wire = " exact prompt\n<recall>citation</recall>\n"
row_id = db.append_message("origin", "user", "question", api_content=wire,
api_content_sources=sources)
db.close()
db = db_factory()
row = db.get_messages("origin")[0]
assert row["api_content"] == wire
assert row["api_content_sources"] == sources
replay = db.get_messages_as_conversation("origin")
assert replay[0]["api_content_sources"] == sources
assert replay[0]["api_content"] == wire
for history in db.get_resume_conversations("origin"):
assert history[0]["api_content_sources"] == sources
db.replace_messages("origin", replay)
assert db.get_messages("origin")[0]["api_content_sources"] == sources
db.create_session("branch", source="cli", parent_session_id="origin",
model_config={"_branched_from": "origin"})
assert db.append_messages_batch("branch", db.get_messages("origin")) == 1
assert db.get_messages_as_conversation("branch")[0]["api_content_sources"] == sources
# Rows arriving after the watermark use the actual pure-SQL clone path.
assert db.archive_and_compact("origin", [], watermark=row_id) == 1
cloned = db.get_messages("origin")[0]
assert cloned["api_content"] == wire
assert cloned["api_content_sources"] == sources
projected = _db_flush_row(SimpleNamespace(_persist_user_message_override="question"),
{"role": "user", "content": wire,
"api_content_sources": json.dumps(sources)}, True)
db.create_session("flush", source="cli")
assert db.append_messages_batch("flush", [projected]) == 1
flushed = db.get_messages_as_conversation("flush")[0]
assert flushed["content"] == "question"
assert flushed["api_content"] == wire
assert flushed["api_content_sources"] == sources
def test_backfill_is_atomic_and_legacy_reconcile_preserves_unknown_metadata(db_factory):
db = db_factory()
db.create_session("legacy", source="cli")
db.append_message("legacy", "user", "question")
path = db.db_path
db.close()
# An older on-disk schema, not a mocked migration or source-text assertion.
with sqlite3.connect(path) as conn:
conn.execute("ALTER TABLE messages DROP COLUMN api_content_sources")
db = db_factory()
assert db.get_messages("legacy")[0]["api_content_sources"] is None
assert db.get_messages_as_conversation("legacy")[0].get("api_content_sources") is None
sources = {"version": 1, "sources": [{"row_id": 4}], "future": [1, False]}
encoded = json.dumps(sources)
with sqlite3.connect(path) as conn:
conn.execute("""CREATE TRIGGER require_atomic_pair BEFORE UPDATE ON messages
WHEN NEW.api_content = 'augmented' AND
NEW.api_content_sources IS NOT '%s'
BEGIN SELECT RAISE(ABORT, 'torn sidecar'); END""" % encoded)
assert db.set_latest_user_api_content("legacy", "question", "augmented", sources=sources) == 1
row = db.get_messages("legacy")[0]
assert (row["api_content"], row["api_content_sources"]) == ("augmented", sources)
with pytest.raises(sqlite3.IntegrityError, match="torn sidecar"):
db.set_latest_user_api_content("legacy", "question", "augmented", sources={"version": 9})
assert db.get_messages("legacy")[0] == row
db.append_message("legacy", "assistant", "answer")
db.append_message("legacy", "user", "different latest question")
assert db.set_latest_user_api_content("legacy", "question", "rejected", sources=sources) == 0
assert db.set_latest_user_api_content("missing", "question", "rejected", sources=sources) == 0
assert db.get_messages("legacy")[0] == row
assert db.get_messages("legacy")[-1]["api_content"] is None
# Existing positional callers still work; replacing a sidecar clears old dependencies.
assert db.set_latest_user_api_content("legacy", "different latest question", "old caller") == 1
assert db.get_messages("legacy")[-1]["api_content_sources"] is None
for opaque in ({}, [], False, 0, {"version": 999, "future": [None]}, "{broken-json"):
assert db.set_latest_user_api_content("legacy", "different latest question", "opaque", sources=opaque) == 1
replay = db.get_messages_as_conversation("legacy")
assert replay[-1]["api_content_sources"] == opaque
db.replace_messages("legacy", replay)
assert db.get_messages("legacy")[-1]["api_content_sources"] == opaque
# Encoding failure cannot update the sidecar before discovering invalid metadata.
before = db.get_messages("legacy")[-1]
with pytest.raises(TypeError):
db.set_latest_user_api_content("legacy", "different latest question", "must not land", sources={object()})
assert db.get_messages("legacy")[-1] == before
@pytest.mark.parametrize("path", ["cli", "gateway_branch", "gateway_transcript", "tui_branch", "tui_seed"])
def test_surface_persistence_serializers_preserve_provenance(db_factory, monkeypatch, path):
db = db_factory()
db.create_session("parent", source="cli")
sources = {"version": 1, "sources": [{"session_id": "source", "row_id": 7}]}
message = {"role": "user", "content": "question", "api_content": "question\nrecall",
"api_content_sources": sources}
target = "child"
if path == "cli":
import cli
from hermes_cli.cli_commands_mixin import CLICommandsMixin
monkeypatch.setattr(cli, "_sync_process_session_id", lambda _: None)
shell = SimpleNamespace(_session_db=db, session_id="parent", model="test",
max_turns=1, reasoning_config={}, agent=None,
conversation_history=[message],
_transfer_session_yolo=lambda *_: None)
CLICommandsMixin._handle_branch_command(shell, "/branch child")
target = shell.session_id
elif path == "gateway_branch":
from gateway.slash_commands_session import _branch_row
db.create_session(target, source="cli")
db.append_messages_batch(target, [_branch_row(message)])
elif path == "gateway_transcript":
from gateway.session_transcript import SessionTranscriptMixin
db.create_session(target, source="cli")
store = SimpleNamespace(_db_for_session_id=lambda _: db)
SessionTranscriptMixin._append_transcript_message(store, target, message)
elif path == "tui_branch":
from tui_gateway import server
monkeypatch.setattr(server, "_resolve_model", lambda: "test")
server._persist_branch(db, target, "parent", "child", [message],
source="cli", cwd=None, profile_name=None)
else:
from tui_gateway import session_workdir
db.create_session(target, source="cli")
monkeypatch.setattr(session_workdir, "_session_db", lambda _: nullcontext(db))
session_workdir._persist_branch_seed({"session_key": target, "parent_session_id": "parent", "seeded": True,
"history_lock": threading.Lock(), "history": [message]})
db.close()
restored = db_factory().get_messages_as_conversation(target)[0]
assert restored["api_content"] == message["api_content"]
assert restored["api_content_sources"] == sources
@pytest.mark.parametrize("role", ["user", "assistant"])
def test_gateway_replay_preserves_opaque_sidecar_dependencies(role):
from gateway.run import _build_replay_entry
message = {"role": role, "content": "question", "api_content": "question\nrecall"}
legacy = _build_replay_entry(role, message["content"], message)
assert legacy["api_content"] == message["api_content"]
assert "api_content_sources" not in legacy
for sources in (
{"version": 1, "sources": [{"session_id": "source", "row_id": 7}]},
None, {}, [], False, 0, "", "{broken-json",
{"version": 999, "future": [None, False]},
):
original = {**message, "api_content_sources": sources}
replay = _build_replay_entry(role, original["content"], original)
assert replay["content"] == original["content"]
assert replay["api_content"] == original["api_content"]
assert replay["api_content_sources"] == sources
assert type(replay["api_content_sources"]) is type(sources)
rewritten = _build_replay_entry(role, "rewritten content", original)
assert "api_content" not in rewritten
assert "api_content_sources" not in rewritten
without_sidecar = _build_replay_entry(
role, original["content"], {**original, "api_content": None}
)
assert "api_content" not in without_sidecar
assert "api_content_sources" not in without_sidecar
assert original == {**message, "api_content_sources": sources}