"""Privacy-controlled local face recognition for SHiRE Vision."""

import hashlib
import json
import os
from pathlib import Path

import cv2
import numpy as np
from cryptography.fernet import Fernet, InvalidToken
from PyQt5.QtCore import QObject, pyqtSignal


class EncryptedFaceStore:
    """Encrypted local biometric templates with no cloud access."""

    def __init__(self, state_dir):
        self.state_dir = Path(state_dir).expanduser()
        self.key_path = self.state_dir / "face-store.key"
        self.store_path = self.state_dir / "identities.enc"

    def _prepare(self):
        self.state_dir.mkdir(
            mode=0o700,
            parents=True,
            exist_ok=True,
        )
        os.chmod(self.state_dir, 0o700)

        if not self.key_path.exists():
            fd = os.open(
                self.key_path,
                os.O_WRONLY | os.O_CREAT | os.O_EXCL,
                0o600,
            )
            try:
                os.write(fd, Fernet.generate_key())
            finally:
                os.close(fd)

        os.chmod(self.key_path, 0o600)

    def _cipher(self):
        self._prepare()
        return Fernet(self.key_path.read_bytes())

    def load(self):
        if not self.store_path.exists():
            return {
                "schema_version": 1,
                "identities": {},
            }

        try:
            raw = self._cipher().decrypt(
                self.store_path.read_bytes()
            )
        except InvalidToken as exc:
            raise RuntimeError(
                "Encrypted face store could not be decrypted."
            ) from exc

        data = json.loads(raw.decode("utf-8"))

        if data.get("schema_version") != 1:
            raise ValueError(
                "Unsupported face-store schema."
            )

        return data

    def save(self, data):
        encrypted = self._cipher().encrypt(
            json.dumps(
                data,
                separators=(",", ":"),
            ).encode("utf-8")
        )

        temporary = self.store_path.with_suffix(".tmp")

        fd = os.open(
            temporary,
            os.O_WRONLY | os.O_CREAT | os.O_TRUNC,
            0o600,
        )
        try:
            os.write(fd, encrypted)
        finally:
            os.close(fd)

        os.replace(temporary, self.store_path)
        os.chmod(self.store_path, 0o600)


