"""Re-print the consultant's T-shirt logo on a tracked position — conservative erase.

Plan JSON per frame: prints [[x,y,s,theta,width,weight],...], box [x0,x1,y0,y1] (template search area around the
first print, in units of its scale), extra [[x0,x1,y0,y1],...] optional windows where stray print fragments are
cleaned, sigma (print softness).

Only three things are ever erased, so hair, skyline, collar and laptop are never touched:
  1. old prints found by template matching the logo, and only their print-coloured pixels;
  2. strictly magenta pixels of the chest marker (compact blobs);
  3. print-coloured pixels in a tight window around the new print (and any `extra` window).
Erased pixels are refilled from the white fabric around them, then the vector logo is printed once, behind hair
and hands, tinted by the light on the shirt."""
import cv2, numpy as np, json, sys, os
src, dst, plan = sys.argv[1], sys.argv[2], sys.argv[3]
LOGO = os.environ.get('LOGO', 'logo_tee.png')          # source/brand/logo_tee.png from the founder's package
P = {int(k): v for k, v in json.load(open(plan)).items()}
RAW = cv2.imread(LOGO, -1)
logo = RAW.astype(np.float32); logo[..., :3] *= logo[..., 3:4] / 255.0   # premultiplied: no fringe after warp/blur
lh, lw = logo.shape[:2]
a_ = RAW[..., 3:4] / 255.0
LG = cv2.cvtColor((RAW[..., :3] * a_ + 255 * (1 - a_)).astype(np.uint8), cv2.COLOR_BGR2GRAY)   # the logo printed on white


def quad(x, y, th, pts):
    c, s_ = np.cos(th), np.sin(th)
    return np.array([[x + c * px - s_ * py, y + s_ * px + c * py] for px, py in pts], np.int32)


def rect(H, W, x, y, th, s, r):
    m = np.zeros((H, W), np.uint8); x0, x1, y0, y1 = [v * s for v in r]
    cv2.fillPoly(m, [quad(x, y, th, [(x0, y0), (x1, y0), (x1, y1), (x0, y1)])], 1); return m


def on_shirt(bb, hsv0, b, r, base):
    """True when the band around a box is white fabric (not blue sky, not a dark scene)."""
    ring = (cv2.dilate(bb, np.ones((11, 11), np.uint8)) - bb) > 0
    return np.median(hsv0[..., 2][ring]) > 185 and np.median(hsv0[..., 1][ring]) < 50 and np.median((b - r)[ring]) < base + 8


