Files
codex-mmo/libexec/mmo_snapshot.py
T

526 lines
21 KiB
Python
Raw Normal View History

2026-08-24 08:11:59 -07:00
#!/usr/bin/env python3
"""Compile validated composition profiles into immutable logical snapshots."""
from __future__ import annotations
import contextlib
import json
import os
import shutil
import stat
import tempfile
from collections.abc import Mapping
from pathlib import Path, PurePosixPath
from typing import Any
from mmo_guidance import compiled_guidance, coordination_capable_agents
from mmo_profiles import resolve_profile
from mmo_util import (
atomic_write_json,
atomic_write_text,
file_lock,
make_tree_read_only,
manifest_for_tree,
read_json,
sha256_bytes,
stable_hash,
state_root,
toml_dumps,
)
from mmo_version import MMO_SCHEMA_VERSION, PACKAGE_VERSION
SWITCHYARD_SCHEMA_VERSION = 1
_VERIFIED_SNAPSHOT_STATS: dict[str, tuple[tuple[str, int, int, int, int, int], ...]] = {}
def _snapshot_fingerprint(
resolved: Mapping[str, Any],
*,
guidance: Mapping[str, str] | None = None,
guidance_schema_version: int = MMO_SCHEMA_VERSION,
) -> dict[str, Any]:
"""Return every semantic input that determines compiled snapshot bytes."""
return {
"snapshot_schema_version": MMO_SCHEMA_VERSION,
"guidance_schema_version": guidance_schema_version,
"package_version": resolved["package_version"],
"resolved": resolved,
"guidance": dict(guidance) if guidance is not None else compiled_guidance(resolved),
}
def snapshots_root() -> Path:
root = state_root() / "snapshots"
root.mkdir(parents=True, exist_ok=True, mode=0o700)
return root
def snapshot_dir(snapshot_hash: str) -> Path:
if len(snapshot_hash) != 64 or any(char not in "0123456789abcdef" for char in snapshot_hash):
raise ValueError("invalid snapshot hash")
path = snapshots_root() / snapshot_hash
try:
mode = path.lstat().st_mode
except FileNotFoundError as exc:
raise FileNotFoundError(f"unknown profile snapshot: {snapshot_hash}") from exc
if not stat.S_ISDIR(mode):
raise RuntimeError(f"snapshot path is not a regular directory: {path}")
return path
def _remove_snapshot_tree(path: Path) -> None:
"""Remove a snapshot tree without following attacker-controlled links."""
try:
mode = path.lstat().st_mode
except FileNotFoundError:
return
if stat.S_ISLNK(mode):
path.unlink()
return
if not stat.S_ISDIR(mode):
raise RuntimeError(f"snapshot path is not a directory: {path}")
for item in path.rglob("*"):
with contextlib.suppress(OSError):
item_mode = item.lstat().st_mode
if not stat.S_ISLNK(item_mode):
os.chmod(item, 0o700 if stat.S_ISDIR(item_mode) else 0o600)
with contextlib.suppress(OSError):
os.chmod(path, 0o700)
shutil.rmtree(path)
def _route_id(snapshot_hash: str, model_key: str) -> str:
return f"mmo-{snapshot_hash[:12]}-{model_key}"
def _switchyard_routes(
resolved: Mapping[str, Any],
) -> tuple[dict[str, Any], dict[str, str], str | None]:
clients: dict[str, Any] = {}
targets: dict[str, Any] = {}
routes: dict[str, Any] = {}
route_ids: dict[str, str] = {}
for route_key, route in resolved["routes"].items():
if route["driver"] != "switchyard":
continue
client: dict[str, Any] = {
"format": route["wire_protocol"],
"base_url": route["base_url"],
"max_retries": int(route.get("max_retries", 1)),
}
credentials = list(route.get("credential_envs", []))
if credentials:
client["api_key_env"] = credentials[0]
if route.get("extra_headers"):
client["extra_headers"] = route["extra_headers"]
clients[route_key] = client
for model_key, model in resolved["models"].items():
route = resolved["routes"][model["route"]]
if route["driver"] != "switchyard":
continue
targets[model_key] = {
"id": model["upstream_id"],
"llm_client": model["route"],
}
extra_body = dict(model.get("extra_body") or {})
route_policy = model.get("route_policy") or route.get("openrouter_policy")
if route_policy is not None:
extra_body["provider"] = route_policy
if extra_body:
targets[model_key]["extra_body"] = extra_body
if not targets:
return (
{
"schema_version": SWITCHYARD_SCHEMA_VERSION,
"llm_clients": clients,
"targets": targets,
"routes": routes,
},
route_ids,
None,
)
route_semantics: dict[str, dict[str, Any]] = {}
for model_key, model in resolved["models"].items():
route = resolved["routes"][model["route"]]
if route["driver"] != "switchyard":
continue
route_semantics[model_key] = {
"type": "passthrough",
"target": model_key,
"context_window": int(model["context_window"]),
"tool_calling": bool(model.get("tool_calling", True)),
"reasoning": model.get("default_reasoning") != "none",
}
# Gateway identity is derived only from transport/model semantics, not from
# the surrounding profile. Profiles with an identical route set can share
# one Switchyard process while retaining independent immutable snapshots.
gateway_hash = stable_hash(
{
"schema_version": SWITCHYARD_SCHEMA_VERSION,
"llm_clients": clients,
"targets": targets,
"routes": route_semantics,
}
)
for model_key, semantics in route_semantics.items():
route_id = _route_id(gateway_hash, model_key)
route_ids[model_key] = route_id
routes[model_key] = {
"id": route_id,
**semantics,
}
return (
{
"schema_version": SWITCHYARD_SCHEMA_VERSION,
"llm_clients": clients,
"targets": targets,
"routes": routes,
},
route_ids,
gateway_hash,
)
def _credential_groups(resolved: Mapping[str, Any], *, drivers: set[str]) -> list[dict[str, Any]]:
groups: list[dict[str, Any]] = []
for route_key, route in sorted(resolved["routes"].items()):
if route["driver"] not in drivers:
continue
alternatives = list(route.get("credential_envs", []))
if not alternatives:
continue
groups.append(
{
"route": route_key,
"driver": route["driver"],
"target_env": alternatives[0],
"alternatives": alternatives,
}
)
return groups
def _json_text(value: Any) -> str:
return json.dumps(value, ensure_ascii=False, indent=2, sort_keys=True, allow_nan=False) + "\n"
def _snapshot_material(
resolved: Mapping[str, Any],
snapshot_hash: str,
*,
package: str,
) -> tuple[dict[str, Any], dict[str, str]]:
"""Derive every manifest field and payload byte from hashed profile data."""
routes, route_ids, gateway_hash = _switchyard_routes(resolved)
gateway_credential_groups = _credential_groups(resolved, drivers={"switchyard"})
direct_credential_groups = _credential_groups(resolved, drivers={"codex_custom"})
credentials = sorted({group["target_env"] for group in gateway_credential_groups})
direct_credentials = sorted({group["target_env"] for group in direct_credential_groups})
manifest = {
"schema_version": MMO_SCHEMA_VERSION,
"guidance_schema_version": MMO_SCHEMA_VERSION,
"snapshot_hash": snapshot_hash,
"logical_hash": resolved["logical_hash"],
"package_version": package,
"profile_id": resolved["profile"]["id"],
"profile_version": resolved["profile"]["version"],
"root_agent": resolved["profile"]["root"],
"orchestration": resolved["coordination"]["orchestration"],
"native_agents": resolved["capabilities"]["native_agents"],
"mcp_agents": resolved["capabilities"]["mcp_agents"],
"warnings": resolved.get("warnings", []),
"gateway_required": bool(routes["routes"]),
"gateway_hash": gateway_hash,
"route_ids": route_ids,
"credential_envs": credentials,
"credential_groups": gateway_credential_groups,
"direct_credential_envs": direct_credentials,
"direct_credential_groups": direct_credential_groups,
"agents": sorted(resolved["agents"]),
"models": sorted(resolved["models"]),
"routes": sorted(resolved["routes"]),
"coordination_capable_agents": coordination_capable_agents(resolved),
}
tool_mcp_servers = resolved.get("tool_mcp_servers", {})
manifest["tool_mcp_servers"] = sorted(tool_mcp_servers)
payload = {
"resolved-profile.json": _json_text(resolved),
"resolved-profile.toml": toml_dumps(
{key: value for key, value in resolved.items() if key != "logical_hash"}
),
}
if routes["routes"]:
payload["routes.toml"] = toml_dumps(routes)
for agent_key, agent in resolved["agents"].items():
payload[f"instructions/{agent_key}.md"] = agent.get("instructions_text", "").rstrip() + "\n"
contract = agent.get("output_contract_schema")
if contract is not None:
payload[f"contracts/{agent_key}.json"] = _json_text(contract)
if resolved.get("smoke") is not None:
payload["smoke.json"] = _json_text(resolved["smoke"])
payload.update(compiled_guidance(resolved))
return manifest, payload
def compile_profile(
profile: str | Path,
*,
bindings: Mapping[str, str] | None = None,
force: bool = False,
) -> dict[str, Any]:
resolved = resolve_profile(profile, bindings=bindings)
return compile_resolved_profile(resolved, force=force)
def compile_resolved_profile(resolved: Mapping[str, Any], *, force: bool = False) -> dict[str, Any]:
"""Compile already validated profile semantics into a content-addressed snapshot.
This is intentionally narrow: evaluation uses it to reduce a validated
profile's delegation graph for matched root-only and ablation controls.
Callers may only pass a complete current-schema document and must recompute its
logical hash after any restriction.
"""
if resolved.get("schema_version") != MMO_SCHEMA_VERSION or isinstance(
resolved.get("schema_version"), bool
):
raise ValueError(f"compile_resolved_profile requires resolved schema v{MMO_SCHEMA_VERSION}")
if not isinstance(resolved.get("logical_hash"), str):
raise ValueError("resolved profile logical_hash is required")
package = resolved.get("package_version")
if not isinstance(package, str) or not package:
raise ValueError("resolved profile package_version is required")
fingerprint = _snapshot_fingerprint(resolved)
snapshot_hash = stable_hash(fingerprint)
destination = snapshots_root() / snapshot_hash
lock = state_root() / ".snapshot-compiler.lock"
with file_lock(lock):
if destination.is_dir() and not force:
return load_snapshot(snapshot_hash)
if destination.exists() or destination.is_symlink():
_remove_snapshot_tree(destination)
temporary = Path(tempfile.mkdtemp(prefix=f".{snapshot_hash[:12]}-", dir=snapshots_root()))
published = False
try:
manifest, payload = _snapshot_material(resolved, snapshot_hash, package=package)
instructions_dir = temporary / "instructions"
contracts_dir = temporary / "contracts"
instructions_dir.mkdir(mode=0o755)
contracts_dir.mkdir(mode=0o755)
for relative, text in payload.items():
atomic_write_text(temporary / relative, text, 0o644)
manifest["payload_files"] = {
relative: sha256_bytes(text.encode("utf-8"))
for relative, text in sorted(payload.items())
}
atomic_write_json(temporary / "manifest.json", manifest, 0o644)
os.replace(temporary, destination)
published = True
make_tree_read_only(destination)
except Exception:
with contextlib.suppress(OSError):
_remove_snapshot_tree(temporary)
if published and (destination.exists() or destination.is_symlink()):
with contextlib.suppress(OSError):
_remove_snapshot_tree(destination)
raise
return load_snapshot(snapshot_hash)
def _snapshot_stat_fingerprint(
directory: Path, expected_files: Mapping[str, str]
) -> tuple[tuple[str, int, int, int, int, int], ...]:
rows: list[tuple[str, int, int, int, int, int]] = []
names = {".", *expected_files}
for relative in expected_files:
parent = Path(relative).parent
while parent != Path("."):
names.add(parent.as_posix())
parent = parent.parent
for relative in sorted(names):
path = directory if relative == "." else directory / relative
try:
status = path.lstat()
except OSError:
rows.append((relative, -1, -1, -1, -1, -1))
continue
rows.append(
(
relative,
int(status.st_size),
int(status.st_mtime_ns),
int(status.st_ctime_ns),
int(status.st_mode),
int(status.st_ino),
)
)
return tuple(rows)
def _stored_snapshot_guidance(
directory: Path,
manifest: Mapping[str, Any],
) -> dict[str, str]:
"""Load guidance bytes whose identity is already bound by the snapshot hash.
Guidance is generated documentation, so its prose may legitimately evolve
between MMO releases without changing an existing immutable snapshot. A
historical snapshot must therefore authenticate its stored guidance rather
than regenerate that prose with the currently installed compiler.
"""
payload_files = manifest.get("payload_files")
if not isinstance(payload_files, Mapping):
raise RuntimeError(f"snapshot payload manifest is missing in {directory}")
guidance: dict[str, str] = {}
for relative in sorted(payload_files):
if not isinstance(relative, str) or not relative.startswith("guidance/"):
continue
parsed = PurePosixPath(relative)
if (
parsed.is_absolute()
or parsed.as_posix() != relative
or any(part in {"", ".", ".."} for part in parsed.parts)
):
raise RuntimeError(f"snapshot guidance path is unsafe in {directory}: {relative!r}")
path = directory.joinpath(*parsed.parts)
try:
status = path.lstat()
except OSError as exc:
raise RuntimeError(f"snapshot guidance is unreadable: {path}") from exc
if not stat.S_ISREG(status.st_mode):
raise RuntimeError(f"snapshot guidance is not a regular file: {path}")
try:
guidance[relative] = path.read_text(encoding="utf-8")
except (OSError, UnicodeError) as exc:
raise RuntimeError(f"snapshot guidance is not valid UTF-8: {path}") from exc
if not guidance:
raise RuntimeError(f"snapshot guidance payload is missing in {directory}")
return guidance
def load_snapshot(snapshot_hash: str) -> dict[str, Any]:
directory = snapshot_dir(snapshot_hash)
for name in ("manifest.json", "resolved-profile.json"):
status = (directory / name).lstat()
if not stat.S_ISREG(status.st_mode):
raise RuntimeError(f"snapshot required file is not regular: {directory / name}")
manifest = read_json(directory / "manifest.json")
if not isinstance(manifest, Mapping):
raise RuntimeError(f"snapshot manifest is not an object in {directory}")
schema_version = manifest.get("schema_version")
if (
not isinstance(schema_version, int)
or isinstance(schema_version, bool)
or schema_version != MMO_SCHEMA_VERSION
):
raise RuntimeError(f"unsupported snapshot schema in {directory}")
if manifest.get("snapshot_hash") != snapshot_hash:
raise RuntimeError(f"snapshot manifest identity mismatch in {directory}")
resolved = read_json(directory / "resolved-profile.json")
if not isinstance(resolved, Mapping):
raise RuntimeError(f"snapshot resolved profile is not an object in {directory}")
resolved_schema = resolved.get("schema_version")
if (
not isinstance(resolved_schema, int)
or isinstance(resolved_schema, bool)
or resolved_schema != MMO_SCHEMA_VERSION
):
raise RuntimeError(f"unsupported resolved profile schema in {directory}")
guidance_schema_version = manifest.get("guidance_schema_version")
if not isinstance(guidance_schema_version, int) or isinstance(guidance_schema_version, bool):
raise RuntimeError(f"snapshot guidance schema is invalid in {directory}")
stored_guidance = _stored_snapshot_guidance(directory, manifest)
fingerprint = _snapshot_fingerprint(
resolved,
guidance=stored_guidance,
guidance_schema_version=guidance_schema_version,
)
expected_hash = stable_hash(fingerprint)
if expected_hash != snapshot_hash:
raise RuntimeError(
f"snapshot content-address mismatch in {directory}: expected {expected_hash}"
)
package_version = resolved.get("package_version")
profile = resolved.get("profile")
if not isinstance(package_version, str) or package_version != PACKAGE_VERSION:
raise RuntimeError(f"snapshot package version must be {PACKAGE_VERSION} in {directory}")
if not isinstance(profile, Mapping) or profile.get("version") != PACKAGE_VERSION:
raise RuntimeError(f"snapshot profile version must be {PACKAGE_VERSION} in {directory}")
expected_manifest, expected_payload = _snapshot_material(
resolved,
snapshot_hash,
package=package_version,
)
expected_manifest["guidance_schema_version"] = guidance_schema_version
expected_payload = {
relative: text
for relative, text in expected_payload.items()
if not relative.startswith("guidance/")
}
expected_payload.update(stored_guidance)
expected_files = {
relative: sha256_bytes(text.encode("utf-8"))
for relative, text in sorted(expected_payload.items())
}
complete_expected_manifest = {
**expected_manifest,
"payload_files": expected_files,
}
if manifest != complete_expected_manifest:
raise RuntimeError(f"snapshot manifest integrity failure in {directory}")
if not expected_files:
raise RuntimeError(f"snapshot payload manifest is missing in {directory}")
fingerprint_files = {"manifest.json": "", **expected_files}
stat_fingerprint = _snapshot_stat_fingerprint(directory, fingerprint_files)
if _VERIFIED_SNAPSHOT_STATS.get(snapshot_hash) != stat_fingerprint:
try:
actual_files = manifest_for_tree(directory, exclude=("manifest.json",))
except ValueError as exc:
raise RuntimeError(f"snapshot payload integrity failure in {directory}: {exc}") from exc
if actual_files != expected_files:
missing = sorted(set(expected_files) - set(actual_files))
unexpected = sorted(set(actual_files) - set(expected_files))
changed = sorted(
key
for key in set(actual_files) & set(expected_files)
if actual_files[key] != expected_files[key]
)
raise RuntimeError(
"snapshot payload integrity failure in "
f"{directory}: missing={missing}, unexpected={unexpected}, changed={changed}"
)
_VERIFIED_SNAPSHOT_STATS[snapshot_hash] = stat_fingerprint
return {
"directory": str(directory),
"manifest": manifest,
"resolved": resolved,
}
def find_snapshot_for_profile(profile_id: str) -> list[dict[str, Any]]:
results: list[dict[str, Any]] = []
for directory in snapshots_root().iterdir():
if not directory.is_dir() or len(directory.name) != 64:
continue
try:
snapshot = load_snapshot(directory.name)
manifest = snapshot["manifest"]
except (OSError, RuntimeError, ValueError, json.JSONDecodeError):
continue
if manifest.get("profile_id") == profile_id:
results.append(
{
"snapshot_hash": directory.name,
"profile_id": profile_id,
"profile_version": manifest.get("profile_version"),
"logical_hash": manifest.get("logical_hash"),
}
)
return sorted(results, key=lambda item: item["snapshot_hash"])