#!/usr/bin/env python3
"""Non-shrinking curvature fairing outside the protected front ornament.

This is a complementary candidate to guided normal filtering.  It uses
Taubin-style positive/negative Laplacian passes, projected predominantly onto
the original normals, with hard boundary/feature/ornament locks and a global
displacement trust region.  Connectivity is never changed.
"""

from __future__ import annotations

import argparse
import json
import time
from pathlib import Path

import numpy as np

import denoise_mesh as dm


VARIANTS = {
    "fair-gentle": {"pairs": 12, "lam": 0.46, "mu": -0.49, "max_pct": 0.0015},
    "fair-balanced": {"pairs": 28, "lam": 0.48, "mu": -0.51, "max_pct": 0.0030},
    "fair-strong": {"pairs": 52, "lam": 0.50, "mu": -0.53, "max_pct": 0.0050},
}


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))
    return parser.parse_args()


def ornament_protection(vertices: np.ndarray) -> np.ndarray:
    y = vertices[:, 1]
    frontness = (float(y.max()) - y) / max(float(y.max() - y.min()), 1e-12)
    return (0.90 * dm.smoothstep(0.66, 0.91, frontness)).astype(np.float32)


def neighbor_weights(vertices: np.ndarray, edges: np.ndarray) -> np.ndarray:
    lengths = np.linalg.norm(vertices[edges[:, 0]] - vertices[edges[:, 1]], axis=1)
    median = float(np.median(lengths))
    weights = median / np.maximum(lengths, median * 0.18)
    return np.clip(weights, 0.20, 4.0).astype(np.float32)


def weighted_laplacian(
    vertices: np.ndarray,
    edges: np.ndarray,
    weights: np.ndarray,
    weight_sum: np.ndarray,
) -> np.ndarray:
    left, right = edges[:, 0], edges[:, 1]
    accumulated = np.zeros_like(vertices, dtype=np.float32)
    for destination, source in ((left, right), (right, left)):
        for axis in range(3):
            accumulated[:, axis] += np.bincount(
                destination,
                weights=weights * vertices[source, axis],
                minlength=len(vertices),
            ).astype(np.float32)
    return accumulated / np.maximum(weight_sum[:, None], dm.EPS) - vertices


def fair(
    original: np.ndarray,
    edges: np.ndarray,
    original_normals: np.ndarray,
    protection: np.ndarray,
    config: dict[str, float | int],
    diagonal: float,
) -> np.ndarray:
    vertices = original.copy()
    weights = neighbor_weights(original, edges)
    left, right = edges[:, 0], edges[:, 1]
    weight_sum = np.bincount(left, weights=weights, minlength=len(original)).astype(np.float32)
    weight_sum += np.bincount(right, weights=weights, minlength=len(original)).astype(np.float32)
    trust = np.float32(diagonal * float(config["max_pct"]))
    movable = (1.0 - protection).astype(np.float32)

    for pair in range(int(config["pairs"])):
        for coefficient in (float(config["lam"]), float(config["mu"])):
            laplacian = weighted_laplacian(vertices, edges, weights, weight_sum)
            normal_amount = np.einsum("ij,ij->i", laplacian, original_normals)
            normal_part = original_normals * normal_amount[:, None]
            tangent_part = laplacian - normal_part
            update = normal_part + np.float32(0.035) * tangent_part
            vertices += np.float32(coefficient) * movable[:, None] * update

            displacement = vertices - original
            # Strongly discourage tangential drift accumulated across passes.
            amount = np.einsum("ij,ij->i", displacement, original_normals)
            normal_displacement = original_normals * amount[:, None]
            tangent_displacement = displacement - normal_displacement
            displacement = normal_displacement + np.float32(0.06) * tangent_displacement
            length = np.linalg.norm(displacement, axis=1)
            displacement *= np.minimum(1.0, trust / np.maximum(length, dm.EPS))[:, None]
            vertices = original + displacement
        if (pair + 1) % 10 == 0 or pair + 1 == int(config["pairs"]):
            displacement = np.linalg.norm(vertices - original, axis=1)
            print(
                f"  pair {pair + 1:02d}/{config['pairs']}: "
                f"p95={np.quantile(displacement, 0.95):.7f}, max={displacement.max():.7f}"
            )
    return vertices.astype(np.float32)


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 = np.asarray(data["vertices"], dtype=np.float32)
    faces = np.asarray(data["faces"], dtype=np.int32)
    diagonal = float(np.linalg.norm(original.max(axis=0) - original.min(axis=0)))
    topology = dm.build_topology(faces, len(original))
    edges = topology["unique_edges"]
    adjacent_faces = topology["adjacent_faces"]
    adjacent_edges = topology["adjacent_edges"]
    boundary_vertices = topology["boundary_vertices"]
    original_face_normals, face_areas, _ = dm.face_geometry(original, faces)
    original_vertex_normals = dm.vertex_normals(original, faces, face_areas, original_face_normals)
    guidance = dm.diffuse_normals(
        original_face_normals, face_areas, adjacent_faces, 18, 0.78, blend=0.82
    )
    feature, guide_distance = dm.feature_protection(
        guidance, adjacent_faces, adjacent_edges, boundary_vertices, len(original)
    )
    protection = np.maximum(feature, ornament_protection(original))
    protection[boundary_vertices] = 1.0
    smooth_mask = guide_distance < 0.22
    mean_edge = dm.mean_edge_length(original, edges)
    original_area = float(face_areas.sum(dtype=np.float64))
    original_volume = abs(dm.signed_volume(original, faces))
    baseline = dm.geometry_metrics(
        original,
        original,
        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),
        "triangles": len(faces),
        "topology": topology["stats"],
        "baseline": baseline,
        "protected_vertex_fraction_over_50pct": float(np.mean(protection > 0.5)),
        "variants": {},
    }
    started = time.perf_counter()
    for name in args.variants:
        config = VARIANTS[name]
        print(f"[{name}] starting")
        candidate = fair(
            original, edges, original_vertex_normals, protection, config, diagonal
        )
        candidate, reverted = dm.revert_unsafe_faces(
            original, candidate, faces, original_face_normals
        )
        output = outdir / f"{name}.npz"
        np.savez(output, vertices=candidate, faces=faces)
        metrics = dm.geometry_metrics(
            original,
            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 / "fairing-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()
