#!/usr/bin/env python3
"""
Reproject an atlas-space eyeliner map into the strip space used by
`region: 'eyeLidUpper'` (src/eye_strip_geometry.ts).

Existing eyeliner art in this repo is painted in the canonical MediaPipe face-UV
layout (against demos/assets/imgs/face/canonical_face_model_uv_visualization.png).
The ribbon wants a flat strip instead: x along the lash line, y across the lid.
Rather than redraw, this walks the strip, works out which point of the canonical
face each strip pixel corresponds to, and samples the atlas there.

The mapping is exact in one respect worth knowing. At runtime the ribbon scales
its width and its v range by the same fold factor, so

    skin offset from the lash margin (mm) = v * strokeWidth * eyeWidth

with no dependence on how open the eye is. v is therefore a fixed physical
offset, which is why a single bake is valid for every lid state.

Requires only numpy + Pillow. Run from the repo root:

    python3 tools/bake_eye_strip.py demos/assets/imgs/face/color_map_eye_liner.png

Options:
    --out PATH        default: alongside the input, suffixed _strip
    --size W H        default: 1024 128
    --stroke F        v=1 offset as a fraction of eye width. Must match the
                      layer's eyeStrip.strokeWidth. Default 0.11 (~2.1 mm).
    --columns N       spine samples; only affects where the wing begins.
                      Default 48, matching the geometry default.
    --corner-u F      default 0.75, snapped to a column exactly as the geometry
                      snaps it.
    --wing-length F   default 0.20 (fraction of eye width)
    --wing-lift F     default 0.50
    --eye WHICH       left | right | average (default average, which cancels any
                      left/right asymmetry in the source art)
    --supersample N   NxN samples per output pixel, default 2
"""

import argparse
import math
import os
import re
import sys

import numpy as np
from PIL import Image

# --- the same landmark rings the geometry uses ----------------------------
# Kept in step with RINGS in src/eye_strip_geometry.ts. The lower-lid pairing is
# mirror-verified: 133/243 ... 33/130 on the left, and the exact mirrors on the
# right, with separations running 1.5-2.8 mm.
RINGS = {
    'upper': [
        {'lash': [133, 173, 157, 158, 159, 160, 161, 246, 33],
         'offset': [243, 190, 56, 28, 27, 29, 30, 247, 130],
         'far': 223},
        {'lash': [362, 398, 384, 385, 386, 387, 388, 466, 263],
         'offset': [463, 414, 286, 258, 257, 259, 260, 467, 359],
         'far': 443},
    ],
    'lower': [
        {'lash': [133, 155, 154, 153, 145, 144, 163, 7, 33],
         'offset': [243, 112, 26, 22, 23, 24, 110, 25, 130],
         'far': 118},
        {'lash': [362, 382, 381, 380, 374, 373, 390, 249, 263],
         'offset': [463, 341, 256, 252, 253, 254, 339, 255, 359],
         'far': 347},
    ],
}
DEFAULT_WING_LENGTH = {'upper': 0.20, 'lower': 0.0}

# Candidate triangles per region, generously: the lid rings plus what surrounds
# them, so a winged tip that leaves the lid still finds surface to land on.
REGION_TRIS = {
    ('upper', 0): [
        7, 163, 144, 145, 153, 154, 155, 25, 110, 24, 23, 22, 26, 112,
        226, 113, 225, 224, 223, 222, 221, 189, 244, 245, 233, 232, 231, 230,
        229, 228, 31, 35, 124, 46, 53, 52, 65, 55, 107, 111, 117, 118, 119,
        120, 121, 128, 114, 143, 156, 70, 63, 105, 66, 9, 8, 193, 122, 6, 168,
    ],
    ('upper', 1): [
        249, 390, 373, 374, 380, 381, 382, 255, 339, 254, 253, 252, 256, 341,
        446, 342, 445, 444, 443, 442, 441, 413, 464, 465, 453, 452, 451, 450,
        449, 448, 261, 265, 353, 276, 283, 282, 295, 285, 336, 340, 346, 347,
        348, 349, 350, 357, 343, 372, 383, 300, 293, 334, 296, 9, 8, 417, 351,
        6, 168,
    ],
}
# the lower lid reuses the same neighbourhood: the rings below the eye and the
# cheek the ribbon drops onto
REGION_TRIS[('lower', 0)] = REGION_TRIS[('upper', 0)]
REGION_TRIS[('lower', 1)] = REGION_TRIS[('upper', 1)]