stats = {}
for i in sorted(P):
    e = P[i]; fr = cv2.imread(f'{src}/f{i:03d}.png'); H, W = fr.shape[:2]
    x, y, s, th, LW = e['prints'][0][:5]
    f = fr.astype(np.float32); b, g, r = f[..., 0], f[..., 1], f[..., 2]; V = f.max(2)
    reg = rect(H, W, x, y, th, s, e['box'])
    br = (reg > 0) & (V > 200)
    base = float(np.median((b - r)[br])) if br.sum() > 50 else 4.0
    vmed = float(np.median(V[br])) if br.sum() > 50 else 240
    hair = (r - b > 10)                                                       # hair and skin are warm; the print never is
    blk = (V < 100) & (r - b < 15)                                           # dark, not brown
    n, lab, st, _ = cv2.connectedComponentsWithStats(cv2.erode(blk.astype(np.uint8), np.ones((3, 3), np.uint8)))
    big = np.isin(lab, [k for k in range(1, n) if st[k][4] > 450 * s * s or max(st[k][2], st[k][3]) > 160 * s])
    bezel = (cv2.dilate(big.astype(np.uint8), np.ones((5, 5), np.uint8)) > 0) & blk   # a long dark shape (laptop bezel); letters are small
    laptops = [st[k][:4] for k in range(1, n) if max(st[k][2], st[k][3]) > 160 * s]      # bounding boxes of long dark frames: text inside is a screen
    printcol = (((b - r - base > 14) & (b > g - 10)) | ((V < vmed - 70) & (b - r > 8))) & ~hair & ~bezel
    fg = np.zeros((H, W), np.uint8)
    hsv0 = cv2.cvtColor(fr, cv2.COLOR_BGR2HSV)
    # 1. whole old prints, by template matching at any size; erase their print-coloured (and greyed-core) pixels
    ys, xs = np.where(reg > 0); rx0, ry0, rx1, ry1 = xs.min(), ys.min(), xs.max() + 1, ys.max() + 1
    sub = cv2.cvtColor(fr, cv2.COLOR_BGR2GRAY)[ry0:ry1, rx0:rx1].copy(); nm = 0
    for _ in range(3):
        best = (0, None)
        for wpx in np.geomspace(36, 260, 34):
            hpx = max(6, int(round(wpx * lh / lw))); wpx = int(round(wpx))
            if hpx >= sub.shape[0] or wpx >= sub.shape[1]: continue
            res = cv2.matchTemplate(sub, cv2.resize(LG, (wpx, hpx), interpolation=cv2.INTER_AREA), cv2.TM_CCOEFF_NORMED)
            _, mv, _, ml = cv2.minMaxLoc(res)
            if mv > best[0]: best = (mv, (ml[0], ml[1], wpx, hpx))
        if best[0] < 0.55: break
        bx, by, bw, bh = best[1]; box = np.zeros_like(reg)
        cv2.rectangle(box, (rx0 + bx - 2, ry0 + by - 2), (rx0 + bx + bw + 2, ry0 + by + bh + 2), 1, -1)
        if on_shirt(box, hsv0, b, r, base):                                   # a 'match' in the skyline is not a print
            fg |= ((box > 0) & (printcol | ((V < vmed - 25) & (b - r > 2) & ~hair & ~bezel))).astype(np.uint8); nm += 1
        sub[max(0, by - 2):by + bh + 2, max(0, bx - 2):bx + bw + 2] = int(np.median(sub))
    # 1b. old prints the template misses (each was warped a little differently): a word-shaped cluster of print-coloured
    #     pixels sitting on white fabric. Hair is warm and the skyline is not ringed by white shirt, so neither qualifies.
    wreg = reg if 'word' not in e else np.max([rect(H, W, x, y, th, s, rr) for rr in e['word']], axis=0)
    pc = ((printcol | ((V < vmed - 45) & (b - r > 2) & ~hair & ~bezel)) & (wreg > 0)).astype(np.uint8)
    word = cv2.morphologyEx(pc, cv2.MORPH_CLOSE, cv2.getStructuringElement(cv2.MORPH_RECT, (int(13 * max(s, .5)) | 1, 3)))
    n, lab, st, _ = cv2.connectedComponentsWithStats(word)
    for k in range(1, n):
        bx_, by_, bw_, bh_, ar_ = st[k]
        if not (30 * s <= bw_ <= 220 * s and 5 <= bh_ <= 45 * s and bw_ >= 2.5 * bh_ and ar_ >= 0.12 * bw_ * bh_): continue
        cxw, cyw = bx_ + bw_ / 2, by_ + bh_ / 2
        if any(lx <= cxw <= lx + lw_ and ly <= cyw <= ly + lh_ for lx, ly, lw_, lh_ in laptops): continue   # laptop screen text
        bb = np.zeros((H, W), np.uint8); cv2.rectangle(bb, (bx_ - 3, by_ - 3), (bx_ + bw_ + 3, by_ + bh_ + 3), 1, -1)
        if on_shirt(bb, hsv0, b, r, base):
            fg |= ((bb > 0) & (printcol | ((V < vmed - 25) & (b - r > 2) & ~hair & ~bezel))).astype(np.uint8); nm += 1
    # 2. the magenta chest marker: strictly magenta, compact, and inside the box (a ribbon always crosses its edge);
    #    then its soft fringe (pink on white, lilac under the cyan sweep, dark magenta where it meets hair)
    core = (r - g > 30) & (b - g > 28)
    if e.get('marker_soft'): core |= (r - g > 10) & (b - g > 8) & (np.abs(r - b) < 18)    # washed-out marker in a light sweep
    mg = cv2.dilate((core & (reg > 0)).astype(np.uint8), np.ones((3, 3), np.uint8))
    edge = cv2.dilate(reg, np.ones((3, 3), np.uint8)) - cv2.erode(reg, np.ones((3, 3), np.uint8))
    n, lab, st, _ = cv2.connectedComponentsWithStats(mg); mk = np.zeros((H, W), np.uint8)
    for k in range(1, n):
        cm = (lab == k).astype(np.uint8)
        if st[k][4] >= 12 and st[k][2] < 160 * s and st[k][3] < 90 * s and not (cm & edge).any(): mk |= cm
    if mk.any():
        near = cv2.dilate(mk, np.ones((int(27 * max(s, .6)) | 1,) * 2, np.uint8)) > 0
        tinge = ((r - g > 6) & (b - g > 4) & (r - b < 20)) | ((b - g > 10) & (r - g > 8))   # pink / purple
        if e.get('marker_soft'): tinge |= (r - g > 2) & (b - g > 0)            # lilac on a cyan-washed shirt (shirt itself has r < g)
        fg |= (mk > 0).astype(np.uint8)
        fg |= (near & tinge).astype(np.uint8)
    # 3. stray print fragments right where the new print goes, and in any extra window
    lhh = LW * lh / lw / 2
    win = rect(H, W, x, y, th, s, [-(LW / 2 + 30), LW / 2 + 30, -(lhh + 14), lhh + 14])
    for rr in e.get('extra', []): win |= rect(H, W, x, y, th, s, rr)
    fg |= ((printcol | ((V < vmed - 40) & (b - r > 2) & ~hair & ~bezel)) & (win > 0)).astype(np.uint8)   # incl. greyed letter cores
    fg = cv2.dilate(fg, np.ones((3, 3), np.uint8)); fg[bezel] = 0
    if fg.any():
        fab = ((fg == 0) & (V > 165) & (cv2.cvtColor(fr, cv2.COLOR_BGR2HSV)[..., 1] < 60)).astype(np.float32); sg = 5 * max(s, .5)
        num = cv2.GaussianBlur(f * fab[..., None], (0, 0), sg); den = cv2.GaussianBlur(fab, (0, 0), sg)[..., None]
        fill = num / np.maximum(den, 1e-4); ok = (den[..., 0] > .05) & (fg > 0)
        out = cv2.inpaint(fr, fg, 4, cv2.INPAINT_TELEA).astype(np.float32); out[ok] = fill[ok]
        fr = np.clip(out, 0, 255).astype(np.uint8)
    f = fr.astype(np.float32)
    hsv = cv2.cvtColor(fr, cv2.COLOR_BGR2HSV).astype(np.float32); Vv, Sa = hsv[..., 2], hsv[..., 1]
    shirt = np.clip((Vv - 130) / 50, 0, 1) * np.clip((70 - Sa) / 35, 0, 1); shirt = cv2.GaussianBlur(shirt, (0, 0), 1.0)
    light = cv2.GaussianBlur(f, (0, 0), 6)
    tint = np.clip(light / np.maximum(light.max(2, keepdims=True), 1), 0, 1) * np.clip(cv2.GaussianBlur(Vv, (0, 0), 6) / 238.0, 0.55, 1.0)[..., None]
    wash = np.clip(cv2.GaussianBlur(Sa, (0, 0), 6) / 110.0, 0, 0.35)[..., None]
    for (px, py, ps, pth, LWp, wgt) in e['prints']:
        if wgt <= 0: continue
        M = cv2.getRotationMatrix2D((lw / 2, lh / 2), -np.degrees(pth), LWp * ps / lw); M[:, 2] += [px - lw / 2, py - lh / 2]
        wl = cv2.GaussianBlur(cv2.warpAffine(logo, M, (W, H), flags=cv2.INTER_AREA, borderValue=(0, 0, 0, 0)), (0, 0), e['sigma'])
        a = (wl[..., 3] / 255.0) * shirt * 0.95 * wgt
        rgb = np.clip(wl[..., :3] / np.maximum(wl[..., 3:4] / 255.0, 1e-3), 0, 255) * tint
        rgb = rgb * (1 - wash) + light * wash
        f = f * (1 - a[..., None]) + rgb * a[..., None]
    cv2.imwrite(f'{dst}/f{i:03d}.png', np.clip(f, 0, 255).astype(np.uint8))
    stats[i] = (nm, int(fg.sum()))
print('done', len(P), 'frames; template hits per frame:', ''.join(str(v[0]) for k, v in sorted(stats.items())))
