#!/usr/bin/env python3
"""ft-eyes: our own eye tracker, live.

Reads the eye-camera frames that ft-eyegrab (the root service frametop-eyegrab, installed by
gaze/tracker/install.sh) keeps in /dev/shm/frametop-eyes-cams, finds each eye's pupil and
glint pair (eyes_pupil.py), turns them into a gaze with the saved calibration and a running
slip estimate (eyes_model.py), and publishes the result in /dev/shm/frametop-eyes-gaze,
where ft-gaze reads it as the source "own". It touches /dev/shm/frametop-eyes-want every
second, and ft-eyegrab copies frames only while someone does.

The gaze service (gaze/ft-gazed) runs it in the dev container, with --watch-stdin (it quits
when its stdin closes), while Eye tracker is Own tracker or the gaze probe uses it. The
probe calibrates and teaches it; the gaze pointer's nudges reach it as clicks, from ft-gazed.
By hand: distrobox enter dev -- python3 gaze/tracker/ft-eyes -v

Calibration: ~/.local/state/frametop/gaze/eyes/calibration.json, made by the gaze probe's
calibration with the tracker toggle on Own tracker (or lab/ft-eyes-score --save CAPTURE).
Each eye's shift since the calibration (eyes_model.Shift, taught by clicks) is kept in
state.json next to it. After a restart, or frames stopping for GAP seconds (the headset
off), the next click starts that eye's shift over, since the headset may sit differently now.

Control socket: abstract datagram "@ft_eyes"; each command gets one reply line.
  status                       JSON: calibration, and per eye its shift, clicks, whether a
                               glint jump is applied, and "reseat" (the next click starts
                               the shift over: the probe asks for a one-dot check then)
  calib-start                  a new calibration: collect dots from now on
  calib-point T0 T1 YAW PITCH  a dot you looked at from T0 to T1 (CLOCK_MONOTONIC_RAW) in
                               that direction (head-relative degrees): "ok N0 N1 SD0 SD1"
                               (frames and spread in px per eye, right first) or "fail WHY"
  calib-fit                    fit the dots, save, and start the shifts and clicks over
  click T YAW PITCH            a click: you were looking there just before T; teaches the
                               shift: "ok DX0 DY0 DX1 DY1" (the shift each eye measured)

Environment, for replays (lab/ft-eyes-e2e): FT_EYES_CAMS, FT_EYES_GAZE, FT_EYES_STATE (the
state folder), FT_EYES_SOCKET (the socket's name). With FT_EYES_CAMS set, the want file isn't
touched.

/dev/shm/frametop-eyes-gaze, 128 bytes, little-endian (mirrored in ft-gaze.cpp):
   0 u32 seq      odd while it's being written: read it before and after, retry if it moved
   4 u32 version  1
   8 f64 t        the newest frame's time, CLOCK_MONOTONIC_RAW seconds
  16 f32 yaw, pitch  the gaze, head-relative degrees (yaw +left, pitch +up), eyes averaged
  24 u32 flags    bit 0: right eye in it, 1: left eye in it, 2: right shift from clicks, 3: left
  28 u32 n        samples published
  32 f32 x4       right yaw, pitch, left yaw, pitch (NaN when that eye isn't seen)
  48 f32 x4       each eye's shift since the calibration, pixels: right x, y, left x, y
  64 f32 x4       pupil centre, pixels: right x, y, left x, y
"""
import json
import math
import mmap
import os
import socket
import struct
import sys
import threading
import time
from collections import deque
from pathlib import Path

import numpy as np

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

CAMS = os.environ.get("FT_EYES_CAMS", "/dev/shm/frametop-eyes-cams")  # overrides for replays
OUT = os.environ.get("FT_EYES_GAZE", "/dev/shm/frametop-eyes-gaze")
WANT = None if "FT_EYES_CAMS" in os.environ else Path("/dev/shm/frametop-eyes-want")
OUT_SIZE = 128
CALIBRATION = Path(os.environ.get("FT_EYES_STATE", Path.home() / ".local/state/frametop/gaze/eyes")) / "calibration.json"
W, H = 512, 400
FRESH = 0.03          # an eye's reading counts toward the output for this long (s)
LOST_EVERY = 3        # while an eye is lost, search the whole frame only every 3rd frame
EYES = ("right", "left")
STATE = CALIBRATION.parent / "state.json"
CLICKS = CALIBRATION.parent / "clicks.jsonl"
SOCKET = "\0" + os.environ.get("FT_EYES_SOCKET", "ft_eyes")
HISTORY = 12.0        # seconds of pupil positions kept per eye, for dots and clicks (the pointer's
                      # clicks come from ft-gazed up to 10 s after the look)