# --- geometry.ts parsing -----------------------------------------------------
def load_geometry(path):
    txt = open(path).read()

    def block(name):
        m = re.search(r'export const ' + name + r' = \[(.*?)\n\];', txt, re.S)
        if not m:
            sys.exit(f'{path}: no {name} export. Rebase ft/delineador first — '
                     'VERTICES only exists there.')
        return m.group(1)

    faces = np.array([int(x) for x in re.findall(r'-?\d+', block('FACES'))])
    uvs = np.array([[float(a), float(b)] for a, b in re.findall(
        r'\[\s*(-?[\d.]+)\s*,\s*(-?[\d.]+)\s*\]', block('UVS'))])
    verts = np.array([[float(a), float(b), float(c)] for a, b, c in re.findall(
        r'\[\s*(-?[\d.]+)\s*,\s*(-?[\d.]+)\s*,\s*(-?[\d.]+)\s*\]',
        block('VERTICES'))])
    return faces.reshape(-1, 3), uvs, verts


# --- three's CatmullRomCurve3, curveType 'centripetal' -----------------------
class CubicPoly:
    __slots__ = ('c0', 'c1', 'c2', 'c3')

    def init(self, x0, x1, t0, t1):
        self.c0, self.c1 = x0, t0
        self.c2 = -3 * x0 + 3 * x1 - 2 * t0 - t1
        self.c3 = 2 * x0 - 2 * x1 + t0 + t1

    def init_nonuniform(self, x0, x1, x2, x3, dt0, dt1, dt2):
        t1 = (x1 - x0) / dt0 - (x2 - x0) / (dt0 + dt1) + (x2 - x1) / dt1
        t2 = (x2 - x1) / dt1 - (x3 - x1) / (dt1 + dt2) + (x3 - x2) / dt2
        self.init(x1, x2, t1 * dt1, t2 * dt1)

    def calc(self, t):
        return self.c0 + self.c1 * t + self.c2 * t * t + self.c3 * t * t * t


class CentripetalCatmullRom:
    """Faithful port of three's CatmullRomCurve3 with curveType 'centripetal',
    including getPointAt (arc-length parameterised) via a 200-division LUT."""

    def __init__(self, points, divisions=200):
        self.points = [np.asarray(p, dtype=float) for p in points]
        self._lengths = None
        self.divisions = divisions

    def get_point(self, t):
        P, l = self.points, len(self.points)
        p = (l - 1) * t
        ip = int(math.floor(p))
        w = p - ip
        if w == 0 and ip == l - 1:
            ip, w = l - 2, 1.0
        p1, p2 = P[ip], P[ip + 1]
        p0 = P[ip - 1] if ip > 0 else p1 - p2 + p1
        p3 = P[ip + 2] if ip + 2 < l else p2 - p1 + p2
        POW = 0.25  # distanceToSquared ** 0.25 == distance ** 0.5
        dt1 = float(np.dot(p1 - p2, p1 - p2)) ** POW
        dt0 = float(np.dot(p0 - p1, p0 - p1)) ** POW
        dt2 = float(np.dot(p2 - p3, p2 - p3)) ** POW
        if dt1 < 1e-4:
            dt1 = 1.0
        if dt0 < 1e-4:
            dt0 = dt1
        if dt2 < 1e-4:
            dt2 = dt1
        out = np.empty(3)
        poly = CubicPoly()
        for k in range(3):
            poly.init_nonuniform(p0[k], p1[k], p2[k], p3[k], dt0, dt1, dt2)
            out[k] = poly.calc(w)
        return out

    def lengths(self):
        if self._lengths is None:
            last = self.get_point(0.0)
            sums, s = [0.0], 0.0
            for i in range(1, self.divisions + 1):
                cur = self.get_point(i / self.divisions)
                s += float(np.linalg.norm(cur - last))
                sums.append(s)
                last = cur
            self._lengths = sums
        return self._lengths

    def u_to_t(self, u):
        L = self.lengths()
        target = u * L[-1]
        lo, hi = 0, len(L) - 1
        while lo <= hi:
            mid = (lo + hi) // 2
            if L[mid] < target:
                lo = mid + 1
            elif L[mid] > target:
                hi = mid - 1
            else:
                return mid / (len(L) - 1)
        i = max(0, hi)
        seg = L[i + 1] - L[i] if i + 1 < len(L) else 0.0
        frac = (target - L[i]) / seg if seg > 0 else 0.0
        return (i + frac) / (len(L) - 1)

    def get_point_at(self, u):
        return self.get_point(self.u_to_t(u))


