#!/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 select
import socket
import struct
import sys
import threading
import time
from collections import deque
from pathlib import Path

# One thread for numpy's BLAS and OpenMP, set before numpy loads: it would start one per core
# (8 here) for small arrays that never need them. eyes_pupil keeps OpenCV to one as well.
for _var in ("OPENBLAS_NUM_THREADS", "OMP_NUM_THREADS"):
    os.environ.setdefault(_var, "1")

import numpy as np  # noqa: E402

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
# Waiting for frames. They come only as a counter in shared memory, so there's nothing to block
# on: sleep until a camera's next frame is due, then look every POLL. ft-eyegrab passes each one
# on within a ms or two of when its camera starts the one after, so they arrive 11.1 ms apart,
# give or take that.
PERIOD = 1 / 90       # s between a camera's frames
EARLY = 0.002         # start looking this long before a frame is due
POLL = 0.001          # then this often until it comes
STALLED = 0.1         # s without a frame: that camera stopped (headset off), and isn't waited for
IDLE_POLL = 0.02      # how often to look while both are stopped


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)
            select.select([sock], [], [], 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 below_steamvr():
    """Nice 10 and SCHED_BATCH, for this thread and those it starts. ft-eyes runs in the dev
    container's podman scope, where the gaze service's unit doesn't reach it, so it ran at
    nice 0 on the cores vrcompositor and vrserver use. Batch also lets a waking ft-eyes wait
    for the running task's turn instead of taking the core: a frame a few ms late costs the
    gaze little, a late compositor frame costs a dropped frame in the headset."""
    try:
        os.setpriority(os.PRIO_PROCESS, 0, max(os.getpriority(os.PRIO_PROCESS, 0), 10))
    except OSError as e:
        print(f"ft-eyes: nice: {e}", file=sys.stderr, flush=True)
    try:
        os.sched_setscheduler(0, os.SCHED_BATCH, os.sched_param(0))
    except (OSError, AttributeError) as e:
        print(f"ft-eyes: SCHED_BATCH: {e}", file=sys.stderr, flush=True)


def main():
    below_steamvr()
    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()
    arrived = [0.0, 0.0]  # when each camera's newest frame was seen (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
            arrived[c] = time.monotonic()
            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)
        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)]
        # Until the next frame is due (a command on the socket wakes us sooner).
        wait = IDLE_POLL
        for c in (0, 1):
            if now - arrived[c] < STALLED:
                wait = min(wait, max(arrived[c] + PERIOD - EARLY - now, POLL))
        select.select([sock], [], [], wait)


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