#!/usr/bin/env python3
"""ft-eyes-e2e: the live path end to end on two recordings, without the headset.

Starts a scratch ft-eyes (its own shared memory, socket, and state folder, so the real
calibration is untouched), then:
  1. plays CALIB into it and sends each of its practice clicks as a calibration dot (the
     300 ms before the press), as the probe's calibration would, and fits;
  2. plays TEST into it and, at each of its practice clicks, scores what ft-eyes was
     publishing in the 300 ms before the press, then sends the click, as the probe does.
So every TEST click is scored with only earlier data, like ft-eyes-score's `clicks` method.

Usage: frame-job -- lab/py lab/ft-eyes-e2e CALIB TEST [--for S] [--dump FILE]
CALIB and TEST are recordings (a bare name is in eyes_lab.CAPTURES; give full paths for
frame-job to copy them to the 7i). Runs at the recorded pace (about the two recordings'
length). --dump saves everything ft-eyes published during TEST (OUT_FIELDS per row) and each
click's score, as a pickle, for looking into the bad clicks.
"""
import json
import mmap
import os
import pickle
import re
import shutil
import socket
import struct
import subprocess
import sys
import tempfile
import time
from pathlib import Path

import numpy as np

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

BEFORE = 0.3  # the probe's fixation window before a press (s)
# ft-eyes' output after its seq and version (ft-eyes' docstring): one row per sample.
OUT_FORMAT = "<dffII12f"
OUT_FIELDS = ("t", "yaw", "pitch", "flags", "n", "r_yaw", "r_pitch", "l_yaw", "l_pitch",
              "r_shift_x", "r_shift_y", "l_shift_x", "l_shift_y", "r_pupil_x", "r_pupil_y", "l_pupil_x", "l_pupil_y")


class Run:
    def __init__(self):
        tag = f"ft-eyes-e2e-{os.getpid()}"
        self.state = Path(tempfile.mkdtemp(prefix=tag + "-"))
        self.env = dict(os.environ, FT_EYES_CAMS=f"/dev/shm/{tag}-cams", FT_EYES_GAZE=f"/dev/shm/{tag}-gaze",
                        FT_EYES_STATE=str(self.state), FT_EYES_SOCKET=tag)
        self.log = self.state / "ft-eyes.log"
        self.trackd = subprocess.Popen([sys.executable, str(TRACKER / "ft-eyes"), "-v"], env=self.env,
                                       stdout=subprocess.DEVNULL, stderr=open(self.log, "w"))
        self.sock = socket.socket(socket.AF_UNIX, socket.SOCK_DGRAM)
        self.sock.bind("\0" + tag + "-client")
        self.sock.settimeout(2)
        self.to = "\0" + tag
        self.gaze = None

    def cmd(self, line):
        for _ in range(50):  # ft-eyes may not have bound its socket yet
            try:
                self.sock.sendto(line.encode(), self.to)
                return self.sock.recv(4096).decode()
            except ConnectionRefusedError:
                time.sleep(0.2)
        raise RuntimeError("ft-eyes never answered")

    def latest_frame_t(self):
        """The newest replayed frame's time, or None before the replay starts."""
        try:
            with open(self.env["FT_EYES_CAMS"], "rb") as f:
                mm = mmap.mmap(f.fileno(), 0, prot=mmap.PROT_READ)
        except (OSError, ValueError):
            return None
        slots, esize = struct.unpack_from("<2I", mm, 16)
        best = None
        for cam in (0, 1):
            n = struct.unpack_from("<Q", mm, 24 + 8 * cam)[0]
            if n:
                t = struct.unpack_from("<d", mm, 64 + (cam * slots + (n - 1) % slots) * esize + 8)[0]
                best = t if best is None else max(best, t)
        return best

    def published(self):
        """ft-eyes' newest output (OUT_FIELDS), or None."""
        if self.gaze is None:
            try:
                with open(self.env["FT_EYES_GAZE"], "rb") as f:
                    self.gaze = mmap.mmap(f.fileno(), 128, prot=mmap.PROT_READ)
            except (OSError, ValueError):
                return None
        for _ in range(3):
            seq = struct.unpack_from("<I", self.gaze, 0)[0]
            v = struct.unpack_from(OUT_FORMAT, self.gaze, 8)
            if not seq & 1 and struct.unpack_from("<I", self.gaze, 0)[0] == seq:
                return v
        return None

    def stage(self, cap, secs, handle):
        """Play `cap` and call handle(click, samples) as each click's press time goes by.
        Returns everything ft-eyes published meanwhile."""
        clicks = pickle.load(open(cap / "features.pkl", "rb"))["clicks"]
        replay = subprocess.Popen([sys.executable, str(LAB / "ft-eyes-replay"), str(cap), self.env["FT_EYES_CAMS"],
                                   "--for", str(secs)], stdout=subprocess.DEVNULL)
        todo, samples = list(clicks), []
        while replay.poll() is None:
            t = self.latest_frame_t()
            v = self.published()
            if v and v[0] > 0 and (not samples or v[0] != samples[-1][0]):
                samples.append(v)
            while t and todo and t > todo[0]["t"] + 0.05:
                handle(todo.pop(0), samples)
            time.sleep(0.003)
        return samples

    def close(self):
        self.trackd.terminate()
        self.trackd.wait()
        for p in (self.env["FT_EYES_CAMS"], self.env["FT_EYES_GAZE"]):
            if os.path.exists(p):
                os.unlink(p)
        shutil.rmtree(self.state, ignore_errors=True)


