Files
hermes-agent/agent/project_recall_send.py

259 lines
12 KiB
Python

"""Prepare canonical model messages; no network or final wire authorization.
Only the host may supply canonical DB, resolved route, selected row references
and a bound consent store. Output still needs the existing agent message/reasoning
conversion and ChatCompletionsTransport.build_kwargs, then a final wire hash gate
on every actual attempt. This optimistic check is NOT an atomic send lease.
"""
from __future__ import annotations
import copy
import json
import math
import time
from pathlib import Path
from agent.project_recall_envelope import (
dependency_ref, render_evidence_text, validate_carrier_render, validate_derived_payload,
)
from agent.project_recall_grants import ProjectRecallGrantStore
from agent.project_recall_guard import (
SafeContextRequired, normalize_recall_route, validate_recall_sources,
)
from agent.project_recall_manifest import PAYLOAD_FIELDS, canonical_payload, validate_manifest_binding
from hermes_constants import get_hermes_home
from tools.session_search_project import ProjectRecall
MAX_DEPTH = 16
MAX_NODES = 256
_JSON_FIELDS = {"tool_calls", "reasoning_details", "codex_reasoning_items", "codex_message_items"}
_MODEL_FIELDS = {"role", "content", "tool_calls", "tool_call_id", "reasoning", "reasoning_content",
"reasoning_details"}
def _payload(db, row):
snapshot = {key: row[key] for key in PAYLOAD_FIELDS}
snapshot["content"] = db._decode_content(row["content"])
for key in _JSON_FIELDS:
if isinstance(snapshot[key], str):
snapshot[key] = json.loads(snapshot[key])
return canonical_payload(snapshot)
def _model_message(db, row):
snapshot = _payload(db, row)
if row["role"] in ("user", "assistant") and row.get("api_content"):
snapshot["content"] = row["api_content"]
_require(row["role"] in ("user", "assistant", "tool", "system", "developer"))
return {key: value for key, value in snapshot.items()
if key in _MODEL_FIELDS and (value is not None or key == "content")}
def _require(condition):
if not condition:
raise SafeContextRequired()
def _check_deadline(deadline):
_require(type(deadline) in (float, int) and math.isfinite(deadline)
and time.monotonic() < deadline)
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)
class _Preparation:
def __init__(self, db, consumer, route, store, deadline):
_check_deadline(deadline)
_require(type(consumer) is str and bool(consumer))
home = get_hermes_home().resolve()
_require(Path(db.db_path).resolve() == home / "state.db")
_require(isinstance(store, ProjectRecallGrantStore)
and store.profile_home == str(home)
and type(store.backend_namespace) is str and bool(store.backend_namespace))
self.db, self.consumer, self.store, self.deadline = db, consumer, store, deadline
self.home, self.backend = str(home), store.backend_namespace
self.rows, self.bindings, self.visiting, self.done = {}, {}, set(), set()
self.heights = {}
self.dependencies = []
self.route = normalize_recall_route(route)
self.recall = ProjectRecall(db, consumer)
_require(self.recall.scope.get("status") == "ready")
rows = self.recall._rows("SELECT user_id FROM sessions WHERE id=?", (consumer,))
_require(len(rows) == 1)
self.principal = rows[0]["user_id"]
with db._read_ctx() as conn:
self.has_manifest = conn.execute(
"SELECT 1 FROM sqlite_master WHERE type='table' AND name='project_recall_manifest'"
).fetchone() is not None
def row(self, sid, mid, *, source=False):
_check_deadline(self.deadline)
_identity(sid, mid)
_require(sid in self.recall.allowed)
where = self.recall._eligible_sql() if source else (
"m.session_id IN (SELECT value FROM json_each(?))")
rows = self.recall._rows(
"SELECT m.* FROM messages m JOIN sessions s ON s.id=m.session_id WHERE "
+ where + " AND m.session_id=? AND m.id=? AND s.user_id IS ?",
(self.recall.members_json, sid, mid, self.principal))
_require(len(rows) == 1)
key = (sid, mid)
_require(key not in self.rows or self.rows[key] == rows[0])
_require(key in self.rows or len(self.rows) < MAX_NODES)
self.rows[key] = rows[0]
return rows[0]
def binding(self, row):
if self.has_manifest:
manifest = validate_manifest_binding(self.db, row["session_id"], row["id"])
else:
_require(row.get("api_content_sources") is None)
manifest = None
key = (row["session_id"], row["id"])
_require(key not in self.bindings or self.bindings[key] == manifest)
self.bindings[key] = manifest
return manifest
def slices(self, evidence):
result = []
for ref in evidence["source_refs"]:
row = self.row(ref["session_id"], ref["message_id"], source=True)
_require(row.get("api_content_sources") is None and not row.get("api_content"))
_require(self.binding(row) is None)
safe = self.recall._safe(row, content_length=None)
if safe is None:
raise SafeContextRequired()
_require(safe["content_hash"] == ref["content_hash"])
start, length = ref["offset"], ref["length"]
_require(start + length <= safe["content_total_chars"])
result.append({**ref, "text": safe["content"][start:start + length]})
return result
def carrier(self, row, envelope):
evidence = envelope["evidence"]
origin = envelope["origin_session_id"]
_require(row["role"] == "user" and evidence["current_session_id"] == origin)
_require(origin in self.recall.allowed and evidence["principal"] == self.principal
and evidence["project_key"] == self.recall.scope["project_key"])
origin_rows = self.recall._rows("SELECT user_id FROM sessions WHERE id=?", (origin,))
_require(len(origin_rows) == 1 and origin_rows[0]["user_id"] == self.principal)
slices = self.slices(evidence)
validate_carrier_render(self.db._decode_content(row["content"]), row["api_content"],
envelope, slices)
validate_recall_sources(self.db, origin, render_evidence_text(slices), evidence,
self.route, self.store.lookup, self.deadline)
def resolve(self, ref):
_check_deadline(self.deadline)
_require(ref["profile_home"] == self.home and ref["backend_namespace"] == self.backend)
_require(ref["origin_session_id"] in self.recall.allowed)
origin = self.recall._rows("SELECT user_id FROM sessions WHERE id=?",
(ref["origin_session_id"],))
_require(len(origin) == 1 and origin[0]["user_id"] == self.principal)
# Resolve identity before reading any message body. Never use the
# manifest's UTF-8 envelope_hash as the envelope ASCII graph hash.
with self.db._read_ctx() as conn:
candidates = conn.execute("""SELECT session_id, message_id FROM project_recall_manifest
WHERE origin_uuid=? AND json_extract(envelope, '$.origin_session_id')=?
AND COALESCE(json_extract(envelope, '$.profile_home'),
json_extract(envelope, '$.evidence.profile_home'))=?
AND COALESCE(json_extract(envelope, '$.backend_namespace'),
json_extract(envelope, '$.evidence.backend_namespace'))=? LIMIT 2""",
(ref["message_key"], ref["origin_session_id"], self.home, self.backend)).fetchall()
_require(len(candidates) == 1)
return candidates[0]["session_id"], candidates[0]["message_id"]
def visit(self, sid, mid, depth=0, expected=None):
_require(depth <= MAX_DEPTH)
key = (sid, mid)
_require(key not in self.visiting)
row = self.row(sid, mid)
manifest = self.binding(row)
if manifest is not None:
envelope = json.loads(row["api_content_sources"])
ref = dependency_ref(envelope)
_require(self.resolve(ref) == key and (expected is None or ref == expected))
_require(manifest["origin_uuid"] == envelope["message_key"]
and manifest["protection_kind"] == envelope["kind"])
if key in self.done:
_require(depth + self.heights[key] <= MAX_DEPTH)
return row
self.visiting.add(key)
height = 0
if envelope["kind"] == "carrier":
self.carrier(row, envelope)
else:
_require(envelope["role"] == row["role"])
# Independent immutable manifest is the dependency witness, not
# the received envelope handed back as its own expected list.
validate_derived_payload(_payload(self.db, row), envelope,
manifest["envelope"]["dependencies"])
for child in envelope["dependencies"]:
child_sid, child_mid = self.resolve(child)
self.visit(child_sid, child_mid, depth + 1, child)
height = max(height, 1 + self.heights[(child_sid, child_mid)])
self.visiting.remove(key)
self.heights[key] = height
self.done.add(key)
self.dependencies.append(ref)
else:
_require(expected is None)
return row
def finish(self):
for sid, mid in list(self.rows):
self.binding(self.row(sid, mid))
for sid, mid in self.done:
row = self.rows[(sid, mid)]
envelope = json.loads(row["api_content_sources"])
_require(self.resolve(dependency_ref(envelope)) == (sid, mid))
if envelope["kind"] == "carrier":
self.carrier(row, envelope)
# Detect sidecar/source mutations observable during the final guards.
for sid, mid in list(self.rows):
self.binding(self.row(sid, mid))
fresh = ProjectRecall(self.db, self.consumer)
_require(fresh.scope.get("status") == "ready"
and fresh.scope["revision"] == self.recall.scope["revision"])
_require(self.store.profile_home == self.home and self.store.backend_namespace == self.backend)
_check_deadline(self.deadline)
def prepare_recall_send(db, consumer_session_id, message_refs, actual_route, grant_store,
deadline) -> dict:
"""Return {messages, dependencies}, or constant SafeContextRequired.
messages are canonical model input, NOT final provider wire payloads.
Preserve reasoning for the agent's provider-specific reasoning echo step;
transport conversion still strips nested call_id/response_item_id. Internal
top-level bookkeeping/provenance never appears in messages.
To persist derived output, pass manifest.canonical_payload(decoded raw DB
PAYLOAD_FIELDS) to inherit_dependencies(..., payload=...), not the projected
messages returned here. Both raw content and all tool/reasoning fields are
bound independently. dependencies is the deduplicated, postorder COMPLETE
protected closure; inherit this exact host-retained list in later outputs.
Limits: 16 dependency edges, 256 total rows including roots, and an absolute
monotonic deadline. Blocking DB calls are checked after return, not cancelled.
The caller must never substitute replay-cleaned text for these raw DB reads.
Re-run for every retry/fallback with the actual resolved route. No grants
are created, no schema installed and no external plugin is consulted here.
"""
try:
_require(type(message_refs) is list and len(message_refs) <= 256)
refs, route = copy.deepcopy(message_refs), copy.deepcopy(actual_route)
prep = _Preparation(db, consumer_session_id, route, grant_store, deadline)
messages = []
for item in refs:
_require(type(item) is dict and set(item) == {"session_id", "message_id"})
row = prep.visit(item["session_id"], item["message_id"])
messages.append(_model_message(db, row))
prep.finish()
_require(message_refs == refs and actual_route == route)
return {"messages": messages, "dependencies": prep.dependencies}
except Exception:
raise SafeContextRequired() from None