Files
stack/packages/mosaic/framework/tools/lease-broker/promote-begin.py
T

400 lines
16 KiB
Python

#!/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())