#!/usr/bin/env python3 """Claude UserPromptSubmit hook for operator-triggered lease promotion.""" from __future__ import annotations import fcntl import importlib.util import json import os import secrets import stat import subprocess import sys import time from collections.abc import Callable, Mapping from pathlib import Path from typing import Final, TextIO _MODULE_DIRECTORY = str(Path(__file__).resolve().parent) if _MODULE_DIRECTORY not in sys.path: sys.path.insert(0, _MODULE_DIRECTORY) from receipt_challenge import receipt_for # noqa: E402 _observer_spec = importlib.util.spec_from_file_location( "mosaic_receipt_observer_client", Path(__file__).resolve().with_name("receipt-observer-client.py") ) if _observer_spec is None or _observer_spec.loader is None: raise RuntimeError("unable to load receipt observer client") _observer_module = importlib.util.module_from_spec(_observer_spec) _observer_spec.loader.exec_module(_observer_module) observer_request = _observer_module.observer_request MAX_FRAME: Final = 64 * 1024 PENDING_MAX_AGE_SECONDS: Final = 60 * 60 PROMOTER_TIMEOUT_SECONDS: Final = 10.0 PROMOTION_PROMPT: Final = "/mosaic-promote" PROMOTER: Final = Path(__file__).resolve().with_name("lease_promote.py") PENDING_DIRECTORY: Final = "mosaic-lease" AUTHORIZATION_DIRECTORY: Final = "authorizations" AUTHORIZATION_TTL_SECONDS: Final = 60 LEASE_TTL_SECONDS: Final = 60 * 60 LOCK_FILE: Final = "promotion.lock" RESULT_FILE: Final = "last-result.json" EXPECTED_BEGIN_KEYS: Final = frozenset( {"ok", "state", "receipt_challenge", "receipt", "binding"} ) EXPECTED_BINDING_KEYS: Final = frozenset( { "compaction_epoch", "request_epoch", "h_source", "h_payload", "runtime_generation", "schema_version", } ) class PromotionAlreadyInProgress(RuntimeError): pass def reject_duplicate_json_keys(pairs: list[tuple[str, object]]) -> dict[str, object]: value: dict[str, object] = {} for key, item in pairs: if key in value: raise ValueError("duplicate promoter JSON key") value[key] = item return value def read_hook_input(stream: object) -> dict[str, object]: raw = getattr(stream, "buffer", stream).read(MAX_FRAME + 1) if not isinstance(raw, bytes) or len(raw) > MAX_FRAME: raise ValueError("invalid UserPromptSubmit input") value = json.loads(raw, object_pairs_hook=reject_duplicate_json_keys) if not isinstance(value, dict): raise ValueError("invalid UserPromptSubmit input") return value def emit_context(stream: TextIO, message: str) -> None: json.dump( { "hookSpecificOutput": { "hookEventName": "UserPromptSubmit", "additionalContext": message, } }, stream, separators=(",", ":"), ) stream.write("\n") def session_pending_name(environ: Mapping[str, str]) -> tuple[Path, str]: runtime_dir = Path(environ["XDG_RUNTIME_DIR"]) session_id = environ["MOSAIC_LEASE_SESSION_ID"] if not runtime_dir.is_absolute(): raise ValueError("XDG_RUNTIME_DIR must be absolute") if len(session_id) != 64 or any(character not in "0123456789abcdef" for character in session_id): raise ValueError("invalid lease session id") return runtime_dir, f"pending-{session_id}" def open_pending_directory(runtime_dir: Path) -> int: directory_flags = ( os.O_RDONLY | getattr(os, "O_CLOEXEC", 0) | getattr(os, "O_DIRECTORY", 0) | getattr(os, "O_NOFOLLOW", 0) ) runtime_descriptor = os.open(runtime_dir, directory_flags) try: runtime_metadata = os.fstat(runtime_descriptor) if ( not stat.S_ISDIR(runtime_metadata.st_mode) or runtime_metadata.st_uid != os.getuid() or stat.S_IMODE(runtime_metadata.st_mode) != 0o700 ): raise ValueError("unsafe XDG runtime directory") try: os.mkdir(PENDING_DIRECTORY, mode=0o700, dir_fd=runtime_descriptor) except FileExistsError: pass descriptor = os.open(PENDING_DIRECTORY, directory_flags, dir_fd=runtime_descriptor) finally: os.close(runtime_descriptor) metadata = os.fstat(descriptor) if ( not stat.S_ISDIR(metadata.st_mode) or metadata.st_uid != os.getuid() or stat.S_IMODE(metadata.st_mode) != 0o700 ): os.close(descriptor) raise ValueError("unsafe promotion pending directory") return descriptor def acquire_lock(directory_descriptor: int) -> int: flags = ( os.O_RDWR | os.O_CREAT | getattr(os, "O_CLOEXEC", 0) | getattr(os, "O_NOFOLLOW", 0) ) descriptor = os.open(LOCK_FILE, flags, 0o600, dir_fd=directory_descriptor) metadata = os.fstat(descriptor) if ( not stat.S_ISREG(metadata.st_mode) or metadata.st_uid != os.getuid() or stat.S_IMODE(metadata.st_mode) != 0o600 ): os.close(descriptor) raise ValueError("unsafe promotion lock file") try: fcntl.flock(descriptor, fcntl.LOCK_EX | fcntl.LOCK_NB) except BlockingIOError as error: os.close(descriptor) raise PromotionAlreadyInProgress() from error return descriptor def sweep_stale_pending(directory_descriptor: int, current_time: float) -> None: cutoff = current_time - PENDING_MAX_AGE_SECONDS removed = False with os.scandir(directory_descriptor) as entries: for candidate in entries: if not ( candidate.name.startswith("pending-") or candidate.name.startswith(".pending-") ): continue try: metadata = candidate.stat(follow_symlinks=False) if metadata.st_mtime < cutoff and not stat.S_ISDIR(metadata.st_mode): os.unlink(candidate.name, dir_fd=directory_descriptor) removed = True except FileNotFoundError: continue if removed: os.fsync(directory_descriptor) def consume_authorization(directory_descriptor: int, session_id: str, wall_clock: float) -> str | None: flags = os.O_RDONLY | getattr(os, "O_CLOEXEC", 0) | getattr(os, "O_DIRECTORY", 0) | getattr(os, "O_NOFOLLOW", 0) try: authorization_descriptor = os.open(AUTHORIZATION_DIRECTORY, flags, dir_fd=directory_descriptor) except FileNotFoundError: return None try: metadata = os.fstat(authorization_descriptor) if not stat.S_ISDIR(metadata.st_mode) or metadata.st_uid != os.getuid() or stat.S_IMODE(metadata.st_mode) != 0o700: raise ValueError("unsafe promotion authorization directory") name = f"{session_id}.auth" try: descriptor = os.open(name, os.O_RDONLY | getattr(os, "O_CLOEXEC", 0) | getattr(os, "O_NOFOLLOW", 0), dir_fd=authorization_descriptor) except FileNotFoundError: return None try: token_metadata = os.fstat(descriptor) if not stat.S_ISREG(token_metadata.st_mode) or token_metadata.st_uid != os.getuid() or stat.S_IMODE(token_metadata.st_mode) != 0o600 or token_metadata.st_size <= 0 or token_metadata.st_size > MAX_FRAME: raise ValueError("unsafe promotion authorization") raw = os.read(descriptor, MAX_FRAME + 1) finally: os.close(descriptor) os.unlink(name, dir_fd=authorization_descriptor) os.fsync(authorization_descriptor) token = json.loads(raw, object_pairs_hook=reject_duplicate_json_keys) if not isinstance(token, dict) or set(token) != {"nonce", "seat", "session_id", "expires_at", "ts"}: return None nonce = token.get("nonce") expires_at = token.get("expires_at") issued_at = token.get("ts") if token.get("session_id") != session_id or not isinstance(token.get("seat"), str) or not isinstance(nonce, str) or len(nonce) != 64 or any(char not in "0123456789abcdef" for char in nonce) or type(expires_at) not in (int, float) or type(issued_at) not in (int, float) or expires_at <= wall_clock or expires_at > issued_at + AUTHORIZATION_TTL_SECONDS: return None return nonce finally: os.close(authorization_descriptor) def write_result(directory_descriptor: int, attempt_id: str, verified: bool, reason: str | None, session_id: str, wall_clock: float) -> None: temporary = f".{RESULT_FILE}.tmp-{secrets.token_hex(8)}" descriptor = os.open(temporary, os.O_WRONLY | os.O_CREAT | os.O_EXCL | getattr(os, "O_CLOEXEC", 0) | getattr(os, "O_NOFOLLOW", 0), 0o600, dir_fd=directory_descriptor) try: os.fchmod(descriptor, 0o600) with os.fdopen(descriptor, "w", encoding="utf-8", closefd=False) as stream: json.dump({"attempt_id": attempt_id, "expires_at_wallclock": wall_clock + LEASE_TTL_SECONDS if verified else None, "reason": reason, "session_id": session_id, "ts": wall_clock, "verified": verified}, stream, separators=(",", ":"), sort_keys=True) stream.flush(); os.fsync(stream.fileno()) os.replace(temporary, RESULT_FILE, src_dir_fd=directory_descriptor, dst_dir_fd=directory_descriptor) os.fsync(directory_descriptor) finally: os.close(descriptor) def write_pending(directory_descriptor: int, name: str, challenge: str) -> None: temporary = f".{name}.tmp-{secrets.token_hex(8)}" flags = ( os.O_WRONLY | os.O_CREAT | os.O_EXCL | getattr(os, "O_CLOEXEC", 0) | getattr(os, "O_NOFOLLOW", 0) ) descriptor = os.open(temporary, flags, 0o600, dir_fd=directory_descriptor) try: os.fchmod(descriptor, 0o600) with os.fdopen(descriptor, "w", encoding="utf-8", closefd=False) as stream: stream.write(challenge) stream.flush() os.fsync(stream.fileno()) os.replace( temporary, name, src_dir_fd=directory_descriptor, dst_dir_fd=directory_descriptor, ) os.fsync(directory_descriptor) except Exception: try: os.unlink(temporary, dir_fd=directory_descriptor) except FileNotFoundError: pass raise finally: os.close(descriptor) def parse_begin_reply( completed: subprocess.CompletedProcess[str], ) -> tuple[str, dict[str, object] | None]: if completed.returncode != 0: return f"PROMOTER_EXIT_{completed.returncode}", None try: value = json.loads( completed.stdout, object_pairs_hook=reject_duplicate_json_keys, ) except (json.JSONDecodeError, RecursionError, TypeError, ValueError): return "INVALID_PROMOTER_REPLY", None if not isinstance(value, dict): return "INVALID_PROMOTER_REPLY", None if value.get("ok") is False and set(value) == {"ok", "code"}: code = value.get("code") return code if isinstance(code, str) and code else "PROMOTION_BEGIN_REFUSED", value if set(value) != EXPECTED_BEGIN_KEYS or value.get("ok") is not True: return "INVALID_PROMOTER_REPLY", None if value.get("state") != "PENDING_VERIFICATION": return "INVALID_PROMOTER_REPLY", None challenge = value.get("receipt_challenge") receipt = value.get("receipt") binding = value.get("binding") if ( not isinstance(challenge, str) or len(challenge) != 64 or any(character not in "0123456789abcdef" for character in challenge) or not isinstance(receipt, str) or not isinstance(binding, dict) or set(binding) != EXPECTED_BINDING_KEYS ): return "INVALID_PROMOTER_REPLY", None integer_fields = ( "compaction_epoch", "request_epoch", "runtime_generation", "schema_version", ) if any(type(binding.get(field)) is not int or binding[field] < 0 for field in integer_fields): return "INVALID_PROMOTER_REPLY", None if not all( isinstance(binding.get(field), str) and len(binding[field]) == 64 and all(character in "0123456789abcdef" for character in binding[field]) for field in ("h_source", "h_payload") ): return "INVALID_PROMOTER_REPLY", None if not secrets.compare_digest( receipt.encode("utf-8"), receipt_for(challenge, binding).encode("utf-8"), ): return "INVALID_PROMOTER_REPLY", None return "", value def main( *, environ: Mapping[str, str] | None = None, stdin: object | None = None, stdout: TextIO | None = None, stderr: TextIO | None = None, run: Callable[..., subprocess.CompletedProcess[str]] = subprocess.run, now: Callable[[], float] = time.time, ) -> int: source_environment = os.environ if environ is None else environ input_stream = sys.stdin if stdin is None else stdin output_stream = sys.stdout if stdout is None else stdout error_stream = sys.stderr if stderr is None else stderr try: hook_input = read_hook_input(input_stream) except (OSError, RecursionError, ValueError, json.JSONDecodeError) as error: print(f"Mosaic promotion trigger ignored invalid hook input: {error}", file=error_stream) return 0 if hook_input.get("prompt") != PROMOTION_PROMPT: return 0 directory_descriptor: int | None = None lock_descriptor: int | None = None try: runtime_dir, pending_name = session_pending_name(source_environment) session_id = source_environment["MOSAIC_LEASE_SESSION_ID"] directory_descriptor = open_pending_directory(runtime_dir) lock_descriptor = acquire_lock(directory_descriptor) wall_clock = now() nonce = consume_authorization(directory_descriptor, session_id, wall_clock) if nonce is None: write_result(directory_descriptor, "0" * 64, False, "NOT_AUTHORIZED", session_id, wall_clock) print("Mosaic promotion denied: NOT_AUTHORIZED.", file=error_stream) return 0 sweep_stale_pending(directory_descriptor, wall_clock) completed = run([sys.executable, "-I", "-S", "-B", str(PROMOTER), "--begin"], check=False, capture_output=True, text=True, env=dict(source_environment), timeout=PROMOTER_TIMEOUT_SECONDS) code, reply = parse_begin_reply(completed) if code or reply is None: write_result(directory_descriptor, nonce, False, code or "PROMOTION_BEGIN_FAILED", session_id, now()) return 0 challenge = str(reply["receipt_challenge"]) observation = observer_request( Path(source_environment["MOSAIC_RECEIPT_OBSERVER_SOCKET"]), {"action": "record_runtime_observation", "session_id": session_id, "runtime_generation": int(source_environment["MOSAIC_RUNTIME_GENERATION"]), "runtime": "claude", "latest_assistant_message": reply["receipt"]}, ) if set(observation) != {"ok"} or observation.get("ok") is not True: write_result(directory_descriptor, challenge, False, "OBSERVATION_REJECTED", session_id, now()) return 0 completion = run([sys.executable, "-I", "-S", "-B", str(PROMOTER), "--complete", challenge], check=False, capture_output=True, text=True, env=dict(source_environment), timeout=PROMOTER_TIMEOUT_SECONDS) try: outcome = json.loads(completion.stdout, object_pairs_hook=reject_duplicate_json_keys) except (json.JSONDecodeError, ValueError): outcome = None if completion.returncode == 0 and isinstance(outcome, dict) and outcome.get("stage") == "promote_lease" and outcome.get("ok") is True and outcome.get("state") == "VERIFIED": write_result(directory_descriptor, challenge, True, None, session_id, now()) else: reason = outcome.get("code") if isinstance(outcome, dict) and isinstance(outcome.get("code"), str) else "PROMOTION_INCOMPLETE" write_result(directory_descriptor, challenge, False, reason, session_id, now()) except PromotionAlreadyInProgress: print("Mosaic promotion denied: PROMOTION_ALREADY_IN_PROGRESS.", file=error_stream) except (KeyError, OSError, RecursionError, ValueError, subprocess.SubprocessError) as error: print(f"Mosaic promotion begin failed: {type(error).__name__}: {error}", file=error_stream) finally: if lock_descriptor is not None: os.close(lock_descriptor) if directory_descriptor is not None: os.close(directory_descriptor) return 0 if __name__ == "__main__": raise SystemExit(main())