GAP = 3.0             # s without frames: the headset was off, and may sit differently now
CLICK_BEFORE = 0.3    # a click's frames: the 300 ms before it (like the probe's fixation)
CALIB_MIN = 15        # frames an eye needs in a calibration dot's window
CALIB_SPREAD = 4.0    # px: more than this and the eye moved during the dot


class Cams:
    """The shared frames. Header and entry layout: ft-eyegrab.c, share_head_t/share_entry_t."""

    def __init__(self):
        fd = os.open(CAMS, os.O_RDONLY)
        try:
            self.ino = os.fstat(fd).st_ino
            self.mm = mmap.mmap(fd, 0, prot=mmap.PROT_READ)
        finally:
            os.close(fd)
        magic, version, w, h, self.slots, self.esize = struct.unpack_from("<6I", self.mm, 0)
        if magic != 0x31434546 or version != 1 or (w, h) != (W, H):
            raise RuntimeError(f"{CAMS}: unexpected header")

    def replaced(self):
        """ft-eyegrab restarted: it makes a new file, and this one is stale."""
        try:
            return os.stat(CAMS).st_ino != self.ino
        except OSError:
            return True

    def count(self, cam):
        return struct.unpack_from("<Q", self.mm, 24 + 8 * cam)[0]

    def tracker(self):
        return struct.unpack_from("<I", self.mm, 40)[0]

    def frame(self, cam, n):
        """Frame n of a camera as (time, array), or None if it was overwritten meanwhile."""
        off = 64 + (cam * self.slots + n % self.slots) * self.esize
        for _ in range(3):
            seq = struct.unpack_from("<Q", self.mm, off)[0]
            if seq & 1:
                continue
            img = np.frombuffer(self.mm, np.uint8, W * H, off + 64).reshape(H, W).copy()
            t, got = struct.unpack_from("<dQ", self.mm, off + 8)
            if struct.unpack_from("<Q", self.mm, off)[0] == seq and got == n:
                return t, img
        return None


class Out:
    def __init__(self):
        fd = os.open(OUT, os.O_RDWR | os.O_CREAT | os.O_NOFOLLOW, 0o600)
        try:
            os.ftruncate(fd, OUT_SIZE)
            self.mm = mmap.mmap(fd, OUT_SIZE)
        finally:
            os.close(fd)
        self.seq = (struct.unpack_from("<I", self.mm, 0)[0] + 1) & ~1  # even: at rest
        self.n = 0

    def write(self, t, gaze, flags, eyes, slips, pupils):
        struct.pack_into("<I", self.mm, 0, self.seq + 1)  # odd while writing
        self.n += 1
        struct.pack_into("<IdffII12f", self.mm, 4, 1, t, gaze[0], gaze[1], flags, self.n,
                         *eyes, *slips, *pupils)
        self.seq = (self.seq + 2) & 0xFFFFFFFE
        struct.pack_into("<I", self.mm, 0, self.seq)


