#!/usr/bin/env python3
"""ft-eyes-score: score our pupil tracker and SteamVR against the gaze probe's practice clicks.

Usage (from gaze/tracker; A and B are recordings: a bare name is in eyes_lab.CAPTURES):
  frame-job -- lab/py lab/ft-eyes-score A        fit and test on A, leave-one-out
  frame-job -- lab/py lab/ft-eyes-score A B      fit on A, test on B
  lab/py lab/ft-eyes-score --save A              fit on A, save it for ft-eyes (on the Frame)
  lab/py lab/ft-eyes-score --clicks A...         on the Frame: copy each capture's practice
      clicks from the probe's log into it (the first scoring does it too)
For frame-job to copy a recording to the 7i, give its full path, e.g.
~/.local/share/frametop/eyes/captures/A.

Each practice click gives a known gaze direction: SteamVR's raw gaze at the press plus the
angle from it to where you released (you were looking there). For each click we take the
frames from just before the press and find each eye's pupil and glint pair.

Methods, each a quadratic fit per eye, both eyes averaged when both are seen:
  pupil   the pupil centre alone. Breaks when the headset slips on the face.
  glint   pupil minus the glint pair's midpoint. Slip moves both alike, so this holds up,
          but the right eye's pair is often off the cornea.
  clicks  the pupil centre minus the shift the earlier clicks measured, with the glints
          only noticing a sudden jump (eyes_model.Shift). The one ft-eyes uses.
  slip    the pupil centre minus a slip estimate. Wherever the pair is seen, the glint
          method gives the gaze, the fit says where the pupil should be for that gaze, and
          the difference is the slip. Slip changes slowly, so the median over the last
          30 seconds applies to every frame, with or without glints (see eyes_model.py).
SteamVR gets the same quadratic fit on its raw gaze.
"""
import json
import pickle
import sys
import time
from pathlib import Path

import numpy as np

sys.path.insert(0, str(Path(__file__).resolve().parent))
from eyes_lab import capture  # noqa: E402
import eyes_model  # noqa: E402
import eyes_pupil  # noqa: E402

PRACTICE = Path.home() / ".local/state/frametop/gaze/practice.jsonl"
BEFORE = (0.25, 0.02)   # frames from 250 ms to 20 ms before the press
EYES = {0: "right", 1: "left"}
MIN_FRAMES = 5
TRACK_EVERY = 9         # the slip track uses every 9th frame per camera (10 a second)
VERSION = 4             # bump when the features change, to rebuild the caches


# --- Features -------------------------------------------------------------------------

def load_index(cap):
    # Skip a half-written last line (the recorder may still be running).
    return np.array([[float(v) for v in l.split()]
                     for l in (cap / "index.txt").read_text().splitlines() if len(l.split()) == 4])


def eye_features(frames):
    """Median pupil centre and glint-pair midpoint over some frames of one eye."""
    ps = [p for p in map(eyes_pupil.find_pupil, frames) if p]
    if len(ps) < MIN_FRAMES:
        return None
    pupil = np.median([(p["x"], p["y"]) for p in ps], axis=0)
    mids = []
    for p in ps:
        pair = eyes_pupil.glint_pair(p)
        if pair:
            mids.append(((pair[0][0] + pair[1][0]) / 2, (pair[0][1] + pair[1][1]) / 2))
    mid = np.median(mids, axis=0) if len(mids) >= 3 else None
    return dict(pupil=pupil, mid=mid)


def practice_log(cap, t0, t1, wall, mono):
    """The probe's practice records during the capture. The first time (on the Frame), cut
    from the probe's log into the capture's practice.jsonl, so the capture carries its own
    clicks (the 7i has no probe log)."""
    own = cap / "practice.jsonl"
    if not own.exists():
        if not PRACTICE.exists():
            sys.exit(f"{own} is missing: run `lab/py lab/ft-eyes-score --clicks {cap.name}` once on the Frame")
        keep = [line for line in open(PRACTICE)
                if t0 - 5 < json.loads(line)["time"] - wall + mono < t1 + 60]
        own.write_text("".join(keep))
    return [json.loads(line) for line in open(own)]


