Files
hermes-agent/agent/project_recall_guard.py
T

221 lines
11 KiB
Python

"""Fail-closed source checks for a future recall send gate; sends nothing.
Currently reuses tools.session_search_project.ProjectRecall's private safety
projection and eligibility SQL. Its scope resolver is optimistic metadata with
cached git probes, NOT fresh filesystem ownership or an atomic send lease.
"""
from __future__ import annotations
import copy
import hashlib
import math
import re
import time
from pathlib import Path
from typing import Any, cast
from agent.project_recall_grants import ProjectRecallGrantStore, normalize_route
from hermes_constants import get_hermes_home
from tools.session_search_project import ProjectRecall
class SafeContextRequired(Exception):
"""The caller must rebuild without protected recalled context."""
def __init__(self):
super().__init__("safe_context_required")
def _require(condition):
if not condition:
raise SafeContextRequired()
def normalize_recall_route(route: Any) -> dict:
"""Validate the explicit chat route with the grant store's canonicalization.
Source: agent.turn_api_request.build_api_request uses the live agent's
provider/base_url/api_mode/model per attempt (fallback can change them).
Only scheme/host case and default ports normalize. HTTP is loopback-only;
path, provider, mode and model stay exact, with no defaults or aliases.
The future caller must obtain this from the final transport, not model args.
"""
_require(type(route) is dict and set(route) == {"provider", "base_url", "api_mode", "model"})
route = cast(dict[str, Any], route)
_require(all(type(v) is str for v in route.values()))
_require(route["api_mode"] == "chat_completions")
try:
return normalize_route(route)
except (ValueError, TypeError):
raise SafeContextRequired() from None
def _deadline(deadline):
_require(type(deadline) in (int, float) and math.isfinite(deadline)
and time.monotonic() < deadline)
def _scope(db, current_session_id, sources):
home = get_hermes_home().resolve()
_require(sources["profile_home"] == str(home)
and Path(db.db_path).resolve() == home / "state.db")
recall = ProjectRecall(db, current_session_id)
_require(recall.scope.get("status") == "ready"
and recall.scope["project_key"] == sources["project_key"])
rows = recall._rows("SELECT user_id FROM sessions WHERE id=?", (current_session_id,))
_require(bool(rows) and rows[0]["user_id"] == sources["principal"])
return recall
def _backend(sources, grant_lookup, backend_namespace):
"""Resolve host authority, never namespace from evidence or lookup results."""
if backend_namespace is not None:
_require(type(backend_namespace) is str and bool(backend_namespace))
store = getattr(grant_lookup, "__self__", None)
if isinstance(store, ProjectRecallGrantStore):
_require(type(store.backend_namespace) is str and bool(store.backend_namespace)
and store.profile_home == sources["profile_home"])
_require(backend_namespace is None or backend_namespace == store.backend_namespace)
backend_namespace = store.backend_namespace
if sources["v"] == 2:
_require(backend_namespace is not None
and sources["backend_namespace"] == backend_namespace)
return backend_namespace
def _grant(sources: dict[str, Any], route, grant_lookup, backend_namespace):
_require(_backend(sources, grant_lookup, backend_namespace) == backend_namespace)
grant: Any = grant_lookup(sources["grant_id"])
_require(_backend(sources, grant_lookup, backend_namespace) == backend_namespace)
required = {"id", "project_key", "principal", "profile_home", "purpose", "route",
"revoked", "revision"}
_require(type(grant) is dict and required <= set(grant))
if backend_namespace is None:
# Real ledger records must not enter the unbound v1 compatibility path.
_require(not {"backend_namespace", "receipt_id", "created_at", "revoked_at"} & set(grant))
else:
_require("backend_namespace" in grant
and grant["backend_namespace"] == backend_namespace)
_require(grant["id"] == sources["grant_id"] and grant["revoked"] is False
and grant.get("revoked_at") is None
and grant["purpose"] == "history_to_chat"
and type(grant["revision"]) is int
and grant["revision"] == sources["grant_revision"]
and all(grant[k] == sources[k] for k in ("project_key", "principal", "profile_home"))
and normalize_recall_route(grant["route"]) == route)
def _source_text(recall, sources, deadline):
parts = []
for ref in sources["source_refs"]:
_deadline(deadline)
rows = recall._rows(
"SELECT m.* FROM messages m JOIN sessions s ON s.id=m.session_id WHERE "
+ recall._eligible_sql() + " AND m.session_id=? AND m.id=? AND s.user_id IS ?",
(recall.members_json, ref["session_id"], ref["message_id"], sources["principal"]),
)
safe = recall._safe(rows[0], content_length=None) if rows else 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"])
parts.append(safe["content"][start:start + length])
return "\n".join(parts)
def _validate(db, current_session_id, api_content, sources: Any, route, grant_lookup, deadline,
backend_namespace):
_deadline(deadline)
_require(type(api_content) is str and bool(api_content))
_require(type(sources) is dict)
fields = {
"v", "profile_home", "project_key", "principal", "current_session_id", "grant_id",
"grant_revision", "route", "source_refs", "sidecar_hash"}
sources = cast(dict[str, Any], sources)
_require(type(sources["v"]) is int and sources["v"] in (1, 2)
and sources["current_session_id"] == current_session_id
and type(sources["grant_revision"]) is int and sources["grant_revision"] > 0)
if sources["v"] == 2:
fields.add("backend_namespace")
_require(type(sources.get("backend_namespace")) is str
and bool(sources["backend_namespace"]))
_require(set(sources) == fields)
_require(all(type(sources[k]) is str and bool(sources[k]) for k in
("profile_home", "project_key", "current_session_id", "grant_id")))
_require(sources["principal"] is None or type(sources["principal"]) is str)
backend_namespace = _backend(sources, grant_lookup, backend_namespace)
route = normalize_recall_route(route)
_require(normalize_recall_route(sources["route"]) == route)
_require(hashlib.sha256(api_content.encode("utf-8")).hexdigest() == sources["sidecar_hash"])
_require(type(sources["source_refs"]) is list and bool(sources["source_refs"]))
for ref in sources["source_refs"]:
_require(type(ref) is dict and set(ref) == {
"session_id", "message_id", "content_hash", "offset", "length"})
ref = cast(dict[str, Any], ref)
_require(type(ref["session_id"]) is str and bool(ref["session_id"])
and type(ref["message_id"]) is int and 0 < ref["message_id"] <= 2**63 - 1
and type(ref["offset"]) is int and 0 <= ref["offset"] <= 2**63 - 1
and type(ref["length"]) is int and 0 < ref["length"] <= 2**63 - 1
and type(ref["content_hash"]) is str
and re.fullmatch(r"[0-9a-f]{64}", ref["content_hash"]) is not None)
recall = _scope(db, current_session_id, sources)
_grant(sources, route, grant_lookup, backend_namespace)
_deadline(deadline)
_require(_source_text(recall, sources, deadline) == api_content)
current = _scope(db, current_session_id, sources)
_require(current.scope["revision"] == recall.scope["revision"])
_require(_source_text(current, sources, deadline) == api_content)
_grant(sources, route, grant_lookup, backend_namespace)
final_scope = _scope(db, current_session_id, sources)
_require(final_scope.scope["revision"] == recall.scope["revision"])
_require(_backend(sources, grant_lookup, backend_namespace) == backend_namespace)
_deadline(deadline)
def validate_recall_sources(db, current_session_id, api_content, sources, route,
grant_lookup, deadline, *, backend_namespace=None) -> None:
"""Validate v1/v2 evidence or raise SafeContextRequired with a constant message.
api_content is ONLY the protected sidecar string, exactly the ordered source
slices joined by one newline. Hashes are SHA256 of UTF-8; offsets/lengths are
Python character indices into full safe-decoded text. No extra original text
is persisted. (None, None) means no protected sidecar; missing sources while
content remains is never safe. Caller must retain that protection marker.
grant_lookup(grant_id) must read current host-authorized state, never model
arguments. Required grant fields are id/project_key/principal/profile_home,
purpose/route/revoked/revision; additional audit fields are preserved. An
active record's revoked_at, if supplied, must be None. deadline is absolute
time.monotonic(); blocking dependencies are not preempted, but a late result
is denied. This check provides no authorization after return and is not yet
integrated with request serialization, retries, middleware, or transport.
v1 sources requires profile_home/project_key/principal/current_session_id,
grant_id/grant_revision (positive integer), route, source_refs, sidecar_hash.
v2 adds a required nonempty backend_namespace. Principal is an exact string
(including empty) or None. The trusted backend_namespace keyword or a bound
ProjectRecallGrantStore supplies authority, never the evidence/record itself.
If both authorities exist they must agree; bound stores must match profile.
A backend-aware lookup must return the same backend_namespace in its record.
Unbound callbacks without a namespace support only legacy v1 non-ledger
records; wrapping store.lookup requires the explicit trusted keyword. v1 is
retained for compatibility, not a backend-bound wire evidence format.
No persistent scope revision is stored: only revisions observed within this
invocation are compared. A later invocation can accept unrelated metadata
additions. Repeated reads detect observable races, not ABA or mutations
after the last read. Grant revocation must monotonically change revision.
"""
if api_content is None and sources is None:
return None
try:
source_snapshot, route_snapshot = copy.deepcopy(sources), copy.deepcopy(route)
_validate(db, current_session_id, api_content, source_snapshot,
route_snapshot, grant_lookup, deadline, backend_namespace)
_require(sources == source_snapshot and route == route_snapshot)
_deadline(deadline)
except Exception:
# Do not log or chain database/lookup exceptions containing source text.
raise SafeContextRequired() from None
return None