#!/usr/bin/env python3
"""Iterative guided-normal denoising with shape and ornament safeguards."""

from __future__ import annotations

import argparse
import json
import time
from pathlib import Path

import numpy as np

import denoise_mesh as dm


VARIANTS = {
    "cycle-gentle": {
        "outer_iterations": 6,
        "normal_iterations": 4,
        "vertex_iterations": 3,
        "sigma": 0.40,
        "step": 0.68,
        "data_pull": 0.0030,
        "max_pct": 0.0010,
    },
    "cycle-balanced": {
        "outer_iterations": 12,
        "normal_iterations": 6,
        "vertex_iterations": 4,
        "sigma": 0.48,
        "step": 0.76,
        "data_pull": 0.0015,
        "max_pct": 0.0020,
    },
    "cycle-strong": {
        "outer_iterations": 20,
        "normal_iterations": 8,
        "vertex_iterations": 5,
        "sigma": 0.56,
        "step": 0.82,
        "data_pull": 0.0008,
        "max_pct": 0.0040,
        "ornament_lock": 0.72,
    },
    "cycle-polish": {
        "outer_iterations": 72,
        "normal_iterations": 12,
        "vertex_iterations": 6,
        "sigma": 0.64,
        "step": 0.86,
        "data_pull": 0.0002,
        "max_pct": 0.0080,
        "ornament_lock": 0.96,
    },
}


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser()
    parser.add_argument("--input", required=True)
    parser.add_argument("--outdir", required=True)
    parser.add_argument("--variants", nargs="+", default=list(VARIANTS), choices=list(VARIANTS))
    parser.add_argument(
        "--ornament-protection",
        choices=("front-y", "none"),
        default="front-y",
        help="Use the original ring's front ornament lock, or rely only on mesh features.",
    )
    return parser.parse_args()


def vertex_area_weights(
    vertex_count: int, faces: np.ndarray, face_areas: np.ndarray
) -> np.ndarray:
    result = np.zeros(vertex_count, dtype=np.float32)
    for corner in range(3):
        result += np.bincount(
            faces[:, corner], weights=face_areas, minlength=vertex_count
        ).astype(np.float32)
    return result / np.float32(3.0)


def ornament_protection(vertices: np.ndarray, strength: float = 0.72) -> np.ndarray:
    """Protect the distinctive front knot while still allowing mild cleanup."""
    y = vertices[:, 1]
    y_min, y_max = float(y.min()), float(y.max())
    normalized_front = (y_max - y) / max(y_max - y_min, 1e-12)
    # Starts in the front quarter and reaches 72% protection at the foremost tip.
    return (strength * dm.smoothstep(0.72, 0.94, normalized_front)).astype(np.float32)


def project_current_vertices(
    current_vertices: np.ndarray,
    original_vertices: np.ndarray,
    faces: np.ndarray,
    target_normals: np.ndarray,
    original_vertex_normals: np.ndarray,
    protection: np.ndarray,
    counts: np.ndarray,
    config: dict[str, float | int],
    diagonal: float,
) -> np.ndarray:
    vertices = current_vertices.copy()
    trust = np.float32(diagonal * float(config["max_pct"]))
    for _ in range(int(config["vertex_iterations"])):
        centroids = (
            vertices[faces[:, 0]] + vertices[faces[:, 1]] + vertices[faces[:, 2]]
        ) / np.float32(3.0)
        correction = np.zeros_like(vertices, dtype=np.float32)
        for corner in range(3):
            indices = faces[:, corner]
            residual = np.einsum("ij,ij->i", target_normals, centroids - vertices[indices])
            projected = target_normals * residual[:, None]
            for axis in range(3):
                correction[:, axis] += np.bincount(
                    indices,
                    weights=projected[:, axis],
                    minlength=len(vertices),
                ).astype(np.float32)
        correction /= counts[:, None]
        correction *= 1.0 - 0.94 * protection[:, None]
        proposed = vertices + np.float32(config["step"]) * correction

        displacement = proposed - original_vertices
        normal_amount = np.einsum("ij,ij->i", displacement, original_vertex_normals)
        normal_part = original_vertex_normals * normal_amount[:, None]
        tangent_part = displacement - normal_part
        displacement = normal_part + np.float32(0.10) * tangent_part
        displacement *= np.float32(1.0 - float(config["data_pull"]))
        length = np.linalg.norm(displacement, axis=1)
        displacement *= np.minimum(1.0, trust / np.maximum(length, dm.EPS))[:, None]
        vertices = original_vertices + displacement
    return vertices.astype(np.float32)