def features(cap):
    """Per-click and per-session features, cached in the capture (numbers only)."""
    cache = cap / "features.pkl"
    if cache.exists():
        f = pickle.loads(cache.read_bytes())
        if f.get("version") == VERSION:
            return f
    wall, mono = map(float, (cap / "clocks.txt").read_text().split())
    idx = load_index(cap)
    frames = np.memmap(cap / "frames.raw", dtype=np.uint8, mode="r").reshape(-1, 400, 512)
    idx = idx[:len(frames)]
    t0, t1 = idx[0, 3], idx[-1, 3]

    clicks = []
    for r in practice_log(cap, t0, t1, wall, mono):
        src = r.get("sources", {})
        s = src.get("mmap1")
        if r.get("mode") != "practice" or not s or "off" not in s:
            continue
        press = r["time"] - r["held_s"] - wall + mono
        if not t0 + BEFORE[0] < press < t1:
            continue
        k = np.where((idx[:, 3] > press - BEFORE[0]) & (idx[:, 3] < press - BEFORE[1]))[0]
        eye = {c: eye_features([frames[i] for i in k if idx[i, 2] == c]) for c in (0, 1)}
        # The truth: a source's gaze at the press plus the angle from it to the release
        # point. The angle is converted with a local linear fit of the screen, so the
        # closer the source, the better: ours when the live tracker was running.
        own = src.get("own") if src.get("own", {}).get("off") else None
        base = own or s
        truth = (base["hy"] + base["off"][0], base["hp"] + base["off"][1])
        clicks.append(dict(t=press, truth=truth, truth_from="own" if own else "mmap1",
                           steam=(s["hy"], s["hp"]),
                           steam_err=float(np.hypot(s["hy"] - truth[0], s["hp"] - truth[1])),
                           live_err=float(np.hypot(*own["off"])) if own else None,
                           press_err=r.get("press_err_deg"), press_source=r.get("source"), eye=eye))

    # The slip track: pupil and pair midpoint on a sample of frames through the session.
    track = {}
    for c in (0, 1):
        rows = []
        for i in np.where(idx[:, 2] == c)[0][::TRACK_EVERY]:
            p = eyes_pupil.find_pupil(frames[i])
            pair = p and eyes_pupil.glint_pair(p)
            if pair:
                rows.append((idx[i, 3], p["x"], p["y"],
                             (pair[0][0] + pair[1][0]) / 2, (pair[0][1] + pair[1][1]) / 2))
        track[c] = np.array(rows).reshape(-1, 5)
    f = dict(version=VERSION, clicks=clicks, track=track, span=(t0, t1))
    cache.write_bytes(pickle.dumps(f))
    return f


# --- Fitting --------------------------------------------------------------------------

class Model:
    """A calibration fitted on some clicks, plus SteamVR's quadratic fit on the same."""

    def __init__(self, clicks):
        self.cal = eyes_model.Calibration.fit(clicks)
        self.steam = eyes_model.Quad([k["steam"] for k in clicks], [k["truth"] for k in clicks])

    def slip(self, c, track, t):
        """Median slip (pixels) over the track in the SLIP_WINDOW seconds before t."""
        tr = track[c]
        if len(tr) == 0:
            return None
        return self.cal.slip(c, tr[(tr[:, 0] < t) & (tr[:, 0] > t - eyes_model.SLIP_WINDOW)])

    def predict(self, method, k, track):
        """Gaze for one click by a method, combining the eyes it has; None if neither."""
        cal, out = self.cal, {}
        for c in (0, 1):
            e = k["eye"][c]
            if e is None or not cal.has("pupil", c):
                continue
            if method == "pupil":
                out[c] = cal.gaze(c, *e["pupil"])
            elif method == "glint" and cal.has("glint", c) and e["mid"] is not None:
                out[c] = cal.fits["glint", c].one(*(e["pupil"] - e["mid"]))
            elif method == "slip":
                s = self.slip(c, track, k["t"])
                if s is not None:
                    out[c] = cal.gaze(c, *e["pupil"], slip=s)
        return cal.combine(out)


# --- Scoring --------------------------------------------------------------------------

METHODS = ("pupil", "glint", "slip")


def report(name, err, total):
    err = np.asarray([e for e in err if e is not None])
    if len(err) == 0:
        print(f"  {name:36s} no clicks")
        return
    print(f"  {name:36s} {len(err):3d}/{total}  median {np.median(err):5.2f}  "
          f"mean {err.mean():5.2f}  90% {np.percentile(err, 90):5.2f} deg")


def score(train, test, track, same):
    """Errors per click for each method; leave-one-out when train and test are the same.
    "clicks" goes through the test clicks in order, as live: each is predicted with the
    shift the earlier ones measured (eyes_model.Shift), then teaches it."""
    errs = {m: [] for m in METHODS + ("clicks", "steam")}
    model = None if same else Model(train)
    shifts = {c: eyes_model.Shift() for c in (0, 1)}
    for i, k in enumerate(test):
        mdl = Model(train[:i] + train[i + 1:]) if same else model
        for m in METHODS:
            g = mdl.predict(m, k, track)
            errs[m].append(None if g is None else float(np.hypot(*(g - k["truth"]))))
        errs["steam"].append(float(np.hypot(*(mdl.steam(k["steam"])[0] - k["truth"]))))
        out = {}
        for c in (0, 1):
            e = k["eye"][c]
            if e is None or not mdl.cal.has("pupil", c):
                continue
            tr = track[c]
            # The glint estimate as live would have had it over the last second, for the hold.
            for back in (eyes_model.JUMP_HOLD, eyes_model.JUMP_HOLD / 2, 0.0):
                te = k["t"] - back
                g = mdl.cal.slip(c, tr[(tr[:, 0] < te) & (tr[:, 0] > te - eyes_model.JUMP_WINDOW)]) if len(tr) else None
                shifts[c].glint(g, te)
            out[c] = mdl.cal.gaze(c, *e["pupil"], slip=shifts[c].value)
            shifts[c].click(mdl.cal.click_shift(c, e["pupil"], k["truth"]), g)
        errs["clicks"].append(float(np.hypot(*(mdl.cal.combine(out) - k["truth"]))) if out else None)
    return errs


