"""Opcode-exact proof for interruption-safe owned descriptor closure."""

from __future__ import annotations

import fcntl
import os
import sys
from concurrent.futures import ThreadPoolExecutor
from pathlib import Path
from threading import Barrier
from types import FrameType, TracebackType
from typing import Literal, Protocol, TypeAlias

from gateway.platforms.nutrition_weekly_reminder_resources import (
    CloseGuard,
    FD_LIFECYCLE_LOCK,
    OwnedDescriptor,
    ResourceRegistrar,
)

_TraceArgument: TypeAlias = (
    None | int | str
    | tuple[type[BaseException], BaseException, TracebackType]
)
_Boundary = tuple[Literal["line", "opcode"], int, int]
_Scenario = Literal["open", "uncertain-open", "uncertain-closed"]


class _Trace(Protocol):
    def __call__(
        self, frame: FrameType, event: str, argument: _TraceArgument, /
    ) -> _Trace | None: ...


def _fds() -> frozenset[int]:
    return frozenset(int(entry.name) for entry in Path("/proc/self/fd").iterdir())


def _is_close_frame(frame: FrameType) -> bool:
    return (
        frame.f_code is OwnedDescriptor.close.__code__
        and frame.f_code.co_filename.endswith(
            "gateway/platforms/nutrition_weekly_reminder_resources.py"
        )
    )


def _owned(scenario: _Scenario) -> tuple[OwnedDescriptor, CloseGuard, int]:
    descriptor = os.open("/dev/null", os.O_RDONLY | os.O_CLOEXEC)
    initial_state: Literal["open", "uncertain"] = (
        "open" if scenario == "open" else "uncertain"
    )
    owned = OwnedDescriptor(os.close, initial_state=initial_state)
    owned.fd = descriptor
    if scenario == "uncertain-closed":
        os.close(descriptor)
    return owned, CloseGuard(
        owned.close, retry=True, pending_failure=owned.take_pending_failure
    ), descriptor


def _close(guard: CloseGuard) -> None:
    with FD_LIFECYCLE_LOCK:
        guard()


def _boundaries(scenario: _Scenario) -> frozenset[_Boundary]:
    owned, guard, _descriptor = _owned(scenario)
    observed: set[_Boundary] = set()

    def collect(
        frame: FrameType, event: str, _argument: _TraceArgument
    ) -> _Trace | None:
        if event == "call" and _is_close_frame(frame):
            frame.f_trace_lines = True
            frame.f_trace_opcodes = True
            return collect
        if _is_close_frame(frame) and event == "line":
            observed.add(("line", frame.f_lineno or 0, frame.f_lasti))
        elif _is_close_frame(frame) and event == "opcode":
            observed.add(("opcode", frame.f_lineno or 0, frame.f_lasti))
        return collect

    try:
        sys.settrace(collect)
        _close(guard)
    finally:
        sys.settrace(None)
        _close(guard)
    assert owned.fd == -1
    return frozenset(observed)


def test_every_owned_fd_close_line_and_opcode_boundary_is_safe() -> None:
    scenarios: tuple[_Scenario, ...] = (
        "open", "uncertain-open", "uncertain-closed"
    )
    campaigns: dict[_Scenario, list[_Boundary]] = {
        scenario: sorted(_boundaries(scenario)) for scenario in scenarios
    }
    assert sum(len(boundaries) for boundaries in campaigns.values()) >= 100
    for scenario, boundaries in campaigns.items():
        for boundary in boundaries:
            for interruption in (KeyboardInterrupt, SystemExit):
                for attempt in (1, 2):
                    assert attempt in (1, 2)
                    baseline = _fds()
                    owned, guard, original_fd = _owned(scenario)
                    pending = interruption()
                    hit = [False]
                    propagated: list[BaseException] = []

                    def interrupt(
                        frame: FrameType, event: str, _argument: _TraceArgument
                    ) -> _Trace | None:
                        if event == "call" and _is_close_frame(frame):
                            frame.f_trace_lines = True
                            frame.f_trace_opcodes = True
                            return interrupt
                        current = (event, frame.f_lineno or 0, frame.f_lasti)
                        if _is_close_frame(frame) and not hit[0] and current == boundary:
                            hit[0] = True
                            raise pending
                        return interrupt

                    try:
                        sys.settrace(interrupt)
                        try:
                            _close(guard)
                        except interruption as error:
                            propagated.append(error)
                        finally:
                            sys.settrace(None)
                        assert hit[0], (scenario, boundary)
                        assert propagated == [pending]
                        assert owned.fd == -1
                        _close(guard)
                        replacement = os.open(
                            "/dev/null", os.O_RDONLY | os.O_CLOEXEC
                        )
                        try:
                            assert replacement == original_fd
                            _close(guard)
                            _ = fcntl.fcntl(replacement, fcntl.F_GETFD)
                        finally:
                            os.close(replacement)
                        assert _fds() == baseline
                        assert len(_fds()) == len(baseline)
                    finally:
                        sys.settrace(None)
                        _close(guard)


def test_repeated_interruptions_preserve_the_first_pending_exception() -> None:
    for first_type, repeated_type in (
        (KeyboardInterrupt, SystemExit), (SystemExit, KeyboardInterrupt)
    ):
        for attempt in (1, 2):
            assert attempt in (1, 2)
            baseline = _fds()
            first = first_type()
            repeated = repeated_type()
            pending = [first, repeated]
            real_close = os.close

            def close_descriptor(descriptor: int) -> None:
                if pending:
                    raise pending.pop(0)
                real_close(descriptor)

            descriptor = os.open("/dev/null", os.O_RDONLY | os.O_CLOEXEC)
            owned = OwnedDescriptor(close_descriptor)
            owned.fd = descriptor
            guard = CloseGuard(
                owned.close, retry=True,
                pending_failure=owned.take_pending_failure,
            )
            propagated: list[BaseException] = []
            try:
                _close(guard)
            except first_type as error:
                propagated.append(error)
            assert propagated == [first]
            assert pending == []
            _close(guard)
            assert _fds() == baseline


def test_concurrent_registrar_close_cannot_hit_reused_descriptor() -> None:
    baseline = _fds()
    workers = 8
    start = Barrier(workers)

    def exercise(_worker: int) -> None:
        _ = start.wait()
        for _ in range(300):
            owner = ResourceRegistrar()
            descriptor = owner.open_fd(
                lambda: os.open("/dev/null", os.O_RDONLY | os.O_CLOEXEC)
            )
            owner.close()
            owner.close()
            replacement_owner = ResourceRegistrar()
            replacement = replacement_owner.open_fd(
                lambda: os.open("/dev/null", os.O_RDONLY | os.O_CLOEXEC)
            )
            owner.close()
            _ = fcntl.fcntl(replacement, fcntl.F_GETFD)
            replacement_owner.close()
            assert descriptor >= 0

    with ThreadPoolExecutor(max_workers=workers) as pool:
        _ = list(pool.map(exercise, range(workers)))
    assert _fds() == baseline
    assert len(_fds()) == len(baseline)