class FaceRecognitionService(QObject):
    """Detects and recognises only explicitly enrolled people."""

    status_changed = pyqtSignal(dict)

    def __init__(
        self,
        policy_path=None,
        model_dir=None,
        state_dir=None,
        similarity_threshold=0.45,
        parent=None,
    ):
        super().__init__(parent)

        root = Path(__file__).resolve().parent.parent

        self.policy_path = Path(
            policy_path
            or root / "config" / "vision_modes.json"
        )
        self.model_dir = Path(
            model_dir
            or root / "models" / "vision" / "face"
        )
        self.state_dir = Path(
            state_dir
            or Path.home()
            / ".local"
            / "share"
            / "shire"
            / "vision"
            / "faces"
        )

        self.policy = json.loads(
            self.policy_path.read_text(
                encoding="utf-8"
            )
        )

        rules = self.policy["face_recognition"]

        if rules["processing_location"] != "local_only":
            raise ValueError(
                "Face processing must remain local."
            )

        if rules["cloud_upload_allowed"]:
            raise ValueError(
                "Cloud face upload must remain prohibited."
            )

        self.enabled = bool(
            rules["enabled_by_default"]
        )
        self.similarity_threshold = float(
            similarity_threshold
        )

        self.store = EncryptedFaceStore(
            self.state_dir
        )
        self.detector = None
        self.recogniser = None

    @staticmethod
    def _normalise(embedding):
        vector = np.asarray(
            embedding,
            dtype=np.float32,
        ).reshape(-1)

        norm = float(np.linalg.norm(vector))

        if norm <= 0:
            raise ValueError(
                "Face embedding has zero magnitude."
            )

        return vector / norm

    @staticmethod
    def _sha256(path):
        digest = hashlib.sha256()

        with Path(path).open("rb") as handle:
            for block in iter(
                lambda: handle.read(1024 * 1024),
                b"",
            ):
                digest.update(block)

        return digest.hexdigest()

    def _load_models(self):
        if (
            self.detector is not None
            and self.recogniser is not None
        ):
            return

        manifest_path = (
            self.model_dir / "MODEL-MANIFEST.json"
        )
        manifest = json.loads(
            manifest_path.read_text(
                encoding="utf-8"
            )
        )

        detector_data = manifest["models"]["detector"]
        recogniser_data = manifest[
            "models"
        ]["recogniser"]

        detector_path = (
            self.model_dir / detector_data["file"]
        )
        recogniser_path = (
            self.model_dir / recogniser_data["file"]
        )

        for path, expected in (
            (detector_path, detector_data["sha256"]),
            (recogniser_path, recogniser_data["sha256"]),
        ):
            if self._sha256(path) != expected:
                raise RuntimeError(
                    f"Model checksum failed: {path.name}"
                )

        self.detector = cv2.FaceDetectorYN_create(
            str(detector_path),
            "",
            (320, 320),
            0.9,
            0.3,
            5000,
        )

        self.recogniser = (
            cv2.FaceRecognizerSF_create(
                str(recogniser_path),
                "",
            )
        )

    def set_enabled(self, enabled):
        self.enabled = bool(enabled)
        self.status_changed.emit(self.snapshot())
        return self.snapshot()

    def enrol_identity(
        self,
        display_name,
        embedding,
        greeting=None,
        approved_by_ray=False,
    ):
        if not approved_by_ray:
            raise PermissionError(
                "Face enrolment requires Ray's approval."
            )

        display_name = " ".join(
            str(display_name).strip().split()
        )

        if not display_name:
            raise ValueError(
                "A display name is required."
            )

        vector = self._normalise(embedding)
        key = display_name.casefold()

        data = self.store.load()
        data["identities"][key] = {
            "display_name": display_name,
            "greeting": (
                str(greeting).strip()
                if greeting
                else f"Hi {display_name}"
            ),
            "embedding": vector.tolist(),
        }
        self.store.save(data)

        return {
            "enrolled": True,
            "display_name": display_name,
        }

    def delete_identity(
        self,
        display_name,
        approved_by_ray=False,
    ):
        if not approved_by_ray:
            raise PermissionError(
                "Face deletion requires Ray's approval."
            )

        key = str(display_name).strip().casefold()
        data = self.store.load()
        removed = data["identities"].pop(
            key,
            None,
        )
        self.store.save(data)

        return removed is not None

    def list_identities(self):
        data = self.store.load()

        return [
            {
                "display_name": item["display_name"],
                "greeting": item["greeting"],
            }
            for item in data["identities"].values()
        ]

    def match_embedding(self, embedding):
        candidate = self._normalise(embedding)
        data = self.store.load()

        best = None

        for item in data["identities"].values():
            known = self._normalise(
                item["embedding"]
            )
            score = float(
                np.dot(candidate, known)
            )

            if best is None or score > best["score"]:
                best = {
                    "display_name": item[
                        "display_name"
                    ],
                    "greeting": item["greeting"],
                    "score": score,
                }

        if (
            best is None
            or best["score"]
            < self.similarity_threshold
        ):
            return {
                "known": False,
                "display_name": None,
                "greeting": None,
                "score": (
                    best["score"]
                    if best is not None
                    else None
                ),
            }

        return {
            "known": True,
            **best,
        }

    def recognise_frame(self, image):
        """Consume a supplied RGB/BGR frame; never open a camera."""

        if not self.enabled:
            return []

        if not isinstance(image, np.ndarray):
            raise TypeError(
                "Face recognition requires a NumPy image."
            )

        if image.ndim != 3 or image.shape[2] != 3:
            raise ValueError(
                "Expected a three-channel image."
            )

        self._load_models()

        height, width = image.shape[:2]
        self.detector.setInputSize(
            (width, height)
        )

        _result, faces = self.detector.detect(image)

        if faces is None:
            return []

        recognised = []

        for face in faces:
            aligned = self.recogniser.alignCrop(
                image,
                face,
            )
            feature = self.recogniser.feature(
                aligned
            )
            match = self.match_embedding(feature)

            recognised.append({
                "box": [
                    int(value)
                    for value in face[:4]
                ],
                **match,
            })

        return recognised

    def snapshot(self):
        return {
            "enabled": self.enabled,
            "processing_location": "local_only",
            "cloud_upload_allowed": False,
            "enrolled_count": len(
                self.list_identities()
            ),
            "camera_owned": False,
        }
