279 lines
12 KiB
Python
279 lines
12 KiB
Python
"""Host-owned recall evidence in the supplied canonical SessionDB.
|
|
|
|
This module sends nothing and is not wired into production writers/transports.
|
|
The caller owns canonical DB selection and envelope version/source authorization.
|
|
This is a database-scoped integrity contract, not a defense against hostile code
|
|
able to alter the database. A successful read is not an atomic send lease.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import hashlib
|
|
import json
|
|
import math
|
|
from uuid import UUID
|
|
|
|
from agent.project_recall_guard import SafeContextRequired
|
|
|
|
|
|
PAYLOAD_FIELDS = (
|
|
"role", "content", "tool_call_id", "tool_calls", "tool_name",
|
|
"effect_disposition", "finish_reason", "reasoning", "reasoning_content",
|
|
"reasoning_details", "codex_reasoning_items", "codex_message_items",
|
|
"observed", "_compressed_summary",
|
|
)
|
|
_JSON_FIELDS = {"tool_calls", "reasoning_details", "codex_reasoning_items", "codex_message_items"}
|
|
|
|
|
|
def _require(condition):
|
|
if not condition:
|
|
raise SafeContextRequired()
|
|
|
|
|
|
def _json_value(value):
|
|
if type(value) is dict:
|
|
_require(all(type(key) is str for key in value))
|
|
for child in value.values():
|
|
_json_value(child)
|
|
elif type(value) is list:
|
|
for child in value:
|
|
_json_value(child)
|
|
else:
|
|
_require(value is None or type(value) in (str, int, bool)
|
|
or (type(value) is float and math.isfinite(value)))
|
|
|
|
|
|
def _json(value):
|
|
try:
|
|
_json_value(value)
|
|
return json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"), allow_nan=False)
|
|
except Exception:
|
|
raise SafeContextRequired() from None
|
|
|
|
|
|
def _hash(value):
|
|
return hashlib.sha256(_json(value).encode("utf-8")).hexdigest()
|
|
|
|
|
|
def canonical_payload(snapshot: dict) -> dict:
|
|
"""v1 host payload: fixed whitelist, missing nullable fields become None.
|
|
|
|
observed/_compressed_summary default to 0. Unknown top-level fields are
|
|
rejected, never silently projected away. Nested JSON is retained in full.
|
|
This is NOT wire normalization: callers must supply decoded canonical DB
|
|
values (including structured tool/reasoning arrays), not provider objects.
|
|
"""
|
|
_require(type(snapshot) is dict and not set(snapshot) - set(PAYLOAD_FIELDS))
|
|
_require(type(snapshot.get("role")) is str and bool(snapshot["role"]))
|
|
result = {key: snapshot.get(key) for key in PAYLOAD_FIELDS}
|
|
for key in ("observed", "_compressed_summary"):
|
|
result[key] = snapshot.get(key, 0)
|
|
return json.loads(_json(result))
|
|
|
|
|
|
def payload_hash(snapshot: dict) -> str:
|
|
return _hash(canonical_payload(snapshot))
|
|
|
|
|
|
def content_hash(role: str, content) -> str:
|
|
"""SHA256(canonical JSON {'role': role, 'content': unsanitized decoded content})."""
|
|
return _hash({"role": role, "content": content})
|
|
|
|
|
|
def _row_payload(db, row, api_content):
|
|
snapshot = {key: row[key] for key in PAYLOAD_FIELDS}
|
|
snapshot["content"] = db._decode_content(row["content"]) if api_content is None else api_content
|
|
for key in _JSON_FIELDS:
|
|
if isinstance(snapshot[key], str):
|
|
snapshot[key] = json.loads(snapshot[key])
|
|
return snapshot
|
|
|
|
|
|
def install_schema(db) -> None:
|
|
"""Opt-in atomic install/upgrade; no constructor hook or foreign-key cascade.
|
|
|
|
Identity moves invalidate bindings and reserve destination IDs permanently,
|
|
including later moves/reuse across sessions. Reservations contain no origin
|
|
evidence and cannot become dependency-resolution candidates.
|
|
"""
|
|
def install(conn):
|
|
conn.execute("""CREATE TABLE IF NOT EXISTS project_recall_manifest (
|
|
session_id TEXT NOT NULL, message_id INTEGER NOT NULL,
|
|
content_hash TEXT NOT NULL, envelope_hash TEXT NOT NULL,
|
|
payload_hash TEXT NOT NULL, origin_uuid TEXT NOT NULL,
|
|
protection_kind TEXT NOT NULL CHECK(protection_kind IN ('carrier', 'derived')),
|
|
envelope TEXT NOT NULL, tombstone INTEGER NOT NULL DEFAULT 0,
|
|
PRIMARY KEY(session_id, message_id)
|
|
)""")
|
|
conn.execute("""CREATE TABLE IF NOT EXISTS project_recall_reserved_ids (
|
|
message_id INTEGER PRIMARY KEY
|
|
)""")
|
|
conn.execute("""CREATE TRIGGER IF NOT EXISTS project_recall_manifest_deleted
|
|
AFTER DELETE ON messages BEGIN
|
|
UPDATE project_recall_manifest SET tombstone=1
|
|
WHERE session_id=OLD.session_id AND message_id=OLD.id;
|
|
END""")
|
|
# REPLACE need not fire DELETE triggers with recursive_triggers disabled.
|
|
conn.execute("""CREATE TRIGGER IF NOT EXISTS project_recall_manifest_reused
|
|
AFTER INSERT ON messages BEGIN
|
|
UPDATE project_recall_manifest SET tombstone=1
|
|
WHERE message_id=NEW.id;
|
|
END""")
|
|
# UPDATE OF id misses SQLite rowid/_rowid_/oid assignments. Replace the
|
|
# old trigger in this same write transaction, including existing DBs.
|
|
conn.execute("DROP TRIGGER IF EXISTS project_recall_manifest_moved")
|
|
conn.execute("""CREATE TRIGGER project_recall_manifest_moved
|
|
AFTER UPDATE ON messages
|
|
WHEN OLD.id IS NOT NEW.id OR OLD.session_id IS NOT NEW.session_id BEGIN
|
|
INSERT INTO project_recall_reserved_ids(message_id)
|
|
SELECT NEW.id
|
|
WHERE EXISTS (SELECT 1 FROM project_recall_manifest WHERE message_id=OLD.id)
|
|
OR EXISTS (SELECT 1 FROM project_recall_reserved_ids WHERE message_id=OLD.id)
|
|
ON CONFLICT(message_id) DO NOTHING;
|
|
UPDATE project_recall_manifest SET tombstone=1
|
|
WHERE message_id=OLD.id OR message_id=NEW.id;
|
|
END""")
|
|
db._execute_write(install)
|
|
|
|
|
|
def _identity(session_id, message_id):
|
|
_require(type(session_id) is str and bool(session_id))
|
|
_require(type(message_id) is int and 0 < message_id <= 2**63 - 1)
|
|
|
|
|
|
def _manifest(conn, session_id, message_id):
|
|
row = conn.execute(
|
|
"SELECT * FROM project_recall_manifest WHERE session_id=? AND message_id=?",
|
|
(session_id, message_id),
|
|
).fetchone()
|
|
if row is None:
|
|
return None
|
|
result = dict(row)
|
|
result["envelope"] = json.loads(result["envelope"])
|
|
return result
|
|
|
|
|
|
def lookup_manifest(db, session_id: str, message_id: int) -> dict | None:
|
|
"""Read original evidence, including tombstones, NOT authorization.
|
|
|
|
Moved IDs may have only an ID reservation: None does not mean ordinary.
|
|
Use validate_manifest_binding to distinguish an ordinary current row.
|
|
"""
|
|
try:
|
|
_identity(session_id, message_id)
|
|
with db._read_ctx() as conn:
|
|
return _manifest(conn, session_id, message_id)
|
|
except Exception:
|
|
raise SafeContextRequired() from None
|
|
|
|
|
|
def validate_manifest_binding(db, session_id: str, message_id: int) -> dict | None:
|
|
"""Validate current canonical row; None means an existing ordinary message."""
|
|
try:
|
|
_identity(session_id, message_id)
|
|
with db._read_ctx() as conn:
|
|
return _validate(conn, db, session_id, message_id)
|
|
except Exception:
|
|
raise SafeContextRequired() from None
|
|
|
|
|
|
def _validate(conn, db, session_id, message_id):
|
|
_require(conn.execute("SELECT 1 FROM project_recall_reserved_ids WHERE message_id=?",
|
|
(message_id,)).fetchone() is None)
|
|
manifest = _manifest(conn, session_id, message_id)
|
|
row = conn.execute("SELECT * FROM messages WHERE session_id=? AND id=?",
|
|
(session_id, message_id)).fetchone()
|
|
_require(row is not None)
|
|
if manifest is None:
|
|
_require(row["api_content_sources"] is None)
|
|
_require(conn.execute("SELECT 1 FROM project_recall_manifest WHERE message_id=?",
|
|
(message_id,)).fetchone() is None)
|
|
return None
|
|
_require(not manifest["tombstone"])
|
|
_require(row["api_content_sources"] is not None)
|
|
envelope = json.loads(row["api_content_sources"])
|
|
_require(envelope == manifest["envelope"] and _hash(envelope) == manifest["envelope_hash"])
|
|
_require((manifest["protection_kind"] == "carrier" and type(row["api_content"]) is str)
|
|
or (manifest["protection_kind"] == "derived" and row["api_content"] is None))
|
|
_require(content_hash(row["role"], db._decode_content(row["content"])) == manifest["content_hash"])
|
|
_require(payload_hash(_row_payload(db, row, row["api_content"])) == manifest["payload_hash"])
|
|
for ref in _source_refs(envelope):
|
|
_identity(ref["session_id"], ref["message_id"])
|
|
_require(conn.execute(
|
|
"SELECT 1 FROM messages WHERE session_id=? AND id=?",
|
|
(ref["session_id"], ref["message_id"]),
|
|
).fetchone() is not None)
|
|
return manifest
|
|
|
|
|
|
def _source_refs(envelope):
|
|
"""Only direct row existence; source authorization/recursive closure is caller-owned."""
|
|
refs = envelope.get("source_refs", [])
|
|
_require(type(refs) is list)
|
|
evidence = envelope.get("evidence")
|
|
if evidence is not None:
|
|
_require(type(evidence) is dict)
|
|
nested = evidence.get("source_refs", [])
|
|
_require(type(nested) is list)
|
|
refs = refs + nested
|
|
return refs
|
|
|
|
|
|
def persist_protected(db, session_id: str, message_id: int, expected_content_hash: str,
|
|
api_content: str | None, envelope: dict) -> dict:
|
|
"""Atomically write sidecar + envelope + independent binding; failures are closed.
|
|
|
|
Envelope requires message_key (canonical UUID string) and kind. Its schema,
|
|
version, payload_hash and all other fields are opaque; optional source_refs
|
|
or evidence.source_refs get direct canonical row-existence checks only.
|
|
Dependencies are recorded verbatim, not traversed. Carrier payload content
|
|
is api_content; derived payload content is canonical decoded row content.
|
|
Existing bindings are immutable; exact retries are allowed, not repairs.
|
|
All hashes are lowercase SHA256 over UTF-8 canonical JSON: ensure_ascii=False,
|
|
sort_keys=True, separators=(',', ':'), allow_nan=False; no text sanitization.
|
|
"""
|
|
try:
|
|
_identity(session_id, message_id)
|
|
_require(type(envelope) is dict)
|
|
encoded = _json(envelope)
|
|
envelope = json.loads(encoded)
|
|
_require(str(UUID(envelope["message_key"])) == envelope["message_key"])
|
|
kind = envelope["kind"]
|
|
_require(kind in ("carrier", "derived"))
|
|
_require((kind == "carrier" and type(api_content) is str and bool(api_content))
|
|
or (kind == "derived" and api_content is None))
|
|
|
|
def persist(conn):
|
|
_require(conn.execute("SELECT 1 FROM project_recall_reserved_ids WHERE message_id=?",
|
|
(message_id,)).fetchone() is None)
|
|
row = conn.execute("SELECT * FROM messages WHERE session_id=? AND id=?",
|
|
(session_id, message_id)).fetchone()
|
|
if row is None:
|
|
raise SafeContextRequired()
|
|
actual = content_hash(row["role"], db._decode_content(row["content"]))
|
|
_require(actual == expected_content_hash)
|
|
bound_payload_hash = payload_hash(_row_payload(db, row, api_content))
|
|
old = _manifest(conn, session_id, message_id)
|
|
if old is not None:
|
|
_require(old["envelope_hash"] == _hash(envelope)
|
|
and old["payload_hash"] == bound_payload_hash
|
|
and row["api_content"] == api_content)
|
|
return _validate(conn, db, session_id, message_id)
|
|
_require(row["api_content_sources"] is None)
|
|
_require(conn.execute("SELECT 1 FROM project_recall_manifest WHERE message_id=?",
|
|
(message_id,)).fetchone() is None)
|
|
changed = conn.execute(
|
|
"UPDATE messages SET api_content=?, api_content_sources=? WHERE session_id=? AND id=?",
|
|
(api_content, encoded, session_id, message_id),
|
|
).rowcount
|
|
_require(changed == 1)
|
|
conn.execute("""INSERT INTO project_recall_manifest
|
|
(session_id, message_id, content_hash, envelope_hash, payload_hash,
|
|
origin_uuid, protection_kind, envelope) VALUES (?, ?, ?, ?, ?, ?, ?, ?)""",
|
|
(session_id, message_id, actual, _hash(envelope), bound_payload_hash,
|
|
envelope["message_key"], kind, encoded))
|
|
return _validate(conn, db, session_id, message_id)
|
|
|
|
return db._execute_write(persist)
|
|
except Exception:
|
|
raise SafeContextRequired() from None |