#!/usr/bin/env python3
"""
Diff demos/glitter-studio/glitter_gen.js against tools/make_glitter_map.py.

The studio generates in the browser so an artist never has to leave the page,
which means the flake generator now exists twice. That is only acceptable while
the two agree: a palette tried in the studio and the PNG produced at the prompt
have to be the SAME map, or the download silently is not what was previewed.
This is what checks that.

Run from the repo root, after touching EITHER side:

    python3 tools/check_glitter_gen.py

Needs Pillow and node. Neither side is allowed its language's RNG — numpy's
PCG64 has no JS equivalent — so both run the same explicit mulberry32 and draw
the same seven values per flake in the same order. See the note at the top of
either file.

TOLERANCE IS ZERO, unlike check_strip_bake.py — there is no supersample collapse
and no accumulation-order difference here, and the one rounding trap
(round-half-to-EVEN vs Math.round's half-up) is written out explicitly on both
sides rather than delegated.

It is worth being precise about WHY zero is achievable, because it is not quite
"identical IEEE operations". `Math.cos`/`Math.sin` are not required to be
correctly rounded, and V8's fdlibm port does differ from glibc's: measured
6703/200000 cos and 6654/200000 sin values differ by 1 ULP between this machine's
python3 and node v22. Both sides call them (the flake tilt, `mag * cos(phi)`).
The maps still come out byte-identical because a 1-ULP double difference is ~1e-16
against an 8-bit quantisation step of 1/255, so crossing a rounding boundary would
take a value landing within 1e-16 of exactly x.5 — which does not happen here and
is checked empirically by every case below. So: treat a diff as a real divergence,
but if one ever appears as a SINGLE texel on one machine and not another, suspect
libm before suspecting the generator.

The Python side is deliberately pure-Python rather than numpy, for the same
reason — a vectorised inner loop would reorder the max-wins test and could not
be mirrored.
"""

import argparse
import json
import os
import subprocess
import sys
import tempfile

from PIL import Image

ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
GEN_JS = os.path.join(ROOT, 'demos', 'glitter-studio', 'glitter_gen.js')
GEN_PY = os.path.join(ROOT, 'tools', 'make_glitter_map.py')

# Small sizes on purpose: the pure-Python reference is O(count * radius^2) and
# these run in a second each. The generator has no size-dependent branch, so a
# 128 tile exercises exactly the code a 1024 one does.
#
# (label, python flags, js key=value options)
CASES = [
    ('default-ish',
     '--size 128 --count 120 --seed 7',
     'size=128 count=120 seed=7'),
    ('fixed palette',
     "--size 128 --count 150 --seed 11 --palette '#ff3d8b,#ffd23d,#3dffd2'",
     'size=128 count=150 seed=11 palette=#ff3d8b,#ffd23d,#3dffd2'),
    ('random palette, 5 hues',
     '--size 128 --count 150 --seed 3 --random-count 5 --saturation 0.7 --value 0.9',
     'size=128 count=150 seed=3 randomCount=5 saturation=0.7 value=0.9'),
    ('domed flakes',
     '--size 128 --count 120 --seed 5 --dome 0.8',
     'size=128 count=120 seed=5 dome=0.8'),
    ('flat, zero slope',
     '--size 128 --count 120 --seed 5 --slope 0.0',
     'size=128 count=120 seed=5 slope=0'),
    ('big soft flakes',
     '--size 128 --count 40 --seed 9 --radius-px 9 --edge 0.9 --radius-jitter 0.8',
     'size=128 count=40 seed=9 radiusPx=9 edge=0.9 radiusJitter=0.8'),
    ('tiny hard flakes',
     '--size 128 --count 400 --seed 2 --radius-px 0.7 --edge 0.05',
     'size=128 count=400 seed=2 radiusPx=0.7 edge=0.05'),
    ('narrow thickness band',
     '--size 128 --count 150 --seed 13 --thickness 0.4 0.45',
     'size=128 count=150 seed=13 thickness=0.4,0.45'),
    ('inverted thickness args',
     '--size 128 --count 150 --seed 13 --thickness 0.45 0.4',
     'size=128 count=150 seed=13 thickness=0.45,0.4'),
    ('dense overlap, occlusion order',
     '--size 64 --count 500 --seed 21 --radius-px 5',
     'size=64 count=500 seed=21 radiusPx=5'),
    ('non-square-friendly seed 0',
     '--size 64 --count 200 --seed 0',
     'size=64 count=200 seed=0'),
    ('large seed, wraps 32 bits',
     '--size 64 --count 200 --seed 4294967290',
     'size=64 count=200 seed=4294967290'),
]


def run(cmd, **kw):
    return subprocess.run(cmd, capture_output=True, text=True, **kw)


def png_bytes(path):
    with Image.open(path) as im:
        return im.convert('RGBA').tobytes()


