#!/usr/bin/python3 -I
"""mood-face-auth — Mood face unlock (YuNet + SFace via OpenCV).

PAM (pam_exec.so quiet, env PAM_TYPE/PAM_USER/PAM_SERVICE): exit 0 when the enrolled face is seen in
>= 2 frames, non-zero otherwise (the password prompt follows). Every "not applicable" case exits
before cv2 is imported. Hard time limit via a watchdog thread + SIGALRM (default action kills).

CLI:  --probe <user>        same checks as the lock screen PAM step (exit code only)
      --test <user>         run the matcher regardless of settings, print JSON
      --enroll              capture a guided enrolment session, JSON lines on stdout
      --selftest <img>...   detect + embed image files, print JSON (pairwise similarity for 2+)
      --cameras             list capture cameras as JSON
"""
import json
import os
import pwd
import signal
import sys
import syslog
import threading
import time

sys.path.insert(0, "/usr/lib/mood/runtime")
import moodface as mf  # noqa: E402  (no cv2 at import)

T0 = time.monotonic()


def log(msg):
    try:
        syslog.openlog("mood-face", 0, syslog.LOG_AUTHPRIV)
        syslog.syslog(syslog.LOG_INFO, msg)
    except Exception:
        pass


def deadline(seconds, what):
    def fire():
        log(f"{what}: timed out after {seconds:.0f}s")
        os._exit(3)
    t = threading.Timer(seconds, fire)
    t.daemon = True
    t.start()
    signal.signal(signal.SIGALRM, signal.SIG_DFL)  # backstop even if a C call holds the GIL
    signal.alarm(int(seconds) + 2)


def typed_password(user):
    """gtklock's face module touches this while the password box has text: skip the camera then."""
    try:
        uid = pwd.getpwnam(user).pw_uid
        st = os.stat(f"/run/user/{uid}/mood-face-typed")
        return st.st_uid == uid and time.time() - st.st_mtime < 600
    except (KeyError, OSError):
        return False


def gate(user, service):
    """-> (cfg, enrolment, camera) or exits non-zero without touching cv2."""
    cfg = mf.config()
    if not cfg.get("enabled"):
        sys.exit(10)
    if not mf.service_allowed(service, cfg):
        sys.exit(11)
    enr = mf.enrolment(user)
    if not enr:
        sys.exit(12)
    cam = mf.pick_camera(cfg)
    if not cam:
        sys.exit(13)
    if service == "gtklock" and typed_password(user):
        sys.exit(14)
    return cfg, enr, cam


def authenticate(user, service, label):
    cfg, enr, cam = gate(user, service)
    deadline(cfg["timeout"] + 3, f"{label} {service} user={user}")
    try:
        r = mf.match_camera(cam, enr["embeddings"], cfg["threshold"], cfg["timeout"] - (time.monotonic() - T0) + 0.5)
    except Exception as e:
        r = {"ok": False, "score": 0, "frames": 0, "error": str(e)[:120]}
    log(f"{label} service={service} user={user} result={'success' if r['ok'] else 'fail'} score={r['score']:.3f} "
        f"frames={r['frames']} t={time.monotonic() - T0:.2f}s{' err=' + r['error'] if r.get('error') else ''}")
    return 0 if r["ok"] else 1


def emit(d):
    sys.stdout.write(json.dumps(d) + "\n")
    sys.stdout.flush()


def preview(eng, img, face, color):
    cv2 = eng.cv2
    h, w = img.shape[:2]
    s = 240 / w
    small = cv2.resize(img, (240, round(h * s)))
    if face is not None:
        x, y, fw, fh = [int(v * s) for v in face[:4]]
        cv2.rectangle(small, (x, y), (x + fw, y + fh), color, 2)
    ok, jpg = cv2.imencode(".jpg", small, [cv2.IMWRITE_JPEG_QUALITY, 60])
    import base64
    return base64.b64encode(jpg.tobytes()).decode() if ok else ""


