"""Durable, privacy-safe acknowledgement gating for Telegram polling."""

from __future__ import annotations

import asyncio
import fcntl
import json
import logging
import os
import secrets
import stat
from collections.abc import Callable, Iterable, Iterator
from dataclasses import dataclass
from contextlib import contextmanager
from contextvars import ContextVar
from pathlib import Path

logger = logging.getLogger(__name__)
_CURRENT_UPDATE_ID: ContextVar[int | None] = ContextVar(
    "telegram_ingress_update_id",
    default=None,
)

_SCHEMA_VERSION = 1
_RECOVERY_SCHEMA_VERSION = 2
_FAILURE_SCHEMA_VERSION = 3
_RECEIPT_STAGE = "receipt"
_RECEIPT_REASON = "handled"
_RECOVERY_STAGE = "recovery_receipt"
_RECOVERY_REASON = "business_commit_reconciled"
_FAILURE_STAGE = "blocked_receipt"
_VALID_REASONS = frozenset(
    {
        "business_commit_reconciled",
        "duplicate_replay",
        "handler_exception",
        "receipt_store_corrupt",
        "receipt_lock_failed",
        "receipt_write_failed",
        "recovery_authority_mismatch",
    }
)


class TelegramIngressReceiptError(RuntimeError):
    """The durable update-ID receipt authority cannot be used safely."""


@dataclass(frozen=True)
class _ReceiptState:
    pending: frozenset[int]
    recovered: dict[int, str]
    failed: dict[int, str]


@dataclass(frozen=True)
class TelegramBusinessRecoveryCandidate:
    update_id: int
    provenance_digest: str
    actor_id: int
    chat_id: int
    topic_id: int
    message_id: int
    callback_data: str


