192 lines
9.5 KiB
Python
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} |