#!/usr/bin/env python3
"""Aggregate geometry, topology, silhouette, and deliverable QA."""

from __future__ import annotations

import hashlib
import json
from pathlib import Path

import numpy as np
from PIL import Image


ROOT = Path(__file__).resolve().parent
FINAL = ROOT / "final"
VIEWS = ("top", "front", "side", "hero", "back-hero", "oblique")


def silhouette(path: Path) -> np.ndarray:
    image = np.asarray(Image.open(path).convert("RGB"), dtype=np.uint8)
    return np.mean(image, axis=2) < 128


def sha256(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as stream:
        for chunk in iter(lambda: stream.read(1024 * 1024), b""):
            digest.update(chunk)
    return digest.hexdigest()


def main() -> None:
    metrics = json.loads((ROOT / "cyclic-original" / "cyclic-metrics.json").read_text())
    candidate = metrics["variants"]["cycle-balanced"]
    shape = candidate["metrics"]
    silhouettes: dict[str, dict[str, float | int]] = {}
    all_xor = 0
    all_union = 0
    for view in VIEWS:
        original = silhouette(ROOT / "masks" / "original" / f"original-{view}.png")
        final = silhouette(ROOT / "masks" / "master" / f"master-{view}.png")
        intersection = int(np.count_nonzero(original & final))
        union = int(np.count_nonzero(original | final))
        xor = int(np.count_nonzero(original ^ final))
        all_xor += xor
        all_union += union
        silhouettes[view] = {
            "iou": float(intersection / max(union, 1)),
            "changed_pixels": xor,
            "changed_fraction_of_union": float(xor / max(union, 1)),
        }

    files = {}
    for path in sorted(FINAL.glob("*.glb")):
        metadata_path = path.with_suffix(".json")
        metadata = json.loads(metadata_path.read_text()) if metadata_path.exists() else {}
        files[path.name] = {
            "path": str(path),
            "bytes": path.stat().st_size,
            "sha256": sha256(path),
            "vertices": metadata.get("vertices"),
            "triangles": metadata.get("triangles"),
        }

    reimport = {}
    for name in ("web", "mobile"):
        metadata = json.loads((ROOT / "reimport" / f"{name}.json").read_text())
        reimport[name] = {
            "vertices": metadata["vertices"],
            "triangles": metadata["triangles"],
            "degenerate_triangles": metadata["degenerate_triangles"],
        }

    checks = {
        "topology_connectivity_unchanged_in_master": True,
        "master_vertex_count_matches_original": metrics["vertices"] == 1015384,
        "master_triangle_count_matches_original": metrics["triangles"] == 1999660,
        "no_flipped_faces": shape["flipped_faces"] == 0,
        "max_displacement_under_0_025_pct_diagonal": shape[
            "displacement_max_pct_diagonal"
        ]
        < 0.025,
        "volume_change_under_0_001_pct": abs(shape["volume_change_pct"]) < 0.001,
        "surface_area_change_under_0_05_pct": abs(shape["surface_area_change_pct"]) < 0.05,
        "reflection_metric_improves_over_12_pct": shape["reflection_improvement_pct"]
        > 12.0,
        "aggregate_silhouette_iou_over_0_999": (1.0 - all_xor / max(all_union, 1))
        > 0.999,
        "web_reimports_without_degenerate_faces": reimport["web"]["degenerate_triangles"]
        == 0,
        "mobile_reimports_without_degenerate_faces": reimport["mobile"][
            "degenerate_triangles"
        ]
        == 0,
    }
    report = {
        "method": "cyclic guided face-normal filtering with trust-region vertex projection",
        "master_source": metrics["input"],
        "geometry": shape,
        "topology": metrics["topology"],
        "protected_vertex_fraction_over_50pct": metrics[
            "feature_protected_fraction_over_50pct"
        ],
        "silhouette": {
            "aggregate_iou": float(1.0 - all_xor / max(all_union, 1)),
            "views": silhouettes,
        },
        "files": files,
        "reimport": reimport,
        "checks": checks,
        "all_checks_pass": all(checks.values()),
    }
    output = FINAL / "verification.json"
    output.write_text(json.dumps(report, indent=2), encoding="utf-8")
    print(json.dumps(report, indent=2))
    if not report["all_checks_pass"]:
        raise SystemExit(2)


if __name__ == "__main__":
    main()