def main(argv):
    args = [a for i, a in enumerate(argv) if not a.startswith("--") and (i == 0 or argv[i - 1] not in ("--for", "--dump"))]
    if len(args) != 2:
        sys.exit(__doc__)
    calib, test = map(capture, args)
    secs = argv[argv.index("--for") + 1] if "--for" in argv else "1e9"
    dump = Path(argv[argv.index("--dump") + 1]) if "--dump" in argv else None
    run = Run()
    try:
        print("calib-start:", run.cmd("calib-start"), flush=True)
        dots = []
        run.stage(calib, secs, lambda k, _s: dots.append(
            run.cmd(f"calib-point {k['t'] - BEFORE} {k['t']} {k['truth'][0]} {k['truth'][1]}")))
        fails = [d for d in dots if not d.startswith("ok")]
        print(f"{calib.name}: {len(dots) - len(fails)} of {len(dots)} clicks taken as dots", flush=True)
        whys = [re.sub(r"[\d.]+", "N", f) for f in fails]
        for why in sorted(set(whys)):
            print(f"  {whys.count(why)} x {why}")
        print("calib-fit:", run.cmd("calib-fit"), flush=True)
        # ft-eyes-replay removes its file at the end; ft-eyes notices within 5 s and waits for the next.
        time.sleep(6)
        scored = []

        def click(k, samples):
            s = np.array([v for v in samples if k["t"] - BEFORE <= v[0] <= k["t"]]).reshape(-1, len(OUT_FIELDS))
            err = float(np.hypot(*(np.median(s[:, 1:3], axis=0) - k["truth"]))) if len(s) >= 5 else None
            reply = run.cmd(f"click {k['t']} {k['truth'][0]} {k['truth'][1]}")
            st = json.loads(run.cmd("status"))["eyes"]
            scored.append((k["t"], err, reply, [(v["clicks"], v["jump"]) for v in st.values()]))

        test_samples = run.stage(test, secs, click)
        if dump:
            with open(dump, "wb") as f:
                pickle.dump({"fields": OUT_FIELDS, "samples": np.array(test_samples, float),
                             "clicks": [dict(t=x[0], err=x[1], reply=x[2], eyes=x[3]) for x in scored]}, f)
            print("dumped to", dump)
        e = np.array([x[1] for x in scored if x[1] is not None])
        print(f"{test.name}: {len(e)} of {len(scored)} clicks scored live", flush=True)
        if len(e):
            print(f"  median {np.median(e):.2f} deg, 90% {np.percentile(e, 90):.2f}, "
                  f"after the first 5: median {np.median(e[5:]):.2f}")
        print("  clicks ft-eyes refused:", sum(not x[2].startswith("ok") for x in scored))
        t0 = scored[0][0] if scored else 0
        print("  clicks that started an eye's shift over after a jump: right {}, left {}".format(
            *(sum(x[3][c][0] == 1 for x in scored[1:]) for c in (0, 1))))
        print("  by time (s): " + ", ".join(
            f"{lo}-{lo + 30}: {np.median(b):.2f}" for lo in range(0, 300, 30)
            if len(b := [x[1] for x in scored if x[1] is not None and lo <= x[0] - t0 < lo + 30])))
        print("  worst: " + ", ".join(
            f"{x[0] - t0:.0f}s {x[1]:.1f} (clicks/jump R {x[3][0][0]}/{x[3][0][1]:d} L {x[3][1][0]}/{x[3][1][1]:d})"
            for x in sorted((x for x in scored if x[1] is not None), key=lambda x: -x[1])[:10]))
        print("status:", run.cmd("status"))
        print("ft-eyes' last report:", run.log.read_text().strip().splitlines()[-1:])
    finally:
        run.close()


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