class Eye:
    def __init__(self, eye):
        self.eye = eye
        self.cal = None
        self.slip = None
        self.shift = eyes_model.Shift()
        self.history = deque()  # (t, x, y, glint mid x, y)
        self.last = None        # the previous pupil (window hint)
        self.gaze = None        # (t, yaw, pitch, x, y)
        self.lost = 0
        self.frames = self.found = 0
        self.work = 0.0
        self.last_t = None      # the previous frame's time

    def use(self, cal, shift=None):
        self.cal = cal if cal is not None and cal.has("pupil", self.eye) else None
        self.slip = eyes_model.SlipTracker(cal, self.eye, window=eyes_model.JUMP_WINDOW) if self.cal else None
        self.shift = shift or eyes_model.Shift()
        self.gaze = None

    def feed(self, t, img):
        self.frames += 1
        if self.last_t is not None and t - self.last_t > GAP:
            self.shift.reseat()
        self.last_t = t
        if self.last is None:
            self.lost += 1
            if self.lost % LOST_EVERY:
                return
        t0 = time.perf_counter()
        p = eyes_pupil.find_pupil(img, self.last)
        self.last = p
        if p is not None:
            self.found += 1
            self.lost = 0
            pair = eyes_pupil.glint_pair(p)
            mid = eyes_model.pair_mid(pair) if pair else (math.nan, math.nan)
            self.history.append((t, p["x"], p["y"], mid[0], mid[1]))
            while self.history and self.history[0][0] < t - HISTORY:
                self.history.popleft()
            if self.cal:
                if pair:
                    self.slip.add(t, (p["x"], p["y"]), mid)
                self.shift.glint(self.slip.get(), t)
                g = self.cal.gaze(self.eye, p["x"], p["y"], self.shift.value)
                self.gaze = (t, float(g[0]), float(g[1]), p["x"], p["y"])
        self.work += time.perf_counter() - t0

    def window(self, t0, t1):
        """Median pupil and glint midpoint over [t0, t1], the frame count, and the spread."""
        rows = np.array([r for r in self.history if t0 <= r[0] <= t1]).reshape(-1, 5)
        if len(rows) == 0:
            return None
        pupil = np.median(rows[:, 1:3], axis=0)
        spread = float(np.median(np.hypot(*(rows[:, 1:3] - pupil).T)))
        mids = rows[~np.isnan(rows[:, 3]), 3:5]
        mid = np.median(mids, axis=0) if len(mids) >= 3 else None
        return dict(pupil=pupil, mid=mid, n=len(rows), spread=spread)


