"""
Direct Poisson solver for the pressure-projection step.

Solves  laplacian(phi) = rhs  on a uniform cell-centred grid with homogeneous
Neumann conditions on all six faces, using a discrete cosine transform.

Why DCT and not multigrid: the second-difference operator on a uniform grid is
diagonalised exactly by the DCT-II basis, so the solve is one forward
transform, one pointwise divide, and one inverse transform. It is a *direct*
solver -- no iteration, no convergence tolerance, no residual to babysit -- and
scipy's pocketfft backend threads it across all cores. Measured 57 ms for
2.4M cells on a 5900X, versus seconds for an iterative solve at the same
accuracy. That single choice is what makes a sub-30-minute 3D run possible.

The price is that the grid must be uniform and the pressure BCs must be all
Neumann. Both are acceptable here: velocity is prescribed on every boundary
(inlet, freestream, outlet), which makes Neumann pressure the physically
consistent choice, and the immersed-boundary method removes any need for a
body-fitted or stretched mesh.

Compatibility condition: an all-Neumann Poisson problem is solvable only if
the RHS integrates to zero, i.e. net mass flux through the boundary is zero.
The solver subtracts the mean of the RHS to enforce this exactly; the caller
is still responsible for correcting outflow so the subtraction stays small.
The solution is defined up to a constant, fixed here by setting the zero mode
to zero (zero-mean pressure).
"""

import numpy as np
import scipy.fft

from .grid import NG


class PoissonDCT:
    def __init__(self, grid, workers=-1):
        self.grid = grid
        self.workers = workers

        nx, ny, nz = grid.nx, grid.ny, grid.nz

        # Eigenvalues of the 1D second-difference operator under DCT-II
        # (homogeneous Neumann):  lam_i = 2*(cos(pi*i/N) - 1) / h^2
        lx = 2.0 * (np.cos(np.pi * np.arange(nx) / nx) - 1.0) / grid.dx**2
        ly = 2.0 * (np.cos(np.pi * np.arange(ny) / ny) - 1.0) / grid.dy**2
        lz = 2.0 * (np.cos(np.pi * np.arange(nz) / nz) - 1.0) / grid.dz**2

        # For nz == 1 this is exactly [0.0], which is what collapses the
        # solver to 2D without a separate code path.
        denom = lx[:, None, None] + ly[None, :, None] + lz[None, None, :]

        # The (0,0,0) mode is the nullspace (a constant shift in pressure).
        # Park it at 1.0 so the divide is safe, then zero the mode after.
        self._null = denom == 0.0
        denom = np.where(self._null, 1.0, denom)
        self.inv_denom = 1.0 / denom

        # Transform only the non-degenerate axes. A DCT along a length-1 axis
        # is mathematically the identity but scipy still walks the array for
        # it, and in 2D (nz=1) that pass costs as much as the two real ones
        # put together.
        self.axes = tuple(a for a, n in enumerate((nx, ny, nz)) if n > 1)

        self._mean_correction = 0.0

    def solve(self, rhs_interior, out=None):
        """
        rhs_interior : (nx, ny, nz) array, no ghosts.
        Returns phi as a ghosted (nx+2, ny+2, nz+2) array with ghosts filled
        by zero-gradient extrapolation, ready for the gradient kernel.
        """
        g = self.grid

        # Enforce the discrete compatibility condition. The residual mean is
        # recorded so run diagnostics can flag a leaking mass balance.
        mean = float(rhs_interior.mean())
        self._mean_correction = mean
        rhs = rhs_interior - mean

        rhs_hat = scipy.fft.dctn(
            rhs, type=2, axes=self.axes, norm="ortho", workers=self.workers
        )
        rhs_hat *= self.inv_denom
        rhs_hat[self._null] = 0.0
        phi_in = scipy.fft.idctn(
            rhs_hat, type=2, axes=self.axes, norm="ortho", workers=self.workers
        )

        if out is None:
            out = g.zeros_p()
        out[NG:-NG, NG:-NG, NG:-NG] = phi_in

        # Zero-gradient ghosts: consistent with the Neumann condition the
        # eigenvalues already assume, so grad(phi).n == 0 on every boundary
        # and the projection cannot inject spurious boundary-normal velocity.
        out[0, :, :] = out[1, :, :]
        out[-1, :, :] = out[-2, :, :]
        out[:, 0, :] = out[:, 1, :]
        out[:, -1, :] = out[:, -2, :]
        out[:, :, 0] = out[:, :, 1]
        out[:, :, -1] = out[:, :, -2]
        return out

    @property
    def last_mean_correction(self):
        """Mean of the last RHS. Large values mean the boundary mass balance
        is not closing -- check the outflow correction before trusting forces."""
        return self._mean_correction
