from __future__ import annotations import argparse from datetime import datetime, timezone from hashlib import sha256 import json import os from pathlib import Path import sys import tempfile from typing import Any from app.db.session import SessionLocal from app.services.data_contract_validation import ( PYTORCH_MODEL_CONTRACT_KEY, PYTORCH_MODEL_CONTRACT_VERSION, ) from app.services.runtime_model_provenance_service import RuntimeModelProvenanceService from app.services.source_registry_service import SourceRegistryService CLAIM_BOUNDARY = ( "Recovered legacy runtime artifact binding only. This receipt does not assert complete human review, " "the historical training commit/container, protected-test independence, national validity or model promotion." ) def _sha256_file(path: Path) -> str: digest = sha256() with path.open("rb") as handle: for chunk in iter(lambda: handle.read(1024 * 1024), b""): digest.update(chunk) return digest.hexdigest() def _required_file(value: str, field: str) -> Path: path = Path(value).expanduser().resolve() if not path.is_file(): raise ValueError(f"{field} does not point to a readable file: {path}") return path def _json_object(path: Path, field: str) -> dict[str, Any]: try: payload = json.loads(path.read_text(encoding="utf-8")) except (OSError, UnicodeDecodeError, json.JSONDecodeError) as exc: raise ValueError(f"{field} is not a readable UTF-8 JSON object: {path}") from exc if not isinstance(payload, dict): raise ValueError(f"{field} must contain a JSON object: {path}") return payload def _require_equal(observed: Any, expected: Any, field: str) -> None: if observed != expected: raise ValueError(f"{field} mismatch: expected {expected!r}, observed {observed!r}") def inspect_evidence(args: argparse.Namespace) -> dict[str, Any]: paths = { "model": _required_file(args.model_path, "--model-path"), "checkpoint": _required_file(args.checkpoint_path, "--checkpoint-path"), "base_model": _required_file(args.base_model_path, "--base-model-path"), "training_summary": _required_file(args.training_summary_path, "--training-summary-path"), "training_args": _required_file(args.training_args_path, "--training-args-path"), "training_results": _required_file(args.training_results_path, "--training-results-path"), "dataset_summary": _required_file(args.dataset_summary_path, "--dataset-summary-path"), "dataset_yaml": _required_file(args.dataset_yaml_path, "--dataset-yaml-path"), } checksums = {name: _sha256_file(path) for name, path in paths.items()} _require_equal(checksums["checkpoint"], checksums["model"], "checkpoint/model SHA-256") training_summary = _json_object(paths["training_summary"], "training summary") dataset_summary = _json_object(paths["dataset_summary"], "dataset summary") _require_equal(training_summary.get("status"), "ok", "training_summary.status") _require_equal(training_summary.get("trained_model_sha256"), checksums["model"], "trained_model_sha256") _require_equal(training_summary.get("base_model_sha256"), checksums["base_model"], "base_model_sha256") _require_equal( training_summary.get("dataset_summary_sha256"), checksums["dataset_summary"], "dataset_summary_sha256", ) _require_equal(training_summary.get("dataset_yaml_sha256"), checksums["dataset_yaml"], "dataset_yaml_sha256") _require_equal(dataset_summary.get("status"), "ok", "dataset_summary.status") class_names = dataset_summary.get("class_names") if not isinstance(class_names, list) or not class_names or not all( isinstance(value, str) and value.strip() for value in class_names ): raise ValueError("dataset_summary.class_names must be a non-empty string list") class_mapping = {str(index): value.strip() for index, value in enumerate(class_names)} evidence = { "paths": {name: str(path) for name, path in paths.items()}, "checksums": checksums, "class_mapping": class_mapping, "training_summary": training_summary, "dataset_summary": { "status": dataset_summary.get("status"), "tile_count": dataset_summary.get("tile_count"), "train_tile_count": dataset_summary.get("train_tile_count"), "val_tile_count": dataset_summary.get("val_tile_count"), "label_count": dataset_summary.get("label_count"), "class_names": class_names, }, } return evidence def _snapshot_key(*, model_id: str, model_sha256: str) -> str: return f"runtime-model-{model_id}-{model_sha256[:24]}" def _manifest_payload( *, args: argparse.Namespace, evidence: dict[str, Any], source_registry_id: str, source_snapshot_id: str, imported_at: str, ) -> dict[str, Any]: checksums = evidence["checksums"] paths = evidence["paths"] payload: dict[str, Any] = { "schema_version": RuntimeModelProvenanceService.MANIFEST_SCHEMA_VERSION, "data_contract": { "key": PYTORCH_MODEL_CONTRACT_KEY, "version": PYTORCH_MODEL_CONTRACT_VERSION, }, "model": { "model_id": args.model_id, "task_type": args.task_type, "sha256": checksums["model"], "model_format": "pytorch", "framework": "ultralytics/pytorch", "class_mapping": evidence["class_mapping"], "source_version": args.source_version, }, "source": { "source_registry_id": source_registry_id, "source_snapshot_id": source_snapshot_id, "source_registry_key": RuntimeModelProvenanceService.SOURCE_REGISTRY_KEY, "source_snapshot_checksum_sha256": checksums["model"], }, "lineage": { "upstream_asset_ids": [ f"base-model:{Path(paths['base_model']).name}", f"dataset-summary:{Path(paths['dataset_summary']).name}", f"dataset-yaml:{Path(paths['dataset_yaml']).name}", f"training-args:{Path(paths['training_args']).name}", f"training-results:{Path(paths['training_results']).name}", f"training-summary:{Path(paths['training_summary']).name}", ], "upstream_checksums_sha256": [ checksums["base_model"], checksums["dataset_summary"], checksums["dataset_yaml"], checksums["training_args"], checksums["training_results"], checksums["training_summary"], ], "transformations": [ { "name": "ultralytics-yolo-finetune", "version": args.framework_version, "checksum_sha256": checksums["training_args"], } ], }, "metadata": { "training_manifest_sha256": checksums["training_summary"], "claim_boundary": CLAIM_BOUNDARY, "evidence_paths": paths, "checkpoint_sha256": checksums["checkpoint"], }, "imported_at": imported_at, } payload["metadata"]["runtime_manifest_sha256"] = RuntimeModelProvenanceService.manifest_self_checksum( payload ) return payload def _write_manifest_atomically(path: Path, payload: dict[str, Any]) -> bool: encoded = (json.dumps(payload, indent=2, sort_keys=True) + "\n").encode("utf-8") if path.exists(): if path.read_bytes() == encoded: return False raise ValueError(f"Refusing to overwrite a different runtime provenance sidecar: {path}") path.parent.mkdir(parents=True, exist_ok=True) temporary: Path | None = None try: with tempfile.NamedTemporaryFile( mode="wb", prefix=f".{path.name}.", suffix=".tmp", dir=path.parent, delete=False, ) as handle: handle.write(encoded) handle.flush() os.fsync(handle.fileno()) temporary = Path(handle.name) os.replace(temporary, path) return True finally: if temporary is not None and temporary.exists(): temporary.unlink() def migrate(args: argparse.Namespace) -> tuple[int, dict[str, Any]]: try: evidence = inspect_evidence(args) except (OSError, ValueError) as exc: return 2, {"status": "evidence_invalid", "message": str(exc), "applied": False} model_sha256 = evidence["checksums"]["model"] snapshot_key = _snapshot_key(model_id=args.model_id, model_sha256=model_sha256) result: dict[str, Any] = { "status": "ready_to_apply", "message": "Recovered evidence is internally checksum-consistent; re-run with --apply to register it.", "applied": False, "model_id": args.model_id, "task_type": args.task_type, "source_version": args.source_version, "model_sha256": model_sha256, "snapshot_key": snapshot_key, "claim_boundary": CLAIM_BOUNDARY, "evidence": evidence, "will_download_models": False, "will_run_inference": False, } if not args.apply: return 0, result db = SessionLocal() manifest_path = RuntimeModelProvenanceService.manifest_path_for_model(evidence["paths"]["model"]) created_manifest = False try: source = SourceRegistryService.ensure_server_owned_source(db, RuntimeModelProvenanceService.SOURCE_REGISTRY_KEY) source.ingest_status = "configured" source.freshness_status = "current" registry_metadata = dict(source.registry_metadata_json or {}) registry_metadata["runtime_model_contract"] = { "key": PYTORCH_MODEL_CONTRACT_KEY, "version": PYTORCH_MODEL_CONTRACT_VERSION, "claim_boundary": CLAIM_BOUNDARY, } source.registry_metadata_json = registry_metadata snapshot = SourceRegistryService.record_snapshot( db, source_key=RuntimeModelProvenanceService.SOURCE_REGISTRY_KEY, snapshot_key=snapshot_key, checksum_sha256=model_sha256, source_version=args.source_version, source_url=f"file://{evidence['paths']['checkpoint']}", crs="not_applicable", units="model_weights", freshness_status="current", ingest_status="ingested", known_limitations=[CLAIM_BOUNDARY], observed_schema={ "model_format": "pytorch", "framework": "ultralytics/pytorch", "class_mapping": evidence["class_mapping"], }, snapshot_metadata={ "evidence_paths": evidence["paths"], "evidence_checksums_sha256": evidence["checksums"], "claim_boundary": CLAIM_BOUNDARY, }, reuse_existing_snapshot=True, ) if manifest_path.exists(): payload = _json_object(manifest_path, "existing runtime provenance sidecar") _require_equal(payload.get("model", {}).get("sha256"), model_sha256, "existing model.sha256") _require_equal( payload.get("model", {}).get("source_version"), args.source_version, "existing model.source_version", ) _require_equal( payload.get("source", {}).get("source_registry_id"), str(source.id), "existing source.source_registry_id", ) _require_equal( payload.get("source", {}).get("source_snapshot_id"), str(snapshot.id), "existing source.source_snapshot_id", ) else: imported_at = datetime.now(timezone.utc).isoformat() payload = _manifest_payload( args=args, evidence=evidence, source_registry_id=str(source.id), source_snapshot_id=str(snapshot.id), imported_at=imported_at, ) created_manifest = _write_manifest_atomically(manifest_path, payload) RuntimeModelProvenanceService.validate_for_runtime( model_path=evidence["paths"]["model"], model_id=args.model_id, task_type=args.task_type, expected_model_version=args.source_version, allowed_frameworks=("ultralytics/pytorch",), ) db.commit() production_evidence = RuntimeModelProvenanceService.validate_for_production_runtime( db=db, model_path=evidence["paths"]["model"], model_id=args.model_id, task_type=args.task_type, expected_model_version=args.source_version, allowed_frameworks=("ultralytics/pytorch",), ) result.update( { "status": "applied", "message": "Runtime model provenance was registered and validated against Postgres.", "applied": True, "manifest_path": str(manifest_path), "manifest_created": created_manifest, "source_registry_id": str(source.id), "source_snapshot_id": str(snapshot.id), "runtime_provenance": production_evidence.as_dict(), } ) return 0, result except Exception as exc: db.rollback() if created_manifest and manifest_path.exists(): manifest_path.unlink() return 3, { **result, "status": "apply_failed", "message": str(exc), "applied": False, } finally: db.close() def parse_args(argv: list[str]) -> argparse.Namespace: parser = argparse.ArgumentParser( description=( "Migrate a surviving legacy YOLO checkpoint into GeoIntel's runtime provenance contract without " "claiming missing training or review evidence." ) ) parser.add_argument("--model-path", required=True) parser.add_argument("--checkpoint-path", required=True) parser.add_argument("--base-model-path", required=True) parser.add_argument("--training-summary-path", required=True) parser.add_argument("--training-args-path", required=True) parser.add_argument("--training-results-path", required=True) parser.add_argument("--dataset-summary-path", required=True) parser.add_argument("--dataset-yaml-path", required=True) parser.add_argument("--model-id", default="yolo-configured") parser.add_argument("--task-type", default="object_detection") parser.add_argument("--source-version", required=True) parser.add_argument("--framework-version", required=True) parser.add_argument("--apply", action="store_true") parser.add_argument("--json", action="store_true") return parser.parse_args(argv) def main(argv: list[str] | None = None) -> int: args = parse_args(argv or sys.argv[1:]) exit_code, payload = migrate(args) if args.json: print(json.dumps(payload, indent=2, sort_keys=True)) else: print(f"status: {payload['status']}") print(f"message: {payload['message']}") print(f"claim boundary: {payload.get('claim_boundary', CLAIM_BOUNDARY)}") return exit_code if __name__ == "__main__": raise SystemExit(main())