#!/usr/bin/env python3
"""Local, dependency-free GLB base-geometry accounting. Python 3.9+.

No network calls, renderer, export conversion, welding, or performance prediction.
Supports embedded GLB 2.0 buffers, FLOAT/VEC3 POSITION, ordinary unsigned indices,
strided accessors, all core primitive modes, and EXT_mesh_gpu_instancing.
Refuses sparse/compressed/quantized geometry and unknown required extensions.
This is a deliberately bounded inspector, not the Khronos glTF Validator.
"""
import argparse
import hashlib
import json
import math
import struct
import sys
from pathlib import Path

VERSION = "1.0.0"
MAX_BYTES = 64 * 1024 * 1024
MAX_ROWS = 2_000_000
MODES = {0: "POINTS", 1: "LINES", 2: "LINE_LOOP", 3: "LINE_STRIP",
         4: "TRIANGLES", 5: "TRIANGLE_STRIP", 6: "TRIANGLE_FAN"}
SAFE_REQUIRED = {
    "EXT_mesh_gpu_instancing", "KHR_materials_unlit", "KHR_materials_pbrSpecularGlossiness",
    "KHR_materials_clearcoat", "KHR_materials_transmission", "KHR_materials_volume",
    "KHR_materials_ior", "KHR_materials_specular", "KHR_materials_sheen",
    "KHR_materials_emissive_strength", "KHR_materials_iridescence",
    "KHR_materials_anisotropy", "KHR_texture_transform", "KHR_texture_basisu",
    "EXT_texture_webp", "KHR_lights_punctual", "KHR_materials_variants",
}


class InspectionError(ValueError):
    pass


def require(ok, message):
    if not ok:
        raise InspectionError(message)


def integer(value, label, minimum=0):
    require(type(value) is int and value >= minimum, f"{label}: invalid integer")
    return value


def at(items, index, label):
    integer(index, label)
    require(index < len(items), f"{label}: index out of range")
    return items[index]


def parse_glb(data):
    require(20 <= len(data) <= MAX_BYTES, "GLB must be 20 bytes to 64 MiB")
    magic, version, length = struct.unpack_from("<4sII", data)
    require(magic == b"glTF" and version == 2, "Expected GLB version 2")
    require(length == len(data), "GLB header length differs from actual bytes")
    chunks = []
    offset = 12
    while offset < length:
        require(offset + 8 <= length, "Truncated chunk header")
        size, kind = struct.unpack_from("<II", data, offset)
        offset += 8
        require(size % 4 == 0 and offset + size <= length, "Invalid chunk bounds/alignment")
        chunks.append((kind, data[offset:offset + size]))
        offset += size
    require(chunks and chunks[0][0] == 0x4E4F534A, "First chunk must be JSON")
    require(sum(k == 0x4E4F534A for k, _ in chunks) == 1, "Multiple JSON chunks")
    require(sum(k == 0x004E4942 for k, _ in chunks) <= 1, "Multiple BIN chunks")
    document = json.loads(chunks[0][1].decode("utf-8"))
    require(document.get("asset", {}).get("version") == "2.0", "Expected glTF 2.0 asset")
    binary = next((b for k, b in chunks if k == 0x004E4942), b"")
    if binary:
        require(len(chunks) > 1 and chunks[1][0] == 0x004E4942, "BIN must follow JSON")
    buffers = document.get("buffers", [])
    require(len(buffers) <= 1, "Only one embedded GLB buffer is supported")
    if buffers:
        require("uri" not in buffers[0], "External/data-URI buffers are unsupported")
        n = integer(buffers[0].get("byteLength"), "buffer byteLength")
        require(n <= len(binary) <= n + 3, "BIN size differs from declared buffer")
        binary = binary[:n]
    else:
        require(not binary, "BIN without buffer declaration")
    return document, binary


class Reader:
    def __init__(self, document, binary):
        self.d, self.binary = document, binary
        self.cache = {}

    def accessor(self, index):
        return at(self.d.get("accessors", []), index, "accessor")

    def decode(self, index):
        if index in self.cache:
            return self.cache[index]
        a = self.accessor(index)
        require("sparse" not in a, "Sparse accessor unsupported; expand it first")
        require(not a.get("normalized", False), "Normalized accessor decoding unsupported")
        formats = {5121: ("B", 1), 5123: ("H", 2), 5125: ("I", 4), 5126: ("f", 4)}
        require(a.get("componentType") in formats, "Unsupported accessor component type")
        sizes = {"SCALAR": 1, "VEC2": 2, "VEC3": 3, "VEC4": 4}
        require(a.get("type") in sizes, "Unsupported accessor shape")
        code, width = formats[a["componentType"]]
        components = sizes[a["type"]]
        n = integer(a.get("count"), "accessor count", 1)
        require(n <= MAX_ROWS, "Accessor exceeds 2 million row safety limit")
        require("bufferView" in a, "Accessor without bufferView unsupported")
        view = at(self.d.get("bufferViews", []), a["bufferView"], "bufferView")
        require(view.get("buffer") == 0, "Expected embedded buffer 0")
        start = integer(view.get("byteOffset", 0), "view offset")
        length = integer(view.get("byteLength"), "view length")
        require(start + length <= len(self.binary), "bufferView exceeds BIN")
        offset = integer(a.get("byteOffset", 0), "accessor offset")
        row_size = width * components
        stride = integer(view.get("byteStride", row_size), "stride", row_size)
        require(stride % width == 0 and offset % width == 0 and start % width == 0,
                "Misaligned accessor")
        require(offset + stride * (n - 1) + row_size <= length, "Accessor exceeds bufferView")
        fmt = "<" + code * components
        rows = [struct.unpack_from(fmt, self.binary, start + offset + i * stride) for i in range(n)]
        require(all(math.isfinite(v) for row in rows for v in row), "Non-finite accessor value")
        self.cache[index] = rows
        return rows


