"""
dpnp + dpctl (SYCL) backend。

主に Intel GPU/CPU などの SYCL device を対象とする。
device は backend 名とは独立に selector で指定する。

Examples
--------
    bk = DpnpBackend(device="gpu")
    bk = DpnpBackend(device="level_zero:gpu:0")
    bk = DpnpBackend(device="cpu")
"""

from __future__ import annotations

import numpy as np

from ..base import Backend


class DpnpBackend(Backend):
    """dpnp/SYCL backend。"""

    name = "dpnp"

    def __init__(self, device: str = "gpu"):
        try:
            import dpnp
            import dpctl
        except Exception as exc:
            raise RuntimeError(
                "dpnp backend is unavailable. Install dpnp and dpctl. "
                f"Original error: {exc}"
            ) from exc

        self._dpnp = dpnp
        self._dpctl = dpctl
        self.device_selector = device

        try:
            # backend が queue を所有することで、asarray/同期を同じ queue に固定する。
            self._queue = dpctl.SyclQueue(device)
        except Exception as exc:
            raise RuntimeError(
                f"Could not create a SYCL queue for device={device!r}: {exc}"
            ) from exc

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

    @property
    def is_accelerator(self) -> bool:
        try:
            return bool(self._queue.sycl_device.is_gpu)
        except Exception:
            return True

    def asarray(self, x, dtype=None, order=None):
        kwargs = {"sycl_queue": self._queue}

        if dtype is not None:
            kwargs["dtype"] = dtype
        if order is not None:
            kwargs["order"] = order

        return self._dpnp.asarray(x, **kwargs)

    def to_numpy(self, x) -> np.ndarray:
        # dpnp.asnumpy() は NumPy ndarray を返す。
        return self._dpnp.asnumpy(x)

    def synchronize(self) -> None:
        # dpctl.SyclQueue.wait() で、この backend が所有する queue の完了を待つ。
        self._queue.wait()

    def fftn(self, x, axes=None):
        return self._dpnp.fft.fftn(x, axes=axes)

    def ifftn(self, x, axes=None):
        return self._dpnp.fft.ifftn(x, axes=axes)

    def fft(self, x, axis=-1):
        return self._dpnp.fft.fft(x, axis=axis)

    def ifft(self, x, axis=-1):
        return self._dpnp.fft.ifft(x, axis=axis)

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

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

    def get_device_name(self) -> str:
        dev = self._queue.sycl_device
        try:
            return str(dev.name)
        except Exception:
            return str(dev)

    def get_device_info(self) -> dict[str, object]:
        info = super().get_device_info()
        dev = self._queue.sycl_device

        info["device_selector"] = self.device_selector
        info["dpnp_version"] = getattr(self._dpnp, "__version__", "unknown")
        info["dpctl_version"] = getattr(self._dpctl, "__version__", "unknown")

        for key, attr in (
            ("vendor", "vendor"),
            ("driver_version", "driver_version"),
        ):
            try:
                info[key] = str(getattr(dev, attr))
            except Exception:
                pass

        try:
            info["sycl_backend"] = str(self._queue.backend)
        except Exception:
            pass

        try:
            info["is_cpu"] = bool(dev.is_cpu)
            info["is_gpu"] = bool(dev.is_gpu)
        except Exception:
            pass

        return info
