"""Private, bounded push receiver for GJC snapshots and local status controls."""

from __future__ import annotations

import errno
import fcntl
import json
import os
import selectors
import socket
import stat
import threading
import time
from dataclasses import dataclass, field
from pathlib import Path
from typing import Any

from .gjc_protocol import (
    MAX_FRAME_BYTES,
    GjcProtocolError,
    _unique_json_object,
    _uuid,
    parse_snapshot,
    valid_clock,
)
from .gjc_tracker import MAX_REGISTRY_BYTES, GjcTracker, ensure_private_directory

MAX_CLIENTS = 32
CLIENT_TIMEOUT = 10.0
CLEAR_REFUSAL_REASONS = frozenset({
    "connected", "process_alive", "process_death_unconfirmed",
})


def _acquire_private_lock(path: Path, label: str) -> int:
    """Hold a stable sidecar inode; never unlink it, including on contention."""
    ensure_private_directory(path.parent)
    fd = os.open(path, os.O_CREAT | os.O_RDWR | os.O_NOFOLLOW, 0o600)
    try:
        info = os.fstat(fd)
        if (not stat.S_ISREG(info.st_mode) or info.st_uid != os.getuid()
                or stat.S_IMODE(info.st_mode) & 0o077):
            raise PermissionError(f"{label} lock must be a private user-owned file")
        try:
            fcntl.flock(fd, fcntl.LOCK_EX | fcntl.LOCK_NB)
        except BlockingIOError as exc:
            raise RuntimeError(f"Another GJC receiver already owns this {label.lower()}") from exc
        return fd
    except BaseException:
        os.close(fd)
        raise


def _decode(data: bytes) -> Any:
    try:
        return json.loads(data, object_pairs_hook=_unique_json_object)
    except (ValueError, UnicodeError, RecursionError) as exc:
        raise GjcProtocolError("Invalid control JSON") from exc


def _control(row: Any) -> tuple[str, str | None]:
    if not isinstance(row, dict):
        raise GjcProtocolError("Invalid control request")
    command = row.get("command")
    keys = {"version", "type", "command"}
    if command == "clear":
        keys.add("launchId")
    if (set(row) != keys or type(row.get("version")) is not int
            or row["version"] != 1 or row.get("type") != "control"
            or command not in ("status", "clear")):
        raise GjcProtocolError("Invalid control request")
    return command, _uuid(row["launchId"], "launchId") if command == "clear" else None


def request_control(
    socket_path: str | Path, command: str = "status", launch_id: str | None = None,
    timeout: float = 2,
) -> dict[str, Any]:
    """Perform one bounded local request; reject malformed or unsolicited responses."""
    if valid_clock(timeout) <= 0:
        raise ValueError("timeout must be positive")
    request: dict[str, Any] = {"version": 1, "type": "control", "command": command}
    if launch_id is not None:
        request["launchId"] = launch_id
    _control(request)
    deadline = time.monotonic() + timeout
    with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as client:
        client.settimeout(timeout)
        try:
            client.connect(str(socket_path))
            client.sendall(json.dumps(request).encode() + b"\n")
        except socket.timeout as exc:
            raise TimeoutError("GJC control response timed out") from exc
        data = bytearray()
        while b"\n" not in data:
            remaining = deadline - time.monotonic()
            if remaining <= 0:
                raise TimeoutError("GJC control response timed out")
            client.settimeout(remaining)
            try:
                chunk = client.recv(min(65536, MAX_REGISTRY_BYTES + 1 - len(data)))
            except socket.timeout as exc:
                raise TimeoutError("GJC control response timed out") from exc
            if not chunk:
                raise GjcProtocolError("Truncated control response")
            data.extend(chunk)
            if len(data) > MAX_REGISTRY_BYTES:
                raise GjcProtocolError("Control response exceeds limit")
    line, extra = bytes(data).split(b"\n", 1)
    if extra:
        raise GjcProtocolError("Unexpected trailing control response")
    row = _decode(line)
    required = {"version", "type", "ok", "status"}
    if command == "clear":
        required.add("cleared")
    rejected = isinstance(row, dict) and row.get("ok") is False
    if rejected:
        required.add("error")
    if (not isinstance(row, dict) or set(row) != required
            or type(row.get("version")) is not int or row["version"] != 1
            or row.get("type") != "response" or type(row.get("ok")) is not bool
            or not isinstance(row.get("status"), dict)
            or (command == "clear" and type(row.get("cleared")) is not bool)):
        raise GjcProtocolError("Invalid control response")
    if rejected:
        error = row["error"]
        if (command != "clear" or row["cleared"] is not False
                or not isinstance(error, dict) or set(error) != {"code", "reason"}
                or error["code"] != "clear_refused"
                or not isinstance(error["reason"], str)
                or error["reason"] not in CLEAR_REFUSAL_REASONS):
            raise GjcProtocolError("Invalid control rejection")
    status_row = row["status"]
    if (set(status_row) != {"version", "maxSlots", "launches", "overflow"}
            or type(status_row["version"]) is not int or status_row["version"] != 1
            or type(status_row["maxSlots"]) is not int
            or not 1 <= status_row["maxSlots"] <= 1024
            or not isinstance(status_row["launches"], list)
            or not isinstance(status_row["overflow"], list)):
        raise GjcProtocolError("Invalid control status")
    # Reuse the persisted display schema, including root identity and bounds.
    validator = GjcTracker(status_row["maxSlots"])
    try:
        validator._restore_payload({"version": 1, "maxSlots": status_row["maxSlots"],
                                    "launches": status_row["launches"], "retired": []})
    except (ValueError, TypeError, KeyError) as exc:
        raise GjcProtocolError("Invalid control display registry") from exc
    if status_row["overflow"] != [r for r in status_row["launches"] if r["slot"] is None]:
        raise GjcProtocolError("Invalid control overflow")
    return row


