#!/usr/bin/env python3
"""
Generate the pair of texture maps a `glitter` layer needs: a field of flat,
randomly tilted, randomly coloured flakes.

Writes TWO pixel-aligned RGBA PNGs. They must be used together — the shader
samples both at the same UV and pairs them by texel:

    <prefix>_normal.png   RGB = tangent-space flake normal, (128,128,255)
                                between flakes
                          A   = flake coverage, 0 between flakes
    <prefix>_color.png    RGB = that flake's own base hue
                          A   = that flake's thin-film thickness, 0..1, which
                                the shader lerps into glitterThicknessRange (nm)

The two channels in the alpha slots are what make the look: coverage decides
WHERE a flake is, thickness decides how its colour SHIFTS as the head turns.
Base hue and thickness are drawn independently per flake, so two neighbouring
flakes of the same hue still travel to different colours.

Flakes are FLAT — one constant normal across the whole flake — not domed. A flat
flake is either catching the light or it is not, which is what makes real glitter
twinkle rather than glow. `--dome` bends them toward a lens if you want a softer,
always-on shimmer instead.

The map TILES. The shader multiplies tUV by glitterScale and the textures are
loaded with RepeatWrapping, so a 1024 map repeated 3-4x over the lips gives far
more flakes than a 4096 one-shot at 1/16 the memory. Flakes that cross an edge
are drawn again on the opposite side so the seam is invisible.

This generator exists TWICE — here, and as demos/glitter-studio/glitter_gen.js,
so an artist can try a palette without opening a shell. That is only tolerable
while the two agree, which `python3 tools/check_glitter_gen.py` checks pixel for
pixel. Change one side, change the other, re-run it.

Agreeing at all required giving up numpy's RNG. default_rng is PCG64 and cannot
be reproduced in JS, so neither side uses its language's generator: both run the
mulberry32 in _Rng below, explicitly, drawing the same seven values per flake in
the same order. Do not "simplify" this back to np.random.

    python3 tools/make_glitter_map.py \
        --out-prefix demos/assets/imgs/face/glitter_festa \
        --palette '#ff3d8b,#ffd23d,#3dffd2,#8b3dff'
"""

import argparse
import colorsys
import math

from PIL import Image


M32 = 0xFFFFFFFF


class _Rng:
    """mulberry32, mirrored from demos/glitter-studio/glitter_gen.js.

    Kept as explicit 32-bit integer arithmetic rather than anything numpy so the
    JS port can be identical. Returns a float in [0, 1).
    """

    __slots__ = ('a',)

    def __init__(self, seed):
        self.a = seed & M32

    def next(self):
        self.a = (self.a + 0x6D2B79F5) & M32
        t = self.a
        t = ((t ^ (t >> 15)) * (t | 1)) & M32
        t = (t ^ ((t + (((t ^ (t >> 7)) * (t | 61)) & M32)) & M32)) & M32
        return ((t ^ (t >> 14)) & M32) / 4294967296.0


def rint(x):
    """numpy's round-half-to-even, written out so the JS port can match it."""
    f = math.floor(x)
    d = x - f
    if d > 0.5:
        return f + 1
    if d < 0.5:
        return f
    return f if f % 2 == 0 else f + 1


def clamp255(v):
    return 0 if v < 0 else (255 if v > 255 else int(v))


def parse_palette(spec):
    """'#rrggbb,#rrggbb' -> list of (r, g, b) in 0..1. 'random' -> None."""
    if spec.strip().lower() == 'random':
        return None
    out = []
    for tok in spec.split(','):
        tok = tok.strip().lstrip('#')
        if len(tok) != 6:
            raise SystemExit(f'bad palette entry {tok!r}, want #rrggbb')
        out.append(tuple(int(tok[i:i + 2], 16) / 255.0 for i in (0, 2, 4)))
    return out


