#!/usr/bin/env python3
"""Render consistent studio and zebra-reflection diagnostics for GLB candidates."""

from __future__ import annotations

import argparse
import math
import sys
from pathlib import Path

import bpy
from mathutils import Vector


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("--outdir", required=True)
    parser.add_argument("--label", required=True)
    parser.add_argument("--environment", required=True)
    parser.add_argument("--roughness", type=float, default=0.025)
    parser.add_argument("--exposure", type=float, default=0.15)
    parser.add_argument("--world-strength", type=float, default=0.82)
    return parser.parse_args(argv)


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


def point_at(obj: bpy.types.Object, target: Vector) -> None:
    obj.rotation_euler = (target - obj.location).to_track_quat("-Z", "Y").to_euler()


def import_object(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 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)
    for polygon in obj.data.polygons:
        polygon.use_smooth = True
    return obj


def bounds(obj: bpy.types.Object) -> tuple[Vector, Vector, Vector, float]:
    points = [obj.matrix_world @ Vector(corner) for corner in obj.bound_box]
    minimum = Vector(tuple(min(point[i] for point in points) for i in range(3)))
    maximum = Vector(tuple(max(point[i] for point in points) for i in range(3)))
    center = (minimum + maximum) * 0.5
    radius = max((maximum - minimum).length * 0.5, 0.1)
    return minimum, maximum, center, radius


def silver_material(roughness: float) -> bpy.types.Material:
    material = bpy.data.materials.new("Diagnostic mirror silver")
    material.use_nodes = True
    bsdf = material.node_tree.nodes.get("Principled BSDF")
    bsdf.inputs["Base Color"].default_value = (0.96, 0.97, 1.0, 1.0)
    bsdf.inputs["Metallic"].default_value = 1.0
    bsdf.inputs["Roughness"].default_value = roughness
    return material


def set_world(environment: Path, strength: float) -> None:
    world = bpy.context.scene.world
    world.use_nodes = True
    nodes = world.node_tree.nodes
    links = world.node_tree.links
    nodes.clear()
    texture = nodes.new("ShaderNodeTexEnvironment")
    texture.image = bpy.data.images.load(str(environment), check_existing=True)
    texture.interpolation = "Linear"
    background = nodes.new("ShaderNodeBackground")
    background.inputs["Strength"].default_value = strength
    output = nodes.new("ShaderNodeOutputWorld")
    links.new(texture.outputs["Color"], background.inputs["Color"])
    links.new(background.outputs["Background"], output.inputs["Surface"])


def setup_scene(
    obj: bpy.types.Object,
    environment: Path,
    roughness: float,
    exposure: float,
    world_strength: float,
) -> None:
    scene = bpy.context.scene
    scene.render.engine = "BLENDER_EEVEE"
    scene.render.resolution_x = 900
    scene.render.resolution_y = 900
    scene.render.resolution_percentage = 100
    scene.render.image_settings.file_format = "PNG"
    scene.render.image_settings.color_mode = "RGB"
    scene.render.film_transparent = False
    scene.render.image_settings.color_depth = "8"
    scene.render.image_settings.color_mode = "RGB"
    scene.render.image_settings.compression = 24
    scene.render.use_file_extension = True
    scene.render.film_transparent = False
    scene.render.image_settings.color_mode = "RGB"
    scene.view_settings.look = "AgX - Medium High Contrast"
    scene.view_settings.exposure = exposure
    set_world(environment, world_strength)

    obj.data.materials.clear()
    obj.data.materials.append(silver_material(roughness))
    minimum, _, center, radius = bounds(obj)

    camera_data = bpy.data.cameras.new("Diagnostic Camera")
    camera = bpy.data.objects.new("Diagnostic Camera", camera_data)
    bpy.context.collection.objects.link(camera)
    camera.data.lens = 72
    camera.data.sensor_width = 36
    scene.camera = camera

    plane_size = radius * 8.0
    bpy.ops.mesh.primitive_plane_add(
        size=plane_size, location=(center.x, center.y, minimum.z - 0.045 * radius)
    )
    floor = bpy.context.object
    floor_material = bpy.data.materials.new("Neutral diagnostic floor")
    floor_material.use_nodes = True
    floor_bsdf = floor_material.node_tree.nodes.get("Principled BSDF")
    floor_bsdf.inputs["Base Color"].default_value = (0.38, 0.40, 0.44, 1.0)
    floor_bsdf.inputs["Roughness"].default_value = 0.32
    floor.data.materials.append(floor_material)


def render_views(obj: bpy.types.Object, output: Path, label: str) -> None:
    scene = bpy.context.scene
    _, _, center, radius = bounds(obj)
    camera = scene.camera
    views = {
        "hero": (Vector((0.0, -2.45, 2.05)), Vector((0.0, -0.06, -0.03))),
        "front": (Vector((0.0, -2.75, 0.48)), Vector((0.0, 0.0, -0.02))),
        "side": (Vector((2.55, -0.30, 1.15)), Vector((0.0, 0.0, -0.02))),
        "under": (Vector((0.0, -2.35, -1.30)), Vector((0.0, 0.0, -0.02))),
    }
    output.mkdir(parents=True, exist_ok=True)
    for view_name, (location_scale, target_scale) in views.items():
        camera.location = center + location_scale * radius
        target = center + target_scale * radius
        point_at(camera, target)
        scene.render.filepath = str(output / f"{label}-{view_name}.png")
        bpy.ops.render.render(write_still=True)


def main() -> None:
    args = parse_args()
    clear_scene()
    obj = import_object(Path(args.input).resolve())
    setup_scene(
        obj,
        Path(args.environment).resolve(),
        args.roughness,
        args.exposure,
        args.world_strength,
    )
    render_views(obj, Path(args.outdir).resolve(), args.label)


if __name__ == "__main__":
    main()