class Tracker:
    def __init__(self):
        self.eyes = [Eye(0), Eye(1)]
        self.cal = None
        self.dots = []  # calibration dots so far, in ft-eyes-score's click form
        self.calibrating = False
        if CALIBRATION.exists():
            self.cal = eyes_model.Calibration.load(CALIBRATION)
            if not self.cal.spread:
                self.cal.spread = spread_from_dots(self.cal)
        shifts = {}
        try:
            d = json.loads(STATE.read_text())
            if self.cal and d.get("calibration") == self.cal.info.get("made"):
                shifts = {int(k): eyes_model.Shift.from_json(v) for k, v in d.get("shift", {}).items()}
        except (OSError, ValueError):
            pass
        for e in self.eyes:
            e.use(self.cal, shifts.get(e.eye))
            # Kept from the last run, but the headset may have been off since: the first
            # click starts the history over (the saved shift is used until then).
            e.shift.reseat()

    def save_state(self):
        d = {"calibration": self.cal.info.get("made") if self.cal else None,
             "shift": {e.eye: e.shift.to_json() for e in self.eyes}}
        tmp = STATE.with_suffix(".tmp")
        tmp.write_text(json.dumps(d))
        tmp.replace(STATE)

    def command(self, line):
        w = line.split()
        if not w:
            return "fail empty"
        if w[0] == "status":
            return json.dumps(self.status())
        if w[0] == "calib-start":
            self.dots, self.calibrating = [], True
            return "ok"
        if w[0] == "calib-point" and len(w) == 5:
            if not self.calibrating:
                return "fail no calibration started"
            t0, t1, yaw, pitch = map(float, w[1:])
            got = {e.eye: e.window(t0, t1) for e in self.eyes}
            for c, g in got.items():
                if g is None or g["n"] < CALIB_MIN:
                    return f"fail the {EYES[c]} eye was seen in only {0 if g is None else g['n']} frames"
                if g["spread"] > CALIB_SPREAD:
                    return f"fail the {EYES[c]} eye moved ({g['spread']:.1f} px)"
            self.dots.append(dict(truth=(yaw, pitch), eye={c: dict(pupil=g["pupil"], mid=g["mid"]) for c, g in got.items()}))
            return "ok {} {} {:.2f} {:.2f}".format(got[0]["n"], got[1]["n"], got[0]["spread"], got[1]["spread"])
        if w[0] == "calib-fit":
            if len(self.dots) < eyes_model.MIN_CLICKS:
                return f"fail only {len(self.dots)} dots (need {eyes_model.MIN_CLICKS})"
            cal = eyes_model.Calibration.fit(self.dots, {"made": time.strftime("%Y-%m-%d %H:%M:%S"),
                                                       "dots": len(self.dots), "from": "probe calibration"})
            if not all(cal.has(n, c) for n in ("pupil", "where") for c in (0, 1)):
                return "fail not enough dots with both eyes"
            errs = [float(np.hypot(*(np.mean([cal.gaze(c, *k["eye"][c]["pupil"]) for c in (0, 1)], axis=0)
                                     - k["truth"]))) for k in self.dots]
            if CALIBRATION.exists():
                CALIBRATION.replace(CALIBRATION.with_name(time.strftime("calibration-%Y%m%d-%H%M%S.json")))
            cal.save(CALIBRATION)
            self.cal, self.calibrating = cal, False
            for e in self.eyes:
                e.use(cal)
            self.save_state()
            with open(CALIBRATION.with_name("calibration-dots.jsonl"), "a") as f:
                for k in self.dots:
                    f.write(json.dumps({"made": cal.info["made"], "truth": k["truth"],
                                        "eye": {c: {"pupil": v["pupil"].tolist(),
                                                    "mid": None if v["mid"] is None else v["mid"].tolist()}
                                                for c, v in k["eye"].items()}}) + "\n")
            return (f"ok {len(self.dots)} dots, fit median {np.median(errs):.2f} deg, eyes "
                    + ", ".join(f"{EYES[c]} {cal.spread[c]:.2f}" for c in sorted(cal.spread)))
        if w[0] == "click" and len(w) == 4:
            if not self.cal:
                return "fail not calibrated"
            t, yaw, pitch = map(float, w[1:])
            out, rec = [], {"time": time.time(), "t": t, "truth": [yaw, pitch], "eyes": {}}
            for e in self.eyes:
                g = e.window(t - CLICK_BEFORE, t)
                if g is None or g["n"] < 5 or not e.cal:
                    out += ["nan", "nan"]
                    continue
                d = self.cal.click_shift(e.eye, g["pupil"], (yaw, pitch))
                before = e.shift.value.tolist()
                e.shift.click(d, e.slip.value if e.slip else None)
                rec["eyes"][e.eye] = {"pupil": g["pupil"].tolist(), "measured": d.tolist(), "before": before,
                                      "after": e.shift.value.tolist()}
                out += [f"{d[0]:.2f}", f"{d[1]:.2f}"]
            self.save_state()
            with open(CLICKS, "a") as f:
                f.write(json.dumps(rec) + "\n")
            return "ok " + " ".join(out)
        return f"fail unknown command {w[0]}"

    def status(self):
        return {"calibration": self.cal.info if self.cal else None, "calibrating": self.calibrating,
                "dots": len(self.dots),
                "eyes": {EYES[e.eye]: {"shift": e.shift.value.tolist(), "clicks": len(e.shift.meas),
                                       "jump": bool(np.any(e.shift.jump)),
                                       "reseat": e.shift.reseated} for e in self.eyes}}


def spread_from_dots(cal):
    """The eyes' fit spreads for a calibration saved without them, from its dots in
    calibration-dots.jsonl ({} if they aren't there: the eyes are then weighted alike)."""
    try:
        lines = CALIBRATION.with_name("calibration-dots.jsonl").read_text().splitlines()
    except OSError:
        return {}
    dots = []
    for line in lines:
        d = json.loads(line)
        if d.get("made") == cal.info.get("made"):
            dots.append(dict(truth=d["truth"], eye={
                int(c): dict(pupil=np.array(v["pupil"]), mid=None if v["mid"] is None else np.array(v["mid"]))
                for c, v in d["eye"].items()}))
    return eyes_model.Calibration.fit(dots).spread if dots else {}


class Want:
    """Touches the want file every second, so ft-eyegrab keeps copying frames."""

    def __init__(self):
        self.at = 0.0

    def __call__(self):
        if WANT is None or time.monotonic() - self.at < 1.0:
            return
        self.at = time.monotonic()
        try:
            fd = os.open(WANT, os.O_WRONLY | os.O_CREAT | os.O_NOFOLLOW | os.O_CLOEXEC, 0o600)
            os.utime(fd)
            os.close(fd)
        except OSError as e:
            print(f"ft-eyes: {WANT}: {e}", file=sys.stderr, flush=True)


