"""Async wrapper around :class:`SoftReadWriteLock` for use with ``asyncio``."""

from __future__ import annotations

import asyncio
import functools
import os
from contextlib import asynccontextmanager
from typing import TYPE_CHECKING, Literal, ParamSpec, TypeVar

from filelock._async import (
    _BackendOutcome,
    _capture_call,
    _drain_future,
    _future_result,
    _raise_cancelled_error,
    _task_owners_for,
    _wait_until_done,
)

from ._sync import SoftReadWriteLock

if TYPE_CHECKING:
    from collections.abc import AsyncGenerator, Callable
    from concurrent import futures
    from types import TracebackType

    from filelock._api import AcquireReturnProxy
    from filelock._lease import LeaseCompromise

_P = ParamSpec("_P")
_R = TypeVar("_R")


class AsyncSoftReadWriteLock:
    """
    Async wrapper around :class:`SoftReadWriteLock` for ``asyncio`` applications.

    The sync class's blocking filesystem operations run on a thread pool via ``loop.run_in_executor()``. Each
    ``asyncio`` task owns its own hold. Tasks share a read lock, and a task asking for the write lock waits until the
    other holders release, with waiting writers ahead of new readers. Upgrading or downgrading a held lock raises
    :class:`RuntimeError`. The underlying :class:`SoftReadWriteLock` handles forks, the heartbeat, and stale eviction.
    Singleton wrappers share one sync lock and its task holds.

    :param lock_file: path to the lock file; the protocol directory lives next to it as ``<lock_file>.rw``
    :param timeout: maximum wait time in seconds; ``-1`` means block indefinitely
    :param blocking: if ``False``, raise :class:`~filelock.Timeout` immediately on contention
    :param is_singleton: if ``True``, reuse existing :class:`SoftReadWriteLock` instances per resolved path
    :param heartbeat_interval: seconds between heartbeat refreshes; default 30 s
    :param stale_threshold: seconds a holder record may stay unchanged before a contender evicts it; defaults to
        ``3 * heartbeat_interval``
    :param poll_interval: seconds between acquire retries under contention; default 0.25 s
    :param on_compromise: called from the heartbeat thread with a :class:`~filelock.LeaseCompromise` when the hold is
        lost
    :param loop: event loop for ``run_in_executor``; ``None`` uses the running loop
    :param executor: executor for ``run_in_executor``; ``None`` uses the default executor

    .. versionadded:: 3.27.0

    """

    def __init__(  # ruff:ignore[too-many-arguments]  # public constructor: one parameter per documented lock option
        self,
        lock_file: str | os.PathLike[str],
        timeout: float = -1,
        *,
        blocking: bool = True,
        is_singleton: bool = True,
        heartbeat_interval: float = 30.0,
        stale_threshold: float | None = None,
        poll_interval: float = 0.25,
        on_compromise: Callable[[LeaseCompromise], None] | None = None,
        loop: asyncio.AbstractEventLoop | None = None,
        executor: futures.Executor | None = None,
    ) -> None:
        self._creator_pid = os.getpid()
        self._lock = SoftReadWriteLock(
            lock_file,
            timeout,
            blocking=blocking,
            is_singleton=is_singleton,
            heartbeat_interval=heartbeat_interval,
            stale_threshold=stale_threshold,
            poll_interval=poll_interval,
            on_compromise=on_compromise,
        )
        self._owners = _task_owners_for(self._lock)
        self._loop = loop
        self._executor = executor

    @property
    def lock_file(self) -> str:
        """The path to the lock file passed to the constructor."""
        return self._lock.lock_file

    @property
    def timeout(self) -> float:
        """The default timeout applied when ``acquire_read`` / ``acquire_write`` is called without one."""
        return self._lock.timeout

    @property
    def blocking(self) -> bool:
        """Whether ``acquire_*`` defaults to blocking; ``False`` makes contention raise immediately."""
        return self._lock.blocking

    @property
    def generation(self) -> int | None:
        """The generation the current hold was granted at, a fencing token; ``None`` when no lock is held."""
        return self._lock.generation

    @property
    def compromise(self) -> LeaseCompromise | None:
        """How the current hold was lost, or ``None`` while it stands."""
        return self._lock.compromise

    @property
    def loop(self) -> asyncio.AbstractEventLoop | None:
        """The event loop used for ``run_in_executor``, or ``None`` for the running loop."""
        return self._loop

    @property
    def executor(self) -> futures.Executor | None:
        """The executor used for ``run_in_executor``, or ``None`` for the default executor."""
        return self._executor

    @asynccontextmanager
    async def read_lock(self, timeout: float | None = None, *, blocking: bool | None = None) -> AsyncGenerator[None]:
        """
        Async context manager that acquires and releases a shared read lock.

        :param timeout: maximum wait time in seconds, or ``None`` to use the instance default
        :param blocking: if ``False``, raise :class:`~filelock.Timeout` immediately; ``None`` uses the instance default

        :raises RuntimeError: if the calling task already holds the write lock
        :raises Timeout: if the lock cannot be acquired within *timeout* seconds

        """
        await self.acquire_read(timeout, blocking=blocking)
        try:
            yield
        finally:
            await self.release()

    @asynccontextmanager
    async def write_lock(self, timeout: float | None = None, *, blocking: bool | None = None) -> AsyncGenerator[None]:
        """
        Async context manager that acquires and releases an exclusive write lock.

        :param timeout: maximum wait time in seconds, or ``None`` to use the instance default
        :param blocking: if ``False``, raise :class:`~filelock.Timeout` immediately; ``None`` uses the instance default

        :raises RuntimeError: if the calling task already holds the read lock
        :raises Timeout: if the lock cannot be acquired within *timeout* seconds

        """
        await self.acquire_write(timeout, blocking=blocking)
        try:
            yield
        finally:
            await self.release()

    async def acquire_read(
        self, timeout: float | None = None, *, blocking: bool | None = None
    ) -> AsyncAcquireSoftReadWriteReturnProxy:
        """
        Acquire a shared read lock.

        See :meth:`SoftReadWriteLock.acquire_read` for reentrancy / upgrade / fork semantics. The blocking work runs
        inside ``run_in_executor`` so other coroutines on the same loop keep progressing while this call waits.

        :param timeout: maximum wait time in seconds, or ``None`` to use the instance default
        :param blocking: if ``False``, raise :class:`~filelock.Timeout` immediately; ``None`` uses the instance default

        :returns: a proxy usable as an async context manager to release the lock

        :raises RuntimeError: if the calling task already holds the write lock, if this instance was invalidated by
            :func:`os.fork`, or if :meth:`close` was called
        :raises Timeout: if the lock cannot be acquired within *timeout* seconds

        """
        self._raise_if_inherited()
        await self._acquire("read", timeout, blocking=blocking)
        return AsyncAcquireSoftReadWriteReturnProxy(lock=self)

    async def acquire_write(
        self, timeout: float | None = None, *, blocking: bool | None = None
    ) -> AsyncAcquireSoftReadWriteReturnProxy:
        """
        Acquire an exclusive write lock.

        See :meth:`SoftReadWriteLock.acquire_write` for the writer-preferring semantics. The blocking work runs inside
        ``run_in_executor``.

        :param timeout: maximum wait time in seconds, or ``None`` to use the instance default
        :param blocking: if ``False``, raise :class:`~filelock.Timeout` immediately; ``None`` uses the instance default

        :returns: a proxy usable as an async context manager to release the lock

        :raises RuntimeError: if the calling task already holds the read lock, if this instance was invalidated by
            :func:`os.fork`, or if :meth:`close` was called
        :raises Timeout: if the lock cannot be acquired within *timeout* seconds

        """
        self._raise_if_inherited()
        await self._acquire("write", timeout, blocking=blocking)
        return AsyncAcquireSoftReadWriteReturnProxy(lock=self)

    async def _acquire(self, mode: Literal["read", "write"], timeout: float | None, *, blocking: bool | None) -> None:
        blocking = self._lock.blocking if blocking is None else blocking
        sync_acquire = self._lock.acquire_read if mode == "read" else self._lock.acquire_write

        async def enter(remaining: float) -> None:
            await self._run_acquire(functools.partial(sync_acquire, remaining, blocking=blocking))

        await self._owners.acquire(
            mode,
            timeout=self._lock.timeout if timeout is None else timeout,
            blocking=blocking,
            lock_file=self.lock_file,
            enter=enter,
        )

    async def release(self, *, force: bool = False) -> None:
        """
        Release one level of the calling task's hold.

        The wrapper releases the lock once no task holds it.

        :param force: if ``True``, drop the calling task's whole hold at any nesting level

        :raises RuntimeError: if the calling task holds no lock and *force* is ``False``

        """
        if self._creator_pid == os.getpid():
            # The task table owns the nesting, so the backend is always one level deep here, and the executor can run
            # this on another worker than the acquire: force=True is the backend's cross-thread release.
            await self._owners.release(
                force=force,
                lock_file=self.lock_file,
                leave=functools.partial(self._run, self._lock.release, force=True),
            )

    async def close(self) -> None:
        """Release any held lock and release the underlying filesystem resources. Idempotent."""
        if self._creator_pid == os.getpid():
            await self._run(self._lock.close)
            self._owners.reset()

    def _raise_if_inherited(self) -> None:
        if self._creator_pid != os.getpid():  # pragma: forked child
            msg = f"AsyncSoftReadWriteLock on {self.lock_file} was inherited across fork; construct a new instance"
            raise RuntimeError(msg)

    async def _run_acquire(self, acquire: Callable[[], AcquireReturnProxy]) -> None:
        # run_in_executor cannot recall work the pool already started, so canceling the caller does not stop the sync
        # acquire: it still commits its claim, sets the hold, and starts the heartbeat, which keeps the record fresh
        # forever so no peer on any host can evict it as stale. Wait the submitted call out and hand the claim back,
        # the way AsyncReadWriteLock does.
        acquire_future = self._submit(acquire)
        try:
            await _wait_until_done(acquire_future)
        except asyncio.CancelledError as cancellation:
            try:
                await _drain_future(acquire_future)
            except BaseException as error:  # ruff:ignore[blind-except]  # reported with the cancellation below
                _raise_cancelled_error(cancellation, error)
            try:
                await _drain_future(self._submit(self._lock.release, force=True))
            except BaseException as error:  # ruff:ignore[blind-except]  # reported with the cancellation below
                _raise_cancelled_error(cancellation, error)
            raise
        _future_result(acquire_future)

    async def _run(self, func: Callable[_P, _R], *args: _P.args, **kwargs: _P.kwargs) -> _R:
        # A canceled release or close is already running on the pool thread; drain it so its outcome is observed
        # instead of finishing unwatched, then let the cancellation through.
        future = self._submit(func, *args, **kwargs)
        try:
            await _wait_until_done(future)
        except asyncio.CancelledError as cancellation:
            try:
                await _drain_future(future)
            except BaseException as error:  # ruff:ignore[blind-except]  # reported with the cancellation below
                _raise_cancelled_error(cancellation, error)
            raise
        return _future_result(future)

    def _submit(
        self, func: Callable[_P, _R], *args: _P.args, **kwargs: _P.kwargs
    ) -> asyncio.Future[_BackendOutcome[_R]]:
        loop = self._loop or asyncio.get_running_loop()
        return loop.run_in_executor(self._executor, _capture_call, functools.partial(func, *args, **kwargs))


class AsyncAcquireSoftReadWriteReturnProxy:
    """Async context-aware object that releases an :class:`AsyncSoftReadWriteLock` on exit."""

    def __init__(self, lock: AsyncSoftReadWriteLock) -> None:
        self.lock = lock

    async def __aenter__(self) -> AsyncSoftReadWriteLock:
        return self.lock

    async def __aexit__(
        self,
        exc_type: type[BaseException] | None,
        exc_value: BaseException | None,
        traceback: TracebackType | None,
    ) -> None:
        await self.lock.release()


__all__ = [
    "AsyncAcquireSoftReadWriteReturnProxy",
    "AsyncSoftReadWriteLock",
]