def random_hues(rng, n, sat, val):
    """Saturated hues spread evenly around the wheel, then jittered.

    Drawn from the SAME stream as the flakes and before any flake draw, so the
    palette and the layout are not independent — reseed to change either.
    """
    out = []
    for i in range(n):
        h = (i / n + rng.next() * (1.0 / max(n, 1))) % 1.0
        out.append(colorsys.hsv_to_rgb(h, sat, val))
    return out


def main():
    ap = argparse.ArgumentParser(
        description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument('--out-prefix', required=True,
                    help='writes <prefix>_normal.png and <prefix>_color.png')
    ap.add_argument('--size', type=int, default=1024,
                    help='square edge in pixels. 1024 tiled beats 4096 flat')
    ap.add_argument('--count', type=int, default=6000,
                    help='number of flakes in one tile')
    ap.add_argument('--radius-px', type=float, default=3.0,
                    help='mean flake radius in pixels. Flakes smaller than a '
                         'screen pixel average away instead of twinkling')
    ap.add_argument('--radius-jitter', type=float, default=0.45,
                    help='fraction of --radius-px the size varies by')
    ap.add_argument('--slope', type=float, default=0.75,
                    help='max flake tilt as tan(angle): 0 = all flakes flat and '
                         'facing out (they all flash at once, which reads as a '
                         'sheet), 1.0 = 45 degrees')
    ap.add_argument('--dome', type=float, default=0.0,
                    help='0 = flat flakes (true twinkle), 1 = lens-shaped '
                         '(always catching something, softer)')
    ap.add_argument('--palette', default='random',
                    help="comma list of #rrggbb, or 'random'")
    ap.add_argument('--saturation', type=float, default=0.85,
                    help='for --palette random')
    ap.add_argument('--value', type=float, default=1.0,
                    help='for --palette random')
    ap.add_argument('--random-count', type=int, default=8,
                    help='how many hues --palette random invents')
    ap.add_argument('--thickness', nargs=2, type=float, default=[0.0, 1.0],
                    metavar=('LO', 'HI'),
                    help='range the per-flake film thickness is drawn from, as '
                         'a 0..1 position inside the layer glitterThicknessRange')
    ap.add_argument('--edge', type=float, default=0.35,
                    help='softness of the flake rim as a fraction of its radius. '
                         'Some softness is needed or the flakes alias')
    ap.add_argument('--seed', type=int, default=7)
    args = ap.parse_args()

    S = args.size
    rng = _Rng(args.seed)

    palette = parse_palette(args.palette)
    if palette is None:
        palette = random_hues(rng, args.random_count, args.saturation, args.value)
    if not palette:
        raise SystemExit('empty palette')

    th_lo = min(args.thickness)
    th_hi = max(args.thickness)

    # Coverage is a MAX, not a sum: overlapping flakes occlude each other, they
    # do not add up into a brighter blob.
    n = S * S
    cov = [0.0] * n
    nx = [0.0] * n
    ny = [0.0] * n
    col = [(0.0, 0.0, 0.0)] * n
    thk = [0.0] * n

    # Seven draws per flake, in this order. glitter_gen.js draws the same seven
    # in the same order; changing the order here changes every generated map.
    for _ in range(args.count):
        cx = rng.next() * S
        cy = rng.next() * S
        jitter = rng.next() * 2.0 - 1.0
        phi = rng.next() * 2.0 * math.pi
        mag_u = rng.next()
        hue_idx = int(rng.next() * len(palette))
        thick = th_lo + rng.next() * (th_hi - th_lo)

        r = args.radius_px * (1.0 + args.radius_jitter * jitter)
        if r < 0.6:
            r = 0.6

        # tan(theta) uniform in AREA, so the tilts spread instead of clustering
        # at the rim — the sqrt is doing real work, not cosmetics.
        mag = args.slope * math.sqrt(mag_u)
        fnx0 = mag * math.cos(phi)
        fny0 = mag * math.sin(phi)
        hue = palette[hue_idx]

        pad = math.ceil(r + 1.0)
        x0 = math.floor(cx) - pad
        x1 = math.floor(cx) + pad + 1
        y0 = math.floor(cy) - pad
        y1 = math.floor(cy) + pad + 1
        edge = max(args.edge * r, 0.5)
        inv_r = 1.0 / max(r, 1e-6)

        for y in range(y0, y1):
            dy = y + 0.5 - cy
            gy = y % S
            for x in range(x0, x1):
                dx = x + 0.5 - cx
                d = math.sqrt(dx * dx + dy * dy)
                a = (r - d) / edge
                a = 0.0 if a < 0 else (1.0 if a > 1 else a)
                a = a * a * (3.0 - 2.0 * a)
                if a <= 0.0:
                    continue
                gx = x % S
                k = gy * S + gx
                # nearer flake wins the texel; ties keep whoever got there first
                if a <= cov[k]:
                    continue
                cov[k] = a
                if args.dome > 0.0:
                    rr = min(d * inv_r, 1.0)
                    nx[k] = fnx0 + args.dome * (dx * inv_r) * rr
                    ny[k] = fny0 + args.dome * (dy * inv_r) * rr
                else:
                    nx[k] = fnx0
                    ny[k] = fny0
                col[k] = hue
                thk[k] = thick

    normal = bytearray(n * 4)
    color = bytearray(n * 4)
    inked = 0
    for k in range(n):
        # Between flakes the normal must read flat. The threshold is a whole
        # 8-bit step, not 0: a rim texel with coverage 0.002 quantises to alpha
        # 0, so the shader would never use its normal, and leaving a stray tilt
        # there only makes the map confusing to inspect.
        if cov[k] < 1.0 / 255.0:
            vx, vy, vz = 0.0, 0.0, 1.0
        else:
            vx, vy, vz = nx[k], ny[k], 1.0
            ln = math.sqrt(vx * vx + vy * vy + vz * vz)
            vx /= ln
            vy /= ln
            vz /= ln
        j = k * 4
        normal[j] = clamp255(rint((vx * 0.5 + 0.5) * 255.0))
        normal[j + 1] = clamp255(rint((vy * 0.5 + 0.5) * 255.0))
        normal[j + 2] = clamp255(rint((vz * 0.5 + 0.5) * 255.0))
        normal[j + 3] = clamp255(rint(cov[k] * 255.0))
        c = col[k]
        color[j] = clamp255(rint(c[0] * 255.0))
        color[j + 1] = clamp255(rint(c[1] * 255.0))
        color[j + 2] = clamp255(rint(c[2] * 255.0))
        color[j + 3] = clamp255(rint(thk[k] * 255.0))
        if cov[k] > 0.02:
            inked += 1

    n_path = f'{args.out_prefix}_normal.png'
    c_path = f'{args.out_prefix}_color.png'
    Image.frombytes('RGBA', (S, S), bytes(normal)).save(n_path)
    Image.frombytes('RGBA', (S, S), bytes(color)).save(c_path)

    hexes = ' '.join('#%02x%02x%02x' % tuple(clamp255(round(v * 255)) for v in p)
                     for p in palette)
    print(f'wrote {n_path}  ({S}x{S}, RGB=flake normal, A=coverage)')
    print(f'wrote {c_path}  ({S}x{S}, RGB=flake hue, A=film thickness)')
    print(f'  {args.count} flakes, mean radius {args.radius_px:.2f} px, '
          f'covering {100.0 * inked / n:.1f}% of the tile')
    print(f'  palette of {len(palette)}: {hexes}')
    print(f'  tilt up to {math.degrees(math.atan(args.slope)):.0f} deg, '
          f'dome {args.dome:.2f}, thickness {args.thickness[0]:.2f}..'
          f'{args.thickness[1]:.2f}, seed {args.seed}')
    print(f'  tiles: set glitter.scale so one tile spans ~1/3 of the region')


if __name__ == '__main__':
    main()
