#!/usr/bin/env python3
"""Synthetic accounting fixtures. No generation outputs or product-quality benchmark.
Run: python3 test-budget.py
Optional: python3 test-budget.py --write-fixtures ./synthetic-fixtures
"""
import copy
import json
import struct
import sys
import unittest
from pathlib import Path
from inspect_glb import InspectionError, inspect_bytes, parse_glb, Reader


def pack(document, binary, uri=None):
    d = copy.deepcopy(document)
    d["buffers"] = [{"byteLength": len(binary)}]
    if uri is not None:
        d["buffers"][0]["uri"] = uri
    text = json.dumps(d, separators=(",", ":")).encode()
    text += b" " * (-len(text) % 4)
    binary += b"\0" * (-len(binary) % 4)
    return (struct.pack("<4sII", b"glTF", 2, 28 + len(text) + len(binary))
            + struct.pack("<II", len(text), 0x4E4F534A) + text
            + struct.pack("<II", len(binary), 0x004E4942) + binary)


def fixture(indexed=True, mode=4, stride=False):
    positions = [(0, 0, 0), (1, 0, 0), (1, 1, 0), (0, 1, 0)]
    indices = [0, 1, 2, 0, 2, 3]
    if mode in (5, 6):
        indices = [0, 1, 2, 3]
    if not indexed:
        positions = [positions[i] for i in indices]
    blob = b"HEAD" if stride else b""
    for position in positions:
        blob += struct.pack("<3f", *position) + (b"PADD" if stride else b"")
    view = {"buffer": 0, "byteLength": len(blob)}
    if stride:
        view["byteStride"] = 16
    accessor = {"bufferView": 0, "byteOffset": 4 if stride else 0,
                "componentType": 5126, "count": len(positions), "type": "VEC3",
                "min": [0, 0, 0], "max": [1, 1, 0]}
    d = {"asset": {"version": "2.0", "generator": "Goblin3D synthetic accounting fixture"},
         "scene": 0, "scenes": [{"nodes": [0]}], "nodes": [{"mesh": 0}],
         "meshes": [{"primitives": [{"attributes": {"POSITION": 0}, "mode": mode, "material": 0}]}],
         "materials": [{"name": "synthetic"}], "accessors": [accessor], "bufferViews": [view]}
    if indexed:
        offset = len(blob)
        blob += struct.pack("<" + "H" * len(indices), *indices)
        d["bufferViews"].append({"buffer": 0, "byteOffset": offset, "byteLength": len(indices) * 2})
        d["accessors"].append({"bufferView": 1, "componentType": 5123,
                               "count": len(indices), "type": "SCALAR"})
        d["meshes"][0]["primitives"][0]["indices"] = 1
    return d, blob


def scenarios():
    d, b = fixture()
    result = {"indexed_quad": pack(d, b)}
    result["nonindexed_quad"] = pack(*fixture(indexed=False))
    result["interleaved_quad"] = pack(*fixture(stride=True))
    result["triangle_strip"] = pack(*fixture(mode=5))
    result["triangle_fan"] = pack(*fixture(mode=6))
    d, b = fixture()
    d["nodes"] = [{"mesh": 0}, {"mesh": 0, "translation": [2, 0, 0]}, {"mesh": 0}]
    d["scenes"] = [{"nodes": [0, 1]}, {"nodes": [2]}]
    result["two_selected_scene_copies"] = pack(d, b)
    d, b = fixture()
    d["meshes"][0]["primitives"].append(copy.deepcopy(d["meshes"][0]["primitives"][0]))
    d["meshes"][0]["primitives"][1]["material"] = 1
    d["materials"].append({"name": "second"})
    result["two_materials_shared_positions"] = pack(d, b)
    d, b = fixture()
    offset = len(b)
    b += struct.pack("<9f", 0, 0, 0, 2, 0, 0, 4, 0, 0)
    d["bufferViews"].append({"buffer": 0, "byteOffset": offset, "byteLength": 36})
    d["accessors"].append({"bufferView": 2, "componentType": 5126, "count": 3, "type": "VEC3"})
    d["extensionsUsed"] = d["extensionsRequired"] = ["EXT_mesh_gpu_instancing"]
    d["nodes"][0]["extensions"] = {"EXT_mesh_gpu_instancing": {"attributes": {"TRANSLATION": 2}}}
    result["three_gpu_instances"] = pack(d, b)
    return result


