259 lines
12 KiB
Python
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 |