"""Durable operational pause state for nutrition-coaching customers."""

from __future__ import annotations

import fcntl
import hashlib
import json
import os
import stat
import tempfile
from collections.abc import Iterator
from contextlib import contextmanager
from datetime import date
from pathlib import Path
from typing import cast

_SCHEMA = "customer-service-state-v1"


class CustomerServiceStateError(RuntimeError):
    """The operational state ledger cannot be trusted or updated."""


class CustomerServiceStateStore:
    """Private, atomic per-customer pause state independent of consent."""

    def __init__(self, path: Path) -> None:
        self.path = path
        self.lock_path = path.with_suffix(f"{path.suffix}.lock")

    def _reject_symlink_ancestors(self) -> None:
        current = self.path.absolute().parent
        while True:
            try:
                metadata = current.lstat()
            except FileNotFoundError:
                pass
            else:
                if stat.S_ISLNK(metadata.st_mode):
                    raise CustomerServiceStateError(
                        "service state ancestor must not be a symlink"
                    )
            if current == current.parent:
                return
            current = current.parent

    @staticmethod
    def _digest(payload: dict[str, object]) -> str:
        encoded = json.dumps(
            payload,
            ensure_ascii=False,
            sort_keys=True,
            separators=(",", ":"),
        ).encode("utf-8")
        return hashlib.sha256(encoded).hexdigest()

    @classmethod
    def _empty(cls) -> dict[str, object]:
        unsigned: dict[str, object] = {
            "schema": _SCHEMA,
            "states": {},
        }
        return {**unsigned, "payload_digest": cls._digest(unsigned)}

    def ensure(self) -> None:
        with self._locked():
            if self.path.exists():
                self._read()
                return
            self._write(self._empty())

    def is_paused(self, customer_key: str) -> bool:
        if not self.path.exists():
            raise CustomerServiceStateError("service state ledger is missing")
        with self._locked():
            payload = self._read()
            states = payload["states"]
            if not isinstance(states, dict):
                raise CustomerServiceStateError("service states are invalid")
            row = cast(dict[str, object], states).get(customer_key)
            if not isinstance(row, dict):
                return False
            return cast(dict[str, object], row).get("state") == "paused"

    def set_paused(
        self,
        customer_key: str,
        *,
        paused: bool,
        updated_on: str,
    ) -> bool:
        if not customer_key:
            raise CustomerServiceStateError("customer key is required")
        try:
            if date.fromisoformat(updated_on).isoformat() != updated_on:
                raise ValueError
        except (TypeError, ValueError) as exc:
            raise CustomerServiceStateError(
                "service state audit date is invalid"
            ) from exc
        with self._locked():
            if not self.path.exists():
                raise CustomerServiceStateError("service state ledger is missing")
            payload = self._read()
            states_value = payload["states"]
            if not isinstance(states_value, dict):
                raise CustomerServiceStateError("service states are invalid")
            states = dict(states_value)
            previous = states.get(customer_key)
            desired = "paused" if paused else "active"
            if isinstance(previous, dict) and previous.get("state") == desired:
                return False
            revision = (
                int(previous.get("revision", 0)) + 1
                if isinstance(previous, dict)
                else 1
            )
            states[customer_key] = {
                "state": desired,
                "revision": revision,
                "updated_on": updated_on,
            }
            unsigned: dict[str, object] = {
                "schema": _SCHEMA,
                "states": states,
            }
            self._write({
                **unsigned,
                "payload_digest": self._digest(unsigned),
            })
            return True

    @contextmanager
    def _locked(self) -> Iterator[None]:
        try:
            self._reject_symlink_ancestors()
            self.path.parent.mkdir(mode=0o700, parents=True, exist_ok=True)
            self._reject_symlink_ancestors()
            self.path.parent.chmod(0o700)
            flags = os.O_CREAT | os.O_RDWR
            if hasattr(os, "O_NOFOLLOW"):
                flags |= os.O_NOFOLLOW
            descriptor = os.open(self.lock_path, flags, 0o600)
            metadata = os.fstat(descriptor)
            if not stat.S_ISREG(metadata.st_mode) or metadata.st_nlink != 1:
                os.close(descriptor)
                raise CustomerServiceStateError("service state lock is unsafe")
            os.fchmod(descriptor, 0o600)
        except OSError as exc:
            raise CustomerServiceStateError("service state lock unavailable") from exc
        try:
            fcntl.flock(descriptor, fcntl.LOCK_EX)
            yield
        finally:
            fcntl.flock(descriptor, fcntl.LOCK_UN)
            os.close(descriptor)

    def _read(self) -> dict[str, object]:
        try:
            metadata = self.path.lstat()
            if (
                not stat.S_ISREG(metadata.st_mode)
                or metadata.st_nlink != 1
                or stat.S_IMODE(metadata.st_mode) != 0o600
            ):
                raise CustomerServiceStateError("service state file is unsafe")
            payload = json.loads(self.path.read_text(encoding="utf-8"))
        except CustomerServiceStateError:
            raise
        except (OSError, TypeError, ValueError, json.JSONDecodeError) as exc:
            raise CustomerServiceStateError("service state ledger is unreadable") from exc
        if not isinstance(payload, dict) or set(payload) != {
            "schema",
            "states",
            "payload_digest",
        }:
            raise CustomerServiceStateError("service state ledger shape is invalid")
        if payload.get("schema") != _SCHEMA or not isinstance(
            payload.get("states"),
            dict,
        ):
            raise CustomerServiceStateError("service state ledger schema is invalid")
        unsigned: dict[str, object] = {
            "schema": payload["schema"],
            "states": payload["states"],
        }
        digest = payload.get("payload_digest")
        if not isinstance(digest, str) or digest != self._digest(unsigned):
            raise CustomerServiceStateError("service state ledger digest is invalid")
        for customer_key, row in payload["states"].items():
            updated_on = row.get("updated_on") if isinstance(row, dict) else None
            try:
                valid_updated_on = (
                    isinstance(updated_on, str)
                    and date.fromisoformat(updated_on).isoformat() == updated_on
                )
            except ValueError:
                valid_updated_on = False
            if (
                not isinstance(customer_key, str)
                or not customer_key
                or not isinstance(row, dict)
                or set(row) != {"state", "revision", "updated_on"}
                or row.get("state") not in {"active", "paused"}
                or not isinstance(row.get("revision"), int)
                or row["revision"] < 1
                or not valid_updated_on
            ):
                raise CustomerServiceStateError("service state row is invalid")
        return payload

    def _write(self, payload: dict[str, object]) -> None:
        data = (
            json.dumps(
                payload,
                ensure_ascii=False,
                sort_keys=True,
                separators=(",", ":"),
            )
            + "\n"
        ).encode("utf-8")
        descriptor = -1
        temporary_path = ""
        try:
            descriptor, temporary_path = tempfile.mkstemp(
                prefix=f".{self.path.name}.",
                dir=self.path.parent,
            )
            os.fchmod(descriptor, 0o600)
            with os.fdopen(descriptor, "wb", closefd=True) as stream:
                descriptor = -1
                stream.write(data)
                stream.flush()
                os.fsync(stream.fileno())
            os.replace(temporary_path, self.path)
            temporary_path = ""
            directory = os.open(self.path.parent, os.O_RDONLY)
            try:
                os.fsync(directory)
            finally:
                os.close(directory)
        except OSError as exc:
            raise CustomerServiceStateError("service state write failed") from exc
        finally:
            if descriptor >= 0:
                os.close(descriptor)
            if temporary_path:
                try:
                    os.unlink(temporary_path)
                except FileNotFoundError:
                    pass