@dataclass
class _Client:
    socket: socket.socket
    buffer: bytearray = field(default_factory=bytearray)
    output: bytes = b""
    producer: str | None = None
    touched: float = field(default_factory=time.monotonic)


class GjcStatusSource:
    def __init__(self, socket_path: str | Path, command: Any, max_slots: int = 12,
                 registry_path: str | Path | None = None) -> None:
        self.socket_path = Path(socket_path)
        self.command = command
        self.tracker = GjcTracker(max_slots, registry_path)
        self._thread: threading.Thread | None = None
        self._stop = threading.Event()
        self._lifecycle = threading.RLock()
        self._error: Exception | None = None
        self._lock_fd: int | None = None
        self._registry_lock_fd: int | None = None
        self._inode: tuple[int, int] | None = None
        self._listener: socket.socket | None = None
        self._selector: selectors.BaseSelector | None = None
        self._clients: dict[socket.socket, _Client] = {}
        self._owners: dict[str, _Client] = {}

    def start(self) -> None:
        with self._lifecycle:
            if self._thread is not None:
                self._check_error()
                return
            ensure_private_directory(self.socket_path.parent)
            self._stop.clear()
            self._error = None
            try:
                registry_path = self.tracker.registry_path
                if registry_path is not None:
                    self._registry_lock_fd = _acquire_private_lock(
                        Path(str(registry_path) + ".lock"), "Registry",
                    )
                    # Construction supports offline reads without owning the registry.
                    # Every writer must refresh after ownership, also on restart.
                    self.tracker = GjcTracker(
                        self.tracker.max_slots, registry_path, self.tracker.lease_seconds,
                    )
                self._lock_fd = _acquire_private_lock(
                    Path(str(self.socket_path) + ".lock"), "Socket",
                )
                if self.socket_path.exists() or self.socket_path.is_symlink():
                    info = self.socket_path.lstat()
                    if not stat.S_ISSOCK(info.st_mode) or info.st_uid != os.getuid():
                        raise PermissionError("Refusing to replace an unowned socket path")
                    with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as probe:
                        probe.settimeout(0.2)
                        try:
                            probe.connect(str(self.socket_path))
                        except OSError as exc:
                            if exc.errno != errno.ECONNREFUSED:
                                raise
                        else:
                            raise RuntimeError("A GJC receiver already owns this socket")
                    self.socket_path.unlink()
                self._listener = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
                self._listener.bind(str(self.socket_path))
                info = self.socket_path.lstat()
                self._inode = info.st_dev, info.st_ino
                os.chmod(self.socket_path, 0o600)
                self._listener.listen(MAX_CLIENTS)
                self._listener.setblocking(False)
                self._selector = selectors.DefaultSelector()
                self._selector.register(self._listener, selectors.EVENT_READ)
                self._thread = threading.Thread(target=self._run, name="gjc-receiver", daemon=True)
                self._thread.start()
            except BaseException:
                self._cleanup()
                self._thread = None
                raise

    def stop(self) -> None:
        with self._lifecycle:
            self._stop.set()
            if self._thread is not None:
                self._thread.join(timeout=5)
                if self._thread.is_alive():
                    raise RuntimeError("GJC receiver did not stop")
                self._thread = None
            self._cleanup()

    def _check_error(self) -> None:
        if self._error is not None:
            raise RuntimeError("GJC receiver failed; restart after fixing local storage") from self._error

    def indicators(self, now: float | None = None):
        self._check_error()
        return self.tracker.indicators(now)

    def status(self, now: float | None = None):
        self._check_error()
        return self.tracker.status(now)

    def _drop(self, client: _Client) -> None:
        self._clients.pop(client.socket, None)
        if self._selector is not None:
            try:
                self._selector.unregister(client.socket)
            except (KeyError, ValueError):
                pass
        client.socket.close()
        if client.producer and self._owners.get(client.producer) is client:
            del self._owners[client.producer]
            self.tracker.disconnect(client.producer)

    def _frame(self, client: _Client, line: bytes) -> None:
        row = _decode(line)
        if isinstance(row, dict) and row.get("type") == "control":
            if client.producer is not None:
                raise GjcProtocolError("Producer connections cannot send controls")
            command, launch_id = _control(row)
            response = {"version": 1, "type": "response", "ok": True}
            if command == "clear":
                cleared, reason = self._clear_disconnected_dead_launch(launch_id)
                response["cleared"] = cleared
                if reason is not None:
                    response["ok"] = False
                    response["error"] = {"code": "clear_refused", "reason": reason}
            response["status"] = self.tracker.status()
            client.output = json.dumps(response, separators=(",", ":")).encode() + b"\n"
            if len(client.output) > MAX_REGISTRY_BYTES:
                raise GjcProtocolError("Control response exceeds limit")
            assert self._selector is not None
            self._selector.modify(client.socket, selectors.EVENT_WRITE, client)
            return
        snapshot = parse_snapshot(row)
        if client.producer is not None and client.producer != snapshot.producer_id:
            raise GjcProtocolError("Producer identity changed on connection")
        if self.tracker.apply(snapshot):
            previous = self._owners.get(snapshot.producer_id)
            client.producer = snapshot.producer_id
            self._owners[snapshot.producer_id] = client
            if previous is not None and previous is not client:
                self._drop(previous)

    def _clear_disconnected_dead_launch(self, launch_id: str) -> tuple[bool, str | None]:
        """Check and clear only on the receiver thread, serialized with producer frames."""
        launch = next((row for row in self.tracker.status()["launches"]
                       if row["launchId"] == launch_id), None)
        if launch is None:
            return False, None
        # An expired display lease does not disconnect a still-owned socket.
        if launch["connected"] or launch["producerId"] in self._owners:
            return False, "connected"
        pid = launch["pid"]
        # Do not let special/group PIDs or platform pid_t overflow become death evidence.
        if type(pid) is not int or not 0 < pid <= 2**31 - 1:
            return False, "process_death_unconfirmed"
        try:
            os.kill(pid, 0)
        except ProcessLookupError:
            pass
        except OSError as exc:
            if exc.errno != errno.ESRCH:
                return False, "process_death_unconfirmed"
        except (OverflowError, ValueError):
            return False, "process_death_unconfirmed"
        else:
            return False, "process_alive"
        return self.tracker.clear(launch_id), None

    def _read(self, client: _Client) -> None:
        data = client.socket.recv(65536)
        if not data:
            self._drop(client)
            return
        client.touched = time.monotonic()
        client.buffer.extend(data)
        while b"\n" in client.buffer:
            line, _, rest = client.buffer.partition(b"\n")
            client.buffer = bytearray(rest)
            if len(line) > MAX_FRAME_BYTES:
                raise GjcProtocolError("Snapshot exceeds frame limit")
            self._frame(client, bytes(line))
            if client.output:
                if client.buffer:
                    raise GjcProtocolError("Only one control per connection is allowed")
                break
        if len(client.buffer) > MAX_FRAME_BYTES:
            raise GjcProtocolError("Snapshot exceeds frame limit")

    def _run(self) -> None:
        assert self._selector is not None
        try:
            while not self._stop.is_set():
                for key, mask in self._selector.select(0.2):
                    if key.fileobj is self._listener:
                        connection, _ = self._listener.accept()
                        connection.setblocking(False)
                        if len(self._clients) >= MAX_CLIENTS:
                            connection.close()
                        else:
                            client = _Client(connection)
                            self._clients[connection] = client
                            self._selector.register(connection, selectors.EVENT_READ, client)
                        continue
                    client = key.data
                    if client.socket not in self._clients:
                        continue
                    try:
                        if mask & selectors.EVENT_WRITE:
                            sent = client.socket.send(client.output[:65536])
                            client.output = client.output[sent:]
                            if not client.output:
                                self._drop(client)
                        else:
                            self._read(client)
                    except BlockingIOError:
                        pass
                    except (GjcProtocolError, ConnectionError):
                        self._drop(client)
                now = time.monotonic()
                for client in list(self._clients.values()):
                    if now - client.touched >= CLIENT_TIMEOUT:
                        self._drop(client)
        except Exception as exc:  # noqa: BLE001 - thread failures must reach the caller
            self._error = exc
        finally:
            # A storage failure may occur after apply updated the model but before
            # the connection was registered as its owner. Invalidate those too.
            for row in self.tracker.status()["launches"]:
                try:
                    self.tracker.disconnect(row["producerId"])
                except (OSError, ValueError) as exc:
                    self._error = self._error or exc
            self._cleanup()

    def _cleanup(self) -> None:
        try:
            self._close_resources()
        finally:
            # Even a filesystem cleanup failure must release writer ownership.
            for attribute in ("_lock_fd", "_registry_lock_fd"):
                fd = getattr(self, attribute)
                if fd is not None:
                    os.close(fd)
                    setattr(self, attribute, None)

    def _close_resources(self) -> None:
        for client in list(self._clients.values()):
            try:
                self._drop(client)
            except Exception as exc:  # noqa: BLE001 - continue closing all owned resources
                self._error = self._error or exc
        if self._selector is not None:
            self._selector.close()
            self._selector = None
        if self._listener is not None:
            self._listener.close()
            self._listener = None
        if self._inode is not None:
            try:
                info = self.socket_path.lstat()
                if (info.st_dev, info.st_ino) == self._inode:
                    self.socket_path.unlink()
            except FileNotFoundError:
                pass
            self._inode = None