class InspectorTests(unittest.TestCase):
    def test_indexed_quad(self):
        r = inspect_bytes(scenarios()["indexed_quad"])
        self.assertEqual(r["stored"]["assembled_triangles"], 2)
        self.assertEqual(r["stored"]["primitive_position_rows"], 4)
        self.assertEqual(r["per_primitive"][0]["distinct_local_position_tuples"], 4)
        self.assertEqual(r["stored"]["referenced_materials"], [0])

    def test_distinct_normals_at_same_position(self):
        d, b = fixture(indexed=False)
        offset = len(b)
        normals = [(0, 0, 1)] * 3 + [(0, 1, 0)] * 3
        b += b"".join(struct.pack("<3f", *n) for n in normals)
        d["bufferViews"].append({"buffer": 0, "byteOffset": offset, "byteLength": 72})
        d["accessors"].append({"bufferView": 1, "componentType": 5126, "count": 6, "type": "VEC3"})
        d["meshes"][0]["primitives"][0]["attributes"]["NORMAL"] = 1
        evidence = inspect_bytes(pack(d, b))["per_primitive"][0]["normal_evidence"]
        self.assertEqual(evidence["distinct_normal_tuples"], 2)
        self.assertEqual(evidence["positions_with_multiple_normal_values"], 2)
        self.assertEqual(evidence["distinct_position_normal_pairs"], 6)

    def test_nonindexed(self):
        r = inspect_bytes(scenarios()["nonindexed_quad"])
        self.assertEqual(r["stored"]["assembled_triangles"], 2)
        self.assertEqual(r["stored"]["primitive_position_rows"], 6)
        self.assertEqual(r["per_primitive"][0]["distinct_local_position_tuples"], 4)

    def test_interleaved_and_accessor_offset(self):
        r = inspect_bytes(scenarios()["interleaved_quad"])
        self.assertEqual(r["per_primitive"][0]["distinct_local_position_tuples"], 4)
        d, b = parse_glb(scenarios()["interleaved_quad"])
        self.assertEqual(Reader(d, b).decode(0)[2], (1, 1, 0))

    def test_strip_and_fan(self):
        for key in ("triangle_strip", "triangle_fan"):
            self.assertEqual(inspect_bytes(scenarios()[key])["stored"]["assembled_triangles"], 2)

    def test_scene_instances(self):
        r = inspect_bytes(scenarios()["two_selected_scene_copies"])
        self.assertEqual(r["stored"]["assembled_triangles"], 2)
        self.assertEqual(r["selected_scene"]["assembled_triangles"], 4)
        self.assertEqual(r["selected_scene"]["mesh_copies"], 2)
        self.assertEqual(inspect_bytes(scenarios()["two_selected_scene_copies"], 1)["selected_scene"]["assembled_triangles"], 2)

    def test_gpu_instances_no_extra_base_copy(self):
        r = inspect_bytes(scenarios()["three_gpu_instances"])
        self.assertEqual(r["stored"]["assembled_triangles"], 2)
        self.assertEqual(r["selected_scene"]["assembled_triangles"], 6)
        self.assertEqual(r["selected_scene"]["mesh_nodes"], 1)
        self.assertEqual(r["selected_scene"]["mesh_copies"], 3)

    def test_shared_accessor_and_materials(self):
        r = inspect_bytes(scenarios()["two_materials_shared_positions"])
        self.assertEqual(r["stored"]["primitives"], 2)
        self.assertEqual(r["stored"]["primitive_position_rows"], 8)
        self.assertEqual(r["stored"]["unique_position_accessor_rows"], 4)
        self.assertEqual(r["per_mesh"][0]["distinct_local_position_tuples"], 4)
        self.assertEqual(r["stored"]["referenced_materials"], [0, 1])

    def test_missing_material_is_not_material_zero(self):
        d, b = fixture()
        del d["meshes"][0]["primitives"][0]["material"]
        r = inspect_bytes(pack(d, b))
        self.assertEqual(r["stored"]["referenced_materials"], [])
        self.assertIsNone(r["per_primitive"][0]["material"])

    def test_no_default_scene_is_not_scene_zero(self):
        d, b = fixture()
        del d["scene"]
        self.assertIsNone(inspect_bytes(pack(d, b))["selected_scene"])
        self.assertEqual(inspect_bytes(pack(d, b), 0)["selected_scene"]["assembled_triangles"], 2)

    def test_unused_mesh_not_counted_in_scene(self):
        d, b = fixture()
        d["meshes"].append(copy.deepcopy(d["meshes"][0]))
        r = inspect_bytes(pack(d, b))
        self.assertEqual(r["stored"]["assembled_triangles"], 4)
        self.assertEqual(r["selected_scene"]["assembled_triangles"], 2)

    def test_default_mode_and_lines(self):
        d, b = fixture()
        del d["meshes"][0]["primitives"][0]["mode"]
        self.assertEqual(inspect_bytes(pack(d, b))["stored"]["assembled_triangles"], 2)
        for mode in (0, 1, 2, 3):
            d["meshes"][0]["primitives"][0]["mode"] = mode
            self.assertEqual(inspect_bytes(pack(d, b))["stored"]["assembled_triangles"], 0)

    def test_bad_file_length(self):
        with self.assertRaises(InspectionError):
            inspect_bytes(scenarios()["indexed_quad"][:-1])

    def test_bad_triangle_count(self):
        d, b = fixture()
        d["accessors"][1]["count"] = 5
        with self.assertRaises(InspectionError):
            inspect_bytes(pack(d, b))

    def test_out_of_range_index(self):
        d, b = fixture()
        b = b[:-2] + struct.pack("<H", 9)
        with self.assertRaises(InspectionError):
            inspect_bytes(pack(d, b))

    def test_buffer_bounds(self):
        d, b = fixture()
        d["accessors"][0]["byteOffset"] = 400
        with self.assertRaises(InspectionError):
            inspect_bytes(pack(d, b))

    def test_scene_cycle(self):
        d, b = fixture()
        d["nodes"][0]["children"] = [0]
        with self.assertRaises(InspectionError):
            inspect_bytes(pack(d, b))

    def test_sparse_refusal(self):
        d, b = fixture()
        d["accessors"][0]["sparse"] = {}
        with self.assertRaises(InspectionError):
            inspect_bytes(pack(d, b))

    def test_compression_and_unknown_required_refusal(self):
        d, b = fixture()
        d["meshes"][0]["primitives"][0]["extensions"] = {"KHR_draco_mesh_compression": {}}
        with self.assertRaises(InspectionError):
            inspect_bytes(pack(d, b))
        d, b = fixture()
        d["extensionsRequired"] = ["EXAMPLE_unknown"]
        with self.assertRaises(InspectionError):
            inspect_bytes(pack(d, b))

    def test_external_buffer_refusal(self):
        d, b = fixture()
        with self.assertRaises(InspectionError):
            inspect_bytes(pack(d, b, uri="external.bin"))

    def test_quantization_refusal(self):
        d, b = fixture()
        d["extensionsRequired"] = ["KHR_mesh_quantization"]
        with self.assertRaises(InspectionError):
            inspect_bytes(pack(d, b))

    def test_degenerate_slots_are_not_removed(self):
        d, b = fixture()
        b = b[:-12] + struct.pack("<6H", 0, 0, 0, 1, 1, 1)
        self.assertEqual(inspect_bytes(pack(d, b))["stored"]["assembled_triangles"], 2)

    def test_attribute_count_mismatch(self):
        d, b = fixture()
        d["meshes"][0]["primitives"][0]["attributes"]["NORMAL"] = 1
        with self.assertRaises(InspectionError):
            inspect_bytes(pack(d, b))

    def test_gpu_instance_count_mismatch(self):
        d, b = parse_glb(scenarios()["three_gpu_instances"])
        d["nodes"][0]["extensions"]["EXT_mesh_gpu_instancing"]["attributes"]["_ID"] = 0
        with self.assertRaises(InspectionError):
            inspect_bytes(pack(d, b))


if __name__ == "__main__":
    if len(sys.argv) == 3 and sys.argv[1] == "--write-fixtures":
        output = Path(sys.argv[2])
        output.mkdir(parents=True, exist_ok=True)
        report = {}
        for name, blob in scenarios().items():
            (output / f"{name}.glb").write_bytes(blob)
            audit = inspect_bytes(blob)
            report[name] = {"sha256": audit["sha256"], "stored": audit["stored"],
                            "selected_scene": audit["selected_scene"]}
        (output / "fixture-results.json").write_text(json.dumps(report, indent=2) + "\n")
        print(f"Wrote {len(report)} synthetic fixtures and their accounting results")
    else:
        unittest.main(verbosity=2)
