#!/usr/bin/env python3
"""Render a high-key silver reference from the exact web GLB.

The reference deliberately changes only shading, lighting, camera, and color
management. It is used to decide whether the remaining defects belong to the
mesh or to the realtime web renderer.
"""

from __future__ import annotations

import argparse
import math
import sys
from pathlib import Path

import bpy
from mathutils import Vector


def arguments() -> 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)
    parser.add_argument("--samples", type=int, default=256)
    return parser.parse_args(argv)


def clear_scene() -> None:
    bpy.ops.object.select_all(action="SELECT")
    bpy.ops.object.delete(use_global=False)
    for datablocks in (
        bpy.data.meshes,
        bpy.data.curves,
        bpy.data.materials,
        bpy.data.cameras,
        bpy.data.lights,
    ):
        for datablock in list(datablocks):
            if datablock.users == 0:
                datablocks.remove(datablock)


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_mesh(path: Path) -> bpy.types.Object:
    bpy.ops.import_scene.gltf(filepath=str(path))
    meshes = [obj for obj in bpy.context.scene.objects if obj.type == "MESH"]
    if not meshes:
        raise RuntimeError(f"No mesh found in {path}")

    bpy.ops.object.select_all(action="DESELECT")
    for obj in meshes:
        obj.select_set(True)
    bpy.context.view_layer.objects.active = meshes[0]
    if len(meshes) > 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 object_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 make_principled_material(
    name: str,
    base_color: tuple[float, float, float, float],
    metallic: float,
    roughness: float,
) -> bpy.types.Material:
    material = bpy.data.materials.new(name)
    material.use_nodes = True
    bsdf = material.node_tree.nodes.get("Principled BSDF")
    bsdf.inputs["Base Color"].default_value = base_color
    bsdf.inputs["Metallic"].default_value = metallic
    bsdf.inputs["Roughness"].default_value = roughness
    return material


def add_area(
    name: str,
    location: Vector,
    target: Vector,
    energy: float,
    size: float,
    color: tuple[float, float, float],
    shape: str = "RECTANGLE",
    size_y: float | None = None,
) -> bpy.types.Object:
    data = bpy.data.lights.new(name=name, type="AREA")
    data.energy = energy
    data.color = color
    data.shape = shape
    data.size = size
    if size_y is not None:
        data.size_y = size_y
    light = bpy.data.objects.new(name, data)
    bpy.context.collection.objects.link(light)
    light.location = location
    point_at(light, target)
    return light


def add_reflection_card(
    name: str,
    location: Vector,
    target: Vector,
    scale: tuple[float, float, float],
    value: float,
) -> bpy.types.Object:
    bpy.ops.mesh.primitive_plane_add(size=2.0, location=location)
    card = bpy.context.object
    card.name = name
    card.scale = scale
    point_at(card, target)
    material = make_principled_material(
        f"{name} material", (value, value, value, 1.0), 0.0, 0.78
    )
    card.data.materials.append(material)
    return card


def configure_cycles(samples: int) -> None:
    scene = bpy.context.scene
    scene.render.engine = "BLENDER_EEVEE"

    # Prefer Cycles when the headless NVIDIA device is available. Eevee remains
    # a deterministic fallback rather than failing the whole research render.
    try:
        preferences = bpy.context.preferences.addons["cycles"].preferences
        for compute_type in ("OPTIX", "CUDA"):
            try:
                preferences.compute_device_type = compute_type
                preferences.get_devices()
                enabled = False
                for device in preferences.devices:
                    device.use = device.type in {"OPTIX", "CUDA"}
                    enabled = enabled or device.use
                if enabled:
                    scene.render.engine = "CYCLES"
                    scene.cycles.device = "GPU"
                    break
            except Exception:
                continue
    except Exception as error:
        print(f"Cycles device discovery failed; using Eevee: {error}")

    if scene.render.engine == "CYCLES":
        scene.cycles.samples = samples
        scene.cycles.use_denoising = True
        scene.cycles.preview_samples = min(64, samples)
        scene.cycles.max_bounces = 8
        scene.cycles.glossy_bounces = 6
        scene.cycles.transparent_max_bounces = 6
        scene.cycles.use_adaptive_sampling = True
        scene.cycles.adaptive_threshold = 0.008
    else:
        scene.render.image_settings.color_depth = "16"
        scene.render.film_transparent = False

    print(f"Reference engine: {scene.render.engine}")


