"""Tests for psd_estimation.py: periodogram, Bartlett, Welch, confidence intervals."""

import numpy as np
import pytest
from scipy.signal import welch as scipy_welch
from psd_estimation import (
    periodogram, bartlett, welch, window_normalization, psd_confidence,
)


class TestPeriodogram:
    def test_shapes(self):
        n = 512
        x = np.random.default_rng(0).standard_normal(n)
        f, psd = periodogram(x, fs=1000)
        assert f.shape == (n // 2 + 1,)
        assert psd.shape == f.shape

    def test_white_noise_flat_psd(self):
        n = 4096
        x = np.random.default_rng(1).standard_normal(n)
        _, psd = periodogram(x, fs=1.0)
        # White noise should have roughly flat PSD (but periodogram is noisy).
        # Check that the mean across frequency is positive and finite.
        assert np.mean(psd) > 0
        assert np.all(np.isfinite(psd))

    def test_sine_produces_peak(self):
        n = 1024
        fs = 1000
        t = np.arange(n) / fs
        x = np.sin(2 * np.pi * 50 * t)
        f, psd = periodogram(x, fs)
        idx_50hz = np.argmin(np.abs(f - 50))
        # Peak at 50 Hz should dominate
        assert psd[idx_50hz] > 10 * np.median(psd)

    def test_dc_bin_not_doubled(self):
        n = 256
        x = np.ones(n)
        _, psd = periodogram(x, fs=1.0, window='boxcar')
        # DC bin should have all power; check it's not zero.
        assert psd[0] > 0

    def test_matches_scipy_even_and_odd_length(self):
        # Regression: odd N has no Nyquist bin, so the last bin must be
        # doubled too (it was not, leaving it a factor 2 low for odd N).
        from scipy.signal import periodogram as scipy_periodogram
        rng = np.random.default_rng(7)
        for n in (256, 257):
            x = rng.standard_normal(n)
            _, psd = periodogram(x, fs=100.0, window='boxcar')
            # detrend=False: scipy mean-removes by default, our module does not
            _, psd_ref = scipy_periodogram(x, fs=100.0, window='boxcar',
                                           detrend=False)
            np.testing.assert_allclose(psd, psd_ref, rtol=1e-10, atol=1e-12)

    def test_windowed_periodogram_changes_peak_width(self):
        n = 1024
        t = np.arange(n) / n
        x = np.sin(2 * np.pi * 20 * t)
        _, psd_boxcar = periodogram(x, fs=n, window='boxcar')
        _, psd_hann = periodogram(x, fs=n, window='hann')
        # Hann window should reduce peak amplitude slightly (scalloping loss)
        peak_box = np.max(psd_boxcar)
        peak_hann = np.max(psd_hann)
        # Hann window broadens the main lobe, reduces peak
        assert peak_hann < peak_box


class TestBartlett:
    def test_shapes(self):
        n = 1024
        x = np.random.default_rng(2).standard_normal(n)
        f, psd = bartlett(x, fs=2000, nperseg=256)
        assert f.shape == (129,)
        assert psd.shape == f.shape

    def test_variance_lower_than_periodogram(self):
        """Bartlett's method should reduce variance vs raw periodogram."""
        n = 8192
        fs = 1000
        nperseg = 256
        rng = np.random.default_rng(3)
        # Generate coloured noise so variance is easier to compare
        x = np.cumsum(rng.standard_normal(n))  # Brownian
        x = x / np.std(x)

        _, p_raw = periodogram(x, fs)
        _, p_bart = bartlett(x, fs, nperseg=nperseg)

        # Variance across frequency bins (proxy for estimator variance)
        # Bartlett should have lower frequency-to-frequency variance
        var_raw = np.var(p_raw[10:-10])   # exclude DC and edges
        var_bart = np.var(p_bart[10:-10])
        assert var_bart < var_raw

    def test_short_signal_raises(self):
        x = np.random.default_rng(0).standard_normal(10)
        with pytest.raises(ValueError):
            bartlett(x, nperseg=256)


class TestWelch:
    def test_shapes(self):
        n = 2048
        x = np.random.default_rng(4).standard_normal(n)
        f, psd = welch(x, fs=1000, nperseg=256)
        assert f.shape == (129,)
        assert psd.shape == f.shape

    def test_against_scipy(self):
        """Compare against scipy.signal.welch with same parameters."""
        n = 8192
        fs = 500
        rng = np.random.default_rng(5)
        x = rng.standard_normal(n)
        nperseg = 256
        noverlap = 128

        f_w, psd_w = welch(x, fs=fs, nperseg=nperseg, noverlap=noverlap, window='hann')
        f_s, psd_s = scipy_welch(x, fs=fs, nperseg=nperseg, noverlap=noverlap,
                                 window='hann', return_onesided=True, scaling='density',
                                 detrend=False)

        # scipy and ours should agree closely (within ~1% relative for white noise).
        rel_err = np.abs(psd_w - psd_s) / (psd_s + 1e-12)
        assert np.median(rel_err) < 0.02

    def test_50_percent_overlap_uses_center_segments(self):
        n = 2048
        x = np.random.default_rng(6).standard_normal(n)
        _, psd1 = welch(x, nperseg=256, noverlap=128)
        # Should not raise and should produce valid PSD
        assert np.all(np.isfinite(psd1))

    def test_no_overlap_is_bartlett(self):
        n = 1024
        x = np.random.default_rng(7).standard_normal(n)
        _, psd_w = welch(x, nperseg=256, noverlap=0, window='boxcar')
        _, psd_b = bartlett(x, nperseg=256, window='boxcar')
        assert np.allclose(psd_w, psd_b)

    def test_reproduces_white_noise_level(self):
        """Welch estimate of unit-variance white noise should average to ~1/fs."""
        n = 32768
        fs = 500
        rng = np.random.default_rng(8)
        x = rng.standard_normal(n)
        _, psd = welch(x, fs=fs, nperseg=1024, noverlap=512, window='hann')
        # Mean across frequency should be ~sigma²/(fs/2) = 1/250 = 0.004
        expected = 1.0 / (fs / 2)
        assert 0.5 * expected < np.mean(psd) < 2.0 * expected


class TestWindowNormalization:
    def test_boxcar_norm(self):
        norm = window_normalization('boxcar', 256)
        assert norm > 0

    def test_hann_norm_larger_than_boxcar(self):
        """Hann window attenuates the signal, so needs larger normalisation."""
        norm_box = window_normalization('boxcar', 256)
        norm_hann = window_normalization('hann', 256)
        assert norm_hann > norm_box


class TestConfidenceIntervals:
    def test_shapes(self):
        n = 4096
        x = np.random.default_rng(0).standard_normal(n)
        _, psd = welch(x, nperseg=256, noverlap=128)
        lo, hi = psd_confidence(psd, n_avg=20, ci=0.95)
        assert lo.shape == psd.shape
        assert hi.shape == psd.shape
        assert np.all(lo < psd)
        assert np.all(hi > psd)

    def test_higher_ci_wider_interval(self):
        _, psd = welch(np.random.default_rng(0).standard_normal(4096), nperseg=256, noverlap=128)
        lo_95, hi_95 = psd_confidence(psd, n_avg=15, ci=0.95)
        lo_99, hi_99 = psd_confidence(psd, n_avg=15, ci=0.99)
        width_95 = np.mean(hi_95 - lo_95)
        width_99 = np.mean(hi_99 - lo_99)
        assert width_99 > width_95

    def test_more_averages_narrower_interval(self):
        _, psd = welch(np.random.default_rng(0).standard_normal(8192), nperseg=256, noverlap=128)
        lo_few, hi_few = psd_confidence(psd, n_avg=5, ci=0.95)
        lo_many, hi_many = psd_confidence(psd, n_avg=50, ci=0.95)
        width_few = np.mean(hi_few - lo_few)
        width_many = np.mean(hi_many - lo_many)
        assert width_few > width_many
