from __future__ import annotations

import functools
import math
from typing import ParamSpec, TYPE_CHECKING, TypeVar

import torch

from . import _dtypes_impl, _util
from ._normalizations import ArrayLike, KeepDims, normalizer, OutArray


if TYPE_CHECKING:
    from collections.abc import Callable, Sequence


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


class LinAlgError(Exception):
    pass


def _atleast_float_1(a: torch.Tensor) -> torch.Tensor:
    if not (a.dtype.is_floating_point or a.dtype.is_complex):
        a = a.to(_dtypes_impl.default_dtypes().float_dtype)
    return a


def _atleast_float_2(
    a: torch.Tensor, b: torch.Tensor
) -> tuple[torch.Tensor, torch.Tensor]:
    dtyp = _dtypes_impl.result_type_impl(a, b)
    if not (dtyp.is_floating_point or dtyp.is_complex):
        dtyp = _dtypes_impl.default_dtypes().float_dtype

    a = _util.cast_if_needed(a, dtyp)
    b = _util.cast_if_needed(b, dtyp)
    return a, b


def linalg_errors(func: Callable[_P, _R]) -> Callable[_P, _R]:
    @functools.wraps(func)
    def wrapped(*args: _P.args, **kwds: _P.kwargs) -> _R:
        try:
            return func(*args, **kwds)
        except torch._C._LinAlgError as e:  # pyrefly: ignore[missing-attribute]  # TODO
            raise LinAlgError(*e.args)  # noqa: B904

    return wrapped


# ### Matrix and vector products ###


@normalizer
@linalg_errors
def matrix_power(a: ArrayLike, n: int) -> torch.Tensor:
    a = _atleast_float_1(a)
    return torch.linalg.matrix_power(a, n)


@normalizer
@linalg_errors
def multi_dot(
    inputs: Sequence[ArrayLike], *, out: OutArray | None = None
) -> torch.Tensor:
    return torch.linalg.multi_dot(inputs)


# ### Solving equations and inverting matrices ###


@normalizer
@linalg_errors
def solve(a: ArrayLike, b: ArrayLike) -> torch.Tensor:
    a, b = _atleast_float_2(a, b)
    return torch.linalg.solve(a, b)