def configure_scene(obj: bpy.types.Object, output: Path, samples: int) -> None:
    scene = bpy.context.scene
    configure_cycles(samples)
    scene.render.resolution_x = 1200
    scene.render.resolution_y = 1200
    scene.render.resolution_percentage = 100
    scene.render.image_settings.file_format = "PNG"
    scene.render.image_settings.color_mode = "RGB"
    scene.render.image_settings.color_depth = "16"
    scene.render.image_settings.compression = 38
    scene.render.filepath = str(output)
    scene.render.film_transparent = False
    scene.render.use_file_extension = True

    scene.view_settings.look = "AgX - Medium High Contrast"
    scene.view_settings.exposure = -0.15
    scene.view_settings.gamma = 1.0

    world = scene.world or bpy.data.worlds.new("TN high-key world")
    scene.world = world
    world.use_nodes = True
    background = world.node_tree.nodes.get("Background")
    background.inputs["Color"].default_value = (0.72, 0.75, 0.76, 1.0)
    background.inputs["Strength"].default_value = 0.26

    obj.data.materials.clear()
    obj.data.materials.append(
        make_principled_material(
            "TN polished sterling silver", (0.955, 0.973, 0.99, 1.0), 1.0, 0.105
        )
    )

    minimum, maximum, center, radius = object_bounds(obj)

    camera_data = bpy.data.cameras.new("Reference camera")
    camera = bpy.data.objects.new("Reference camera", camera_data)
    bpy.context.collection.objects.link(camera)
    camera.data.lens = 78
    camera.data.sensor_width = 36
    camera.data.dof.use_dof = False
    camera.location = center + Vector((0.05, -2.6, 2.08)) * radius
    point_at(camera, center + Vector((0.0, -0.04, -0.04)) * radius)
    scene.camera = camera

    floor_z = minimum.z - 0.055 * radius
    bpy.ops.mesh.primitive_plane_add(size=radius * 12.0, location=(center.x, center.y, floor_z))
    floor = bpy.context.object
    floor.name = "Warm white sweep"
    floor.data.materials.append(
        make_principled_material("Warm white floor", (0.78, 0.79, 0.78, 1.0), 0.0, 0.36)
    )

    target = center + Vector((0.0, 0.0, 0.02)) * radius
    add_area(
        "Large left strip",
        center + Vector((-2.6, -2.2, 3.35)) * radius,
        target,
        250.0,
        radius * 3.3,
        (0.94, 0.97, 1.0),
        size_y=radius * 0.65,
    )
    add_area(
        "Large right strip",
        center + Vector((2.85, -1.0, 2.35)) * radius,
        target,
        220.0,
        radius * 2.8,
        (1.0, 0.955, 0.90),
        size_y=radius * 0.52,
    )
    add_area(
        "Top silk",
        center + Vector((0.1, 0.5, 4.4)) * radius,
        target,
        310.0,
        radius * 3.9,
        (0.94, 0.97, 1.0),
        shape="DISK",
    )
    add_area(
        "Front fill",
        center + Vector((0.0, -4.1, 0.35)) * radius,
        target,
        115.0,
        radius * 2.25,
        (1.0, 0.97, 0.93),
        shape="DISK",
    )

    # Two controlled dark reflections separate the polished edges without
    # turning the whole product charcoal.
    add_reflection_card(
        "Left edge card",
        center + Vector((-2.65, -0.1, 1.05)) * radius,
        target,
        (radius * 0.2, radius * 1.65, 1.0),
        0.006,
    )
    add_reflection_card(
        "Right edge card",
        center + Vector((2.7, 0.85, 1.2)) * radius,
        target,
        (radius * 0.17, radius * 1.45, 1.0),
        0.012,
    )

    if scene.render.engine == "CYCLES":
        for mesh in scene.objects:
            if mesh.type == "MESH":
                mesh.visible_shadow = True

    maximum_dimension = max(maximum.x - minimum.x, maximum.y - minimum.y, maximum.z - minimum.z)
    print(
        f"Web reference mesh: {len(obj.data.vertices):,} vertices, "
        f"{len(obj.data.polygons):,} faces, extent {maximum_dimension:.6f}"
    )


def main() -> None:
    args = arguments()
    input_path = Path(args.input).resolve()
    output_path = Path(args.output).resolve()
    output_path.parent.mkdir(parents=True, exist_ok=True)
    clear_scene()
    obj = import_mesh(input_path)
    configure_scene(obj, output_path, args.samples)
    bpy.ops.render.render(write_still=True)
    print(f"Reference written: {output_path}")


if __name__ == "__main__":
    main()