def unit(v):
    n = np.linalg.norm(v)
    return v / n if n > 1e-12 else np.zeros(3)


def eye_frame(verts, lash_ids, crease_ids, brow_id):
    """Averaged browward direction and eye-plane normal — the same construction
    EyeStripGeometry.update uses, and for the same reason: per column the crease
    offset is nearly parallel to the spine at the canthi."""
    lash_c = verts[lash_ids].mean(axis=0)
    up = verts[crease_ids].mean(axis=0) - lash_c
    if np.dot(up, verts[brow_id] - lash_c) < 0:
        up = -up
    up = unit(up)
    axis = verts[lash_ids[-1]] - verts[lash_ids[0]]
    n = unit(np.cross(axis, up))
    if n[2] < 0:
        n = -n
    up = unit(up - n * np.dot(up, n))
    return up, n


def strip_points(u, v, curve, up, n, eye_width, args, corner_col, columns):
    """3D point on the canonical model for strip coordinates (u, v), arrays."""
    corner_u = corner_col / (columns - 1)
    band = args.stroke * eye_width
    wing = args.wing_length * eye_width

    # spine sample and tangent, mirroring buildEye: arc length on the eye,
    # straight extrapolation on the wing
    su = np.clip(u / corner_u, 0.0, 1.0)
    on_wing = u > corner_u

    # tip frame, for the wing
    tip_p = curve.get_point_at(1.0)
    d = 1e-3
    tip_t = unit(curve.get_point_at(1.0) - curve.get_point_at(1.0 - d))
    tip_a = unit(np.cross(n, tip_t))
    if np.dot(tip_a, up) < 0:
        tip_a = -tip_a
    wing_t = unit(tip_t + tip_a * args.wing_lift)

    out = np.empty((u.size, 3))
    # cache spine evaluations: many pixels share a column
    uniq, inv = np.unique(np.round(su, 6), return_inverse=True)
    P = np.empty((uniq.size, 3))
    A = np.empty((uniq.size, 3))
    for k, s in enumerate(uniq):
        P[k] = curve.get_point_at(float(s))
        t = unit(curve.get_point_at(min(1.0, float(s) + d))
                 - curve.get_point_at(max(0.0, float(s) - d)))
        a = unit(np.cross(n, t))
        if np.dot(a, up) < 0:
            a = -a
        A[k] = a
    Pi, Ai = P[inv], A[inv]

    along = np.where(on_wing, (u - corner_u) / max(1e-9, 1.0 - corner_u) * wing, 0.0)
    base = np.where(on_wing[:, None],
                    tip_p[None, :] + wing_t[None, :] * along[:, None],
                    Pi)
    across = np.where(on_wing[:, None], tip_a[None, :], Ai)
    out = base + across * (v * band)[:, None]
    return out


