#!/usr/bin/env python3
"""Boilerplate generators for DCTL files (DaVinci Color Transform Language).

DCTL is a C-like, GPU-compiled pixel shader language used by Resolve as
programmable color transforms, the ResolveFX DCTL plugin, DCTL transitions,
and custom ACES IDT/ODT transforms. See docs/notes/dctl-notes.md for the full
spec; this module just emits ready-to-install source.

After installing a DCTL, call project_settings(action='refresh_luts') to make
Resolve re-scan its LUT directory. ACES DCTLs are scanned only at startup.
"""

from typing import Any, Dict, List, Optional


def header(name: str, kind: str) -> str:
    """Marker comment used by the dctl tool's `list` action."""
    return (
        f"// @mcp-dctl name={name} kind={kind}\n"
        "// Generated by davinci-resolve-mcp. The marker above only affects\n"
        "// how `dctl list` reports this file.\n\n"
    )


# UI-parameter type → pretty-printed DCTL keyword.
DCTL_UI_TYPES = {
    "float": "DCTLUI_SLIDER_FLOAT",
    "int": "DCTLUI_SLIDER_INT",
    "value": "DCTLUI_VALUE_BOX",
    "checkbox": "DCTLUI_CHECK_BOX",
    "combo": "DCTLUI_COMBO_BOX",
    "color": "DCTLUI_COLOR_PICKER",
}


def _render_ui_params(params: List[Dict[str, Any]]) -> str:
    """Render a list of param dicts as a `DEFINE_UI_PARAMS` block per param.

    Each dict supports:
        name   (required) — identifier referenced inside transform()
        label  (optional, defaults to name)
        type   'float'|'int'|'value'|'checkbox'|'combo'|'color' (default 'float')
        default, min, max, step  numeric defaults appropriate per type
        tooltip  optional string
    """
    if not params:
        return ""
    lines = []
    for prm in params:
        pname = prm["name"]
        label = prm.get("label", pname)
        ptype = prm.get("type", "float")
        if ptype not in DCTL_UI_TYPES:
            raise ValueError(f"Invalid DCTL UI type '{ptype}'. "
                             f"Valid: {sorted(DCTL_UI_TYPES.keys())}")
        kw = DCTL_UI_TYPES[ptype]

        if ptype in ("float", "int", "value"):
            default = prm.get("default", 1.0)
            pmin = prm.get("min", 0.0)
            pmax = prm.get("max", 4.0)
            step = prm.get("step", 0.01 if ptype != "int" else 1)
            lines.append(
                f"DEFINE_UI_PARAMS({pname}, {label}, {kw}, "
                f"{default}, {pmin}, {pmax}, {step})"
            )
        elif ptype == "checkbox":
            default = int(bool(prm.get("default", 0)))
            lines.append(f"DEFINE_UI_PARAMS({pname}, {label}, {kw}, {default})")
        elif ptype == "combo":
            default = prm.get("default", 0)
            options = prm.get("options", [])
            opts = ", ".join(options)
            lines.append(f"DEFINE_UI_PARAMS({pname}, {label}, {kw}, {default}, {opts})")
        elif ptype == "color":
            r = prm.get("r", 1.0)
            g = prm.get("g", 1.0)
            b = prm.get("b", 1.0)
            lines.append(f"DEFINE_UI_PARAMS({pname}, {label}, {kw}, {r}, {g}, {b})")

        if "tooltip" in prm:
            lines.append(f'DEFINE_UI_TOOLTIP({pname}, "{prm["tooltip"]}")')
    return "\n".join(lines) + "\n\n"


def transform(name: str, options: Optional[Dict[str, Any]] = None) -> str:
    """Per-pixel transform DCTL with optional UI sliders.

    options:
        params: list of UI param dicts (see _render_ui_params).
                Default exposes a single Gain slider.
        body:   transform body. Receives p_R, p_G, p_B and any UI params by
                their declared name. Default: applies Gain uniformly.

    The generated DCTL uses the float3 entry point (no alpha).
    """
    options = options or {}
    params = options.get("params")
    if params is None:
        params = [{"name": "p_Gain", "label": "Gain", "type": "float",
                   "default": 1.0, "min": 0.0, "max": 4.0, "step": 0.01}]
    body = options.get("body") or (
        "    return make_float3(p_R * p_Gain, p_G * p_Gain, p_B * p_Gain);"
    )
    ui_block = _render_ui_params(params)
    return header(name, "transform") + f'''{ui_block}__DEVICE__ float3 transform(int p_Width, int p_Height, int p_X, int p_Y,
                              float p_R, float p_G, float p_B)
{{
{body}
}}
'''


