"""
CuPy backend。

CuPy は NVIDIA CUDA と AMD ROCm の双方を NumPy/SciPy 互換 API で扱える。
このクラスでは CUDA/ROCm の違いを利用側コードに露出させず、
CuPy backend + device_id として扱う。
"""

from __future__ import annotations

import numpy as np

from ..base import Backend


class CupyBackend(Backend):
    """CuPy backend。"""

    name = "cupy"

    def __init__(self, device_id: int = 0):
        try:
            import cupy as cp
        except Exception as exc:
            raise RuntimeError(
                "CuPy backend is unavailable. Install CuPy for your CUDA/ROCm "
                f"environment. Original error: {exc}"
            ) from exc

        self._cp = cp
        self.device_id = int(device_id)

        try:
            ndev = cp.cuda.runtime.getDeviceCount()
            if ndev <= self.device_id:
                raise RuntimeError(
                    f"Requested device_id={self.device_id}, "
                    f"but only {ndev} CuPy device(s) are visible."
                )
            self._device = cp.cuda.Device(self.device_id)
            self._device.use()
        except Exception as exc:
            raise RuntimeError(
                f"CuPy device initialization failed: {exc}"
            ) from exc

    @property
    def xp(self):
        return self._cp

    @property
    def is_accelerator(self) -> bool:
        return True

    def asarray(self, x, dtype=None, order=None):
        kwargs = {}
        if dtype is not None:
            kwargs["dtype"] = dtype
        if order is not None:
            kwargs["order"] = order

        with self._device:
            return self._cp.asarray(x, **kwargs)

    def to_numpy(self, x) -> np.ndarray:
        with self._device:
            return self._cp.asnumpy(x)

    def synchronize(self) -> None:
        with self._device:
            self._cp.cuda.get_current_stream().synchronize()

    def fftn(self, x, axes=None):
        with self._device:
            return self._cp.fft.fftn(x, axes=axes)

    def ifftn(self, x, axes=None):
        with self._device:
            return self._cp.fft.ifftn(x, axes=axes)

    def fft(self, x, axis=-1):
        with self._device:
            return self._cp.fft.fft(x, axis=axis)

    def ifft(self, x, axis=-1):
        with self._device:
            return self._cp.fft.ifft(x, axis=axis)

    def fft2(self, x, axes=(-2, -1)):
        with self._device:
            return self._cp.fft.fft2(x, axes=axes)

    def ifft2(self, x, axes=(-2, -1)):
        with self._device:
            return self._cp.fft.ifft2(x, axes=axes)

    def get_device_name(self) -> str:
        props = self._cp.cuda.runtime.getDeviceProperties(self.device_id)
        name = props["name"]
        if isinstance(name, bytes):
            name = name.decode(errors="replace")
        return str(name)

    def get_device_info(self) -> dict[str, object]:
        info = super().get_device_info()
        info["device_id"] = self.device_id
        info["cupy_version"] = self._cp.__version__

        try:
            runtime = self._cp.cuda.runtime
            info["runtime_version"] = runtime.runtimeGetVersion()
            info["driver_version"] = runtime.driverGetVersion()
        except Exception:
            pass

        try:
            with self._device:
                free_b, total_b = self._cp.cuda.runtime.memGetInfo()
            info["memory_total_gib"] = total_b / 1024**3
            info["memory_free_gib"] = free_b / 1024**3
        except Exception:
            pass

        # CuPy は ROCm 環境でも cupy.cuda.* API を互換層として使う。
        # is_hip が存在する版では実行基盤をログに残す。
        try:
            is_hip = bool(getattr(self._cp.cuda.runtime, "is_hip"))
            info["platform"] = "rocm/hip" if is_hip else "cuda"
        except Exception:
            info["platform"] = "cupy"

        return info