def project_uv(Q, tris, uvs, verts, origin, e1, e2, n):
    """Atlas UV for each point in Q, by projecting along the eye-plane normal.

    Nearest-point-on-mesh was the obvious approach and it is wrong here. The
    ribbon is built in the eye plane while the lid is curved, so its points sit
    slightly off the surface; a point that far off projects inside the PLANES of
    two adjacent triangles, and those two give different UVs for it. The pick
    then flips back and forth along the stroke, which showed up as ~32 scallops
    of about 80 um along the top edge of the first bake.

    Projecting along a single fixed direction is unambiguous: it is a ray cast,
    expressed as a 2D containment test in the eye-plane basis. The lid does not
    fold in that projection, so at most one front-facing triangle contains any
    given sample.
    """
    # front-facing triangles only, so the temple wrapping around cannot claim a
    # sample that belongs to the lid
    fn = np.cross(verts[tris[:, 1]] - verts[tris[:, 0]],
                  verts[tris[:, 2]] - verts[tris[:, 0]])
    facing = (fn @ n) > 0
    tris = tris[facing]

    def to2d(P):
        d = P - origin
        return np.stack([d @ e1, d @ e2], axis=-1)

    V2 = to2d(verts)
    Q2 = to2d(Q)

    A = V2[tris[:, 0]]
    E1 = V2[tris[:, 1]] - A
    E2 = V2[tris[:, 2]] - A
    det = E1[:, 0] * E2[:, 1] - E1[:, 1] * E2[:, 0]
    det = np.where(np.abs(det) < 1e-20, 1e-20, det)

    best_uv = np.zeros((Q.shape[0], 2))
    best_pen = np.full(Q.shape[0], -np.inf)   # most-interior wins
    hit = np.zeros(Q.shape[0], dtype=bool)

    for t in range(tris.shape[0]):
        W = Q2 - A[t]
        b1 = (W[:, 0] * E2[t, 1] - W[:, 1] * E2[t, 0]) / det[t]
        b2 = (E1[t, 0] * W[:, 1] - E1[t, 1] * W[:, 0]) / det[t]
        b0 = 1.0 - b1 - b2
        # signed depth inside the triangle; > 0 means strictly contained
        pen = np.minimum(np.minimum(b0, b1), b2)
        take = pen > best_pen
        if not take.any():
            continue
        uv = (uvs[tris[t, 0]][None, :]
              + b1[:, None] * (uvs[tris[t, 1]] - uvs[tris[t, 0]])[None, :]
              + b2[:, None] * (uvs[tris[t, 2]] - uvs[tris[t, 0]])[None, :])
        best_uv = np.where(take[:, None], uv, best_uv)
        best_pen = np.where(take, pen, best_pen)
        hit |= pen >= -1e-9
    return best_uv, hit


def sample_bilinear(img, uv):
    """img: HxWx4 float. uv in canonical atlas space, v measured top-down (the
    convention of canonical_face_model_uv_visualization.png, which is how every
    face texture in this repo is authored)."""
    H, W = img.shape[:2]
    x = np.clip(uv[:, 0] * (W - 1), 0, W - 1)
    y = np.clip(uv[:, 1] * (H - 1), 0, H - 1)
    x0, y0 = np.floor(x).astype(int), np.floor(y).astype(int)
    x1, y1 = np.minimum(x0 + 1, W - 1), np.minimum(y0 + 1, H - 1)
    fx, fy = (x - x0)[:, None], (y - y0)[:, None]
    return ((img[y0, x0] * (1 - fx) + img[y0, x1] * fx) * (1 - fy)
            + (img[y1, x0] * (1 - fx) + img[y1, x1] * fx) * fy)


