#!/usr/bin/env python3
"""Extract a GLB's joined triangle mesh to a compact NumPy archive.

Run with Blender:
    blender -b --python extract_mesh.py -- --input model.glb --output mesh.npz
"""

from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path

import bpy
import numpy as np


def parse_args() -> argparse.Namespace:
    argv = sys.argv[sys.argv.index("--") + 1 :] if "--" in sys.argv else []
    parser = argparse.ArgumentParser()
    parser.add_argument("--input", required=True)
    parser.add_argument("--output", required=True)
    return parser.parse_args(argv)


def clear_scene() -> None:
    bpy.ops.object.select_all(action="SELECT")
    bpy.ops.object.delete(use_global=False)


def import_joined(path: Path) -> bpy.types.Object:
    bpy.ops.import_scene.gltf(filepath=str(path))
    objects = [obj for obj in bpy.context.scene.objects if obj.type == "MESH"]
    if not objects:
        raise RuntimeError(f"No mesh found in {path}")
    bpy.ops.object.select_all(action="DESELECT")
    for obj in objects:
        obj.select_set(True)
    bpy.context.view_layer.objects.active = objects[0]
    if len(objects) > 1:
        bpy.ops.object.join()
    obj = bpy.context.view_layer.objects.active
    bpy.ops.object.transform_apply(location=True, rotation=True, scale=True)
    obj.data.validate(verbose=False, clean_customdata=False)
    obj.data.update()
    return obj


def main() -> None:
    args = parse_args()
    source = Path(args.input).resolve()
    output = Path(args.output).resolve()
    output.parent.mkdir(parents=True, exist_ok=True)

    clear_scene()
    obj = import_joined(source)
    mesh = obj.data
    mesh.calc_loop_triangles()

    vertices = np.empty(len(mesh.vertices) * 3, dtype=np.float32)
    mesh.vertices.foreach_get("co", vertices)
    vertices = vertices.reshape((-1, 3))

    faces = np.empty(len(mesh.loop_triangles) * 3, dtype=np.int32)
    mesh.loop_triangles.foreach_get("vertices", faces)
    faces = faces.reshape((-1, 3))

    # Uncompressed NPZ avoids spending minutes compressing a multi-million-face mesh.
    np.savez(output, vertices=vertices, faces=faces)

    edge_a = vertices[faces[:, 1]] - vertices[faces[:, 0]]
    edge_b = vertices[faces[:, 2]] - vertices[faces[:, 0]]
    double_area = np.linalg.norm(np.cross(edge_a, edge_b), axis=1)
    metadata = {
        "input": str(source),
        "output": str(output),
        "object_name": obj.name,
        "vertices": int(vertices.shape[0]),
        "triangles": int(faces.shape[0]),
        "degenerate_triangles": int(np.count_nonzero(double_area < 1e-12)),
        "bounds_min": vertices.min(axis=0).astype(float).tolist(),
        "bounds_max": vertices.max(axis=0).astype(float).tolist(),
    }
    output.with_suffix(".json").write_text(
        json.dumps(metadata, indent=2), encoding="utf-8"
    )
    print(json.dumps(metadata, indent=2))


if __name__ == "__main__":
    main()