def compare(label, a, b, which, verbose):
    """Returns the number of differing bytes, and prints a summary."""
    if len(a) != len(b):
        print(f'  FAIL {label} [{which}]: size differs, {len(a)} vs {len(b)}')
        return max(len(a), len(b))
    diff = [i for i in range(len(a)) if a[i] != b[i]]
    if not diff:
        return 0
    worst = max(abs(a[i] - b[i]) for i in diff)
    pct = 100.0 * len(diff) / len(a)
    print(f'  FAIL {label} [{which}]: {len(diff)} of {len(a)} bytes differ '
          f'({pct:.4f}%), worst {worst}')
    if verbose:
        for i in diff[:8]:
            texel, chan = divmod(i, 4)
            print(f'    texel {texel} channel {"RGBA"[chan]}: py {a[i]} vs js {b[i]}')
    return len(diff)


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--verbose', action='store_true',
                    help='print the first differing texels of each failure')
    ap.add_argument('--only', help='run only cases whose label contains this')
    args = ap.parse_args()

    if run(['node', '--version']).returncode != 0:
        print('node not found; cannot drive glitter_gen.js', file=sys.stderr)
        return 2
    for path in (GEN_JS, GEN_PY):
        if not os.path.exists(path):
            print(f'missing {path}', file=sys.stderr)
            return 2

    cases = [c for c in CASES if not args.only or args.only in c[0]]
    failures = 0

    if not args.only:
        failures += check_primitives(args.verbose)

    with tempfile.TemporaryDirectory() as tmp:
        for label, pyflags, jsopts in cases:
            py_prefix = os.path.join(tmp, 'py')
            js_prefix = os.path.join(tmp, 'js')

            p = run(['python3', GEN_PY, '--out-prefix', py_prefix, *_split(pyflags)])
            if p.returncode != 0:
                print(f'  FAIL {label}: python tool exited {p.returncode}\n{p.stderr}')
                failures += 1
                continue

            j = run(['node', GEN_JS, js_prefix, *jsopts.split()])
            if j.returncode != 0:
                print(f'  FAIL {label}: node exited {j.returncode}\n{j.stderr}')
                failures += 1
                continue

            bad = 0
            for which in ('normal', 'color'):
                a = png_bytes(f'{py_prefix}_{which}.png')
                with open(f'{js_prefix}_{which}.raw', 'rb') as f:
                    b = f.read()
                bad += compare(label, a, b, which, args.verbose)

            # the palette is reported by both and is part of the contract
            with open(f'{js_prefix}.json') as f:
                js_meta = json.load(f)
            py_palette = _palette_from_stdout(p.stdout)
            if py_palette != js_meta['palette']:
                print(f'  FAIL {label} [palette]: {py_palette} vs {js_meta["palette"]}')
                bad += 1

            if bad:
                failures += 1
            else:
                print(f'  ok   {label}')

    total = len(cases) + (0 if args.only else 1)
    print()
    if failures:
        print(f'{failures} of {total} cases DIFFER — the studio and the tool have drifted.')
        return 1
    print(f'all {total} cases identical (zero bytes differ).')
    return 0


def check_primitives(verbose):
    """Compare the shared PRNG, rounding and HSV directly.

    Pixel diffing does NOT cover these. Measured: the only rint tie a real map
    ever produces is 127.5, where half-to-even and half-up agree — so a broken
    tie-break passes every image case. The PRNG is worse: a divergence there
    changes the whole map, which looks like a different-but-plausible map rather
    than an obvious failure. Both are checked head-on here.
    """
    p = run(['node', GEN_JS, '--probe'])
    if p.returncode != 0:
        print(f'  FAIL primitives: node exited {p.returncode}\n{p.stderr}')
        return 1
    js = json.loads(p.stdout)

    sys.path.insert(0, os.path.join(ROOT, 'tools'))
    import make_glitter_map as py

    bad = 0

    rng = py._Rng(12345)
    py_stream = [rng.next() for _ in range(len(js['stream']))]
    if py_stream != js['stream']:
        first = next(i for i, (a, b) in enumerate(zip(py_stream, js['stream'])) if a != b)
        print(f'  FAIL primitives [mulberry32]: streams diverge at draw {first}: '
              f'py {py_stream[first]!r} vs js {js["stream"][first]!r}')
        bad += 1

    ties = [-2.5, -1.5, -0.5, 0.5, 1.5, 2.5, 3.5, 126.5, 127.5, 128.5,
            254.5, 0.25, 0.75, -0.25, 1e-9, 254.9999999]
    py_rint = [py.rint(v) for v in ties]
    if py_rint != js['rint']:
        for v, a, b in zip(ties, py_rint, js['rint']):
            if a != b:
                print(f'  FAIL primitives [rint]: rint({v}) py {a} vs js {b}')
        bad += 1

    import colorsys
    hsv_in = [(0, 0.85, 1), (0.1, 0.85, 1), (0.5, 0.85, 1), (0.99, 0.85, 1),
              (0.0, 0.0, 0.4), (0.3333333333, 1, 1)]
    py_hsv = [list(colorsys.hsv_to_rgb(*a)) for a in hsv_in]
    if py_hsv != js['hsv']:
        print(f'  FAIL primitives [hsvToRgb]: {py_hsv} vs {js["hsv"]}')
        bad += 1

    if bad:
        return 1
    print('  ok   primitives (mulberry32 stream, rint ties, hsvToRgb)')
    return 0


def _split(flags):
    """Split a flag string, honouring the single quotes used around palettes."""
    import shlex
    return shlex.split(flags)


def _palette_from_stdout(text):
    for line in text.splitlines():
        line = line.strip()
        if line.startswith('palette of '):
            return line.split(':', 1)[1].split()
    return []


if __name__ == '__main__':
    sys.exit(main())