def transform_alpha(name: str, options: Optional[Dict[str, Any]] = None) -> str:
    """Per-pixel transform DCTL with alpha (Resolve 19.1+).

    options:
        params:     list of UI param dicts (default: empty)
        alpha_mode: 'straight' | 'premultiply' (default 'straight')
        body:       transform body returning a float4. Default passes through.

    Sets the appropriate `DEFINE_DCTL_ALPHA_MODE_*` tag.
    """
    options = options or {}
    params = options.get("params") or []
    alpha_mode = options.get("alpha_mode", "straight")
    if alpha_mode not in ("straight", "premultiply"):
        raise ValueError(f"Invalid alpha_mode '{alpha_mode}'. "
                         "Valid: straight, premultiply")
    mode_tag = ("DEFINE_DCTL_ALPHA_MODE_STRAIGHT"
                if alpha_mode == "straight"
                else "DEFINE_DCTL_ALPHA_MODE_PREMULTIPLY")
    body = options.get("body") or (
        "    return make_float4(p_R, p_G, p_B, p_A);"
    )
    ui_block = _render_ui_params(params)
    return header(name, "transform_alpha") + f'''{ui_block}{mode_tag}

__DEVICE__ float4 transform(int p_Width, int p_Height, int p_X, int p_Y,
                              float p_R, float p_G, float p_B, float p_A)
{{
{body}
}}
'''


def transition(name: str, options: Optional[Dict[str, Any]] = None) -> str:
    """DCTL transition. Blends From and To clips using TRANSITION_PROGRESS.

    options:
        body: transition body returning a float4. Default linear cross-dissolve.

    The generated DCTL uses texture-based From and To inputs and reads the
    global TRANSITION_PROGRESS value (0.0 to 1.0).
    """
    options = options or {}
    body = options.get("body") or '''    float fromR = _tex2D(p_FromTexR, p_X, p_Y);
    float fromG = _tex2D(p_FromTexG, p_X, p_Y);
    float fromB = _tex2D(p_FromTexB, p_X, p_Y);
    float fromA = _tex2D(p_FromTexA, p_X, p_Y);
    float toR = _tex2D(p_ToTexR, p_X, p_Y);
    float toG = _tex2D(p_ToTexG, p_X, p_Y);
    float toB = _tex2D(p_ToTexB, p_X, p_Y);
    float toA = _tex2D(p_ToTexA, p_X, p_Y);

    float t = TRANSITION_PROGRESS;
    return make_float4(
        fromR * (1.0f - t) + toR * t,
        fromG * (1.0f - t) + toG * t,
        fromB * (1.0f - t) + toB * t,
        fromA * (1.0f - t) + toA * t);'''
    return header(name, "transition") + f'''__DEVICE__ float4 transition(
    int p_Width, int p_Height, int p_X, int p_Y,
    __TEXTURE__ p_FromTexR, __TEXTURE__ p_FromTexG,
    __TEXTURE__ p_FromTexB, __TEXTURE__ p_FromTexA,
    __TEXTURE__ p_ToTexR, __TEXTURE__ p_ToTexG,
    __TEXTURE__ p_ToTexB, __TEXTURE__ p_ToTexA)
{{
{body}
}}
'''


def matrix(name: str, options: Optional[Dict[str, Any]] = None) -> str:
    """3x3 color matrix transform.

    options:
        matrix: nested 3x3 list of floats. Default = identity.
                Order: [[Rr, Rg, Rb], [Gr, Gg, Gb], [Br, Bg, Bb]]

    Embeds the matrix as constants and applies it as `out = M * in`.
    """
    options = options or {}
    m = options.get("matrix") or [[1.0, 0.0, 0.0],
                                  [0.0, 1.0, 0.0],
                                  [0.0, 0.0, 1.0]]
    if (not isinstance(m, list) or len(m) != 3
            or any(not isinstance(row, list) or len(row) != 3 for row in m)):
        raise ValueError("matrix option must be a 3x3 nested list")
    rows = ",\n        ".join(
        f"{{{', '.join(_f(x) for x in row)}}}" for row in m
    )
    return header(name, "matrix") + f'''__CONSTANT__ float M[3][3] = {{
        {rows}
}};

__DEVICE__ float3 transform(int p_Width, int p_Height, int p_X, int p_Y,
                              float p_R, float p_G, float p_B)
{{
    float r = M[0][0] * p_R + M[0][1] * p_G + M[0][2] * p_B;
    float g = M[1][0] * p_R + M[1][1] * p_G + M[1][2] * p_B;
    float b = M[2][0] * p_R + M[2][1] * p_G + M[2][2] * p_B;
    return make_float3(r, g, b);
}}
'''