class TelegramIngressReceiptStore:
    """Persist only pending update-ID receipts until Telegram acknowledges them."""

    def __init__(self, path: Path) -> None:
        self._path = Path(path)
        self._lock_path = self._path.with_suffix(self._path.suffix + ".lock")

    def contains(self, update_id: int) -> bool:
        with self._locked() as directory_fd:
            state = self._read_unlocked(directory_fd)
            return update_id in state.pending or update_id in state.recovered

    def record(self, update_id: int) -> None:
        with self._locked() as directory_fd:
            state = self._read_unlocked(directory_fd)
            if update_id in state.pending or update_id in state.recovered:
                return
            self._write_unlocked(
                directory_fd,
                _ReceiptState(state.pending | {update_id}, state.recovered, state.failed),
            )

    def record_recovered(
        self,
        update_id: int,
        provenance_digest: str,
        *,
        heal_handler_failure: bool = False,
    ) -> None:
        """Retain a terminal receipt for authenticated post-commit recovery."""
        if (
            type(update_id) is not int
            or update_id < 0
            or len(provenance_digest) != 64
            or any(character not in "0123456789abcdef" for character in provenance_digest)
        ):
            raise ValueError("invalid Telegram recovery receipt")
        with self._locked() as directory_fd:
            state = self._read_unlocked(directory_fd)
            failure = state.failed.get(update_id)
            if failure is not None and (
                not heal_handler_failure or failure != "handler_exception"
            ):
                raise TelegramIngressReceiptError(failure)
            existing = state.recovered.get(update_id)
            if existing is not None:
                if existing != provenance_digest:
                    raise TelegramIngressReceiptError("receipt_store_corrupt")
                return
            recovered = dict(state.recovered)
            recovered[update_id] = provenance_digest
            failed = dict(state.failed)
            failed.pop(update_id, None)
            self._write_unlocked(
                directory_fd,
                _ReceiptState(
                    frozenset(item for item in state.pending if item != update_id),
                    recovered,
                    failed,
                ),
            )

    def record_failed(self, update_id: int, reason_code: str) -> None:
        """Persist a handler failure so restart recovery cannot overwrite it."""
        if type(update_id) is not int or update_id < 0 or reason_code != "handler_exception":
            return
        with self._locked() as directory_fd:
            state = self._read_unlocked(directory_fd)
            if update_id in state.recovered:
                raise TelegramIngressReceiptError("receipt_store_corrupt")
            existing = state.failed.get(update_id)
            if existing is not None:
                if existing != reason_code:
                    raise TelegramIngressReceiptError("receipt_store_corrupt")
                return
            failed = dict(state.failed)
            failed[update_id] = reason_code
            self._write_unlocked(
                directory_fd,
                _ReceiptState(state.pending, state.recovered, failed),
            )

    def failure_reason(self, update_id: int) -> str | None:
        with self._locked() as directory_fd:
            return self._read_unlocked(directory_fd).failed.get(update_id)

    def recovered_provenance(self, update_id: int) -> str | None:
        with self._locked() as directory_fd:
            return self._read_unlocked(directory_fd).recovered.get(update_id)

    def forget_before(self, offset: int) -> None:
        with self._locked() as directory_fd:
            state = self._read_unlocked(directory_fd)
            retained = frozenset(
                update_id
                for update_id in state.pending
                if update_id >= offset
            )
            if retained != state.pending:
                self._write_unlocked(
                    directory_fd,
                    _ReceiptState(retained, state.recovered, state.failed),
                )

    @contextmanager
    def _locked(self) -> Iterator[int]:
        directory_fd = -1
        descriptor = -1
        try:
            directory_fd = self._open_authority_directory()
            self._validate_optional_regular(
                self._lock_path.name,
                directory_fd,
                error_code="receipt_lock_failed",
            )
            descriptor = os.open(
                self._lock_path.name,
                self._open_flags(os.O_CREAT | os.O_RDWR | os.O_NONBLOCK),
                0o600,
                dir_fd=directory_fd,
            )
            self._validate_regular(
                os.fstat(descriptor),
                error_code="receipt_lock_failed",
            )
            fcntl.flock(descriptor, fcntl.LOCK_EX)
        except OSError as exc:
            if descriptor >= 0:
                os.close(descriptor)
            if directory_fd >= 0:
                os.close(directory_fd)
            raise TelegramIngressReceiptError("receipt_lock_failed") from exc
        except TelegramIngressReceiptError:
            if descriptor >= 0:
                os.close(descriptor)
            if directory_fd >= 0:
                os.close(directory_fd)
            raise
        try:
            yield directory_fd
        finally:
            fcntl.flock(descriptor, fcntl.LOCK_UN)
            os.close(descriptor)
            os.close(directory_fd)

    def _open_authority_directory(self) -> int:
        parent = self._path.parent
        if parent.is_absolute():
            directory_fd = os.open(
                os.sep,
                os.O_RDONLY | os.O_DIRECTORY,
            )
            components = parent.parts[1:]
        else:
            directory_fd = os.open(
                ".",
                os.O_RDONLY | os.O_DIRECTORY,
            )
            components = parent.parts
        try:
            for component in components:
                if component in {"", "."}:
                    continue
                if component == "..":
                    raise TelegramIngressReceiptError("receipt_lock_failed")
                try:
                    child_fd = os.open(
                        component,
                        self._open_flags(os.O_RDONLY | os.O_DIRECTORY),
                        dir_fd=directory_fd,
                    )
                except FileNotFoundError:
                    try:
                        os.mkdir(component, 0o700, dir_fd=directory_fd)
                    except FileExistsError:
                        pass
                    child_fd = os.open(
                        component,
                        self._open_flags(os.O_RDONLY | os.O_DIRECTORY),
                        dir_fd=directory_fd,
                    )
                try:
                    if not stat.S_ISDIR(os.fstat(child_fd).st_mode):
                        raise TelegramIngressReceiptError("receipt_lock_failed")
                except BaseException:
                    os.close(child_fd)
                    raise
                os.close(directory_fd)
                directory_fd = child_fd
        except BaseException:
            os.close(directory_fd)
            raise
        return directory_fd

    @staticmethod
    def _open_flags(flags: int) -> int:
        no_follow = getattr(os, "O_NOFOLLOW", None)
        if type(no_follow) is not int:
            raise TelegramIngressReceiptError("receipt_lock_failed")
        return flags | no_follow

    @staticmethod
    def _validate_regular(
        metadata: os.stat_result,
        *,
        error_code: str,
    ) -> None:
        if (
            not stat.S_ISREG(metadata.st_mode)
            or metadata.st_nlink != 1
            or stat.S_IMODE(metadata.st_mode) != 0o600
        ):
            raise TelegramIngressReceiptError(error_code)

    def _validate_optional_regular(
        self,
        name: str,
        directory_fd: int,
        *,
        error_code: str,
    ) -> None:
        try:
            metadata = os.lstat(name, dir_fd=directory_fd)
        except FileNotFoundError:
            return
        except OSError as exc:
            raise TelegramIngressReceiptError(error_code) from exc
        self._validate_regular(metadata, error_code=error_code)

    def _read_unlocked(self, directory_fd: int) -> _ReceiptState:
        try:
            metadata = os.lstat(self._path.name, dir_fd=directory_fd)
        except FileNotFoundError:
            return _ReceiptState(frozenset(), {}, {})
        except OSError as exc:
            raise TelegramIngressReceiptError("receipt_store_corrupt") from exc
        self._validate_regular(
            metadata,
            error_code="receipt_store_corrupt",
        )
        descriptor = -1
        try:
            descriptor = os.open(
                self._path.name,
                self._open_flags(os.O_RDONLY | os.O_NONBLOCK),
                dir_fd=directory_fd,
            )
            self._validate_regular(
                os.fstat(descriptor),
                error_code="receipt_store_corrupt",
            )
            chunks: list[bytes] = []
            while chunk := os.read(descriptor, 65536):
                chunks.append(chunk)
            raw: object = json.loads(b"".join(chunks).decode("utf-8"))
        except TelegramIngressReceiptError:
            raise
        except (OSError, UnicodeDecodeError, json.JSONDecodeError) as exc:
            raise TelegramIngressReceiptError("receipt_store_corrupt") from exc
        finally:
            if descriptor >= 0:
                os.close(descriptor)
        if not isinstance(raw, dict):
            raise TelegramIngressReceiptError("receipt_store_corrupt")
        version = raw.get("version")
        expected_keys = (
            {"receipts", "version"}
            if version == _SCHEMA_VERSION
            else (
                {"receipts", "terminal_receipts", "version"}
                if version == _RECOVERY_SCHEMA_VERSION
                else {"receipts", "terminal_failures", "terminal_receipts", "version"}
            )
        )
        if (
            version not in {_SCHEMA_VERSION, _RECOVERY_SCHEMA_VERSION, _FAILURE_SCHEMA_VERSION}
            or set(raw) != expected_keys
            or not isinstance(raw.get("receipts"), dict)
        ):
            raise TelegramIngressReceiptError("receipt_store_corrupt")

        receipts: set[int] = set()
        for raw_update_id, record in raw["receipts"].items():
            if (
                not isinstance(raw_update_id, str)
                or not raw_update_id.isdecimal()
                or raw_update_id != str(int(raw_update_id))
                or not isinstance(record, dict)
                or record
                != {"reason_code": _RECEIPT_REASON, "stage": _RECEIPT_STAGE}
            ):
                raise TelegramIngressReceiptError("receipt_store_corrupt")
            update_id = int(raw_update_id)
            if update_id < 0:
                raise TelegramIngressReceiptError("receipt_store_corrupt")
            receipts.add(update_id)
        recovered: dict[int, str] = {}
        raw_recovered = raw.get("terminal_receipts", {})
        if not isinstance(raw_recovered, dict):
            raise TelegramIngressReceiptError("receipt_store_corrupt")
        for raw_update_id, record in raw_recovered.items():
            if (
                not isinstance(raw_update_id, str)
                or not raw_update_id.isdecimal()
                or raw_update_id != str(int(raw_update_id))
                or not isinstance(record, dict)
                or set(record) != {"provenance_digest", "reason_code", "stage"}
                or record.get("reason_code") != _RECOVERY_REASON
                or record.get("stage") != _RECOVERY_STAGE
                or not isinstance(record.get("provenance_digest"), str)
            ):
                raise TelegramIngressReceiptError("receipt_store_corrupt")
            update_id = int(raw_update_id)
            provenance_digest = record["provenance_digest"]
            if (
                update_id < 0
                or len(provenance_digest) != 64
                or any(character not in "0123456789abcdef" for character in provenance_digest)
                or update_id in receipts
            ):
                raise TelegramIngressReceiptError("receipt_store_corrupt")
            recovered[update_id] = provenance_digest
        failed: dict[int, str] = {}
        raw_failed = raw.get("terminal_failures", {})
        if not isinstance(raw_failed, dict):
            raise TelegramIngressReceiptError("receipt_store_corrupt")
        for raw_update_id, record in raw_failed.items():
            if (
                not isinstance(raw_update_id, str)
                or not raw_update_id.isdecimal()
                or raw_update_id != str(int(raw_update_id))
                or record != {"reason_code": "handler_exception", "stage": _FAILURE_STAGE}
            ):
                raise TelegramIngressReceiptError("receipt_store_corrupt")
            update_id = int(raw_update_id)
            if update_id < 0 or update_id in receipts or update_id in recovered:
                raise TelegramIngressReceiptError("receipt_store_corrupt")
            failed[update_id] = "handler_exception"
        return _ReceiptState(frozenset(receipts), recovered, failed)

    def _write_unlocked(
        self,
        directory_fd: int,
        state: _ReceiptState,
    ) -> None:
        payload: dict[str, object] = {
            "version": (
                _FAILURE_SCHEMA_VERSION
                if state.failed
                else (_RECOVERY_SCHEMA_VERSION if state.recovered else _SCHEMA_VERSION)
            ),
            "receipts": {
                str(update_id): {
                    "stage": _RECEIPT_STAGE,
                    "reason_code": _RECEIPT_REASON,
                }
                for update_id in sorted(state.pending)
            },
        }
        if state.recovered or state.failed:
            payload["terminal_receipts"] = {
                str(update_id): {
                    "stage": _RECOVERY_STAGE,
                    "reason_code": _RECOVERY_REASON,
                    "provenance_digest": provenance_digest,
                }
                for update_id, provenance_digest in sorted(state.recovered.items())
            }
        if state.failed:
            payload["terminal_failures"] = {
                str(update_id): {
                    "stage": _FAILURE_STAGE,
                    "reason_code": reason_code,
                }
                for update_id, reason_code in sorted(state.failed.items())
            }
        temporary_name = (
            f".{self._path.name}.receipt-{secrets.token_hex(16)}.tmp"
        )
        descriptor = -1
        try:
            self._validate_optional_regular(
                self._path.name,
                directory_fd,
                error_code="receipt_store_corrupt",
            )
            descriptor = os.open(
                temporary_name,
                self._open_flags(os.O_CREAT | os.O_EXCL | os.O_WRONLY),
                0o600,
                dir_fd=directory_fd,
            )
            self._validate_regular(
                os.fstat(descriptor),
                error_code="receipt_write_failed",
            )
            encoded = json.dumps(
                payload,
                indent=2,
                sort_keys=True,
            ).encode("utf-8")
            with os.fdopen(descriptor, "wb") as temporary_file:
                descriptor = -1
                if temporary_file.write(encoded) != len(encoded):
                    raise OSError("partial receipt write")
                temporary_file.flush()
                os.fsync(temporary_file.fileno())
            os.replace(
                temporary_name,
                self._path.name,
                src_dir_fd=directory_fd,
                dst_dir_fd=directory_fd,
            )
            os.fsync(directory_fd)
        except TelegramIngressReceiptError:
            raise
        except OSError as exc:
            raise TelegramIngressReceiptError("receipt_write_failed") from exc
        finally:
            if descriptor >= 0:
                os.close(descriptor)
            try:
                os.unlink(temporary_name, dir_fd=directory_fd)
            except FileNotFoundError:
                pass
            except OSError:
                pass