def enroll():
    """Guided capture: 4 straight, 3 turned left, 3 turned right (with fallbacks), JSON lines out."""
    cfg = mf.config()
    cam = mf.pick_camera(cfg)
    if not cam:
        emit({"error": "No camera found"})
        return 2
    deadline(40, "enroll")
    eng = mf.Engine(score=0.8)
    cap = mf.open_camera(cam)
    if not cap:
        emit({"error": "The camera is busy or can't be opened"})
        return 2
    want = {"c": 4, "l": 3, "r": 3}
    got = {"c": [], "l": [], "r": []}
    tips = {"c": "Look at the screen", "l": "Turn slightly left", "r": "Turn slightly right"}
    start, last_keep, last_pv = time.monotonic(), 0, 0
    try:
        while True:
            now = time.monotonic()
            el = now - start
            done = sum(min(len(got[k]), want[k]) for k in want)
            if done >= 10 and el >= 4 or el > 25:
                break
            ok, img = cap.read()
            if not ok or img is None:
                time.sleep(0.05)
                continue
            img = eng.prep(img)
            faces = eng.detect(img)
            need = next((k for k in ("c", "l", "r") if len(got[k]) < want[k]), "c")
            hint, face, color = tips[need], None, (80, 80, 255)
            if not faces:
                hint = "Move into the frame" if img.mean() > 45 else "Find more light"
            elif len(faces) > 1 and faces[1][2] * faces[1][3] > 0.35 * faces[0][2] * faces[0][3]:
                hint, face = "Only one face please", faces[0]
            else:
                face = faces[0]
                if face[2] < img.shape[1] * 0.14:
                    hint = "Come a bit closer"
                elif face[14] < 0.88:
                    hint = "Hold still" if img.mean() > 45 else "Find more light"
                elif now - last_keep >= 0.22:
                    y = mf.yaw(face)
                    b = "c" if abs(y) < 0.12 else ("l" if y > 0 else "r")
                    if el > 12 and len(got[need]) < want[need]:
                        b = need  # can't turn that far? take what we get
                    if len(got[b]) < want[b] + 2:
                        v = eng.embed(img, face)
                        if all(float(v @ u) < 0.985 for u in got[b]):
                            got[b].append(v)
                            last_keep = now
                            color = (120, 220, 90)
            if now - last_pv >= 0.2:
                last_pv = now
                pct = round(min(100, sum(min(len(got[k]), want[k]) for k in want) * 10))
                emit({"progress": pct, "hint": hint, "preview": preview(eng, img, face, color), "face": face is not None})
    finally:
        cap.release()
    embs = [v for k in got for v in got[k]]
    if len(embs) < 6:
        emit({"error": "Couldn't see your face clearly enough. Try facing a window or turning on a light."})
        return 1
    import numpy as np
    m = np.mean(embs, axis=0)
    m /= np.linalg.norm(m)
    embs = [v for v in embs if float(v @ m) > 0.55][:mf.MAX_EMB]  # drop stray frames (someone walking past)
    emit({"done": True, "progress": 100, "embeddings": [[round(float(x), 6) for x in v] for v in embs], "camera": cam})
    return 0


def selftest(paths):
    import cv2
    eng = mf.Engine()
    out, vecs = [], []
    for p in paths:
        img = cv2.imread(p)
        if img is None:
            out.append({"image": p, "error": "unreadable"})
            vecs.append(None)
            continue
        img = eng.prep(img)
        faces = eng.detect(img)
        r = {"image": p, "faces": len(faces)}
        if faces:
            f = faces[0]
            v = eng.embed(img, f)
            vecs.append(v)
            r.update(box=[round(float(x)) for x in f[:4]], score=round(float(f[14]), 3), yaw=round(mf.yaw(f), 3), embedding=[round(float(x), 5) for x in v])
        else:
            vecs.append(None)
        out.append(r)
    res = {"results": out, "ms": round((time.monotonic() - T0) * 1000)}
    if len(paths) > 1:
        res["similarity"] = [[None if a is None or b is None else round(float(a @ b), 4) for b in vecs] for a in vecs]
    print(json.dumps(res))
    return 0


def main():
    a = sys.argv[1:]
    if not a:
        if os.environ.get("PAM_TYPE") != "auth":
            return 1
        return authenticate(os.environ.get("PAM_USER", ""), os.environ.get("PAM_SERVICE", ""), "pam")
    if a[0] == "--probe" and len(a) == 2:
        return authenticate(a[1], "gtklock", "probe")
    if a[0] == "--test" and len(a) == 2:
        cfg = mf.config()
        enr, cam = mf.enrolment(a[1]), mf.pick_camera(cfg)
        if not enr or not cam:
            print(json.dumps({"ok": False, "error": "No camera found" if enr else "No face enrolled"}))
            return 1
        deadline(cfg["timeout"] + 5, "test")
        r = mf.match_camera(cam, enr["embeddings"], cfg["threshold"], cfg["timeout"])
        r["threshold"] = cfg["threshold"]
        log(f"test user={a[1]} result={'success' if r['ok'] else 'fail'} score={r['score']:.3f}")
        print(json.dumps(r))
        return 0 if r["ok"] else 1
    if a[0] == "--enroll":
        return enroll()
    if a[0] == "--selftest" and len(a) > 1:
        return selftest(a[1:])
    if a[0] == "--cameras":
        print(json.dumps(mf.cameras()))
        return 0
    print(__doc__.strip(), file=sys.stderr)
    return 2


if __name__ == "__main__":
    sys.exit(main())
