brulee-1's picture
Add files using upload-large-folder tool
6debdcc verified
Raw
History Blame Contribute Delete
12.3 kB
"""Verify the integrity and provenance of an archived research repository."""
from __future__ import annotations
import argparse
import csv
import hashlib
import json
import logging
from pathlib import Path
from typing import Any
LOGGER = logging.getLogger(__name__)
ARCHIVE_MANIFEST = Path("results/archive_manifest.json")
RAW_ROWS = Path("results/raw_rows.jsonl")
WEIGHT_MANIFEST = Path("weights/manifest.csv")
ARCHIVE_DIRS = ("images", "results", "samples", "weights")
def _sha256(path: Path) -> str:
digest = hashlib.sha256()
with path.open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _safe_relative(root: Path, raw_path: str) -> Path:
candidate = Path(raw_path)
if candidate.is_absolute():
raise ValueError(f"path containment violation: {raw_path}")
raw_candidate = root / candidate
if raw_candidate.is_symlink():
raise ValueError(f"path is a symlink: {raw_path}")
resolved = raw_candidate.resolve()
try:
resolved.relative_to(root.resolve())
except ValueError as exc:
raise ValueError(f"path containment violation: {raw_path}") from exc
return resolved
def _relative_path(root: Path, path: Path) -> str:
return path.relative_to(root).as_posix()
def _validate_entry(root: Path, entry: dict[str, Any], *, label: str) -> Path:
raw_path = entry.get("path", entry.get("artifact_path"))
if not isinstance(raw_path, str) or not raw_path:
raise ValueError(f"{label}: missing path")
path = _safe_relative(root, raw_path)
if path.is_symlink():
raise ValueError(f"{label}: symlink is not allowed: {raw_path}")
if not path.is_file():
raise ValueError(f"{label}: listed file is missing: {raw_path}")
expected_bytes = entry.get("bytes", entry.get("size"))
expected_sha = entry.get("sha256", entry.get("sha"))
if not isinstance(expected_bytes, int) or path.stat().st_size != expected_bytes:
raise ValueError(f"{label}: byte count mismatch: {raw_path}")
if not isinstance(expected_sha, str) or _sha256(path) != expected_sha:
raise ValueError(f"{label}: SHA-256 mismatch: {raw_path}")
return path
def _manifest_entries(value: Any) -> list[dict[str, Any]]:
if isinstance(value, list):
entries = value
elif isinstance(value, dict):
entries = value.get("files", value.get("artifacts"))
else:
entries = None
if not isinstance(entries, list) or not all(isinstance(entry, dict) for entry in entries):
raise ValueError("archive manifest must contain a files/artifacts list")
return entries
def _read_weight_manifest(root: Path) -> list[dict[str, Any]]:
path = _safe_relative(root, str(WEIGHT_MANIFEST))
if path.is_symlink() or not path.is_file():
raise ValueError("weights/manifest.csv is missing or is a symlink")
with path.open(newline="", encoding="utf-8") as handle:
rows = list(csv.DictReader(handle))
if not rows:
raise ValueError("weights/manifest.csv contains no entries")
normalized: list[dict[str, Any]] = []
for index, row in enumerate(rows, 1):
try:
byte_count = int(row.get("bytes", row.get("size", "")))
except ValueError as exc:
raise ValueError(f"weights manifest row {index}: invalid bytes") from exc
normalized.append(
{"path": row.get("path", row.get("artifact_path")), "bytes": byte_count, "sha256": row.get("sha256", row.get("sha"))}
)
return normalized
def _load_jsonl_generation_rows(root: Path, conditions: list[str], prompt_count: int) -> list[dict[str, Any]]:
raw_path = _safe_relative(root, str(RAW_ROWS))
if raw_path.is_symlink() or not raw_path.is_file():
raise ValueError("results/raw_rows.jsonl is missing or is a symlink")
rows: list[dict[str, Any]] = []
seen: set[tuple[str, str]] = set()
try:
lines = raw_path.read_text(encoding="utf-8").splitlines()
except (OSError, UnicodeDecodeError) as exc:
raise ValueError("results/raw_rows.jsonl is not valid UTF-8") from exc
for line_number, line in enumerate(lines, 1):
if not line.strip():
continue
try:
row = json.loads(line)
except json.JSONDecodeError as exc:
raise ValueError(f"raw_rows.jsonl line {line_number} is not valid JSON") from exc
if not isinstance(row, dict):
raise TypeError(f"raw_rows.jsonl line {line_number} is not a JSON object")
if row.get("row_type") != "generation":
continue
condition = row.get("condition")
prompt_id = row.get("prompt_id")
if not isinstance(condition, str) or condition not in conditions:
raise ValueError(f"generation row {line_number}: unknown condition")
if not isinstance(prompt_id, str) or not prompt_id:
raise ValueError(f"generation row {line_number}: missing prompt_id")
key = (condition, prompt_id)
if key in seen:
raise ValueError(f"duplicate generation row: {condition}/{prompt_id}")
seen.add(key)
rows.append(row)
expected_rows = len(conditions) * prompt_count
if len(rows) != expected_rows:
raise ValueError(f"generation row count mismatch: expected {expected_rows}, got {len(rows)}")
prompt_ids = {str(row["prompt_id"]) for row in rows}
if len(prompt_ids) != prompt_count:
raise ValueError(f"prompt count mismatch: expected {prompt_count}, got {len(prompt_ids)}")
return rows
def _validate_generation_provenance(root: Path, rows: list[dict[str, Any]], archive_paths: set[str]) -> int:
successful = 0
for row in rows:
condition = str(row["condition"])
prompt_id = str(row["prompt_id"])
label = f"generation {condition}/{prompt_id}"
provenance = (row.get("sample_path"), row.get("sample_hash"), row.get("sample_sha256"), row.get("sample_bytes"))
if row.get("success"):
sample_raw, sample_hash, sample_sha256, sample_bytes = provenance
if not isinstance(sample_raw, str) or not sample_raw:
raise ValueError(f"{label}: successful row has no sample path")
sample_path = _safe_relative(root, sample_raw)
if sample_path.is_symlink() or not sample_path.is_file():
raise ValueError(f"{label}: successful sample is missing or is a symlink")
actual_bytes = sample_path.stat().st_size
actual_hash = _sha256(sample_path)
if not isinstance(sample_bytes, int) or actual_bytes != sample_bytes:
raise ValueError(f"{label}: sample byte count mismatch")
hashes = [value for value in (sample_hash, sample_sha256) if value is not None]
if not hashes or any(not isinstance(value, str) or value != actual_hash for value in hashes):
raise ValueError(f"{label}: sample SHA-256 mismatch")
relative = _relative_path(root, sample_path)
if relative not in archive_paths:
raise ValueError(f"{label}: sample is not listed in archive manifest: {relative}")
successful += 1
elif any(value is not None for value in provenance):
raise ValueError(f"{label}: failed row contains sample provenance")
return successful
def _reject_unlisted_tree_files(root: Path, archive_paths: set[str]) -> None:
for directory in ARCHIVE_DIRS:
for path in (root / directory).rglob("*"):
if path.is_symlink():
raise ValueError(f"symlink found in archived tree: {_relative_path(root, path)}")
if path.is_file():
relative = _relative_path(root, path)
if relative == ARCHIVE_MANIFEST.as_posix():
continue
if relative not in archive_paths:
raise ValueError(f"unlisted file in archived tree: {relative}")
def verify_archive(root: str | Path) -> dict[str, int]:
"""Verify archive files, raw generation provenance, containment, and symlink safety."""
archive_root = Path(root).resolve()
if not archive_root.is_dir():
raise ValueError(f"archive root is not a directory: {archive_root}")
for directory in ARCHIVE_DIRS:
path = _safe_relative(archive_root, directory)
if path.is_symlink() or not path.is_dir():
raise ValueError(f"archive directory missing or symlinked: {directory}")
manifest_path = _safe_relative(archive_root, str(ARCHIVE_MANIFEST))
if manifest_path.is_symlink() or not manifest_path.is_file():
raise ValueError("results/archive_manifest.json is missing or is a symlink")
try:
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
except (OSError, UnicodeDecodeError, json.JSONDecodeError) as exc:
raise ValueError("archive manifest is not valid UTF-8 JSON") from exc
if not isinstance(manifest, dict):
raise TypeError("archive manifest must be a JSON object")
profile = manifest.get("profile")
prompt_count = manifest.get("prompt_count")
conditions = manifest.get("conditions")
if not isinstance(profile, str) or not profile:
raise ValueError("archive manifest profile is missing")
if not isinstance(prompt_count, int) or prompt_count <= 0:
raise ValueError("archive manifest prompt_count is invalid")
if not isinstance(conditions, list) or not conditions or not all(isinstance(item, str) for item in conditions):
raise ValueError("archive manifest conditions are invalid")
if len(set(conditions)) != len(conditions):
raise ValueError("archive manifest conditions contain duplicates")
entries = _manifest_entries(manifest)
archive_paths: set[str] = set()
checked = 0
for index, entry in enumerate(entries, 1):
path = _validate_entry(archive_root, entry, label=f"archive manifest row {index}")
relative = _relative_path(archive_root, path)
if relative in archive_paths:
raise ValueError(f"duplicate archive path: {relative}")
archive_paths.add(relative)
checked += 1
for required in ("images", "results", "samples", "weights"):
if not any(path == required or path.startswith(required + "/") for path in archive_paths):
raise ValueError(f"archive manifest must list a file under {required}")
for required in (RAW_ROWS.as_posix(), WEIGHT_MANIFEST.as_posix()):
if required not in archive_paths:
raise ValueError(f"archive manifest must list {required}")
weight_entries = _read_weight_manifest(archive_root)
weight_paths: set[str] = set()
for index, entry in enumerate(weight_entries, 1):
path = _validate_entry(archive_root, entry, label=f"weights manifest row {index}")
relative = _relative_path(archive_root, path)
if not relative.startswith("weights/"):
raise ValueError(f"weights manifest path is outside weights/: {relative}")
if relative in weight_paths:
raise ValueError(f"duplicate weight path: {relative}")
weight_paths.add(relative)
generation_rows = _load_jsonl_generation_rows(archive_root, conditions, prompt_count)
successful_samples = _validate_generation_provenance(archive_root, generation_rows, archive_paths)
_reject_unlisted_tree_files(archive_root, archive_paths)
result = {
"checked_files": checked + len(weight_entries),
"checked_weight_entries": len(weight_entries),
"generation_rows": len(generation_rows),
"successful_samples": successful_samples,
}
LOGGER.info("archive verified: %s", result)
return result
def main() -> int:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--root", type=Path, default=Path("."), help="archive repository root")
args = parser.parse_args()
logging.basicConfig(level=logging.INFO, format="%(levelname)s %(message)s")
report = verify_archive(args.root)
print(json.dumps(report, sort_keys=True))
return 0
if __name__ == "__main__":
raise SystemExit(main())