"""Post-pass: remove stray copies of the logo the package printed off the shirt (on an arm, a wall) in a few frames.
usage: strays.py dir plan.json a0 a1 [min_score]  — in place. The new print (from the plan) is masked out of the search."""
import cv2, numpy as np, json, sys, os
d, plan, a0, a1 = sys.argv[1], sys.argv[2], int(sys.argv[3]), int(sys.argv[4]); TH = float(sys.argv[5]) if len(sys.argv) > 5 else 0.5
P = {int(k): v for k, v in json.load(open(plan)).items()}
RAW = cv2.imread(os.environ.get('LOGO', 'logo_tee.png'), -1); lh, lw = RAW.shape[:2]; a_ = RAW[..., 3:4] / 255.0
LG = cv2.cvtColor((RAW[..., :3] * a_ + 255 * (1 - a_)).astype(np.uint8), cv2.COLOR_BGR2GRAY)
for i in range(a0, a1 + 1):
    fr = cv2.imread(f'{d}/f{i:03d}.png'); g = cv2.cvtColor(fr, cv2.COLOR_BGR2GRAY); H, W = g.shape; gm = g.copy()
    if i in P:
        x, y = P[i]['print_at']; w = P[i]['print_w']; gm[int(y - w * .25):int(y + w * .25), int(x - w * .7):int(x + w * .7)] = 128
    sub = cv2.resize(gm, None, fx=.5, fy=.5, interpolation=cv2.INTER_AREA)       # half size: 4x faster, prints are big enough
    fg = np.zeros((H, W), np.uint8); hits = []
    for _ in range(8):
        best = (0, None)
        for wpx in np.geomspace(24, 210, 22):
            hpx = max(4, int(round(wpx * lh / lw))); wpx = int(round(wpx))
            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] < TH: break
        sx, sy, sw, sh_ = best[1]; sub[max(0, sy - 2):sy + sh_ + 2, max(0, sx - 2):sx + sw + 2] = 128
        bx, by, bw, bh = [2 * v for v in best[1]]; p = 5
        box = g[max(0, by - p):by + bh + p, max(0, bx - p):bx + bw + p].astype(np.float32)
        bg = cv2.GaussianBlur(cv2.dilate(box, np.ones((9, 9), np.uint8)), (0, 0), 6)        # the surface under the print
        fg[max(0, by - p):by + bh + p, max(0, bx - p):bx + bw + p] |= (box < bg - 14).astype(np.uint8)
        hits.append((round(best[0], 2), bx, by, bw))
    if fg.any():
        cv2.imwrite(f'{d}/f{i:03d}.png', cv2.inpaint(fr, cv2.dilate(fg, np.ones((5, 5), np.uint8)), 5, cv2.INPAINT_TELEA))
    print(i, hits)