def main():
    ap = argparse.ArgumentParser(description=__doc__,
                                 formatter_class=argparse.RawDescriptionHelpFormatter)
    ap.add_argument('source')
    ap.add_argument('--out')
    ap.add_argument('--geometry', default='src/geometry.ts')
    ap.add_argument('--size', nargs=2, type=int, default=[1024, 128])
    ap.add_argument('--stroke', type=float, default=0.11)
    ap.add_argument('--columns', type=int, default=48)
    ap.add_argument('--corner-u', type=float, default=0.75)
    ap.add_argument('--wing-length', type=float, default=None)
    ap.add_argument('--wing-lift', type=float, default=0.50)
    ap.add_argument('--eye', choices=('left', 'right', 'average'), default='average')
    ap.add_argument('--supersample', type=int, default=2)
    ap.add_argument('--region', choices=('upper', 'lower'), default='upper',
                    help="which lash line to bake onto; must match the layer's "
                         "ThreeMeshConfig.region")
    ap.add_argument('--emit', choices=('color', 'mask'), default='color',
                    help="'color' keeps the sampled RGBA, for colorMapDir. "
                         "'mask' writes the sampled ALPHA into RGB and sets A to "
                         "255, for alphaMapDir — the shader reads alphaMap on the "
                         "GREEN channel, so this is what lets the layer's `color` "
                         "drive the tint. Every eyeliner asset in this repo has "
                         "RGB = pure black inside the stroke, which is why `color` "
                         "does nothing through colorMapDir.")
    ap.add_argument('--smooth-u', type=float, default=0.0,
                    help='gaussian sigma in output pixels, along u only. The '
                         'atlas assets in this repo carry a 1-13 texel stroke, '
                         'so their edge is quantised to whole texels and that '
                         'quantisation transfers as a ~30-period wobble of about '
                         '70 um on skin. It is not artistic intent. A synthetic '
                         'smooth source bakes with zero wobble, so anything you '
                         'see here came from the art. 6 suits these assets; '
                         '0 (default) keeps the bake strictly faithful.')
    args = ap.parse_args()
    if args.wing_length is None:
        args.wing_length = DEFAULT_WING_LENGTH[args.region]

    tris, uvs, verts = load_geometry(args.geometry)
    mm = 120.0 / float(np.linalg.norm(verts[10] - verts[152]))

    src = np.asarray(Image.open(args.source).convert('RGBA'), dtype=np.float64)
    W, H = args.size
    ss = max(1, args.supersample)
    columns = max(4, args.columns)
    corner_col = min(columns - 2, max(1, round((columns - 1) * args.corner_u)))
    corner_u = corner_col / (columns - 1)

    # supersampled strip coordinates, v = 0 at the lash margin. Row H-1 of the
    # image is v = 0: the shader reads tUV = uv and three's default flipY = true
    # puts uv.v = 0 on the image's last row.
    off = (np.arange(ss) + 0.5) / ss
    xs = (np.arange(W)[:, None] + off[None, :]).ravel() / W
    ys = (np.arange(H)[:, None] + off[None, :]).ravel() / H
    U, Vt = np.meshgrid(xs, ys, indexing='xy')
    Vv = 1.0 - Vt                      # image top row -> v = 1
    u_flat, v_flat = U.ravel(), Vv.ravel()

    accum = np.zeros((u_flat.size, 4))
    per_eye = []
    eyes = (['left', 'right'] if args.eye == 'average' else [args.eye])
    for which in eyes:
        idx = 0 if which == 'left' else 1
        rings = RINGS[args.region][idx]
        lash, crease, brow = rings['lash'], rings['offset'], rings['far']
        region = set(lash) | set(crease) | set(REGION_TRIS[(args.region, idx)])

        up, n = eye_frame(verts, lash, crease, brow)
        curve = CentripetalCatmullRom([verts[i] for i in lash])
        eye_width = float(np.linalg.norm(verts[lash[-1]] - verts[lash[0]]))

        cand = tris[np.all(np.isin(tris, list(region)), axis=1)]
        Q = strip_points(u_flat, v_flat, curve, up, n, eye_width, args,
                         corner_col, columns)
        # eye-plane basis for the projection
        axis = verts[lash[-1]] - verts[lash[0]]
        e1 = unit(axis - n * np.dot(axis, n))
        e2 = np.cross(n, e1)
        origin = verts[lash].mean(axis=0)
        uv, inside = project_uv(Q, cand, uvs, verts, origin, e1, e2, n)
        rgba = sample_bilinear(src, uv)
        per_eye.append(rgba)
        accum += rgba
        print(f'  {which:5s}: {cand.shape[0]} candidate triangles, '
              f'{100.0 * inside.mean():.1f}% of samples hit one, '
              f'band v=1 at {args.stroke * eye_width * mm:.2f} mm')

    out = accum / len(eyes)
    if len(per_eye) == 2:
        d = np.abs(per_eye[0] - per_eye[1]).max()
        print(f'  max left/right difference in the source art: {d:.1f}/255'
              f'{"  (asymmetric — check --eye left vs right)" if d > 24 else ""}')

    if args.emit == 'mask':
        # alphaMap is read on the GREEN channel (DynamicMaterialShader.ts:1523),
        # so put the mask in RGB and make the image itself opaque.
        m = out[:, 3:4]
        out = np.concatenate([m, m, m, np.full_like(m, 255.0)], axis=1)

    # collapse the supersamples
    out = out.reshape(H, ss, W, ss, 4).mean(axis=(1, 3))

    if args.smooth_u > 0:
        sig = args.smooth_u
        r = max(1, int(math.ceil(3 * sig)))
        k = np.exp(-0.5 * (np.arange(-r, r + 1) / sig) ** 2)
        k /= k.sum()
        pad = np.pad(out, ((0, 0), (r, r), (0, 0)), mode='edge')
        out = np.stack([np.apply_along_axis(
            lambda m: np.convolve(m, k, mode='valid'), 1, pad[:, :, c])
            for c in range(4)], axis=2)
        print(f'  smoothed along u with sigma {sig} px')
    img = Image.fromarray(np.clip(np.rint(out), 0, 255).astype(np.uint8), 'RGBA')

    dst = args.out or (os.path.splitext(args.source)[0] + '_strip.png')
    img.save(dst)

    a = out[:, :, 1] if args.emit == 'mask' else out[:, :, 3]
    rows = np.where(a.max(axis=1) > 8)[0]
    cols = np.where(a.max(axis=0) > 8)[0]
    print(f'\nwrote {dst}  ({W}x{H}, {args.emit}, {args.region} lid)')
    if rows.size:
        v_lo = 1.0 - (rows.max() + 1) / H
        v_hi = 1.0 - rows.min() / H
        print(f'  ink occupies v {v_lo:.3f}..{v_hi:.3f}  '
              f'({v_lo * args.stroke * 19.4:.2f}..{v_hi * args.stroke * 19.4:.2f} mm '
              f'from the lash margin on a 19.4 mm eye)')
        print(f'  ink occupies u {cols.min() / W:.3f}..{(cols.max() + 1) / W:.3f}'
              f'   (outer canthus at u={corner_u:.3f}, wing beyond it)')
        print(f'  eyeStrip.strokeWidth must be {args.stroke} on the layer that '
              f'uses this, or v will not mean what it means here.')
        if v_hi > 0.99:
            print('  Ink reaches the very top of the strip. If you later lower '
                  'foldOpenFrac below 1.0 to\n        enable the fold law, the top '
                  'of the stroke will be hidden with the eye open —\n        raise '
                  '--stroke and re-bake to leave headroom for the unfolded reveal.')
    else:
        print('  WARNING: output is empty. The source may not be atlas-space art, '
              'or its ink sits outside the lid band.')


if __name__ == '__main__':
    main()