def denoise_variant(
    original_vertices: np.ndarray,
    faces: np.ndarray,
    adjacent_faces: np.ndarray,
    fixed_guidance: np.ndarray,
    original_vertex_normals: np.ndarray,
    protection: np.ndarray,
    counts: np.ndarray,
    config: dict[str, float | int],
    diagonal: float,
) -> np.ndarray:
    vertices = original_vertices.copy()
    for outer in range(int(config["outer_iterations"])):
        face_normals, face_areas, _ = dm.face_geometry(vertices, faces)
        target_normals = dm.diffuse_normals(
            face_normals,
            face_areas,
            adjacent_faces,
            int(config["normal_iterations"]),
            float(config["sigma"]),
            fixed_guidance=fixed_guidance,
            blend=0.88,
        )
        vertices = project_current_vertices(
            vertices,
            original_vertices,
            faces,
            target_normals,
            original_vertex_normals,
            protection,
            counts,
            config,
            diagonal,
        )
        if (outer + 1) % 4 == 0 or outer + 1 == int(config["outer_iterations"]):
            displacement = np.linalg.norm(vertices - original_vertices, axis=1)
            print(
                f"  cycle {outer + 1:02d}/{config['outer_iterations']}: "
                f"p95={np.quantile(displacement, 0.95):.7f}, max={displacement.max():.7f}"
            )
    return vertices


def main() -> None:
    args = parse_args()
    source = Path(args.input).resolve()
    outdir = Path(args.outdir).resolve()
    outdir.mkdir(parents=True, exist_ok=True)
    data = np.load(source)
    original_vertices = np.asarray(data["vertices"], dtype=np.float32)
    faces = np.asarray(data["faces"], dtype=np.int32)
    diagonal = float(
        np.linalg.norm(original_vertices.max(axis=0) - original_vertices.min(axis=0))
    )
    topology = dm.build_topology(faces, len(original_vertices))
    adjacent_faces = topology["adjacent_faces"]
    adjacent_edges = topology["adjacent_edges"]
    boundary_vertices = topology["boundary_vertices"]
    original_face_normals, face_areas, _ = dm.face_geometry(original_vertices, faces)
    original_vertex_normals = dm.vertex_normals(
        original_vertices, faces, face_areas, original_face_normals
    )
    fixed_guidance = dm.diffuse_normals(
        original_face_normals, face_areas, adjacent_faces, 18, 0.78, blend=0.82
    )
    feature, guide_distance = dm.feature_protection(
        fixed_guidance,
        adjacent_faces,
        adjacent_edges,
        boundary_vertices,
        len(original_vertices),
    )
    smooth_mask = guide_distance < 0.22
    counts = np.zeros(len(original_vertices), dtype=np.float32)
    for corner in range(3):
        counts += np.bincount(faces[:, corner], minlength=len(original_vertices)).astype(np.float32)
    counts = np.maximum(counts, 1.0)
    original_area = float(face_areas.sum(dtype=np.float64))
    original_volume = abs(dm.signed_volume(original_vertices, faces))
    mean_edge = dm.mean_edge_length(original_vertices, topology["unique_edges"])
    baseline = dm.geometry_metrics(
        original_vertices,
        original_vertices,
        faces,
        original_face_normals,
        adjacent_faces,
        smooth_mask,
        diagonal,
        mean_edge,
        original_area,
        original_volume,
    )
    report: dict[str, object] = {
        "input": str(source),
        "vertices": len(original_vertices),
        "triangles": len(faces),
        "topology": topology["stats"],
        "baseline": baseline,
        "feature_protected_fraction_over_50pct": None,
        "variants": {},
    }
    started = time.perf_counter()
    for name in args.variants:
        config = VARIANTS[name]
        if args.ornament_protection == "front-y":
            protection = np.maximum(
                feature,
                ornament_protection(
                    original_vertices, float(config.get("ornament_lock", 0.72))
                ),
            )
        else:
            protection = feature.copy()
        protection[boundary_vertices] = 1.0
        report["feature_protected_fraction_over_50pct"] = float(
            np.mean(protection > 0.5)
        )
        print(f"[{name}] starting")
        candidate = denoise_variant(
            original_vertices,
            faces,
            adjacent_faces,
            fixed_guidance,
            original_vertex_normals,
            protection,
            counts,
            config,
            diagonal,
        )
        candidate, reverted = dm.revert_unsafe_faces(
            original_vertices, candidate, faces, original_face_normals
        )
        output = outdir / f"{name}.npz"
        np.savez(output, vertices=candidate, faces=faces)
        metrics = dm.geometry_metrics(
            original_vertices,
            candidate,
            faces,
            original_face_normals,
            adjacent_faces,
            smooth_mask,
            diagonal,
            mean_edge,
            original_area,
            original_volume,
        )
        metrics["normal_tv_improvement_pct"] = float(
            100.0 * (1.0 - metrics["smooth_region_normal_tv_rms"] / baseline["smooth_region_normal_tv_rms"])
        )
        metrics["reflection_improvement_pct"] = float(
            100.0 * (1.0 - metrics["reflection_field_rms"] / baseline["reflection_field_rms"])
        )
        report["variants"][name] = {
            "config": config,
            "output": str(output),
            "metrics": metrics,
            "reverted_vertices_for_flip_safety": reverted,
        }
        print(
            f"[{name}] reflection={metrics['reflection_improvement_pct']:.2f}% "
            f"volume={metrics['volume_change_pct']:.4f}% flips={metrics['flipped_faces']}"
        )
    report["elapsed_seconds"] = time.perf_counter() - started
    (outdir / "cyclic-metrics.json").write_text(
        json.dumps(report, indent=2), encoding="utf-8"
    )
    print(f"Done in {report['elapsed_seconds']:.1f}s")


if __name__ == "__main__":
    main()