@normalizer
@linalg_errors
def lstsq(
    a: ArrayLike, b: ArrayLike, rcond: float | None = None
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
    a, b = _atleast_float_2(a, b)
    # NumPy is using gelsd: https://github.com/numpy/numpy/blob/v1.24.0/numpy/linalg/umath_linalg.cpp#L3991
    # on CUDA, only `gels` is available though, so use it instead
    driver = "gels" if a.is_cuda or b.is_cuda else "gelsd"
    return torch.linalg.lstsq(a, b, rcond=rcond, driver=driver)


@normalizer
@linalg_errors
def inv(a: ArrayLike) -> torch.Tensor:
    a = _atleast_float_1(a)
    result = torch.linalg.inv(a)
    return result


@normalizer
@linalg_errors
def pinv(a: ArrayLike, rcond: float = 1e-15, hermitian: bool = False) -> torch.Tensor:
    a = _atleast_float_1(a)
    return torch.linalg.pinv(a, rtol=rcond, hermitian=hermitian)


@normalizer
@linalg_errors
def tensorsolve(
    a: ArrayLike, b: ArrayLike, axes: Sequence[int] | None = None
) -> torch.Tensor:
    a, b = _atleast_float_2(a, b)
    return torch.linalg.tensorsolve(a, b, dims=axes)


@normalizer
@linalg_errors
def tensorinv(a: ArrayLike, ind: int = 2) -> torch.Tensor:
    a = _atleast_float_1(a)
    return torch.linalg.tensorinv(a, ind=ind)


# ### Norms and other numbers ###


@normalizer
@linalg_errors
def det(a: ArrayLike) -> torch.Tensor:
    a = _atleast_float_1(a)
    return torch.linalg.det(a)


@normalizer
@linalg_errors
def slogdet(a: ArrayLike) -> tuple[torch.Tensor, torch.Tensor]:
    a = _atleast_float_1(a)
    return torch.linalg.slogdet(a)


@normalizer
@linalg_errors
def cond(x: ArrayLike, p: int | str | None = None) -> torch.Tensor:
    x = _atleast_float_1(x)

    # check if empty
    # cf: https://github.com/numpy/numpy/blob/v1.24.0/numpy/linalg/linalg.py#L1744
    if x.numel() == 0 and math.prod(x.shape[-2:]) == 0:
        raise LinAlgError("cond is not defined on empty arrays")

    result = torch.linalg.cond(x, p=p)

    # Convert nans to infs (numpy does it in a data-dependent way, depending on
    # whether the input array has nans or not)
    # XXX: NumPy does this: https://github.com/numpy/numpy/blob/v1.24.0/numpy/linalg/linalg.py#L1744
    return torch.where(torch.isnan(result), float("inf"), result)


@normalizer
@linalg_errors
def matrix_rank(
    a: ArrayLike, tol: float | None = None, hermitian: bool = False
) -> torch.Tensor | int:
    a = _atleast_float_1(a)

    if a.ndim < 2:
        return int((a != 0).any())

    if tol is None:
        # follow https://github.com/numpy/numpy/blob/v1.24.0/numpy/linalg/linalg.py#L1885
        atol = 0
        rtol = max(a.shape[-2:]) * torch.finfo(a.dtype).eps
    else:
        atol, rtol = tol, 0
    return torch.linalg.matrix_rank(a, atol=atol, rtol=rtol, hermitian=hermitian)


@normalizer
@linalg_errors
def norm(
    x: ArrayLike,
    ord: int | float | str | None = None,
    axis: int | tuple[int, ...] | None = None,
    keepdims: KeepDims = False,
) -> torch.Tensor:
    x = _atleast_float_1(x)
    return torch.linalg.norm(x, ord=ord, dim=axis)


# ### Decompositions ###


@normalizer
@linalg_errors
def cholesky(a: ArrayLike) -> torch.Tensor:
    a = _atleast_float_1(a)
    return torch.linalg.cholesky(a)


@normalizer
@linalg_errors
def qr(
    a: ArrayLike, mode: str = "reduced"
) -> tuple[torch.Tensor, torch.Tensor] | torch.Tensor:
    a = _atleast_float_1(a)
    result = torch.linalg.qr(a, mode=mode)
    if mode == "r":
        # match NumPy
        return result.R
    return result


@normalizer
@linalg_errors
def svd(
    a: ArrayLike,
    full_matrices: bool = True,
    compute_uv: bool = True,
    hermitian: bool = False,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor] | torch.Tensor:
    a = _atleast_float_1(a)
    if not compute_uv:
        return torch.linalg.svdvals(a)

    # NB: ignore the hermitian= argument (no pytorch equivalent)
    result = torch.linalg.svd(a, full_matrices=full_matrices)
    return result


# ### Eigenvalues and eigenvectors ###


@normalizer
@linalg_errors
def eig(a: ArrayLike) -> tuple[torch.Tensor, torch.Tensor]:
    a = _atleast_float_1(a)
    w, vt = torch.linalg.eig(a)

    if not a.is_complex() and w.is_complex() and (w.imag == 0).all():
        w = w.real
        vt = vt.real
    return w, vt


@normalizer
@linalg_errors
def eigh(a: ArrayLike, UPLO: str = "L") -> tuple[torch.Tensor, torch.Tensor]:
    a = _atleast_float_1(a)
    return torch.linalg.eigh(a, UPLO=UPLO)


@normalizer
@linalg_errors
def eigvals(a: ArrayLike) -> torch.Tensor:
    a = _atleast_float_1(a)
    result = torch.linalg.eigvals(a)
    if not a.is_complex() and result.is_complex() and (result.imag == 0).all():
        result = result.real
    return result


@normalizer
@linalg_errors
def eigvalsh(a: ArrayLike, UPLO: str = "L") -> torch.Tensor:
    a = _atleast_float_1(a)
    return torch.linalg.eigvalsh(a, UPLO=UPLO)