def inspect_bytes(data, scene=None):
    d, binary = parse_glb(data)
    required = set(d.get("extensionsRequired", []))
    require(not required - SAFE_REQUIRED,
            "Unsupported required extension(s): " + ", ".join(sorted(required - SAFE_REQUIRED)))
    # Optional extensions also cannot silently replace the geometry being counted.
    def check_extensions(value):
        if isinstance(value, dict):
            forbidden = set(value.get("extensions", {})) & {
                "KHR_draco_mesh_compression", "EXT_meshopt_compression", "KHR_meshopt_compression", "KHR_mesh_quantization",
                "MSFT_lod", "EXT_mesh_manifold", "EXT_node_visibility",
            }
            require(not forbidden, "Unsupported geometry/selection extension: " + ", ".join(sorted(forbidden)))
            for child in value.values():
                check_extensions(child)
        elif isinstance(value, list):
            for child in value:
                check_extensions(child)
    check_extensions(d)
    reader = Reader(d, binary)
    rows, mesh_totals = [], []
    all_position_accessors, used_materials = set(), set()
    meshes = d.get("meshes", [])
    for mi, mesh in enumerate(meshes):
        totals = {"triangles": 0, "primitive_position_rows": 0, "primitive_count": 0}
        mesh_positions = set()
        for pi, primitive in enumerate(mesh.get("primitives", [])):
            attrs = primitive.get("attributes", {})
            require("POSITION" in attrs, "Primitive without POSITION unsupported")
            ai = attrs["POSITION"]
            a = reader.accessor(ai)
            require(a.get("type") == "VEC3" and a.get("componentType") == 5126,
                    "POSITION must be unquantized FLOAT VEC3")
            position_rows = reader.decode(ai)
            for semantic, attr in attrs.items():
                require(reader.accessor(attr).get("count") == len(position_rows),
                        f"{semantic}: mismatched attribute count")
            indices = None
            if "indices" in primitive:
                ia = reader.accessor(primitive["indices"])
                require(ia.get("type") == "SCALAR" and ia.get("componentType") in (5121, 5123, 5125),
                        "Indices must be unsigned SCALAR")
                indices = [r[0] for r in reader.decode(primitive["indices"])]
                maximum = {5121: 255, 5123: 65535, 5125: 4294967295}[ia["componentType"]]
                require(all(i < len(position_rows) and i != maximum for i in indices),
                        "Index out of range or reserved primitive-restart value")
            n = len(indices) if indices is not None else len(position_rows)
            mode = primitive.get("mode", 4)
            require(type(mode) is int and mode in MODES, "Unsupported primitive mode")
            minimum = {0: 1, 1: 2, 2: 2, 3: 2, 4: 3, 5: 3, 6: 3}[mode]
            require(n >= minimum, "Insufficient elements for primitive mode")
            require(mode != 4 or n % 3 == 0, "TRIANGLES element count must divide by 3")
            require(mode != 1 or n % 2 == 0, "LINES element count must divide by 2")
            triangles = n // 3 if mode == 4 else n - 2 if mode in (5, 6) else 0
            material = primitive.get("material")
            if material is not None:
                at(d.get("materials", []), material, "material")
                used_materials.add(material)
            all_position_accessors.add(ai)
            unique_positions = set(position_rows)
            mesh_positions.update(unique_positions)
            normal_evidence = None
            if "NORMAL" in attrs:
                normal_accessor = reader.accessor(attrs["NORMAL"])
                require(normal_accessor.get("type") == "VEC3" and normal_accessor.get("componentType") == 5126,
                        "NORMAL must be unquantized FLOAT VEC3")
                normals = reader.decode(attrs["NORMAL"])
                normals_by_position = {}
                for position, normal in zip(position_rows, normals):
                    normals_by_position.setdefault(position, set()).add(normal)
                normal_evidence = {
                    "distinct_normal_tuples": len(set(normals)),
                    "distinct_position_normal_pairs": len(set(zip(position_rows, normals))),
                    "positions_with_multiple_normal_values": sum(len(values) > 1 for values in normals_by_position.values()),
                    "normal_variants_per_distinct_position": sorted(len(values) for values in normals_by_position.values()),
                }
            rows.append({"mesh": mi, "primitive": pi, "mode": MODES[mode],
                         "position_accessor": ai, "position_rows": len(position_rows),
                         "distinct_local_position_tuples": len(unique_positions),
                         "indexed": indices is not None, "element_count": n,
                         "assembled_triangles": triangles, "material": material,
                         "normal_evidence": normal_evidence})
            totals["triangles"] += triangles
            totals["primitive_position_rows"] += len(position_rows)
            totals["primitive_count"] += 1
        totals["distinct_local_position_tuples"] = len(mesh_positions)
        mesh_totals.append(totals)
    chosen_scene = d.get("scene") if scene is None else scene
    active = None
    if chosen_scene is not None:
        selected = at(d.get("scenes", []), chosen_scene, "scene")
        active = {"scene": chosen_scene, "mesh_nodes": 0, "mesh_copies": 0,
                  "assembled_triangles": 0, "primitive_position_rows": 0,
                  "primitive_copies": 0, "referenced_materials": []}
        nodes = d.get("nodes", [])
        seen, scene_materials = set(), set()
        stack = list(reversed(selected.get("nodes", [])))
        while stack:
            ni = stack.pop()
            node = at(nodes, ni, "node")
            require(ni not in seen, "Repeated/cyclic node in selected scene")
            seen.add(ni)
            stack.extend(reversed(node.get("children", [])))
            instancing = node.get("extensions", {}).get("EXT_mesh_gpu_instancing")
            if "mesh" not in node:
                require(instancing is None, "GPU instancing requires a mesh node")
                continue
            mi = node["mesh"]
            total = at(mesh_totals, mi, "node mesh")
            copies = 1
            if instancing is not None:
                instance_attrs = instancing.get("attributes", {})
                require(bool(instance_attrs), "GPU instancing needs attribute accessors")
                counts = {integer(reader.accessor(i).get("count"), "instance count", 1)
                          for i in instance_attrs.values()}
                require(len(counts) == 1, "GPU instance attribute counts disagree")
                copies = counts.pop()
            active["mesh_nodes"] += 1
            active["mesh_copies"] += copies
            active["assembled_triangles"] += total["triangles"] * copies
            active["primitive_position_rows"] += total["primitive_position_rows"] * copies
            active["primitive_copies"] += total["primitive_count"] * copies
            scene_materials.update(row["material"] for row in rows
                                   if row["mesh"] == mi and row["material"] is not None)
        active["referenced_materials"] = sorted(scene_materials)
    return {"inspector_version": VERSION, "byte_length": len(data),
            "sha256": hashlib.sha256(data).hexdigest(),
            "extensions_used": d.get("extensionsUsed", []),
            "extensions_required": d.get("extensionsRequired", []),
            "stored": {"meshes": len(meshes), "primitives": len(rows),
                       "assembled_triangles": sum(x["assembled_triangles"] for x in rows),
                       "primitive_position_rows": sum(x["position_rows"] for x in rows),
                       "unique_position_accessor_rows": sum(reader.accessor(i)["count"] for i in all_position_accessors),
                       "referenced_materials": sorted(used_materials),
                       "declared_materials": len(d.get("materials", []))},
            "selected_scene": active, "per_primitive": rows, "per_mesh": mesh_totals,
            "limitations": [
                "Base geometry only; no animation, skinning, morph, texture or memory-cost evaluation.",
                "Assembled triangle slots include degenerate triangles; no area test is performed.",
                "Position rows are attribute entries, not unique topological vertices or vertex-shader invocations.",
                "Distinct positions use exact local float tuples; matching coordinates do not prove vertices may be welded.",
                "Per-primitive row totals can count a shared POSITION accessor more than once.",
                "Scene counts include every selected-scene copy before culling, runtime LOD, batching or render passes.",
                "Primitive counts are not measured draw calls; expanded rows are not allocated GPU memory.",
                "No default scene means selected_scene=null; pass --scene explicitly if desired.",
                "Not a full glTF validator; run Khronos glTF Validator separately.",
            ]}


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("file", type=Path)
    parser.add_argument("--scene", type=int, help="Explicit scene index; otherwise use declared default")
    args = parser.parse_args()
    try:
        require(args.file.stat().st_size <= MAX_BYTES, "File exceeds 64 MiB safety limit")
        result = inspect_bytes(args.file.read_bytes(), args.scene)
        print(json.dumps(result, indent=2, allow_nan=False))
    except (InspectionError, OSError, ValueError, KeyError, TypeError, AttributeError, RecursionError, struct.error) as error:
        print(f"Inspection refused: {error}", file=sys.stderr)
        return 2
    return 0


if __name__ == "__main__":
    sys.exit(main())