def wait_for_cams(sock, tracker, want):
    while True:
        want()
        try:
            return Cams()
        except (OSError, ValueError, RuntimeError):
            serve(sock, tracker)
            time.sleep(0.2)


def serve(sock, tracker):
    while True:
        try:
            data, addr = sock.recvfrom(512)
        except BlockingIOError:
            return
        try:
            reply = tracker.command(data.decode(errors="replace").strip())
        except Exception as ex:  # a bad command must not take the tracker down
            reply = f"fail {type(ex).__name__}: {ex}"
        if addr:
            try:
                sock.sendto(reply.encode(), addr)
            except OSError:
                pass


def main():
    verbose = "-v" in sys.argv
    if "--watch-stdin" in sys.argv:
        # Run by ft-gazed through distrobox, which doesn't pass a stop on: quit when our
        # stdin (its pipe) closes.
        def watch():
            while sys.stdin.buffer.read(4096):
                pass
            os._exit(0)
        threading.Thread(target=watch, daemon=True).start()
    want = Want()
    CALIBRATION.parent.mkdir(parents=True, exist_ok=True)
    tracker = Tracker()
    eyes = tracker.eyes
    out = Out()
    sock = socket.socket(socket.AF_UNIX, socket.SOCK_DGRAM)
    sock.bind(SOCKET)
    sock.setblocking(False)
    cal = tracker.cal
    print("ft-eyes: " + (f"calibration from {cal.info.get('made')} ({cal.info.get('from', cal.info.get('capture'))})"
                           if cal else "not calibrated: run the probe's calibration with the Own tracker")
          + f"; waiting for {CAMS}", file=sys.stderr, flush=True)
    cams = wait_for_cams(sock, tracker, want)
    print("ft-eyes: frames found, tracking", file=sys.stderr, flush=True)
    seen = [cams.count(0), cams.count(1)]
    report = time.monotonic()
    while True:
        serve(sock, tracker)
        want()
        new = False
        for c in (0, 1):
            n = cams.count(c)
            if n == seen[c]:
                continue
            seen[c] = n
            got = cams.frame(c, n - 1)  # only the newest: never fall behind
            if got:
                eyes[c].feed(*got)
                new = True
        if new:
            latest = max((e.gaze[0] for e in eyes if e.gaze), default=None)
            use = [e for e in eyes if e.gaze and latest - e.gaze[0] < FRESH]
            if use:
                yaw, pitch = tracker.cal.combine({e.eye: e.gaze[1:3] for e in use})
                flags, per, shifts, pupils = 0, [], [], []
                for i, e in enumerate(eyes):
                    fresh = e in use
                    flags |= (1 << i) if fresh else 0
                    flags |= (4 << i) if e.shift.meas else 0
                    per += [e.gaze[1], e.gaze[2]] if fresh else [math.nan, math.nan]
                    shifts += [float(v) for v in e.shift.value]
                    pupils += [e.gaze[3], e.gaze[4]] if fresh else [math.nan, math.nan]
                out.write(latest, (yaw, pitch), flags, per, shifts, pupils)
        else:
            time.sleep(0.001)
        now = time.monotonic()
        if now - report >= 5:
            if verbose:
                parts = []
                for e in eyes:
                    s = e.shift.value
                    parts.append(f"{EYES[e.eye]} {e.frames / 5:.0f} fps, found {e.found / max(e.frames, 1):.0%}, "
                                 f"{e.work / max(e.frames, 1) * 1000:.2f} ms/frame, shift ({s[0]:+.1f},{s[1]:+.1f})"
                                 f" from {len(e.shift.meas)} clicks" + (", jump" if np.any(e.shift.jump) else ""))
                    e.frames = e.found = 0
                    e.work = 0.0
                print("ft-eyes: " + "; ".join(parts), file=sys.stderr, flush=True)
            report = now
            if cams.replaced():
                print("ft-eyes: frames went away; waiting", file=sys.stderr, flush=True)
                cams = wait_for_cams(sock, tracker, want)
                seen = [cams.count(0), cams.count(1)]


if __name__ == "__main__":
    try:
        main()
    except KeyboardInterrupt:
        pass