def summary(f, label):
    cl = f["clicks"]
    T = np.array([k["truth"] for k in cl])
    print(f"{label}: {len(cl)} clicks over {f['span'][1] - f['span'][0]:.0f} s, gaze yaw "
          f"{T[:, 0].min():.0f}..{T[:, 0].max():.0f}, pitch {T[:, 1].min():.0f}..{T[:, 1].max():.0f}")
    for c in (0, 1):
        n = sum(k["eye"][c] is not None for k in cl)
        g = sum(k["eye"][c] is not None and k["eye"][c]["mid"] is not None for k in cl)
        print(f"  {EYES[c]} eye: pupil before {n} clicks, glint pair before {g}; "
              f"slip track {len(f['track'][c])} frames with the pair")


CALIBRATION = Path.home() / ".local/state/frametop/gaze/eyes/calibration.json"


def main(args):
    if args and args[0] == "--clicks":
        for cap in map(capture, args[1:]):
            wall, mono = map(float, (cap / "clocks.txt").read_text().split())
            lines = (cap / "index.txt").read_text().splitlines()
            ts = [float(l.split()[3]) for l in (lines[0], lines[-1])]
            print(f"{cap}: {len(practice_log(cap, *ts, wall, mono))} practice records")
        return
    if args and args[0] == "--save":
        cap = capture(args[1])
        f = features(cap)
        cal = eyes_model.Calibration.fit(f["clicks"], {"capture": cap.name, "clicks": len(f["clicks"]),
                                                     "made": time.strftime("%Y-%m-%d %H:%M")})
        cal.save(CALIBRATION)
        print(f"saved {CALIBRATION}: {sorted(f'{n} {EYES[e]}' for n, e in cal.fits)}")
        return
    ca = capture(args[0])
    a = features(ca)
    summary(a, ca.name)
    if len(args) > 1:
        cb = capture(args[1])
        b = features(cb)
        summary(b, cb.name)
        print(f"\nFit on {ca.name}, tested on {cb.name}:")
        test, track, errs = b["clicks"], b["track"], score(a["clicks"], b["clicks"], b["track"], False)
    else:
        print("\nLeave-one-out within the session:")
        test, track, errs = a["clicks"], a["track"], score(a["clicks"], a["clicks"], a["track"], True)
    n = len(test)
    report("SteamVR raw", [k["steam_err"] for k in test], n)
    steam_press = [k["press_err"] for k in test if k["press_source"] != "own"]
    report("SteamVR + probe's live correction", steam_press, len(steam_press))
    live = [k["live_err"] for k in test if k["live_err"] is not None]
    if live:
        report("ours live (ft-eyes, as the probe saw it)", live, n)
        own_press = [k["press_err"] for k in test if k["press_source"] == "own"]
        report("ours live + probe's live correction", own_press, len(own_press))
    report("SteamVR + quadratic fit", errs["steam"], n)
    for m in METHODS:
        report(f"ours, {m}", errs[m], n)
    report("ours, clicks (shift from earlier clicks)", errs["clicks"], n)
    # Like for like: the clicks every method scored.
    common = [i for i in range(n) if all(errs[m][i] is not None for m in METHODS)]
    print(f"\nSame {len(common)} clicks for every method:")
    report("SteamVR + quadratic fit", [errs["steam"][i] for i in common], len(common))
    for m in METHODS:
        report(f"ours, {m}", [errs[m][i] for i in common], len(common))
    if len(args) == 1:
        for c in (0, 1):
            tr = track[c]
            if len(tr) < 20:
                continue
            mdl = Model(a["clicks"])
            ss = [mdl.slip(c, track, t) for t in np.linspace(tr[0, 0] + eyes_model.SLIP_WINDOW, tr[-1, 0], 6)]
            print(f"  {EYES[c]} eye slip estimate through the session (px): "
                  + " ".join(f"({s[0]:+.1f},{s[1]:+.1f})" for s in ss if s is not None))


if __name__ == "__main__":
    main(sys.argv[1:] or ["practice1"])