def update_id_from_update(update: object) -> int:
    update_id = getattr(update, "update_id", None)
    if type(update_id) is not int or update_id < 0:
        raise TelegramIngressReceiptError("receipt_store_corrupt")
    return update_id


def set_current_telegram_update(update: object) -> None:
    """Expose the PTB update ID to content-free nested callback diagnostics."""
    _CURRENT_UPDATE_ID.set(update_id_from_update(update))


def current_telegram_update_id() -> int | None:
    return _CURRENT_UPDATE_ID.get()


def log_telegram_ingress(
    stage: str,
    update_id: int,
    *,
    reason_code: str | None = None,
) -> None:
    """Log a bounded receipt event without message, callback, or actor data."""
    if reason_code is None:
        logger.info("telegram_ingress stage=%s update_id=%s", stage, update_id)
        return
    if reason_code not in _VALID_REASONS:
        raise ValueError("invalid Telegram ingress receipt reason")
    logger.info(
        "telegram_ingress stage=%s update_id=%s reason_code=%s",
        stage,
        update_id,
        reason_code,
    )


class TelegramPollingReceiptGate:
    """Hold PTB's next offset request until every earlier update is receipted."""

    def __init__(
        self,
        store: TelegramIngressReceiptStore,
        *,
        on_blocked: Callable[[int, str], None],
    ) -> None:
        self._store = store
        self._on_blocked = on_blocked
        self._ready: dict[int, asyncio.Event] = {}
        self._outcome: dict[int, bool] = {}
        self._reason: dict[int, str] = {}
        self._processing: set[int] = set()
        self._business_recoveries: dict[int, TelegramBusinessRecoveryCandidate] = {}

    def register_business_recovery(
        self,
        candidate: TelegramBusinessRecoveryCandidate,
    ) -> None:
        existing = self._business_recoveries.get(candidate.update_id)
        if existing is not None and existing != candidate:
            raise TelegramIngressReceiptError("receipt_store_corrupt")
        self._business_recoveries[candidate.update_id] = candidate

    def captured(self, updates: Iterable[object]) -> None:
        for update in updates:
            update_id = update_id_from_update(update)
            self._ready.setdefault(update_id, asyncio.Event())

    async def begin(self, update: object) -> bool:
        """Return whether the already-receipted update must skip business handlers."""
        update_id = update_id_from_update(update)
        set_current_telegram_update(update)
        ready = self._ready.setdefault(update_id, asyncio.Event())
        try:
            failure = self._store.failure_reason(update_id)
            recovery = self._business_recoveries.get(update_id)
            if failure is not None:
                if (
                    failure == "handler_exception"
                    and recovery is not None
                    and self._matches_business_recovery(update, recovery)
                ):
                    self.reconcile_business_commit(
                        update_id=update_id,
                        provenance_digest=recovery.provenance_digest,
                        heal_handler_failure=True,
                    )
                else:
                    self.failed(update, failure)
                return True
            if self._store.contains(update_id):
                self._resolve(update_id, True)
                log_telegram_ingress(
                    "receipt",
                    update_id,
                    reason_code="duplicate_replay",
                )
                return True
            if recovery is not None:
                if not self._matches_business_recovery(update, recovery):
                    self.failed(update, "recovery_authority_mismatch")
                    return True
                self.reconcile_business_commit(
                    update_id=update_id,
                    provenance_digest=recovery.provenance_digest,
                )
                return True
        except TelegramIngressReceiptError as exc:
            self.failed(update, str(exc))
            return True

        if update_id in self._processing:
            await ready.wait()
            return True
        self._processing.add(update_id)
        return False

    def completed(self, update: object) -> None:
        update_id = update_id_from_update(update)
        if self._outcome.get(update_id) is False:
            return
        try:
            self._store.record(update_id)
        except TelegramIngressReceiptError as exc:
            self.failed(update, str(exc))
            return
        self._resolve(update_id, True)
        log_telegram_ingress("receipt", update_id)

    def failed(self, update: object, reason_code: str) -> None:
        update_id = update_id_from_update(update)
        if self._outcome.get(update_id) is True:
            return
        self._processing.discard(update_id)
        self._reason[update_id] = reason_code
        try:
            self._store.record_failed(update_id, reason_code)
        except TelegramIngressReceiptError as exc:
            self._reason[update_id] = str(exc)
        self._resolve(update_id, False)
        log_telegram_ingress("receipt", update_id, reason_code=reason_code)

    def reconcile_business_commit(
        self,
        *,
        update_id: int,
        provenance_digest: str,
        heal_handler_failure: bool = False,
    ) -> None:
        """Receipt an authenticated business commit after UI recovery completed."""
        if self._outcome.get(update_id) is False:
            raise TelegramIngressReceiptError(
                self._reason.get(update_id, "handler_exception")
            )
        self._store.record_recovered(
            update_id,
            provenance_digest,
            heal_handler_failure=heal_handler_failure,
        )
        self._processing.discard(update_id)
        self._resolve(update_id, True)
        log_telegram_ingress(
            _RECOVERY_STAGE,
            update_id,
            reason_code=_RECOVERY_REASON,
        )

    async def permits_offset(self, offset: int, *, running: bool) -> bool:
        pending = tuple(
            update_id
            for update_id in sorted(self._ready)
            if update_id < offset
        )
        if pending and not running:
            return False
        for update_id in pending:
            await self._ready[update_id].wait()
            if self._outcome.get(update_id) is not True:
                self._on_blocked(
                    update_id,
                    self._reason.get(update_id, "handler_exception"),
                )
                return False
        return True

    def acknowledged_through(self, offset: int) -> None:
        try:
            self._store.forget_before(offset)
        except TelegramIngressReceiptError:
            return
        for update_id in tuple(self._ready):
            if update_id < offset:
                self._ready.pop(update_id, None)
                self._outcome.pop(update_id, None)
                self._reason.pop(update_id, None)
                self._processing.discard(update_id)

    @staticmethod
    def _matches_business_recovery(
        update: object,
        candidate: TelegramBusinessRecoveryCandidate,
    ) -> bool:
        query = getattr(update, "callback_query", None)
        message = getattr(query, "message", None)
        actor = getattr(query, "from_user", None)
        chat = getattr(message, "chat", None)
        topic_id = getattr(message, "message_thread_id", None)
        if topic_id is None:
            topic_id = 0
        return bool(
            getattr(query, "data", None) == candidate.callback_data
            and getattr(actor, "id", None) == candidate.actor_id
            and getattr(chat, "id", getattr(message, "chat_id", None)) == candidate.chat_id
            and topic_id == candidate.topic_id
            and getattr(message, "message_id", None) == candidate.message_id
        )

    def _resolve(self, update_id: int, succeeded: bool) -> None:
        self._outcome[update_id] = succeeded
        self._ready.setdefault(update_id, asyncio.Event()).set()


class ReceiptGatedTelegramBot:
    """PTB bot proxy which does not send an advancing offset before a receipt."""

    def __init__(
        self,
        bot: object,
        gate: TelegramPollingReceiptGate,
        *,
        is_running: Callable[[], bool],
    ) -> None:
        self._bot = bot
        self._gate = gate
        self._is_running = is_running

    def __getattr__(self, name: str) -> object:
        return getattr(self._bot, name)

    async def get_updates(
        self,
        *args: object,
        **kwargs: object,
    ) -> object:
        offset = kwargs.get("offset")
        if type(offset) is int and offset > 0:
            if not await self._gate.permits_offset(
                offset,
                running=self._is_running(),
            ):
                return []
        get_updates = getattr(self._bot, "get_updates")
        updates = await get_updates(*args, **kwargs)
        if not isinstance(updates, (list, tuple)):
            raise TelegramIngressReceiptError("receipt_store_corrupt")
        self._gate.captured(updates)
        if type(offset) is int and offset > 0:
            self._gate.acknowledged_through(offset)
        return updates
