diff --git a/solstone/think/brain_cli.py b/solstone/think/brain_cli.py index 2474ab49c..6ace31cc6 100644 --- a/solstone/think/brain_cli.py +++ b/solstone/think/brain_cli.py @@ -34,15 +34,20 @@ from solstone.think.providers.brain_state import ( BrainEvidenceComponent, BrainProbeOutcome, BrainStateConflictError, + BrainStateExpectedFingerprintStaleError, BrainStateInspection, BrainStateRecord, + abandon_brain_prerequisite_renewal, abandon_brain_refresh, + begin_brain_prerequisite_renewal, begin_brain_refresh, brain_state_path, build_active_brain_fingerprint, + finish_brain_prerequisite_renewal, finish_brain_refresh, inspect_brain_state, probe_brain_refresh_lease_held, + read_active_brain_fingerprint_sha256, runtime_phase_reason, ) from solstone.think.providers.runtime_health import ( @@ -331,6 +336,21 @@ def _expected_fingerprint_matches(expected: str) -> bool: return _current_bundled_runtime_fingerprint(config) == expected +def _expected_active_fingerprint_matches(expected: str) -> bool: + try: + return read_active_brain_fingerprint_sha256() == expected + except Exception: + return False + + +def _expected_refresh_fingerprint_matches( + args: argparse.Namespace, expected: str +) -> bool: + if getattr(args, "expected_active_fingerprint", False): + return _expected_active_fingerprint_matches(expected) + return _expected_fingerprint_matches(expected) + + def _runtime_diagnostic( reason_code: str, *, @@ -422,6 +442,7 @@ def _spp_prerequisite(now: datetime) -> tuple[BrainEvidenceComponent, str | None try: recheck_confidential_attestation() except AttestationStaleError: + # Compatibility for verifier implementations that report stale directly. return _failed_component(now, "attestation_expired"), "attestation_expired" except AttestationFailedError: return _failed_component(now, "attestation_rejected"), "attestation_rejected" @@ -620,10 +641,82 @@ def _render_stale_expected(args: argparse.Namespace) -> int: return brain_exit_code(refresh_outcome="stale_expected_fingerprint") +def _refresh_args_for_active_fallback(args: argparse.Namespace) -> argparse.Namespace: + return argparse.Namespace( + json=args.json, + expected_fingerprint=args.expected_fingerprint, + expected_active_fingerprint=True, + expect_active_fingerprint_absent=False, + ) + + +def _run_renew_prerequisites(args: argparse.Namespace) -> int: + now = _now() + expected = args.expected_fingerprint + if expected and not _expected_active_fingerprint_matches(expected): + return _render_stale_expected(args) + + begin = begin_brain_prerequisite_renewal( + now, + expected_fingerprint_sha256=expected, + run_id=uuid.uuid4().hex, + ) + if begin["status"] == "busy": + inspection = inspect_brain_state(now) + view = _view_from_inspection(inspection, now) + busy = _transient_view( + "busy", + active_lane=view.active_lane, + active_provider=view.active_provider, + active_model=view.active_model, + fingerprint_sha256=view.fingerprint_sha256, + ) + render(busy, json_output=args.json) + return brain_exit_code(refresh_outcome="busy") + if begin["status"] == "unsafe": + return _run_refresh(_refresh_args_for_active_fallback(args)) + + permit = begin["permit"] + try: + lane_prerequisites, _reason = _spp_prerequisite(now) + except Exception: + LOG.exception("brain prerequisite renewal probe failed") + try: + record = abandon_brain_prerequisite_renewal( + permit, + "probe_internal_error", + now, + ) + except BrainStateConflictError: + lost = _transient_view("lost_fence") + render(lost, json_output=args.json) + return brain_exit_code(refresh_outcome="lost_fence") + else: + try: + record = finish_brain_prerequisite_renewal( + permit, + lane_prerequisites, + now, + ) + except BrainStateConflictError: + lost = _transient_view("lost_fence") + render(lost, json_output=args.json) + return brain_exit_code(refresh_outcome="lost_fence") + + inspection = inspect_brain_state(now) + view = _view_from_inspection(inspection, now) + render(view, json_output=args.json) + return brain_exit_code(aggregate_state=record["aggregate_state"]) + + def _run_refresh(args: argparse.Namespace) -> int: now = _now() expected = args.expected_fingerprint - if expected and not _expected_fingerprint_matches(expected): + expected_active = getattr(args, "expected_active_fingerprint", False) + expected_absent = getattr(args, "expect_active_fingerprint_absent", False) + if expected and expected_absent: + return _render_stale_expected(args) + if expected and not _expected_refresh_fingerprint_matches(args, expected): return _render_stale_expected(args) inspection = inspect_brain_state(now) @@ -631,11 +724,24 @@ def _run_refresh(args: argparse.Namespace) -> int: if view.reason_code == "configuration_invalid": render(view, json_output=args.json) return brain_exit_code(aggregate_state=view.aggregate_state) - if expected and view.aggregate_state == "ready": + if ( + expected + and not expected_active + and not expected_absent + and view.aggregate_state == "ready" + ): render(view, json_output=args.json) return brain_exit_code(aggregate_state=view.aggregate_state) - permit = begin_brain_refresh(now, run_id=uuid.uuid4().hex) + try: + permit = begin_brain_refresh( + now, + run_id=uuid.uuid4().hex, + expected_active_fingerprint_sha256=expected if expected_active else None, + expect_active_fingerprint_absent=expected_absent, + ) + except BrainStateExpectedFingerprintStaleError: + return _render_stale_expected(args) if permit is None: busy = False try: @@ -719,6 +825,25 @@ def build_parser() -> argparse.ArgumentParser: "--expected-fingerprint", help="Only refresh if the bundled runtime fingerprint still matches", ) + refresh_parser.add_argument( + "--expected-active-fingerprint", + action="store_true", + help="Interpret --expected-fingerprint as the active brain fingerprint", + ) + refresh_parser.add_argument( + "--expect-active-fingerprint-absent", + action="store_true", + help="Only refresh if no active brain fingerprint exists yet", + ) + renew_parser = subparsers.add_parser( + "renew-prerequisites", + help=argparse.SUPPRESS, + ) + renew_parser.add_argument("--json", action="store_true", help=argparse.SUPPRESS) + renew_parser.add_argument( + "--expected-fingerprint", + help=argparse.SUPPRESS, + ) return parser @@ -727,6 +852,8 @@ def _dispatch(args: argparse.Namespace, parser: argparse.ArgumentParser) -> int: return _run_status(args) if args.subcommand == "refresh": return _run_refresh(args) + if args.subcommand == "renew-prerequisites": + return _run_renew_prerequisites(args) parser.print_help() return 2 diff --git a/solstone/think/cortex.py b/solstone/think/cortex.py index 73351e9bb..a762d0805 100644 --- a/solstone/think/cortex.py +++ b/solstone/think/cortex.py @@ -26,14 +26,20 @@ import subprocess import sys import threading import time +import uuid +from collections.abc import Callable from datetime import datetime, timezone from pathlib import Path from typing import Any, Dict, Optional from solstone.think.callosum import CallosumConnection from solstone.think.models import calc_agent_cost -from solstone.think.providers.brain_state import inspect_brain_state +from solstone.think.providers.brain_state import ( + inspect_brain_state, + read_active_brain_fingerprint_sha256, +) from solstone.think.runner import _atomic_symlink +from solstone.think.services.spp_attest.cadence import TPM_HEARTBEAT_INTERVAL from solstone.think.talent import get_output_path from solstone.think.talents import TALENT_EXECUTION_MODULE from solstone.think.utils import get_journal, get_rev, now_ms @@ -100,10 +106,546 @@ class TalentProcess: return +SPP_RENEWAL_ATTEMPT_BOUND_S = 120.0 +SPP_REFRESH_OBSERVATION_BOUND_S = 300.0 +SPP_RENEWAL_RETRY_DELAYS_S = (5.0, 10.0, 20.0, 40.0, 60.0) +SPP_RENEWAL_PROACTIVE_MARGIN_S = ( + SPP_RENEWAL_ATTEMPT_BOUND_S + SPP_RENEWAL_RETRY_DELAYS_S[0] +) +SPP_RENEWAL_ACK_TIMEOUT_S = 15.0 +SPP_RENEWAL_MAX_WAIT_S = 60.0 +assert SPP_RENEWAL_PROACTIVE_MARGIN_S < TPM_HEARTBEAT_INTERVAL.total_seconds() / 2 + + +class SppRenewalController: + """Unattended SPP prerequisite renewal controller for Cortex.""" + + def __init__( + self, + *, + callosum: CallosumConnection, + stop_event: threading.Event, + logger: logging.Logger, + clock: Callable[[], datetime], + wait: Callable[[float], bool], + journal_path: Path, + ) -> None: + self.callosum = callosum + self.stop_event = stop_event + self.logger = logger + self.clock = clock + self.wait = wait + self.journal_path = journal_path + self._pending_ref: str | None = None + self._pending_action: str | None = None + self._pending_fingerprint: str | None = None + self._pending_expect_fingerprint_absent = False + self._pending_observed_at: datetime | None = None + self._pending_expires_at: datetime | None = None + self._ack_deadline: datetime | None = None + self._running_ref: str | None = None + self._running_action: str | None = None + self._running_fingerprint: str | None = None + self._running_expect_fingerprint_absent = False + self._running_observed_at: datetime | None = None + self._running_expires_at: datetime | None = None + self._running_deadline: datetime | None = None + self._successor_after_ref: str | None = None + self._successor_deadline: datetime | None = None + self._retry_index = 0 + self._retry_after: datetime | None = None + self._last_mode: str | None = None + + def run(self) -> None: + while not self.stop_event.is_set(): + try: + delay = self.step() + except Exception as exc: + now = self._now() + self._handle_step_exception(exc, now) + delay = self._seconds_until( + self._retry_after, now, default=SPP_RENEWAL_RETRY_DELAYS_S[0] + ) + self.wait(max(0.0, min(delay, SPP_RENEWAL_MAX_WAIT_S))) + + def step(self) -> float: + now = self._now() + try: + return self._step(now) + except Exception as exc: + self._handle_step_exception(exc, now) + return self._seconds_until( + self._retry_after, now, default=SPP_RENEWAL_RETRY_DELAYS_S[0] + ) + + def _step(self, now: datetime) -> float: + if self._spp_disabled(now): + self._clear_demand() + if self._last_mode != "disabled": + self._log("disabled", reason="non_spp_lane") + self._last_mode = "disabled" + return 30.0 + if self._pending_ref is not None: + if self._ack_deadline is not None and now >= self._ack_deadline: + self._log("failed", reason="start_ack_timeout", ref=self._pending_ref) + self._clear_pending() + self._schedule_retry(now) + return self._seconds_until(self._ack_deadline, now, default=5.0) + if self._running_ref is not None: + if self._running_deadline is not None and now >= self._running_deadline: + self._log( + "stale", reason="running_observation_timeout", ref=self._running_ref + ) + self._clear_running() + self._schedule_retry(now) + return self._seconds_until( + self._retry_after, now, default=SPP_RENEWAL_RETRY_DELAYS_S[0] + ) + return self._seconds_until(self._running_deadline, now, default=5.0) + if self._successor_after_ref is not None: + if self._successor_deadline is not None and now < self._successor_deadline: + return self._seconds_until( + self._successor_deadline, + now, + default=5.0, + ) + self._log( + "stale", + reason="successor_observation_timeout", + active_ref=self._successor_after_ref, + ) + self._clear_successor() + if self._retry_after is not None: + if now < self._retry_after: + return self._seconds_until(self._retry_after, now, default=5.0) + self._retry_after = None + + plan = self._plan(now) + if plan["action"] == "disabled": + self._clear_demand() + if self._last_mode != "disabled": + self._log("disabled", reason=plan.get("reason")) + self._last_mode = "disabled" + return 30.0 + self._last_mode = plan["action"] + if plan["action"] == "checking": + return 5.0 + if plan["action"] == "wait": + return float(plan["delay"]) + if plan["action"] in {"renew", "refresh"}: + if self._send_request(plan): + return SPP_RENEWAL_ACK_TIMEOUT_S + return self._seconds_until( + self._retry_after, + now, + default=SPP_RENEWAL_RETRY_DELAYS_S[0], + ) + return 30.0 + + def _spp_disabled(self, now: datetime) -> bool: + inspection = inspect_brain_state(now, journal_path=self.journal_path) + return inspection["projection"]["active_lane"] != "spp" + + def handle_supervisor_message(self, message: dict[str, Any]) -> None: + if message.get("tract") != "supervisor": + return + event = message.get("event") + ref = message.get("ref") + now = self._now() + if event == "started" and ref == self._pending_ref: + self._running_ref = self._pending_ref + self._running_action = self._pending_action + self._running_fingerprint = self._pending_fingerprint + self._running_expect_fingerprint_absent = ( + self._pending_expect_fingerprint_absent + ) + self._running_observed_at = self._pending_observed_at + self._running_expires_at = self._pending_expires_at + self._running_deadline = datetime.fromtimestamp( + now.timestamp() + self._observation_bound(self._pending_action), + tz=timezone.utc, + ) + self._log("in_flight", ref=self._running_ref, action=self._running_action) + self._clear_pending(keep_retry=True) + return + if event == "skipped" and ref == self._pending_ref: + active_ref = message.get("active_ref") + self._log( + "in_flight", + reason=str(message.get("reason") or "skipped"), + ref=ref, + active_ref=str(active_ref) if active_ref else None, + ) + self._successor_after_ref = str(active_ref) if active_ref else None + self._successor_deadline = ( + datetime.fromtimestamp( + now.timestamp() + SPP_REFRESH_OBSERVATION_BOUND_S, + tz=timezone.utc, + ) + if self._successor_after_ref is not None + else None + ) + self._clear_pending(keep_retry=True) + if self._successor_after_ref is None: + self._schedule_retry(now) + return + if event == "stopped": + if ref == self._running_ref: + self._verify_running_result(self._exit_code(message), now) + return + if ref == self._successor_after_ref: + self._clear_successor() + + def _plan(self, now: datetime) -> dict[str, Any]: + inspection = inspect_brain_state(now, journal_path=self.journal_path) + projection = inspection["projection"] + if projection["active_lane"] != "spp": + return {"action": "disabled", "reason": "non_spp_lane"} + if projection["aggregate_state"] == "checking": + return {"action": "checking"} + + fingerprint = read_active_brain_fingerprint_sha256( + journal_path=self.journal_path + ) + if fingerprint is None: + return { + "action": "refresh", + "fingerprint": None, + "expect_fingerprint_absent": True, + } + record = inspection["record"] + component = None + if record is not None: + component = record["evidence"].get("lane_prerequisites") + observed_at = None + expires_at = None + if isinstance(component, dict): + observed_at = self._parse_time(component.get("observed_at")) + expires_at = self._parse_time(component.get("expires_at")) + if ( + projection["aggregate_state"] != "ready" + or not isinstance(record, dict) + or record.get("fingerprint_sha256") != fingerprint + or not isinstance(component, dict) + or component.get("status") != "ok" + ): + return { + "action": "refresh", + "fingerprint": fingerprint, + "observed_at": observed_at, + "expires_at": expires_at, + } + + if observed_at is None or expires_at is None: + return { + "action": "refresh", + "fingerprint": fingerprint, + "observed_at": observed_at, + "expires_at": expires_at, + } + if now >= expires_at: + return { + "action": "refresh", + "fingerprint": fingerprint, + "observed_at": observed_at, + "expires_at": expires_at, + } + renew_at = expires_at.timestamp() - SPP_RENEWAL_PROACTIVE_MARGIN_S + delay = renew_at - now.timestamp() + if delay > 0: + self._log("scheduled", delay_s=round(delay, 3)) + return {"action": "wait", "delay": delay} + return { + "action": "renew", + "fingerprint": fingerprint, + "observed_at": observed_at, + "expires_at": expires_at, + } + + def _send_request(self, plan: dict[str, Any]) -> bool: + action = str(plan["action"]) + fingerprint = plan.get("fingerprint") + ref = f"spp-renewal-{uuid.uuid4().hex}" + if action == "renew": + cmd = [ + "journal", + "brain", + "renew-prerequisites", + "--json", + "--expected-fingerprint", + str(fingerprint), + ] + else: + expect_absent = bool(plan.get("expect_fingerprint_absent")) + if fingerprint is None and not expect_absent: + self._log("failed", reason="fingerprint_unavailable", action=action) + self._schedule_retry(self._now()) + return False + cmd = ["journal", "brain", "refresh", "--json"] + if expect_absent: + cmd.append("--expect-active-fingerprint-absent") + else: + cmd.extend( + [ + "--expected-fingerprint", + str(fingerprint), + "--expected-active-fingerprint", + ] + ) + now = self._now() + try: + self.callosum.emit( + "supervisor", + "request", + cmd=cmd, + ref=ref, + scheduler_name="spp-renewal", + ) + except Exception as exc: + self._log("failed", reason=type(exc).__name__, action=action) + self._schedule_retry(now) + return False + self._pending_ref = ref + self._pending_action = action + self._pending_fingerprint = str(fingerprint) if fingerprint else None + self._pending_expect_fingerprint_absent = bool( + plan.get("expect_fingerprint_absent") + ) + self._pending_observed_at = plan.get("observed_at") + self._pending_expires_at = plan.get("expires_at") + self._ack_deadline = datetime.fromtimestamp( + now.timestamp() + SPP_RENEWAL_ACK_TIMEOUT_S, tz=timezone.utc + ) + self._log("in_flight", ref=ref, action=action) + return True + + def _verify_running_result(self, exit_code: int, now: datetime) -> None: + action = self._running_action + fingerprint = self._running_fingerprint + expect_absent = self._running_expect_fingerprint_absent + previous_observed = self._running_observed_at + previous_expires = self._running_expires_at + ref = self._running_ref + self._clear_running() + if action == "renew" and self._persisted_spp_prerequisite_verified( + fingerprint, + previous_observed, + previous_expires, + now, + require_ready=False, + ): + self._retry_index = 0 + self._retry_after = None + self._log("verified", ref=ref) + return + if ( + action == "refresh" + and exit_code == 0 + and ( + self._persisted_spp_prerequisite_verified( + fingerprint, + previous_observed, + previous_expires, + now, + require_ready=True, + ) + if not expect_absent + else self._persisted_spp_absence_bootstrap_verified( + previous_observed, + previous_expires, + now, + ) + ) + ): + self._retry_index = 0 + self._retry_after = None + self._log("verified", ref=ref, action="refresh") + return + self._log("failed", ref=ref, action=action, exit_code=exit_code) + self._schedule_retry(now) + + def _persisted_spp_prerequisite_verified( + self, + fingerprint: str | None, + previous_observed: datetime | None, + previous_expires: datetime | None, + now: datetime, + *, + require_ready: bool, + ) -> bool: + if fingerprint is None: + return False + try: + active_fingerprint = read_active_brain_fingerprint_sha256( + journal_path=self.journal_path + ) + inspection = inspect_brain_state(now, journal_path=self.journal_path) + except Exception: + return False + if active_fingerprint != fingerprint: + return False + projection = inspection["projection"] + if projection["active_lane"] != "spp": + return False + if require_ready and projection["aggregate_state"] != "ready": + return False + record = inspection["record"] + if record is None or record["active_lane"] != "spp": + return False + if record["fingerprint_sha256"] != fingerprint: + return False + component = record["evidence"].get("lane_prerequisites") + if not isinstance(component, dict) or component.get("status") != "ok": + return False + observed_at = self._parse_time(component.get("observed_at")) + expires_at = self._parse_time(component.get("expires_at")) + return ( + observed_at is not None + and expires_at is not None + and (previous_observed is None or observed_at > previous_observed) + and (previous_expires is None or expires_at > previous_expires) + ) + + def _persisted_spp_absence_bootstrap_verified( + self, + previous_observed: datetime | None, + previous_expires: datetime | None, + now: datetime, + ) -> bool: + try: + active_fingerprint = read_active_brain_fingerprint_sha256( + journal_path=self.journal_path + ) + inspection = inspect_brain_state(now, journal_path=self.journal_path) + except Exception: + return False + if active_fingerprint is None: + return False + projection = inspection["projection"] + if ( + projection["active_lane"] != "spp" + or projection["aggregate_state"] != "ready" + ): + return False + record = inspection["record"] + if record is None or record["active_lane"] != "spp": + return False + if record["fingerprint_sha256"] != active_fingerprint: + return False + component = record["evidence"].get("lane_prerequisites") + if not isinstance(component, dict) or component.get("status") != "ok": + return False + observed_at = self._parse_time(component.get("observed_at")) + expires_at = self._parse_time(component.get("expires_at")) + return ( + observed_at is not None + and expires_at is not None + and (previous_observed is None or observed_at > previous_observed) + and (previous_expires is None or expires_at > previous_expires) + ) + + def _schedule_retry(self, now: datetime) -> None: + delay = SPP_RENEWAL_RETRY_DELAYS_S[ + min(self._retry_index, len(SPP_RENEWAL_RETRY_DELAYS_S) - 1) + ] + self._retry_index += 1 + self._retry_after = datetime.fromtimestamp( + now.timestamp() + delay, tz=timezone.utc + ) + self._log("retrying", delay_s=delay) + + def _handle_step_exception(self, exc: Exception, now: datetime) -> None: + self._log("failed", reason=type(exc).__name__) + self._clear_pending(keep_retry=True) + self._clear_running() + self._clear_successor() + self._schedule_retry(now) + + def _observation_bound(self, action: str | None) -> float: + if action == "refresh": + return SPP_REFRESH_OBSERVATION_BOUND_S + return SPP_RENEWAL_ATTEMPT_BOUND_S + + def _exit_code(self, message: dict[str, Any]) -> int: + value = message.get("exit_code") + if value is None: + return -1 + try: + return int(value) + except (TypeError, ValueError): + return -1 + + def _clear_pending(self, *, keep_retry: bool = False) -> None: + self._pending_ref = None + self._pending_action = None + self._pending_fingerprint = None + self._pending_expect_fingerprint_absent = False + self._pending_observed_at = None + self._pending_expires_at = None + self._ack_deadline = None + if not keep_retry: + self._retry_after = None + + def _clear_running(self) -> None: + self._running_ref = None + self._running_action = None + self._running_fingerprint = None + self._running_expect_fingerprint_absent = False + self._running_observed_at = None + self._running_expires_at = None + self._running_deadline = None + + def _clear_successor(self) -> None: + self._successor_after_ref = None + self._successor_deadline = None + + def _clear_demand(self) -> None: + self._clear_pending() + self._clear_running() + self._clear_successor() + self._retry_index = 0 + self._retry_after = None + + def _now(self) -> datetime: + return self.clock().astimezone(timezone.utc) + + def _seconds_until( + self, deadline: datetime | None, now: datetime, *, default: float + ) -> float: + if deadline is None: + return default + return max(0.0, deadline.timestamp() - now.timestamp()) + + def _parse_time(self, value: object) -> datetime | None: + if not isinstance(value, str): + return None + try: + return datetime.fromisoformat(value.replace("Z", "+00:00")).astimezone( + timezone.utc + ) + except ValueError: + return None + + def _log(self, event: str, **fields: object) -> None: + safe = " ".join( + f"{key}={value}" + for key, value in sorted(fields.items()) + if value is not None + ) + suffix = f" {safe}" if safe else "" + self.logger.info("event=spp_renewal_%s%s", event, suffix) + + class CortexService: """Callosum-based talent process manager.""" - def __init__(self, journal_path: Optional[str] = None): + def __init__( + self, + journal_path: Optional[str] = None, + *, + clock: Callable[[], datetime] | None = None, + wait: Callable[[float], bool] | None = None, + ): self.journal_path = Path(journal_path or get_journal()) self.talents_dir = self.journal_path / "talents" self.talents_dir.mkdir(parents=True, exist_ok=True) @@ -117,6 +659,10 @@ class CortexService: self.spawn_queue: queue.Queue = queue.Queue() self._pending_spawns: int = 0 self._spawn_worker: threading.Thread | None = None + self._clock = clock or (lambda: datetime.now(timezone.utc)) + self._wait = wait or self.stop_event.wait + self._spp_renewal_controller: SppRenewalController | None = None + self._spp_renewal_worker: threading.Thread | None = None # Callosum connection for receiving requests and broadcasting events self.callosum = CallosumConnection(defaults={"rev": get_rev()}) @@ -264,6 +810,7 @@ class CortexService: daemon=True, ) self._spawn_worker.start() + self._start_spp_renewal_controller() self.logger.info("Cortex service started, listening for talent requests") @@ -285,16 +832,38 @@ class CortexService: def _should_request_brain_refresh(self) -> bool: try: - inspection = inspect_brain_state(datetime.now(timezone.utc)) + inspection = inspect_brain_state( + self._clock(), journal_path=self.journal_path + ) except Exception: return True projection = inspection["projection"] + if projection["active_lane"] == "spp": + return False if projection["aggregate_state"] in {"checking", "ready"}: return False return not projection["runtime_transition_in_progress"] + def _start_spp_renewal_controller(self) -> None: + self._spp_renewal_controller = SppRenewalController( + callosum=self.callosum, + stop_event=self.stop_event, + logger=self.logger, + clock=self._clock, + wait=self._wait, + journal_path=self.journal_path, + ) + self._spp_renewal_worker = threading.Thread( + target=self._spp_renewal_controller.run, + name="cortex-spp-renewal", + daemon=True, + ) + self._spp_renewal_worker.start() + def _handle_callosum_message(self, message: Dict[str, Any]) -> None: """Handle incoming Callosum messages (callback).""" + if self._spp_renewal_controller is not None: + self._spp_renewal_controller.handle_supervisor_message(message) # Filter for cortex tract and request event if message.get("tract") != "cortex" or message.get("event") != "request": return @@ -907,6 +1476,9 @@ class CortexService: if self.callosum: self.callosum.stop() + if self._spp_renewal_worker is not None: + self._spp_renewal_worker.join(timeout=2.0) + # Let the spawn worker finish its current item and exit (~2s bound; the # worker's 0.5s get-timeout guarantees it observes stop_event promptly). if self._spawn_worker is not None: diff --git a/solstone/think/models.py b/solstone/think/models.py index 05d133595..f74420a32 100644 --- a/solstone/think/models.py +++ b/solstone/think/models.py @@ -292,8 +292,9 @@ class AttestationStaleError(AttestationNotVerifiedError): # Attestation failures are non-retryable. AttestationFailedError is raised by -# solstone.think.services.spp_attest.composite.verify_composite; AttestationStaleError -# is reserved for the follow-on verifier when a session's cadence windows lapse. +# solstone.think.services.spp_attest.composite.verify_composite. AttestationStaleError +# remains part of the verifier contract; the process-local egress transport rotates +# stale live sessions inline before forwarding bytes. _CONFIDENTIAL_ATTESTATION_VERIFIER: Callable[[dict[str, Any]], None] | None = None diff --git a/solstone/think/providers/brain_state.py b/solstone/think/providers/brain_state.py index 8b096c3ca..175d43d46 100644 --- a/solstone/think/providers/brain_state.py +++ b/solstone/think/providers/brain_state.py @@ -523,6 +523,15 @@ class BrainRuntimeFailureResult(TypedDict): error: str | None +BrainPrerequisiteRenewalStatus = Literal["started", "busy", "unsafe"] + + +class BrainPrerequisiteRenewalBeginResult(TypedDict): + status: BrainPrerequisiteRenewalStatus + permit: NotRequired["BrainRefreshPermit"] + reason: NotRequired[str] + + class BrainStateValidationError(ValueError): """Raised when a persisted brain state record violates the closed schema.""" @@ -536,6 +545,10 @@ class BrainStateConflictError(RuntimeError): """Raised when a stale refresh permit attempts to finalize.""" +class BrainStateExpectedFingerprintStaleError(BrainStateConflictError): + """Raised when a fenced refresh observes a different active fingerprint state.""" + + @dataclass class BrainRefreshPermit: """Held active-brain refresh permit.""" @@ -1885,18 +1898,39 @@ def begin_brain_refresh( now: datetime, *, run_id: str | None = None, + expected_active_fingerprint_sha256: str | None = None, + expect_active_fingerprint_absent: bool = False, journal_path: str | Path | None = None, ) -> BrainRefreshPermit | None: now = _utc(now) + if expected_active_fingerprint_sha256 is not None: + try: + _validate_hex( + expected_active_fingerprint_sha256, + "expected_active_fingerprint_sha256", + ) + except BrainStateValidationError as exc: + raise BrainStateExpectedFingerprintStaleError(str(exc)) from exc + if ( + expected_active_fingerprint_sha256 is not None + and expect_active_fingerprint_absent + ): + raise ValueError( + "expected active fingerprint and expected absence are mutually exclusive" + ) + expected_contract = ( + expected_active_fingerprint_sha256 is not None + or expect_active_fingerprint_absent + ) try: config = read_journal_config(journal_path) except (CorruptConfigError, OSError): return None lane, provider, model = _derive_lane(config) path = brain_state_path(journal_path=journal_path) - if lane is None: + if lane is None and not expected_contract: return None - if lane == "none": + if lane == "none" and not expected_contract: _begin_nonrefresh_record( now, path=path, @@ -1908,18 +1942,72 @@ def begin_brain_refresh( if lease is None: return None try: - try: - key = _load_or_generate_fingerprint_key(journal_path=journal_path) - except Exception: - lease.release() - return None - fingerprint = build_active_brain_fingerprint(config, hmac_key=key) - if fingerprint["active_lane"] is None or not fingerprint["ok"]: - lease.release() - return None run_id = run_id or uuid.uuid4().hex expires_at = now + CHECKING_TTL with hold_lock(path, mode=BRAIN_FILE_MODE): + try: + config = read_journal_config(journal_path) + except (CorruptConfigError, OSError): + lease.release() + return None + lane, provider, model = _derive_lane(config) + if lane is None: + if expected_contract: + raise BrainStateExpectedFingerprintStaleError( + "active brain fingerprint is unavailable" + ) + lease.release() + return None + if lane == "none": + if expected_contract: + raise BrainStateExpectedFingerprintStaleError( + "active brain fingerprint is unavailable" + ) + _begin_nonrefresh_record( + now, + path=path, + active_provider=provider, + active_model=model, + ) + lease.release() + return None + try: + key = _load_existing_fingerprint_key(journal_path=journal_path) + if expect_active_fingerprint_absent: + if key is not None: + raise BrainStateExpectedFingerprintStaleError( + "active brain fingerprint is present" + ) + key = _load_or_generate_fingerprint_key(journal_path=journal_path) + elif expected_active_fingerprint_sha256 is not None: + if key is None: + raise BrainStateExpectedFingerprintStaleError( + "active brain fingerprint is absent" + ) + else: + key = _load_or_generate_fingerprint_key(journal_path=journal_path) + except BrainStateExpectedFingerprintStaleError: + raise + except Exception: + lease.release() + return None + assert key is not None + fingerprint = build_active_brain_fingerprint(config, hmac_key=key) + if ( + expected_active_fingerprint_sha256 is not None + and fingerprint["fingerprint_sha256"] + != expected_active_fingerprint_sha256 + ): + raise BrainStateExpectedFingerprintStaleError( + "active brain fingerprint changed" + ) + if fingerprint["active_lane"] is None or not fingerprint["ok"]: + if expected_contract: + raise BrainStateExpectedFingerprintStaleError( + "active brain fingerprint is unavailable" + ) + lease.release() + return None current = _read_record_unlocked(path) revision = _next_revision(current) marker_seen = _runtime_failure_marker_id(current) @@ -1999,6 +2087,52 @@ def read_active_brain_fingerprint_sha256( return fingerprint["fingerprint_sha256"] +def _component_ok_unexpired( + component: BrainEvidenceComponent | None, + component_name: str, + now: datetime, +) -> bool: + if component is None or component["status"] != "ok": + return False + try: + expires_at = _parse_timestamp( + component.get("expires_at"), f"evidence.{component_name}.expires_at" + ) + except BrainStateValidationError: + return False + return now < expires_at + + +def _safe_prerequisite_renewal_evidence( + current: BrainStateRecord, + now: datetime, +) -> BrainEvidenceRecord | None: + if current["active_lane"] != "spp": + return None + if _record_timestamp_invalid(current, now): + return None + evidence = current["evidence"] + preserved: dict[str, BrainEvidenceComponent | None] = dict(evidence) + for component_name in ("configuration", "generate", "cogitate"): + if not _component_ok_unexpired(preserved[component_name], component_name, now): + return None + return cast(BrainEvidenceRecord, preserved) + + +def _prerequisite_renewal_begin_result( + status: BrainPrerequisiteRenewalStatus, + *, + permit: BrainRefreshPermit | None = None, + reason: str | None = None, +) -> BrainPrerequisiteRenewalBeginResult: + result: BrainPrerequisiteRenewalBeginResult = {"status": status} + if permit is not None: + result["permit"] = permit + if reason is not None: + result["reason"] = reason + return result + + def _record_from_evidence( *, evidence: BrainEvidenceRecord, @@ -2082,6 +2216,208 @@ def _assert_finish_allowed( raise BrainStateConflictError("brain runtime failure marker changed") +def begin_brain_prerequisite_renewal( + now: datetime, + *, + expected_fingerprint_sha256: str | None = None, + run_id: str | None = None, + journal_path: str | Path | None = None, +) -> BrainPrerequisiteRenewalBeginResult: + """Begin a fenced SPP prerequisite-only renewal. + + This is deliberately narrower than ``begin_brain_refresh``: it preserves + same-fingerprint model evidence and refuses to start unless replacing only + ``lane_prerequisites`` can be made safe. + """ + + try: + now = _utc(now) + except ValueError as exc: + return _prerequisite_renewal_begin_result("unsafe", reason=str(exc)) + if expected_fingerprint_sha256 is not None: + try: + _validate_hex(expected_fingerprint_sha256, "expected_fingerprint_sha256") + except BrainStateValidationError: + return _prerequisite_renewal_begin_result( + "unsafe", reason="fingerprint_mismatch" + ) + path = brain_state_path(journal_path=journal_path) + lease = acquire_file_lease(brain_refresh_lease_path(journal_path=journal_path)) + if lease is None: + return _prerequisite_renewal_begin_result("busy", reason="lease_held") + try: + try: + _config, key, fingerprint = _load_fingerprint_for_write( + journal_path=journal_path + ) + except (CorruptConfigError, OSError, BrainStateValidationError) as exc: + lease.release() + return _prerequisite_renewal_begin_result("unsafe", reason=str(exc)) + if key is None or fingerprint is None or not fingerprint["ok"]: + lease.release() + return _prerequisite_renewal_begin_result( + "unsafe", reason="fingerprint_not_available" + ) + if fingerprint["active_lane"] != "spp": + lease.release() + return _prerequisite_renewal_begin_result("unsafe", reason="non_spp_lane") + fingerprint_sha = fingerprint["fingerprint_sha256"] + if fingerprint_sha is None or ( + expected_fingerprint_sha256 is not None + and fingerprint_sha != expected_fingerprint_sha256 + ): + lease.release() + return _prerequisite_renewal_begin_result( + "unsafe", reason="fingerprint_mismatch" + ) + + run_id = run_id or uuid.uuid4().hex + expires_at = now + CHECKING_TTL + with hold_lock(path, mode=BRAIN_FILE_MODE): + try: + current = _read_record_unlocked(path) + except ( + OSError, + BrainStateValidationError, + MalformedDataError, + json.JSONDecodeError, + ): + lease.release() + return _prerequisite_renewal_begin_result( + "unsafe", reason="brain_record_unavailable" + ) + if current is None or current["fingerprint_sha256"] != fingerprint_sha: + lease.release() + return _prerequisite_renewal_begin_result( + "unsafe", reason="brain_record_missing" + ) + evidence = _safe_prerequisite_renewal_evidence(current, now) + if evidence is None: + lease.release() + return _prerequisite_renewal_begin_result( + "unsafe", reason="unsafe_evidence" + ) + revision = _next_revision(current) + marker_seen = _runtime_failure_marker_id(current) + checking: BrainCheckingRecord = { + "run_id": run_id, + "started_at": _iso(now), + "expires_at": _iso(expires_at), + "fingerprint_sha256": fingerprint_sha, + "checking_revision": revision, + "runtime_failure_marker_seen": marker_seen, + } + record = _record( + revision=revision, + aggregate_state="checking", + reason_code="brain_check_in_progress", + active_lane="spp", + active_provider=fingerprint["active_provider"], + active_model=fingerprint["active_model"], + fingerprint_sha256=fingerprint_sha, + checking=checking, + evidence=evidence, + runtime_failure_marker=current["runtime_failure_marker"], + diagnostic={}, + now=now, + ) + _write_record(path, record) + return _prerequisite_renewal_begin_result( + "started", + permit=BrainRefreshPermit( + run_id=run_id, + started_at=now, + expires_at=expires_at, + fingerprint_sha256=fingerprint_sha, + checking_revision=revision, + runtime_failure_marker_seen=marker_seen, + lease=lease, + ), + ) + except BaseException: + lease.release() + raise + + +def _validate_lane_prerequisite_component( + component: Mapping[str, Any], +) -> BrainEvidenceComponent: + validated = _validate_component( + component, + "lane_prerequisites", + component_name="lane_prerequisites", + ) + if validated is None or validated["status"] == "not_attempted": + raise BrainStateValidationError( + "lane_prerequisites.status", + "prerequisite renewal requires ok or declared failure", + ) + return validated + + +def finish_brain_prerequisite_renewal( + permit: BrainRefreshPermit, + lane_prerequisites: Mapping[str, Any], + now: datetime, + *, + journal_path: str | Path | None = None, +) -> BrainStateRecord: + now = _utc(now) + path = brain_state_path(journal_path=journal_path) + try: + component = _validate_lane_prerequisite_component(lane_prerequisites) + _config, key, fingerprint = _load_fingerprint_for_write( + journal_path=journal_path + ) + if key is None or fingerprint is None: + raise BrainStateConflictError("brain fingerprint key is unavailable") + with hold_lock(path, mode=BRAIN_FILE_MODE): + current = _read_record_unlocked(path) + _assert_finish_allowed(permit, current, now) + if fingerprint["fingerprint_sha256"] != permit.fingerprint_sha256: + raise BrainStateConflictError("brain fingerprint changed") + if fingerprint["active_lane"] != "spp": + raise BrainStateConflictError("brain lane changed") + assert current is not None + evidence = _safe_prerequisite_renewal_evidence(current, now) + if evidence is None: + raise BrainStateConflictError("brain prerequisite evidence is unsafe") + evidence["lane_prerequisites"] = component + record = _record_from_evidence( + evidence=evidence, + fingerprint=fingerprint, + revision=_next_revision(current), + now=now, + checking=None, + runtime_failure_marker=None, + ) + return _write_record(path, record) + finally: + permit.release() + + +def abandon_brain_prerequisite_renewal( + permit: BrainRefreshPermit, + reason_code: BrainReasonCode, + now: datetime, + *, + diagnostic: Mapping[str, BrainDiagnosticValue] | None = None, + journal_path: str | Path | None = None, +) -> BrainStateRecord: + component = _component( + _component_status_for_reason(reason_code), + _utc(now), + reason_code=reason_code, + diagnostic=diagnostic, + ) + return finish_brain_prerequisite_renewal( + permit, + component, + now, + journal_path=journal_path, + ) + + def finish_brain_refresh( permit: BrainRefreshPermit, outcome: BrainProbeOutcome, @@ -2320,23 +2656,29 @@ __all__ = [ "BrainLaneId", "BrainProbeOutcome", "BrainProjection", + "BrainPrerequisiteRenewalBeginResult", + "BrainPrerequisiteRenewalStatus", "BrainReasonCode", "BrainRefreshPermit", "BrainRuntimeFailureComponent", "BrainRuntimeFailureMarker", "BrainRuntimeFailureResult", + "BrainStateExpectedFingerprintStaleError", "BrainStateConflictError", "BrainStateInspection", "BrainStateRecord", "BrainStateValidationError", "abandon_brain_refresh", + "abandon_brain_prerequisite_renewal", "brain_fingerprint_key_path", "brain_refresh_lease_path", "brain_state_path", "begin_brain_refresh", + "begin_brain_prerequisite_renewal", "build_active_brain_fingerprint", "derive_active_brain_lane", "finish_brain_refresh", + "finish_brain_prerequisite_renewal", "inspect_brain_state", "project_brain_state", "read_active_brain_fingerprint_sha256", diff --git a/solstone/think/services/spp_attest/nvgpu/appraise.py b/solstone/think/services/spp_attest/nvgpu/appraise.py index 5e5e35861..ae450afbb 100644 --- a/solstone/think/services/spp_attest/nvgpu/appraise.py +++ b/solstone/think/services/spp_attest/nvgpu/appraise.py @@ -28,6 +28,7 @@ from solstone.think.services.spp_attest.snp import AppraisalStep from solstone.think.services.spp_attest.tlv import GpuEnvelope log = logging.getLogger(__name__) +NVATTEST_TIMEOUT_S = 60.0 def appraise_gpu_leg( @@ -85,7 +86,15 @@ def appraise_gpu_leg( capture_output=True, text=True, check=False, + timeout=NVATTEST_TIMEOUT_S, ) + except subprocess.TimeoutExpired as exc: + _log_gpu_appraisal_failure( + "gpu_appraisal_failed", + exception_class=type(exc).__name__, + stderr=exc.stderr, + ) + raise GpuAppraisalError("gpu_appraisal_failed") from exc except OSError as exc: _log_gpu_appraisal_failure( "nvattest_unavailable", diff --git a/solstone/think/services/spp_transport.py b/solstone/think/services/spp_transport.py index 5a34db8d1..5afa4af5c 100644 --- a/solstone/think/services/spp_transport.py +++ b/solstone/think/services/spp_transport.py @@ -18,7 +18,7 @@ from urllib.parse import urlsplit from OpenSSL import SSL -from solstone.think.models import AttestationFailedError, AttestationStaleError +from solstone.think.models import AttestationFailedError from solstone.think.providers.nvattest_install import ( ensure_nvattest_installed, nvattest_cache_ready, @@ -260,7 +260,7 @@ def _establish_and_record_locked( ) -def _reuse_or_raise_stale_locked(now: datetime) -> bool: +def _reuse_or_teardown_stale_locked(now: datetime) -> bool: state = spp.get_attestation_state() if ( state.session is not None @@ -274,9 +274,6 @@ def _reuse_or_raise_stale_locked(now: datetime) -> bool: and _transport_live_locked() ): _teardown_locked() - raise AttestationStaleError( - "the confidential attestation cadence lapsed (attestation_stale)" - ) return False @@ -286,13 +283,13 @@ def verify_confidential_attestation(block: dict[str, Any]) -> None: now = datetime.now(timezone.utc) with _LOCK: _CONFIDENTIAL_BLOCK = dict(block) - if _reuse_or_raise_stale_locked(now): + if _reuse_or_teardown_stale_locked(now): return nvattest_dir = _ensure_nvattest_for_attestation(block) with _LOCK: _CONFIDENTIAL_BLOCK = dict(block) - if _reuse_or_raise_stale_locked(now): + if _reuse_or_teardown_stale_locked(now): return _establish_and_record_locked(block, now, nvattest_dir=nvattest_dir) diff --git a/tests/services/test_spp_attest_cadence.py b/tests/services/test_spp_attest_cadence.py index 22cee16f6..702a1ac57 100644 --- a/tests/services/test_spp_attest_cadence.py +++ b/tests/services/test_spp_attest_cadence.py @@ -88,6 +88,25 @@ def test_attestation_session_verified_before_all_windows() -> None: ) +def test_attestation_session_verified_one_microsecond_before_each_boundary() -> None: + assert ( + _session(tpm_heartbeat_at=NOW - TPM_HEARTBEAT_INTERVAL).status( + NOW - timedelta(microseconds=1) + ) + == "verified" + ) + assert ( + _session(gpu_reattest_at=NOW - GPU_REATTEST_INTERVAL).status( + NOW - timedelta(microseconds=1) + ) + == "verified" + ) + assert ( + _session(started_at=NOW - SESSION_CAP).status(NOW - timedelta(microseconds=1)) + == "verified" + ) + + def test_attestation_session_stale_when_tpm_heartbeat_lapses() -> None: session = _session(tpm_heartbeat_at=NOW - TPM_HEARTBEAT_INTERVAL) diff --git a/tests/services/test_spp_attest_nvgpu.py b/tests/services/test_spp_attest_nvgpu.py index f9e2dc979..ebb346a8e 100644 --- a/tests/services/test_spp_attest_nvgpu.py +++ b/tests/services/test_spp_attest_nvgpu.py @@ -396,6 +396,45 @@ def test_gpu_appraisal_failure_log_uses_bounded_stderr_digest( assert "collector detail" not in message +def test_nvattest_timeout_is_bounded_redacted_and_cleans_evidence( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + caplog: pytest.LogCaptureFixture, +) -> None: + nvattest_dir = _fake_nvattest_dir(tmp_path) + observed: dict[str, Any] = {} + marker = "timeout-secret-nonce_from_ar" + + def timeout_run(argv, **kwargs): + observed["timeout"] = kwargs["timeout"] + evidence_path = Path(argv[argv.index("--gpu-evidence-file") + 1]) + observed["evidence_path"] = evidence_path + assert evidence_path.exists() + raise subprocess.TimeoutExpired(argv, kwargs["timeout"], stderr=marker) + + monkeypatch.setattr(appraise_module.subprocess, "run", timeout_run) + caplog.set_level(logging.WARNING, logger=appraise_module.log.name) + + with pytest.raises(GpuAppraisalError) as exc_info: + appraise_module.appraise_gpu_leg( + _envelope(), + _owner_nonce(), + nvattest_dir=nvattest_dir, + ) + + assert exc_info.value.reason == "gpu_appraisal_failed" + assert observed["timeout"] == appraise_module.NVATTEST_TIMEOUT_S + assert not observed["evidence_path"].exists() + messages = [ + record.getMessage() + for record in caplog.records + if "event=nvattest_gpu_appraisal_failed" in record.getMessage() + ] + assert len(messages) == 1 + assert "exception=TimeoutExpired" in messages[0] + assert marker not in messages[0] + + def test_bool_false_returncode_rejects( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, diff --git a/tests/services/test_spp_transport.py b/tests/services/test_spp_transport.py index d0c633064..8f122b83f 100644 --- a/tests/services/test_spp_transport.py +++ b/tests/services/test_spp_transport.py @@ -4,6 +4,7 @@ from __future__ import annotations import json +import threading import time from datetime import datetime, timedelta, timezone from pathlib import Path @@ -13,7 +14,7 @@ from unittest.mock import Mock import pytest from solstone.think.journal_io.locking import hold_lock -from solstone.think.models import AttestationFailedError, AttestationStaleError +from solstone.think.models import AttestationFailedError from solstone.think.providers import nvattest_install from solstone.think.services import spp, spp_transport from solstone.think.services.spp_attest.cadence import AttestationSession @@ -124,7 +125,7 @@ def _stale_session(verdict: object) -> AttestationSession: ) -def test_verify_confidential_attestation_reuses_then_raises_stale_once_then_reestablishes( +def test_verify_confidential_attestation_reuses_then_rotates_stale_session_inline( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: @@ -145,14 +146,6 @@ def test_verify_confidential_attestation_reuses_then_raises_stale_once_then_rees assert spp.get_attestation_state().session is not None spp.record_attestation_verified(_stale_session(verdict)) - with pytest.raises(AttestationStaleError): - spp_transport.verify_confidential_attestation(block) - - state = spp.get_attestation_state() - assert state.session is not None - assert state.session.status(datetime.now(timezone.utc)) == "stale" - assert establish.call_count == 1 - spp_transport.verify_confidential_attestation(block) assert establish.call_count == 2 @@ -161,6 +154,76 @@ def test_verify_confidential_attestation_reuses_then_raises_stale_once_then_rees spp.get_attestation_state().session.status(datetime.now(timezone.utc)) == "verified" ) + spp_transport.verify_confidential_attestation(block) + assert establish.call_count == 2 + + +def test_concurrent_stale_rotation_establishes_one_replacement( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + block = _write_confidential_config(tmp_path, monkeypatch) + _patch_listener(monkeypatch) + old_verdict = object() + old_channel = _FakeChannel(old_verdict) + spp.record_attestation_verified(_stale_session(old_verdict)) + with spp_transport._LOCK: + spp_transport._LISTENER = _FakeListener() + spp_transport._LISTENER_THREAD = _AliveThread() + spp_transport._FORWARDER_BASE_URL = "http://127.0.0.1:1111" + spp_transport._POOL[:] = [old_channel] + + ensure_barrier = threading.Barrier(2) + ensure_calls = 0 + + def delayed_ensure(_block): + nonlocal ensure_calls + ensure_calls += 1 + ensure_barrier.wait(timeout=2) + return Path("/tmp/solstone-nvattest-test") + + establish_started = threading.Event() + release_establish = threading.Event() + + def delayed_establish(*_args, **kwargs): + establish_started.set() + assert release_establish.wait(timeout=2) + return _FakeChannel(object(), epoch=kwargs["epoch"]) + + establish = Mock(side_effect=delayed_establish) + monkeypatch.setattr( + spp_transport, + "_ensure_nvattest_for_attestation", + delayed_ensure, + ) + monkeypatch.setattr(spp_transport, "establish_attested_channel", establish) + + errors: list[BaseException] = [] + + def verify() -> None: + try: + spp_transport.verify_confidential_attestation(block) + except BaseException as exc: # pragma: no cover - surfaced below + errors.append(exc) + + first = threading.Thread(target=verify) + second = threading.Thread(target=verify) + first.start() + second.start() + assert establish_started.wait(timeout=2) + assert first.is_alive() + assert second.is_alive() + release_establish.set() + first.join(timeout=2) + second.join(timeout=2) + + assert errors == [] + assert not first.is_alive() + assert not second.is_alive() + assert ensure_calls == 2 + assert establish.call_count == 1 + assert old_channel.closed is True + assert spp.get_attestation_state().session is not None def test_confidential_egress_base_url_returns_forwarder_not_configured_endpoint( @@ -238,6 +301,50 @@ def test_confidential_forwarder_base_url_requires_verified_forwarder( spp_transport.confidential_forwarder_base_url() +def test_stale_rotation_failure_returns_no_forwarder_and_later_recovers( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + block = _write_confidential_config(tmp_path, monkeypatch) + _patch_listener(monkeypatch) + old_verdict = object() + spp.record_attestation_verified(_stale_session(old_verdict)) + old_channel = _FakeChannel(old_verdict) + spp_transport._POOL[:] = [old_channel] + establish = Mock( + side_effect=[ + RatlsChannelError("gateway_unreachable"), + _FakeChannel(object()), + ] + ) + monkeypatch.setattr(spp_transport, "establish_attested_channel", establish) + pump = Mock(return_value="local_closed") + monkeypatch.setattr(spp_transport, "_pump", pump) + + with pytest.raises(AttestationFailedError): + spp_transport.confidential_egress_base_url(block["endpoint_url"]) + assert old_channel.closed is True + assert spp_transport._FORWARDER_BASE_URL is None + assert spp_transport._LISTENER is None + assert spp_transport._POOL == [] + assert spp_transport._ACTIVE == set() + assert spp.get_attestation_state().session is None + assert spp.get_attestation_state().failure is not None + local = _FakeLocal() + spp_transport._handle_loopback_connection(local) + assert local.closed is True + pump.assert_not_called() + + assert ( + spp_transport.confidential_egress_base_url(block["endpoint_url"]) + == "http://127.0.0.1:4567" + ) + assert establish.call_count == 2 + assert spp.get_attestation_state().session is not None + assert spp_transport._POOL + assert all(channel is not old_channel for channel in spp_transport._POOL) + + def test_confidential_probe_status_reads_state_without_attestation( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_brain_cli.py b/tests/test_brain_cli.py index a146409d2..77e53dedb 100644 --- a/tests/test_brain_cli.py +++ b/tests/test_brain_cli.py @@ -8,7 +8,7 @@ import hashlib import json import logging import sys -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from pathlib import Path from types import SimpleNamespace from typing import Any @@ -145,7 +145,12 @@ def _health_snapshot(journal: Path) -> dict[str, bytes]: def _args(**kwargs: Any) -> argparse.Namespace: - defaults = {"json": False, "expected_fingerprint": None} + defaults = { + "json": False, + "expected_fingerprint": None, + "expected_active_fingerprint": False, + "expect_active_fingerprint_absent": False, + } defaults.update(kwargs) return argparse.Namespace(**defaults) @@ -641,6 +646,104 @@ def test_refresh_prerequisite_failure_skips_inference_and_commits_not_attempted( assert record["evidence"]["cogitate"]["reason_code"] == "local_runtime_not_ready" +def test_renew_prerequisites_updates_spp_only_without_model_probes( + brain_journal: Path, + monkeypatch: pytest.MonkeyPatch, + capsys: pytest.CaptureFixture[str], +) -> None: + config = _endpoint_config(confidential=True) + prior_now = NOW - timedelta(minutes=1) + _write_ready_record(brain_journal, config, now=prior_now) + before = json.loads(brain_state_path(journal_path=brain_journal).read_text()) + expected = brain_state_module.read_active_brain_fingerprint_sha256( + journal_path=brain_journal + ) + assert expected is not None + _forbid_provider_calls(monkeypatch) + monkeypatch.setattr( + brain_cli, + "_spp_prerequisite", + lambda now: ( + brain_cli._ok_component(now, expires_at=now + timedelta(minutes=10)), + None, + ), + ) + + code = brain_cli._run_renew_prerequisites( + _args(json=True, expected_fingerprint=expected) + ) + payload = json.loads(capsys.readouterr().out) + record = json.loads(brain_state_path(journal_path=brain_journal).read_text()) + + assert code == 0 + assert payload["aggregate_state"] == "ready" + assert record["evidence"]["configuration"] == before["evidence"]["configuration"] + assert record["evidence"]["generate"] == before["evidence"]["generate"] + assert record["evidence"]["cogitate"] == before["evidence"]["cogitate"] + assert record["evidence"]["lane_prerequisites"]["observed_at"] == NOW.isoformat() + + +def test_renew_prerequisites_busy_does_not_fallback( + brain_journal: Path, + monkeypatch: pytest.MonkeyPatch, + capsys: pytest.CaptureFixture[str], +) -> None: + _write_ready_record(brain_journal, _endpoint_config(confidential=True)) + lease = acquire_file_lease(brain_refresh_lease_path(journal_path=brain_journal)) + assert lease is not None + monkeypatch.setattr( + brain_cli, + "_run_refresh", + lambda _args: pytest.fail("busy must not fallback to full refresh"), + ) + _forbid_provider_calls(monkeypatch) + try: + code = brain_cli._run_renew_prerequisites(_args(json=True)) + finally: + lease.release() + payload = json.loads(capsys.readouterr().out) + + assert code == 3 + assert payload["reason_code"] == "busy" + + +def test_renew_prerequisites_unsafe_dispatches_full_refresh( + brain_journal: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _write_config(brain_journal, _cloud_config()) + calls: list[argparse.Namespace] = [] + + def fake_refresh(args: argparse.Namespace) -> int: + calls.append(args) + return 7 + + monkeypatch.setattr(brain_cli, "_run_refresh", fake_refresh) + + assert brain_cli._run_renew_prerequisites(_args(json=True)) == 7 + assert len(calls) == 1 + assert calls[0].expected_active_fingerprint is True + + +def test_renew_prerequisites_expected_mismatch_writes_nothing( + brain_journal: Path, + monkeypatch: pytest.MonkeyPatch, + capsys: pytest.CaptureFixture[str], +) -> None: + _write_ready_record(brain_journal, _endpoint_config(confidential=True)) + before = _health_snapshot(brain_journal) + _forbid_provider_calls(monkeypatch) + + code = brain_cli._run_renew_prerequisites( + _args(json=True, expected_fingerprint="f" * 64) + ) + payload = json.loads(capsys.readouterr().out) + + assert code == 3 + assert payload["reason_code"] == "stale_expected_fingerprint" + assert _health_snapshot(brain_journal) == before + + def test_refresh_cloud_missing_key_skips_inference( brain_journal: Path, monkeypatch: pytest.MonkeyPatch, @@ -965,6 +1068,130 @@ def test_expected_fingerprint_mismatch_or_non_bundled_exits_stale( assert payload["reason_code"] == "stale_expected_fingerprint" +def test_refresh_expected_active_fingerprint_mismatch_writes_nothing( + brain_journal: Path, + monkeypatch: pytest.MonkeyPatch, + capsys: pytest.CaptureFixture[str], +) -> None: + original = _endpoint_config(confidential=True) + changed = _endpoint_config(confidential=True) + changed["providers"]["local"]["served_model_id"] = "new-served-model" + changed["services"]["confidential"]["served_model_id"] = "new-served-model" + _write_ready_record(brain_journal, original) + expected = brain_state_module.read_active_brain_fingerprint_sha256( + journal_path=brain_journal + ) + assert expected is not None + before = _health_snapshot(brain_journal) + _write_config(brain_journal, changed) + _forbid_provider_calls(monkeypatch) + + code = brain_cli._run_refresh( + _args( + json=True, + expected_fingerprint=expected, + expected_active_fingerprint=True, + ) + ) + payload = json.loads(capsys.readouterr().out) + + assert code == 3 + assert payload["reason_code"] == "stale_expected_fingerprint" + assert _health_snapshot(brain_journal) == before + + +def test_refresh_expected_active_fingerprint_race_before_begin_writes_nothing( + brain_journal: Path, + monkeypatch: pytest.MonkeyPatch, + capsys: pytest.CaptureFixture[str], +) -> None: + original = _endpoint_config(confidential=True) + changed = _endpoint_config(confidential=True) + changed["providers"]["local"]["served_model_id"] = "race-served-model" + changed["services"]["confidential"]["served_model_id"] = "race-served-model" + _write_ready_record(brain_journal, original) + expected = brain_state_module.read_active_brain_fingerprint_sha256( + journal_path=brain_journal + ) + assert expected is not None + before = _health_snapshot(brain_journal) + real_match = brain_cli._expected_refresh_fingerprint_matches + + def match_then_switch(args: argparse.Namespace, value: str) -> bool: + assert real_match(args, value) + _write_config(brain_journal, changed) + return True + + monkeypatch.setattr( + brain_cli, + "_expected_refresh_fingerprint_matches", + match_then_switch, + ) + _forbid_provider_calls(monkeypatch) + + code = brain_cli._run_refresh( + _args( + json=True, + expected_fingerprint=expected, + expected_active_fingerprint=True, + ) + ) + payload = json.loads(capsys.readouterr().out) + + assert code == 3 + assert payload["reason_code"] == "stale_expected_fingerprint" + assert _health_snapshot(brain_journal) == before + + +def test_refresh_expected_absent_fingerprint_bootstraps_spp( + brain_journal: Path, + monkeypatch: pytest.MonkeyPatch, + capsys: pytest.CaptureFixture[str], +) -> None: + _write_config(brain_journal, _endpoint_config(confidential=True)) + assert not brain_fingerprint_key_path(journal_path=brain_journal).exists() + monkeypatch.setattr( + brain_cli, + "_spp_prerequisite", + lambda now: ( + brain_cli._ok_component(now, expires_at=now + timedelta(minutes=10)), + None, + ), + ) + generate_calls, cogitate_calls = _patch_generate_and_cogitate(monkeypatch) + + code = brain_cli._run_refresh( + _args(json=True, expect_active_fingerprint_absent=True) + ) + payload = json.loads(capsys.readouterr().out) + + assert code == 0 + assert payload["aggregate_state"] == "ready" + assert payload["lane"] == "spp" + assert brain_fingerprint_key_path(journal_path=brain_journal).exists() + assert len(generate_calls) == 1 + assert len(cogitate_calls) == 1 + + +def test_refresh_expected_absent_fingerprint_stale_when_key_appears( + brain_journal: Path, + monkeypatch: pytest.MonkeyPatch, + capsys: pytest.CaptureFixture[str], +) -> None: + _write_ready_record(brain_journal, _endpoint_config(confidential=True)) + before = _health_snapshot(brain_journal) + _forbid_provider_calls(monkeypatch) + + code = brain_cli._run_refresh( + _args(json=True, expect_active_fingerprint_absent=True) + ) + payload = json.loads(capsys.readouterr().out) + + assert code == 3 + assert payload["reason_code"] == "stale_expected_fingerprint" + assert _health_snapshot(brain_journal) == before + + def test_expected_fingerprint_match_proceeds_and_ready_short_circuits( brain_journal: Path, monkeypatch: pytest.MonkeyPatch, diff --git a/tests/test_brain_state.py b/tests/test_brain_state.py index c3a5a29c8..fcf539c72 100644 --- a/tests/test_brain_state.py +++ b/tests/test_brain_state.py @@ -27,12 +27,16 @@ from solstone.think.providers.brain_state import ( DEFAULT_READY_EVIDENCE_TTL, BrainProbeOutcome, BrainStateConflictError, + BrainStateExpectedFingerprintStaleError, BrainStateValidationError, + abandon_brain_prerequisite_renewal, abandon_brain_refresh, + begin_brain_prerequisite_renewal, begin_brain_refresh, brain_fingerprint_key_path, brain_state_path, build_active_brain_fingerprint, + finish_brain_prerequisite_renewal, finish_brain_refresh, inspect_brain_state, project_brain_state, @@ -1052,6 +1056,392 @@ def test_runtime_failure_ingress_rejects_stale_fingerprint_without_write( assert _read_raw_record(tmp_path)["revision"] == prior_revision +def test_begin_refresh_expected_active_fingerprint_stales_after_switch( + tmp_path: Path, +) -> None: + original = _spp_config(account_id="acct-a") + _write_ready_record(tmp_path, original) + expected = _current_fingerprint(tmp_path, original) + before_record = brain_state_path(journal_path=tmp_path).read_bytes() + before_key = brain_fingerprint_key_path(journal_path=tmp_path).read_bytes() + _write_config(tmp_path, _spp_config(account_id="acct-b")) + + with pytest.raises(BrainStateExpectedFingerprintStaleError): + begin_brain_refresh( + NOW + timedelta(seconds=1), + expected_active_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + + assert brain_state_path(journal_path=tmp_path).read_bytes() == before_record + assert brain_fingerprint_key_path(journal_path=tmp_path).read_bytes() == before_key + + +def test_begin_refresh_expected_absent_fingerprint_bootstraps_key( + tmp_path: Path, +) -> None: + _write_config(tmp_path, _spp_config()) + assert not brain_fingerprint_key_path(journal_path=tmp_path).exists() + + permit = begin_brain_refresh( + NOW, + expect_active_fingerprint_absent=True, + journal_path=tmp_path, + ) + + assert permit is not None + assert brain_fingerprint_key_path(journal_path=tmp_path).exists() + record = finish_brain_refresh( + permit, + _ready_outcome(NOW + timedelta(seconds=1)), + NOW + timedelta(seconds=1), + journal_path=tmp_path, + ) + assert record["aggregate_state"] == "ready" + assert record["active_lane"] == "spp" + + +def test_begin_refresh_expected_absent_fingerprint_stales_when_key_exists( + tmp_path: Path, +) -> None: + _write_ready_record(tmp_path, _spp_config()) + before_record = brain_state_path(journal_path=tmp_path).read_bytes() + before_key = brain_fingerprint_key_path(journal_path=tmp_path).read_bytes() + + with pytest.raises(BrainStateExpectedFingerprintStaleError): + begin_brain_refresh( + NOW + timedelta(seconds=1), + expect_active_fingerprint_absent=True, + journal_path=tmp_path, + ) + + assert brain_state_path(journal_path=tmp_path).read_bytes() == before_record + assert brain_fingerprint_key_path(journal_path=tmp_path).read_bytes() == before_key + + +def test_prerequisite_renewal_preserves_same_fingerprint_model_evidence( + tmp_path: Path, +) -> None: + config = _spp_config() + _write_ready_record(tmp_path, config) + before = _read_raw_record(tmp_path) + expected = _current_fingerprint(tmp_path, config) + begin = begin_brain_prerequisite_renewal( + NOW + timedelta(minutes=1), + expected_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + + assert begin["status"] == "started" + checking = _read_raw_record(tmp_path) + assert checking["aggregate_state"] == "checking" + assert checking["evidence"]["generate"] == before["evidence"]["generate"] + assert checking["evidence"]["cogitate"] == before["evidence"]["cogitate"] + + finish_now = NOW + timedelta(minutes=2) + record = finish_brain_prerequisite_renewal( + begin["permit"], + _component(finish_now, expires_at=finish_now + timedelta(minutes=10)), + finish_now, + journal_path=tmp_path, + ) + + assert record["aggregate_state"] == "ready" + assert record["revision"] == before["revision"] + 2 + assert record["evidence"]["configuration"] == before["evidence"]["configuration"] + assert record["evidence"]["generate"] == before["evidence"]["generate"] + assert record["evidence"]["cogitate"] == before["evidence"]["cogitate"] + assert ( + record["evidence"]["lane_prerequisites"]["observed_at"] + == finish_now.isoformat() + ) + + +def test_prerequisite_renewal_refuses_expired_model_evidence(tmp_path: Path) -> None: + config = _spp_config() + _write_ready_record(tmp_path, config) + raw = _read_raw_record(tmp_path) + raw["evidence"]["generate"]["expires_at"] = (NOW - timedelta(seconds=1)).isoformat() + atomic_replace(brain_state_path(journal_path=tmp_path), json.dumps(raw), mode=0o600) + expected = _current_fingerprint(tmp_path, config) + + result = begin_brain_prerequisite_renewal( + NOW, + expected_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + + assert result["status"] == "unsafe" + assert _read_raw_record(tmp_path)["revision"] == raw["revision"] + + +def test_prerequisite_renewal_reports_busy_when_refresh_lease_is_held( + tmp_path: Path, +) -> None: + config = _spp_config() + _write_ready_record(tmp_path, config) + holder = begin_brain_refresh(NOW + timedelta(seconds=1), journal_path=tmp_path) + assert holder is not None + try: + result = begin_brain_prerequisite_renewal( + NOW + timedelta(seconds=2), + expected_fingerprint_sha256=_current_fingerprint(tmp_path, config), + journal_path=tmp_path, + ) + finally: + holder.release() + + assert result["status"] == "busy" + + +def test_prerequisite_renewal_refuses_non_spp_lane(tmp_path: Path) -> None: + config = _cloud_config() + _write_ready_record(tmp_path, config) + before = brain_state_path(journal_path=tmp_path).read_bytes() + + result = begin_brain_prerequisite_renewal( + NOW + timedelta(seconds=1), + expected_fingerprint_sha256=_current_fingerprint(tmp_path, config), + journal_path=tmp_path, + ) + + assert result["status"] == "unsafe" + assert brain_state_path(journal_path=tmp_path).read_bytes() == before + + +@pytest.mark.parametrize( + "mutation", + ( + lambda raw: raw["evidence"].__setitem__("generate", None), + lambda raw: raw["evidence"]["generate"].__setitem__("status", "failed"), + lambda raw: raw["evidence"]["generate"].__setitem__("expires_at", "not-time"), + ), +) +def test_prerequisite_renewal_refuses_missing_malformed_or_non_ok_model_evidence( + tmp_path: Path, + mutation, +) -> None: + config = _spp_config() + _write_ready_record(tmp_path, config) + raw = _read_raw_record(tmp_path) + mutation(raw) + atomic_replace(brain_state_path(journal_path=tmp_path), json.dumps(raw), mode=0o600) + before = brain_state_path(journal_path=tmp_path).read_bytes() + + result = begin_brain_prerequisite_renewal( + NOW + timedelta(seconds=1), + expected_fingerprint_sha256=_current_fingerprint(tmp_path, config), + journal_path=tmp_path, + ) + + assert result["status"] == "unsafe" + assert brain_state_path(journal_path=tmp_path).read_bytes() == before + + +def test_prerequisite_renewal_expected_fingerprint_mismatch_is_no_write( + tmp_path: Path, +) -> None: + config = _spp_config() + _write_ready_record(tmp_path, config) + before = brain_state_path(journal_path=tmp_path).read_bytes() + + result = begin_brain_prerequisite_renewal( + NOW + timedelta(seconds=1), + expected_fingerprint_sha256="0" * 64, + journal_path=tmp_path, + ) + + assert result["status"] == "unsafe" + assert result["reason"] == "fingerprint_mismatch" + assert brain_state_path(journal_path=tmp_path).read_bytes() == before + + +def test_prerequisite_renewal_recovers_orphaned_checking_after_lease_released( + tmp_path: Path, +) -> None: + config = _spp_config() + _write_ready_record(tmp_path, config) + expected = _current_fingerprint(tmp_path, config) + first = begin_brain_prerequisite_renewal( + NOW + timedelta(seconds=1), + expected_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + assert first["status"] == "started" + first["permit"].release() + orphaned = _read_raw_record(tmp_path) + assert orphaned["aggregate_state"] == "checking" + + second = begin_brain_prerequisite_renewal( + NOW + timedelta(seconds=2), + expected_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + assert second["status"] == "started" + recovered = _read_raw_record(tmp_path) + assert recovered["revision"] == orphaned["revision"] + 1 + assert recovered["checking"]["run_id"] != orphaned["checking"]["run_id"] + second["permit"].release() + + +def test_prerequisite_renewal_expired_permit_conflicts_and_releases( + tmp_path: Path, +) -> None: + config = _spp_config() + _write_ready_record(tmp_path, config) + expected = _current_fingerprint(tmp_path, config) + begin = begin_brain_prerequisite_renewal( + NOW, + expected_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + assert begin["status"] == "started" + permit = begin["permit"] + + with pytest.raises(BrainStateConflictError): + finish_brain_prerequisite_renewal( + permit, + _component( + permit.expires_at, expires_at=permit.expires_at + timedelta(minutes=10) + ), + permit.expires_at, + journal_path=tmp_path, + ) + + retry = begin_brain_prerequisite_renewal( + NOW + timedelta(seconds=1), + expected_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + assert retry["status"] == "started" + retry["permit"].release() + + +def test_prerequisite_renewal_conflicts_on_revision_drift_and_releases( + tmp_path: Path, +) -> None: + config = _spp_config() + _write_ready_record(tmp_path, config) + expected = _current_fingerprint(tmp_path, config) + begin = begin_brain_prerequisite_renewal( + NOW, + expected_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + assert begin["status"] == "started" + raw = _read_raw_record(tmp_path) + raw["revision"] += 1 + raw["checking"]["checking_revision"] = raw["revision"] + atomic_replace(brain_state_path(journal_path=tmp_path), json.dumps(raw), mode=0o600) + + with pytest.raises(BrainStateConflictError): + finish_brain_prerequisite_renewal( + begin["permit"], + _component(NOW + timedelta(seconds=1)), + NOW + timedelta(seconds=1), + journal_path=tmp_path, + ) + + retry = begin_brain_prerequisite_renewal( + NOW + timedelta(seconds=2), + expected_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + assert retry["status"] == "started" + retry["permit"].release() + + +def test_prerequisite_renewal_conflicts_on_runtime_marker_drift_and_releases( + tmp_path: Path, +) -> None: + config = _spp_config() + _write_ready_record(tmp_path, config) + expected = _current_fingerprint(tmp_path, config) + begin = begin_brain_prerequisite_renewal( + NOW, + expected_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + assert begin["status"] == "started" + raw = _read_raw_record(tmp_path) + raw["checking"]["runtime_failure_marker_seen"] = "changed-marker" + atomic_replace(brain_state_path(journal_path=tmp_path), json.dumps(raw), mode=0o600) + + with pytest.raises(BrainStateConflictError): + finish_brain_prerequisite_renewal( + begin["permit"], + _component(NOW + timedelta(seconds=1)), + NOW + timedelta(seconds=1), + journal_path=tmp_path, + ) + + retry = begin_brain_prerequisite_renewal( + NOW + timedelta(seconds=2), + expected_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + assert retry["status"] == "started" + retry["permit"].release() + + +def test_prerequisite_renewal_conflicts_on_fingerprint_drift(tmp_path: Path) -> None: + config = _spp_config(account_id="acct-a") + _write_ready_record(tmp_path, config) + expected = _current_fingerprint(tmp_path, config) + begin = begin_brain_prerequisite_renewal( + NOW, + expected_fingerprint_sha256=expected, + journal_path=tmp_path, + ) + assert begin["status"] == "started" + _write_config(tmp_path, _spp_config(account_id="acct-b")) + + with pytest.raises(BrainStateConflictError): + finish_brain_prerequisite_renewal( + begin["permit"], + _component(NOW, expires_at=NOW + timedelta(minutes=10)), + NOW, + journal_path=tmp_path, + ) + + retry = begin_brain_prerequisite_renewal( + NOW + timedelta(seconds=1), + expected_fingerprint_sha256=_current_fingerprint( + tmp_path, _spp_config(account_id="acct-b") + ), + journal_path=tmp_path, + ) + assert retry["status"] == "unsafe" + + +def test_prerequisite_renewal_abandon_records_failure_without_losing_model_evidence( + tmp_path: Path, +) -> None: + config = _spp_config() + _write_ready_record(tmp_path, config) + before = _read_raw_record(tmp_path) + begin = begin_brain_prerequisite_renewal( + NOW, + expected_fingerprint_sha256=_current_fingerprint(tmp_path, config), + journal_path=tmp_path, + ) + assert begin["status"] == "started" + + record = abandon_brain_prerequisite_renewal( + begin["permit"], + "probe_internal_error", + NOW + timedelta(seconds=1), + journal_path=tmp_path, + ) + + assert record["reason_code"] == "probe_internal_error" + assert record["evidence"]["lane_prerequisites"]["reason_code"] == ( + "probe_internal_error" + ) + assert record["evidence"]["generate"] == before["evidence"]["generate"] + assert record["evidence"]["cogitate"] == before["evidence"]["cogitate"] + + def test_key_replacement_invalidates_prior_ready_record(tmp_path: Path) -> None: _write_ready_record(tmp_path, _cloud_config()) inspection = inspect_brain_state(NOW, journal_path=tmp_path) diff --git a/tests/test_brain_state_multiprocess.py b/tests/test_brain_state_multiprocess.py index 4fb8c36ce..a4ac85f80 100644 --- a/tests/test_brain_state_multiprocess.py +++ b/tests/test_brain_state_multiprocess.py @@ -3,15 +3,22 @@ from __future__ import annotations +import hashlib import json import os import subprocess import sys import time -from datetime import datetime, timezone +from datetime import datetime, timedelta, timezone from pathlib import Path -from solstone.think.providers.brain_state import inspect_brain_state +from solstone.think.models import LOCAL_MODEL +from solstone.think.providers.brain_state import ( + DEFAULT_READY_EVIDENCE_TTL, + begin_brain_refresh, + finish_brain_refresh, + inspect_brain_state, +) NOW = datetime(2026, 1, 2, 3, 4, 5, tzinfo=timezone.utc) @@ -36,6 +43,70 @@ def _write_config(journal: Path) -> None: ) +def _write_spp_config(journal: Path) -> None: + credential = "endpoint-secret" + path = journal / "config" / "journal.json" + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text( + json.dumps( + { + "providers": { + "active": {"provider": "local", "model": LOCAL_MODEL}, + "local": { + "endpoint_url": "https://brain.example.test/v1", + "served_model_id": "served-model", + "credential": credential, + }, + }, + "services": { + "confidential": { + "enabled_at": NOW.isoformat(), + "account_id": "acct-a", + "endpoint_url": "https://brain.example.test/v1", + "served_model_id": "served-model", + "credential_created_at": NOW.isoformat(), + "credential_fingerprint_sha256": hashlib.sha256( + credential.encode("utf-8") + ).hexdigest(), + "prior_active": { + "provider": "google", + "model": "gemini-flash-latest", + }, + "prior_local_endpoint": None, + } + }, + "env": {}, + } + ), + encoding="utf-8", + ) + + +def _component(now: datetime) -> dict[str, str]: + return { + "status": "ok", + "observed_at": now.isoformat(), + "expires_at": (now + DEFAULT_READY_EVIDENCE_TTL).isoformat(), + } + + +def _write_spp_ready_record(journal: Path) -> None: + _write_spp_config(journal) + permit = begin_brain_refresh(NOW, journal_path=journal) + assert permit is not None + finish_brain_refresh( + permit, + { + "configuration": _component(NOW), + "lane_prerequisites": _component(NOW), + "generate": _component(NOW), + "cogitate": _component(NOW), + }, + NOW, + journal_path=journal, + ) + + def test_refresh_permit_excludes_contender_and_crash_releases(tmp_path: Path) -> None: _write_config(tmp_path) ready = tmp_path / "ready" @@ -107,3 +178,72 @@ else: if holder.poll() is None: holder.terminate() holder.wait(timeout=5) + + +def test_prerequisite_renewal_permit_excludes_contender_and_crash_releases( + tmp_path: Path, +) -> None: + _write_spp_ready_record(tmp_path) + ready = tmp_path / "ready" + holder_code = f""" +import pathlib +import time +from datetime import datetime +from solstone.think.providers.brain_state import begin_brain_prerequisite_renewal +now = datetime.fromisoformat({(NOW + timedelta(seconds=1)).isoformat()!r}) +result = begin_brain_prerequisite_renewal(now, journal_path={str(tmp_path)!r}) +assert result["status"] == "started", result +pathlib.Path({str(ready)!r}).write_text("ready") +while True: + time.sleep(0.05) +""" + contender_code = f""" +from datetime import datetime +from solstone.think.providers.brain_state import begin_brain_prerequisite_renewal +now = datetime.fromisoformat({(NOW + timedelta(seconds=2)).isoformat()!r}) +result = begin_brain_prerequisite_renewal(now, journal_path={str(tmp_path)!r}) +print(result["status"], flush=True) +if result["status"] == "started": + result["permit"].release() +""" + holder = subprocess.Popen( + [sys.executable, "-c", holder_code], + cwd=Path.cwd(), + env=_env(tmp_path), + stdout=subprocess.PIPE, + stderr=subprocess.PIPE, + text=True, + ) + try: + deadline = time.monotonic() + 10 + while not ready.exists() and time.monotonic() < deadline: + time.sleep(0.05) + assert ready.exists() + + busy = subprocess.run( + [sys.executable, "-c", contender_code], + cwd=Path.cwd(), + env=_env(tmp_path), + capture_output=True, + text=True, + check=True, + ) + assert busy.stdout.strip() == "busy" + + holder.terminate() + stdout, stderr = holder.communicate(timeout=10) + assert holder.returncode is not None, (stdout, stderr) + + free = subprocess.run( + [sys.executable, "-c", contender_code], + cwd=Path.cwd(), + env=_env(tmp_path), + capture_output=True, + text=True, + check=True, + ) + assert free.stdout.strip() == "started" + finally: + if holder.poll() is None: + holder.terminate() + holder.wait(timeout=5) diff --git a/tests/test_cortex.py b/tests/test_cortex.py index bfd3659ca..4862f63a6 100644 --- a/tests/test_cortex.py +++ b/tests/test_cortex.py @@ -10,6 +10,7 @@ import subprocess import sys import threading import time +from datetime import datetime, timedelta, timezone from pathlib import Path from unittest.mock import MagicMock, patch @@ -68,6 +69,125 @@ def _completed_path(journal_path: Path, use_id: str, name: str = "chat") -> Path return journal_path / "talents" / name.replace(":", "--") / f"{use_id}.jsonl" +class _FakeClock: + def __init__(self, now: datetime): + self.now = now + self.waits: list[float] = [] + + def __call__(self) -> datetime: + return self.now + + def advance(self, seconds: float) -> None: + self.now += timedelta(seconds=seconds) + + def wait(self, seconds: float) -> bool: + self.waits.append(seconds) + self.advance(seconds) + return False + + +class _FakeCallosum: + def __init__(self) -> None: + self.emitted: list[tuple[tuple, dict]] = [] + + def emit(self, *args, **kwargs) -> None: + self.emitted.append((args, kwargs)) + + +def _spp_inspector(state: dict): + def inspect(now: datetime, *, journal_path=None): + lane = state.get("lane", "spp") + aggregate = state.get("aggregate", "ready") + fingerprint = state.get("fingerprint", "a" * 64) + record_fingerprint = state.get("record_fingerprint", fingerprint) + observed = state.get("observed", now) + expires = state.get("expires", now + timedelta(minutes=10)) + record = None + if lane == "spp" and state.get("record_present", True): + component = None + if state.get("component_present", True): + component = { + "status": state.get("component_status", "ok"), + "observed_at": observed.isoformat(), + "expires_at": expires.isoformat(), + } + reason = state.get("component_reason") + if reason is not None: + component["reason_code"] = reason + record = { + "active_lane": "spp", + "fingerprint_sha256": record_fingerprint, + "evidence": {"lane_prerequisites": component}, + } + return { + "status": "ok", + "path": "/redacted", + "record": record, + "projection": { + "aggregate_state": aggregate, + "active_lane": lane, + "fingerprint_sha256": fingerprint, + "runtime_transition_in_progress": False, + }, + "reason_code": None, + "error": None, + } + + return inspect + + +def _make_spp_controller( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, + *, + state: dict, + clock: _FakeClock | None = None, + logger: MagicMock | None = None, + callosum: _FakeCallosum | None = None, +): + from solstone.think import cortex + + clock = clock or _FakeClock(datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc)) + callosum = callosum or _FakeCallosum() + logger = logger or MagicMock() + monkeypatch.setattr(cortex, "inspect_brain_state", _spp_inspector(state)) + monkeypatch.setattr( + cortex, + "read_active_brain_fingerprint_sha256", + lambda *, journal_path=None: state.get("fingerprint", "a" * 64), + ) + controller = cortex.SppRenewalController( + callosum=callosum, + stop_event=threading.Event(), + logger=logger, + clock=clock, + wait=clock.wait, + journal_path=tmp_path, + ) + return controller, clock, callosum, logger + + +def _commands(callosum: _FakeCallosum) -> list[list[str]]: + return [kwargs["cmd"] for _args, kwargs in callosum.emitted] + + +def _assert_fenced_refresh_command(command: list[str], fingerprint: str) -> None: + assert command[:4] == ["journal", "brain", "refresh", "--json"] + assert "--expected-active-fingerprint" in command + expected_index = command.index("--expected-fingerprint") + assert command[expected_index + 1] == fingerprint + + +def _assert_absence_fenced_refresh_command(command: list[str]) -> None: + assert command == [ + "journal", + "brain", + "refresh", + "--json", + "--expect-active-fingerprint-absent", + ] + + @pytest.fixture def mock_journal(tmp_path, monkeypatch): """Set up a temporary journal directory.""" @@ -175,6 +295,884 @@ def test_start_starts_spawn_worker_and_stays_resident(cortex_service, monkeypatc service_thread.join(timeout=1) +def test_spp_renewal_controller_replaces_prerequisites_for_seventy_virtual_minutes( + tmp_path, + monkeypatch, +): + from solstone.think import cortex + + start = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + clock = _FakeClock(start) + state = { + "fingerprint": "a" * 64, + "observed": start, + "expires": start + timedelta(minutes=10), + } + monkeypatch.setattr(cortex, "inspect_brain_state", _spp_inspector(state)) + monkeypatch.setattr( + cortex, + "read_active_brain_fingerprint_sha256", + lambda *, journal_path=None: state["fingerprint"], + ) + callosum = _FakeCallosum() + controller = cortex.SppRenewalController( + callosum=callosum, + stop_event=threading.Event(), + logger=MagicMock(), + clock=clock, + wait=clock.wait, + journal_path=tmp_path, + ) + + committed: list[tuple[datetime, datetime]] = [] + while clock.now < start + timedelta(minutes=72): + delay = controller.step() + if delay > 0: + clock.advance(delay) + assert state["expires"] > clock.now + controller.step() + ref = callosum.emitted[-1][1]["ref"] + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "started", "ref": ref} + ) + state["observed"] = clock.now + state["expires"] = clock.now + timedelta(minutes=10) + committed.append((state["observed"], state["expires"])) + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "stopped", "ref": ref, "exit_code": 0} + ) + + commands = _commands(callosum) + assert commands + assert len(committed) >= 9 + assert clock.now >= start + timedelta(minutes=72) + assert start + timedelta(minutes=72) > start + timedelta(minutes=60) + assert any(observed >= start + timedelta(minutes=30) for observed, _ in committed) + assert any(observed >= start + timedelta(minutes=60) for observed, _ in committed) + assert all( + later[0] > earlier[0] and later[1] > earlier[1] + for earlier, later in zip(committed, committed[1:]) + ) + assert all( + command[:3] == ["journal", "brain", "renew-prerequisites"] + for command in commands + ) + assert not any( + command[:3] == ["journal", "brain", "refresh"] for command in commands + ) + assert not any( + "generate" in command or "cogitate" in command for command in commands + ) + + +def test_spp_renewal_controller_tracks_skipped_active_ref_successor( + tmp_path, + monkeypatch, +): + from solstone.think import cortex + + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + clock = _FakeClock(now) + state = { + "fingerprint": "b" * 64, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + } + monkeypatch.setattr(cortex, "inspect_brain_state", _spp_inspector(state)) + monkeypatch.setattr( + cortex, + "read_active_brain_fingerprint_sha256", + lambda *, journal_path=None: state["fingerprint"], + ) + callosum = _FakeCallosum() + controller = cortex.SppRenewalController( + callosum=callosum, + stop_event=threading.Event(), + logger=MagicMock(), + clock=clock, + wait=clock.wait, + journal_path=tmp_path, + ) + + controller.step() + first_ref = callosum.emitted[-1][1]["ref"] + controller.handle_supervisor_message( + { + "tract": "supervisor", + "event": "skipped", + "ref": first_ref, + "active_ref": "active-brain", + "reason": "still_running", + } + ) + controller.step() + assert len(callosum.emitted) == 1 + + controller.handle_supervisor_message( + { + "tract": "supervisor", + "event": "stopped", + "ref": "active-brain", + "exit_code": 0, + } + ) + controller.step() + assert len(callosum.emitted) == 2 + assert callosum.emitted[-1][1]["ref"] != first_ref + + +def test_spp_renewal_controller_times_out_missed_successor_stop( + tmp_path, + monkeypatch, +): + from solstone.think import cortex + + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + state = { + "fingerprint": "b" * 64, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + } + controller, clock, callosum, _logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=state, + clock=_FakeClock(now), + ) + + controller.step() + first_ref = callosum.emitted[-1][1]["ref"] + controller.handle_supervisor_message( + { + "tract": "supervisor", + "event": "skipped", + "ref": first_ref, + "active_ref": "active-brain", + "reason": "still_running", + } + ) + clock.advance(cortex.SPP_REFRESH_OBSERVATION_BOUND_S + 1) + + controller.step() + + assert controller._successor_after_ref is None + assert len(callosum.emitted) == 2 + assert callosum.emitted[-1][1]["ref"] != first_ref + + +def test_spp_renewal_controller_accepts_full_refresh_exit_zero_without_retry( + tmp_path, + monkeypatch, +): + from solstone.think import cortex + + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + state = { + "aggregate": "unknown", + "fingerprint": "e" * 64, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + } + controller, clock, callosum, _logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=state, + clock=_FakeClock(now), + ) + + controller.step() + _assert_fenced_refresh_command(callosum.emitted[-1][1]["cmd"], state["fingerprint"]) + ref = callosum.emitted[-1][1]["ref"] + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "started", "ref": ref} + ) + assert controller._running_deadline is not None + assert ( + controller._running_deadline - clock.now + ).total_seconds() == cortex.SPP_REFRESH_OBSERVATION_BOUND_S + + clock.advance(1) + state["aggregate"] = "ready" + state["observed"] = clock.now + state["expires"] = clock.now + timedelta(minutes=10) + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "stopped", "ref": ref, "exit_code": 0} + ) + + assert controller._retry_after is None + assert controller._running_ref is None + delay = controller.step() + assert len(callosum.emitted) == 1 + assert delay > 0 + + +@pytest.mark.parametrize("case", ("unchanged", "unhealthy", "wrong_fingerprint")) +def test_spp_renewal_controller_rejects_exit_zero_without_new_ready_refresh_proof( + tmp_path, + monkeypatch, + case, +): + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + original_fingerprint = "e" * 64 + state = { + "aggregate": "unknown", + "fingerprint": original_fingerprint, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + } + controller, clock, callosum, logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=state, + clock=_FakeClock(now), + ) + + controller.step() + _assert_fenced_refresh_command(callosum.emitted[-1][1]["cmd"], original_fingerprint) + ref = callosum.emitted[-1][1]["ref"] + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "started", "ref": ref} + ) + + clock.advance(1) + if case == "unchanged": + state["aggregate"] = "ready" + elif case == "unhealthy": + state["aggregate"] = "unhealthy" + state["component_status"] = "failed" + state["component_reason"] = "attestation_rejected" + state["observed"] = clock.now + state["expires"] = clock.now + timedelta(minutes=10) + else: + state["aggregate"] = "ready" + state["fingerprint"] = "f" * 64 + state["record_fingerprint"] = "f" * 64 + state["observed"] = clock.now + state["expires"] = clock.now + timedelta(minutes=10) + + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "stopped", "ref": ref, "exit_code": 0} + ) + + assert controller._retry_after is not None + logged = "\n".join( + " ".join(str(part) for part in call.args) for call in logger.info.call_args_list + ) + assert "verified" not in logged + assert "failed" in logged + + +def test_spp_renewal_controller_running_timeout_clears_and_retries( + tmp_path, + monkeypatch, +): + from solstone.think import cortex + + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + state = { + "fingerprint": "f" * 64, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + } + controller, clock, callosum, _logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=state, + clock=_FakeClock(now), + ) + + controller.step() + ref = callosum.emitted[-1][1]["ref"] + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "started", "ref": ref} + ) + assert controller._running_deadline is not None + assert ( + controller._running_deadline - clock.now + ).total_seconds() == cortex.SPP_RENEWAL_ATTEMPT_BOUND_S + clock.advance(cortex.SPP_RENEWAL_ATTEMPT_BOUND_S + 1) + delay = controller.step() + + assert controller._running_ref is None + assert controller._retry_after is not None + assert delay == 5.0 + + clock.advance(delay) + controller.step() + assert len(callosum.emitted) == 2 + assert callosum.emitted[-1][1]["ref"] != ref + + +@pytest.mark.parametrize( + "state_update", + ( + {"record_present": False}, + {"component_status": "failed", "component_reason": "attestation_not_verified"}, + {"record_fingerprint": "0" * 64}, + {"expires": datetime(2026, 7, 24, 11, 59, tzinfo=timezone.utc)}, + ), +) +def test_spp_renewal_controller_unsafe_spp_records_fall_back_to_full_refresh( + tmp_path, + monkeypatch, + state_update, +): + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + state = { + "fingerprint": "9" * 64, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + **state_update, + } + _controller, _clock, callosum, _logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=state, + clock=_FakeClock(now), + ) + + _controller.step() + + _assert_fenced_refresh_command(callosum.emitted[-1][1]["cmd"], state["fingerprint"]) + + +def test_spp_renewal_controller_bootstraps_with_absence_fenced_refresh( + tmp_path, + monkeypatch, +): + from solstone.think import cortex + + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + state = { + "aggregate": "unknown", + "fingerprint": None, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + } + controller, clock, callosum, logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=state, + clock=_FakeClock(now), + ) + + assert controller.step() == cortex.SPP_RENEWAL_ACK_TIMEOUT_S + _assert_absence_fenced_refresh_command(callosum.emitted[-1][1]["cmd"]) + ref = callosum.emitted[-1][1]["ref"] + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "started", "ref": ref} + ) + clock.advance(1) + state["aggregate"] = "ready" + state["fingerprint"] = "a" * 64 + state["record_fingerprint"] = "a" * 64 + state["observed"] = clock.now + state["expires"] = clock.now + timedelta(minutes=10) + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "stopped", "ref": ref, "exit_code": 0} + ) + + assert controller._retry_after is None + logged = "\n".join(str(call.args) for call in logger.info.call_args_list) + assert "verified" in logged + + +def test_spp_renewal_controller_absence_fenced_refresh_retries_when_key_appears( + tmp_path, + monkeypatch, +): + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + state = { + "aggregate": "unknown", + "fingerprint": None, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + } + controller, clock, callosum, _logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=state, + clock=_FakeClock(now), + ) + + controller.step() + _assert_absence_fenced_refresh_command(callosum.emitted[-1][1]["cmd"]) + ref = callosum.emitted[-1][1]["ref"] + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "started", "ref": ref} + ) + state["fingerprint"] = "b" * 64 + state["record_fingerprint"] = "b" * 64 + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "stopped", "ref": ref, "exit_code": 3} + ) + + assert controller._retry_after is not None + clock.advance((controller._retry_after - clock.now).total_seconds()) + controller.step() + _assert_fenced_refresh_command(callosum.emitted[-1][1]["cmd"], "b" * 64) + + +def test_spp_renewal_controller_restarts_from_ready_vs_stale_record( + tmp_path, + monkeypatch, +): + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + ready_state = { + "fingerprint": "1" * 64, + "observed": now, + "expires": now + timedelta(minutes=10), + } + ready, _clock, ready_callosum, _logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=ready_state, + clock=_FakeClock(now), + ) + + delay = ready.step() + assert ready_callosum.emitted == [] + assert delay > 0 + + stale_state = { + "fingerprint": "1" * 64, + "observed": now - timedelta(minutes=20), + "expires": now - timedelta(seconds=1), + } + stale, _clock, stale_callosum, _logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=stale_state, + clock=_FakeClock(now), + ) + + stale.step() + _assert_fenced_refresh_command( + stale_callosum.emitted[-1][1]["cmd"], + stale_state["fingerprint"], + ) + + +def test_spp_renewal_controller_pending_fingerprint_switch_is_refenced( + tmp_path, + monkeypatch, +): + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + first_fingerprint = "2" * 64 + second_fingerprint = "3" * 64 + state = { + "fingerprint": first_fingerprint, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + } + controller, clock, callosum, _logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=state, + clock=_FakeClock(now), + ) + + controller.step() + first_ref = callosum.emitted[-1][1]["ref"] + assert first_fingerprint in callosum.emitted[-1][1]["cmd"] + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "started", "ref": first_ref} + ) + state["fingerprint"] = second_fingerprint + state["record_fingerprint"] = second_fingerprint + state["observed"] = clock.now + state["expires"] = clock.now + timedelta(seconds=30) + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "stopped", "ref": first_ref, "exit_code": 3} + ) + assert controller._retry_after is not None + + clock.advance((controller._retry_after - clock.now).total_seconds()) + controller.step() + + assert len(callosum.emitted) == 2 + assert second_fingerprint in callosum.emitted[-1][1]["cmd"] + + +def test_spp_renewal_controller_refresh_fallback_is_refenced_after_switch( + tmp_path, + monkeypatch, +): + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + first_fingerprint = "2" * 64 + second_fingerprint = "3" * 64 + state = { + "aggregate": "unknown", + "fingerprint": first_fingerprint, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + } + controller, clock, callosum, _logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=state, + clock=_FakeClock(now), + ) + + controller.step() + first_cmd = callosum.emitted[-1][1]["cmd"] + _assert_fenced_refresh_command(first_cmd, first_fingerprint) + first_ref = callosum.emitted[-1][1]["ref"] + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "started", "ref": first_ref} + ) + state["fingerprint"] = second_fingerprint + state["record_fingerprint"] = second_fingerprint + state["observed"] = clock.now + state["expires"] = clock.now + timedelta(seconds=30) + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "stopped", "ref": first_ref, "exit_code": 3} + ) + assert controller._retry_after is not None + + clock.advance((controller._retry_after - clock.now).total_seconds()) + controller.step() + + assert len(callosum.emitted) == 2 + _assert_fenced_refresh_command(callosum.emitted[-1][1]["cmd"], second_fingerprint) + + +def test_spp_renewal_controller_clock_jumps_do_not_duplicate( + tmp_path, + monkeypatch, +): + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + state = { + "fingerprint": "4" * 64, + "observed": now, + "expires": now + timedelta(minutes=10), + } + controller, clock, callosum, _logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=state, + clock=_FakeClock(now), + ) + + delay = controller.step() + assert delay >= 0 + assert callosum.emitted == [] + + clock.advance(delay + 1) + controller.step() + assert callosum.emitted[-1][1]["cmd"][:3] == [ + "journal", + "brain", + "renew-prerequisites", + ] + + backward_state = { + "fingerprint": "5" * 64, + "observed": now, + "expires": now + timedelta(minutes=10), + } + backward, backward_clock, backward_callosum, _logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=backward_state, + clock=_FakeClock(now), + ) + backward_delay = backward.step() + backward_clock.advance(-600) + later_delay = backward.step() + + assert backward_delay >= 0 + assert later_delay >= 0 + assert backward_callosum.emitted == [] + + +def test_spp_renewal_controller_recovers_from_inspection_and_fingerprint_errors( + tmp_path, + monkeypatch, +): + from solstone.think import cortex + + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + state = { + "fingerprint": "6" * 64, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + } + clock = _FakeClock(now) + callosum = _FakeCallosum() + logger = MagicMock() + inspect_calls = 0 + + def flaky_inspect(current: datetime, *, journal_path=None): + nonlocal inspect_calls + inspect_calls += 1 + if inspect_calls == 1: + raise OSError("SECRET-SENTINEL path") + return _spp_inspector(state)(current, journal_path=journal_path) + + monkeypatch.setattr(cortex, "inspect_brain_state", flaky_inspect) + monkeypatch.setattr( + cortex, + "read_active_brain_fingerprint_sha256", + lambda *, journal_path=None: state["fingerprint"], + ) + controller = cortex.SppRenewalController( + callosum=callosum, + stop_event=threading.Event(), + logger=logger, + clock=clock, + wait=clock.wait, + journal_path=tmp_path, + ) + + assert controller.step() == 5.0 + clock.advance(5.0) + controller.step() + assert callosum.emitted[-1][1]["cmd"][:3] == [ + "journal", + "brain", + "renew-prerequisites", + ] + + fingerprint_calls = 0 + + def flaky_fingerprint(*, journal_path=None): + nonlocal fingerprint_calls + fingerprint_calls += 1 + if fingerprint_calls == 1: + raise OSError("SECRET-SENTINEL path") + return state["fingerprint"] + + callosum = _FakeCallosum() + logger = MagicMock() + monkeypatch.setattr(cortex, "inspect_brain_state", _spp_inspector(state)) + monkeypatch.setattr( + cortex, "read_active_brain_fingerprint_sha256", flaky_fingerprint + ) + controller = cortex.SppRenewalController( + callosum=callosum, + stop_event=threading.Event(), + logger=logger, + clock=clock, + wait=clock.wait, + journal_path=tmp_path, + ) + + assert controller.step() == 5.0 + clock.advance(5.0) + controller.step() + assert callosum.emitted[-1][1]["cmd"][:3] == [ + "journal", + "brain", + "renew-prerequisites", + ] + + logged = "\n".join( + " ".join(str(part) for part in call.args) for call in logger.info.call_args_list + ) + assert "SECRET-SENTINEL" not in logged + assert str(tmp_path) not in logged + + +def test_spp_renewal_controller_run_contains_unexpected_step_exception( + tmp_path, + monkeypatch, +): + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + state = { + "fingerprint": "8" * 64, + "observed": now, + "expires": now + timedelta(minutes=10), + } + controller, _clock, _callosum, logger = _make_spp_controller( + tmp_path, + monkeypatch, + state=state, + clock=_FakeClock(now), + ) + controller.step = MagicMock(side_effect=RuntimeError("SECRET-SENTINEL")) + + def stop_after_wait(_seconds: float) -> bool: + controller.stop_event.set() + return True + + controller.wait = stop_after_wait + controller.run() + + logged = "\n".join( + " ".join(str(part) for part in call.args) for call in logger.info.call_args_list + ) + assert "RuntimeError" in logged + assert "retrying" in logged + assert "SECRET-SENTINEL" not in logged + + +def test_spp_renewal_controller_retries_with_cap_and_clears_on_disable( + tmp_path, + monkeypatch, +): + from solstone.think import cortex + + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + clock = _FakeClock(now) + state = { + "fingerprint": "c" * 64, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + } + monkeypatch.setattr(cortex, "inspect_brain_state", _spp_inspector(state)) + monkeypatch.setattr( + cortex, + "read_active_brain_fingerprint_sha256", + lambda *, journal_path=None: state["fingerprint"], + ) + controller = cortex.SppRenewalController( + callosum=_FakeCallosum(), + stop_event=threading.Event(), + logger=MagicMock(), + clock=clock, + wait=clock.wait, + journal_path=tmp_path, + ) + + delays: list[float] = [] + for _ in range(6): + controller.step() + clock.advance(cortex.SPP_RENEWAL_ACK_TIMEOUT_S + 1) + controller.step() + assert controller._retry_after is not None + delays.append((controller._retry_after - clock.now).total_seconds()) + clock.advance(delays[-1]) + assert delays == [5.0, 10.0, 20.0, 40.0, 60.0, 60.0] + + controller.step() + assert controller._pending_ref is not None + state["lane"] = "byo-cloud" + controller.step() + assert controller._pending_ref is None + assert controller._retry_after is None + + +def test_spp_renewal_controller_logs_lifecycle_events_without_secrets( + tmp_path, + monkeypatch, +): + from solstone.think import cortex + + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + clock = _FakeClock(now) + state = { + "lane": "byo-cloud", + "fingerprint": "d" * 64, + "observed": now - timedelta(minutes=9), + "expires": now + timedelta(seconds=30), + "secret": "SECRET-SENTINEL", + } + monkeypatch.setattr(cortex, "inspect_brain_state", _spp_inspector(state)) + monkeypatch.setattr( + cortex, + "read_active_brain_fingerprint_sha256", + lambda *, journal_path=None: state["fingerprint"], + ) + logger = MagicMock() + callosum = _FakeCallosum() + controller = cortex.SppRenewalController( + callosum=callosum, + stop_event=threading.Event(), + logger=logger, + clock=clock, + wait=clock.wait, + journal_path=tmp_path, + ) + + assert controller.step() == 30.0 + assert callosum.emitted == [] + + state["lane"] = "spp" + state["expires"] = clock.now + timedelta(minutes=10) + scheduled_delay = controller.step() + assert scheduled_delay > 0 + clock.advance(scheduled_delay) + controller.step() + assert callosum.emitted[-1][1]["cmd"][:3] == [ + "journal", + "brain", + "renew-prerequisites", + ] + verified_ref = callosum.emitted[-1][1]["ref"] + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "started", "ref": verified_ref} + ) + clock.advance(1) + state["observed"] = clock.now + state["expires"] = clock.now + timedelta(minutes=10) + controller.handle_supervisor_message( + { + "tract": "supervisor", + "event": "stopped", + "ref": verified_ref, + "exit_code": 0, + } + ) + + state["expires"] = clock.now + timedelta(seconds=30) + controller.step() + failed_ref = callosum.emitted[-1][1]["ref"] + clock.advance(cortex.SPP_RENEWAL_ACK_TIMEOUT_S + 1) + controller.step() + clock.advance(5.0) + controller.step() + stale_ref = callosum.emitted[-1][1]["ref"] + assert stale_ref != failed_ref + controller.handle_supervisor_message( + {"tract": "supervisor", "event": "started", "ref": stale_ref} + ) + clock.advance(cortex.SPP_RENEWAL_ATTEMPT_BOUND_S + 1) + controller.step() + + logged = "\n".join( + " ".join(str(part) for part in call.args) for call in logger.info.call_args_list + ) + for event in ( + "disabled", + "scheduled", + "in_flight", + "verified", + "failed", + "stale", + "retrying", + ): + assert event in logged + assert "SECRET-SENTINEL" not in logged + assert str(tmp_path) not in logged + assert "SECRET-SENTINEL" not in json.dumps(callosum.emitted) + + +def test_cortex_stop_wakes_and_joins_spp_renewal_worker( + mock_journal, + monkeypatch, +): + from solstone.think import cortex + from solstone.think.cortex import CortexService + + now = datetime(2026, 7, 24, 12, 0, tzinfo=timezone.utc) + state = {"lane": "byo-cloud", "fingerprint": "7" * 64} + monkeypatch.setattr(cortex, "inspect_brain_state", _spp_inspector(state)) + service = CortexService(str(mock_journal), clock=lambda: now) + service.callosum = MagicMock() + + service._start_spp_renewal_controller() + assert service._spp_renewal_worker is not None + deadline = time.monotonic() + 1.0 + while not service._spp_renewal_worker.is_alive() and time.monotonic() < deadline: + time.sleep(0.01) + assert service._spp_renewal_worker.is_alive() + + service.stop() + + assert not service._spp_renewal_worker.is_alive() + + def test_handle_request_dedups_existing_active_file( cortex_service, mock_journal, monkeypatch ): diff --git a/tests/test_no_implicit_cloud.py b/tests/test_no_implicit_cloud.py index 83e18fc6a..d631810c6 100644 --- a/tests/test_no_implicit_cloud.py +++ b/tests/test_no_implicit_cloud.py @@ -823,8 +823,11 @@ def test_confidential_stt_attestation_failure_blocks_remote_audio_egress( establish.assert_called_once() -def test_confidential_stt_stale_session_defers_before_egress(tmp_path, monkeypatch): +def test_confidential_stt_stale_session_replacement_failure_defers_before_egress( + tmp_path, monkeypatch +): from solstone.think.services import spp, spp_transport + from solstone.think.services.spp_attest.ratls.channel import RatlsChannelError _empty_journal(tmp_path, monkeypatch) config = _confidential_config(provider_pins=False) @@ -835,6 +838,8 @@ def test_confidential_stt_stale_session_defers_before_egress(tmp_path, monkeypat spp_transport._FORWARDER_BASE_URL = "http://127.0.0.1:4567" spp.record_attestation_verified(_stale_session(object())) mocks = _install_stt_backend_mocks(monkeypatch, confidential_result=None) + establish = Mock(side_effect=RatlsChannelError("gateway_unreachable")) + monkeypatch.setattr(spp_transport, "establish_attested_channel", establish) httpx_post = Mock(side_effect=AssertionError("audio egress attempted")) monkeypatch.setattr("httpx.post", httpx_post) @@ -843,9 +848,10 @@ def test_confidential_stt_stale_session_defers_before_egress(tmp_path, monkeypat with pytest.raises(ConfidentialTranscribeDeferral) as exc_info: transcribe("confidential", _stt_audio(), 16000, {}) - assert exc_info.value.reason_code == "attestation_stale" + assert exc_info.value.reason_code == "attestation_unreachable" mocks["parakeet"].assert_not_called() httpx_post.assert_not_called() + establish.assert_called_once() def test_confidential_stt_setting_off_gate_blocks_confidential_only( diff --git a/tests/test_supervisor.py b/tests/test_supervisor.py index f072f80fa..e3ce97ea3 100644 --- a/tests/test_supervisor.py +++ b/tests/test_supervisor.py @@ -588,6 +588,7 @@ def test_get_command_name(): assert get(["journal", "indexer", "--rescan"]) == "indexer" assert get(["sol", "insight", "20240101"]) == "insight" assert get(["journal", "think", "--day", "20240101"]) == "daily" + assert get(["journal", "brain", "renew-prerequisites"]) == "brain" assert get(["journal", "maintenance", "list"]) == "maintenance" assert get(["journal", "maintenance", "run", "foo:bar"]) == "maintenance:foo:bar" assert get(["journal", "maintenance", "run", "baz:qux"]) == "maintenance:baz:qux"