def lut_apply(name: str, options: Optional[Dict[str, Any]] = None) -> str:
    """DCTL that loads an external .cube LUT and applies it.

    options:
        lut_path: relative or absolute path to a .cube file. The path is
                  resolved relative to the .dctl file's location.
                  Default: 'YourLut.cube' (placeholder — user must edit).
        params:   optional UI param list (mix amount, etc.).
        body:     transform body. Default applies the LUT directly.

    Reference: docs/notes/dctl-notes.md → "LUTs Inside DCTL". Use APPLY_LUT(r, g, b,
    name) to apply either an external cube or an inline DEFINE_CUBE_LUT.
    """
    options = options or {}
    lut_path = options.get("lut_path", "YourLut.cube")
    params = options.get("params") or [
        {"name": "p_Mix", "label": "Mix", "type": "float",
         "default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01}
    ]
    body = options.get("body") or '''    float lutR = p_R, lutG = p_G, lutB = p_B;
    APPLY_LUT(lutR, lutG, lutB, ExternalLut);
    return make_float3(
        p_R + (lutR - p_R) * p_Mix,
        p_G + (lutG - p_G) * p_Mix,
        p_B + (lutB - p_B) * p_Mix);'''
    ui_block = _render_ui_params(params)
    return header(name, "lut_apply") + f'''DEFINE_LUT(ExternalLut, {lut_path})

{ui_block}__DEVICE__ float3 transform(int p_Width, int p_Height, int p_X, int p_Y,
                              float p_R, float p_G, float p_B)
{{
{body}
}}
'''


def aces_idt(name: str, options: Optional[Dict[str, Any]] = None) -> str:
    """ACES Input Device Transform. Installs to ACES Transforms/IDT/.

    options:
        parametric: bool — if true, emits the V1 parametric template stub
                    (DEFINE_ACES_PARAM). If false, non-parametric stub.
                    Default false.
        body:       transform body. Default passes through.

    ACES DCTLs are scanned at Resolve startup, NOT through RefreshLUTList().
    Install requires a Resolve restart. See docs/notes/dctl-notes.md → "DCTL And ACES".
    """
    options = options or {}
    parametric = bool(options.get("parametric", False))
    body = options.get("body") or (
        "    return make_float3(p_R, p_G, p_B);"
    )
    aces_def = (
        "DEFINE_ACES_PARAM(IS_PARAMETRIC_ACES_TRANSFORM: 1)"
        if parametric
        else "DEFINE_ACES_PARAM(IS_PARAMETRIC_ACES_TRANSFORM: 0)"
    )
    return header(name, "aces_idt") + f'''{aces_def}

__DEVICE__ float3 transform(int p_Width, int p_Height, int p_X, int p_Y,
                              float p_R, float p_G, float p_B)
{{
{body}
}}
'''


def aces_odt(name: str, options: Optional[Dict[str, Any]] = None) -> str:
    """ACES Output Device Transform. Installs to ACES Transforms/ODT/.

    Identical structure to aces_idt; the directory placement is what tells
    Resolve which side of the pipeline this transform belongs to.
    """
    options = options or {}
    parametric = bool(options.get("parametric", False))
    body = options.get("body") or (
        "    return make_float3(p_R, p_G, p_B);"
    )
    aces_def = (
        "DEFINE_ACES_PARAM(IS_PARAMETRIC_ACES_TRANSFORM: 1)"
        if parametric
        else "DEFINE_ACES_PARAM(IS_PARAMETRIC_ACES_TRANSFORM: 0)"
    )
    return header(name, "aces_odt") + f'''{aces_def}

__DEVICE__ float3 transform(int p_Width, int p_Height, int p_X, int p_Y,
                              float p_R, float p_G, float p_B)
{{
{body}
}}
'''


def kernel(name: str, options: Optional[Dict[str, Any]] = None) -> str:
    """Bare-bones DCTL with structured TODO comments.

    Useful when the user wants to write the DCTL by hand but wants the
    boilerplate (header, signature, alpha note) generated for them.
    """
    options = options or {}
    return header(name, "kernel") + f'''// TODO: Replace with your transform implementation.
// Reference: docs/notes/dctl-notes.md
//
// Useful globals:
//   __RESOLVE_VER_MAJOR__, __RESOLVE_VER_MINOR__
//   DEVICE_IS_CUDA / DEVICE_IS_OPENCL / DEVICE_IS_METAL
//   TIMELINE_FRAME_INDEX  (defaults to 1 when used as a LUT)
//
// Float literals must use the `f` suffix: 1.2f not 1.2

__DEVICE__ float3 transform(int p_Width, int p_Height, int p_X, int p_Y,
                              float p_R, float p_G, float p_B)
{{
    // TODO
    return make_float3(p_R, p_G, p_B);
}}
'''


def _f(x: float) -> str:
    """Format a float with the trailing 'f' suffix DCTL requires."""
    s = f"{float(x):.6g}"
    if "." not in s and "e" not in s and "n" not in s:
        s += ".0"
    return s + "f"


TEMPLATES = {
    "transform": transform,
    "transform_alpha": transform_alpha,
    "transition": transition,
    "matrix": matrix,
    "kernel": kernel,
    "lut_apply": lut_apply,
    "aces_idt": aces_idt,
    "aces_odt": aces_odt,
}

# Maps each template kind to the install category, which the dctl tool uses
# to pick the right install directory (regular LUT vs. ACES Transforms).
KIND_CATEGORY = {
    "transform": "lut",
    "transform_alpha": "lut",
    "transition": "lut",
    "matrix": "lut",
    "kernel": "lut",
    "lut_apply": "lut",
    "aces_idt": "aces_idt",
    "aces_odt": "aces_odt",
}
