"""NumPy + SciPy FFT backend."""

from __future__ import annotations

import numpy as np
from scipy import fft as scipy_fft

from ..base import Backend
from ..device_info import get_cpu_name


class NumpyBackend(Backend):
    """
    CPU backend。

    Parameters
    ----------
    workers:
        scipy.fft の worker 数。
        -1 なら利用可能 worker をすべて使用する。
    """

    name = "numpy"

    def __init__(self, workers: int = -1):
        self.workers = workers

    @property
    def xp(self):
        return np

    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
        return np.asarray(x, **kwargs)

    def to_numpy(self, x) -> np.ndarray:
        return np.asarray(x)

    def synchronize(self) -> None:
        # CPU backend は同期実行なので何もしない。
        return None

    def fftn(self, x, axes=None):
        return scipy_fft.fftn(x, axes=axes, workers=self.workers)

    def ifftn(self, x, axes=None):
        return scipy_fft.ifftn(x, axes=axes, workers=self.workers)

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

    def ifft(self, x, axis=-1):
        return scipy_fft.ifft(x, axis=axis, workers=self.workers)

    def fft2(self, x, axes=(-2, -1)):
        return scipy_fft.fft2(x, axes=axes, workers=self.workers)

    def ifft2(self, x, axes=(-2, -1)):
        return scipy_fft.ifft2(x, axes=axes, workers=self.workers)

    def get_device_name(self) -> str:
        return get_cpu_name()

    def get_device_info(self) -> dict[str, object]:
        info = super().get_device_info()
        info["fft_workers"] = self.workers
        return info
