"""Tests for the SpectralAnalyzer and DSP utilities."""

from __future__ import annotations

from pathlib import Path
from unittest.mock import MagicMock

import numpy as np
import pytest

from analyzer.context import AnalysisContext
from analyzer.default_registry import build_default_registry
from analyzer.pipeline import AnalysisPipeline
from analyzer.result import AnalysisResult
from analyzer.spectral.spectral_analyzer import SpectralAnalyzer
from config.settings import Settings
from conftest import init_test_db
from core.exceptions import AnalysisError
from database.repositories import ProjectRepository, TrackRepository
from decoder.decoded_audio import DecodedAudio
from dsp.config import FFT_SIZE, HOP_LENGTH, WINDOW
from dsp.spectrum import (
    compute_magnitude_spectrum,
    compute_power_spectrum,
    compute_stft,
    hann_window,
)
from storage.analysis_repository import AnalysisRepository
from storage.feature_repository import FeatureRepository

SAMPLE_RATE = 44100


def _make_decoded_audio(
    samples: np.ndarray,
    sample_rate: int = SAMPLE_RATE,
    duration: float | None = None,
    bit_depth: int | None = 16,
) -> DecodedAudio:
    """Create a DecodedAudio from a NumPy array."""
    n_frames = samples.shape[0]
    channels = samples.shape[1] if samples.ndim > 1 else 1
    if samples.ndim == 1:
        samples = samples.reshape(-1, 1)
    if duration is None:
        duration = n_frames / sample_rate
    return DecodedAudio(
        samples=samples,
        sample_rate=sample_rate,
        channels=channels,
        duration=duration,
        bit_depth=bit_depth,
    )


def _make_context(samples: np.ndarray, tmp_path: Path) -> AnalysisContext:
    """Create an AnalysisContext from a samples array."""
    track = MagicMock()
    track.id = "test-track-id"
    track.relative_path = "test.wav"
    audio = _make_decoded_audio(samples)
    settings = Settings(
        project_name="Test",
        database_url=f"sqlite:///{tmp_path / 'test.db'}",
        log_level="DEBUG",
        cache_directory=str(tmp_path / "cache"),
        data_directory=str(tmp_path / "data"),
        output_directory=str(tmp_path / "output"),
    )
    return AnalysisContext(
        track=track,
        decoded_audio=audio,
        settings=settings,
    )


def _analyze(samples: np.ndarray, tmp_path: Path) -> AnalysisResult:
    """Run SpectralAnalyzer on the given samples."""
    context = _make_context(samples, tmp_path)
    analyzer = SpectralAnalyzer()
    return analyzer.analyze(context)


def _sine_wave(
    freq: float,
    duration: float = 1.0,
    sample_rate: int = SAMPLE_RATE,
) -> np.ndarray:
    """Generate a sine wave at the given frequency."""
    t = np.arange(int(sample_rate * duration)) / sample_rate
    return (0.5 * np.sin(2.0 * np.pi * freq * t)).astype(np.float32)


def _white_noise(
    duration: float = 1.0,
    sample_rate: int = SAMPLE_RATE,
    seed: int = 42,
) -> np.ndarray:
    """Generate white noise."""
    rng = np.random.default_rng(seed)
    n = int(sample_rate * duration)
    return (0.5 * rng.standard_normal(n)).astype(np.float32)


# --- DSP utility tests ---


def test_hann_window() -> None:
    """Verify Hann window has correct shape and properties."""
    w = hann_window(FFT_SIZE)
    assert w.shape == (FFT_SIZE,)
    assert w[0] >= 0.0
    assert w[-1] >= 0.0
    assert np.all(w >= 0.0)
    assert np.all(w <= 1.0)


def test_compute_stft() -> None:
    """Verify STFT output shape."""
    samples = _sine_wave(440.0, duration=0.5)
    window = hann_window(FFT_SIZE)
    stft = compute_stft(samples, FFT_SIZE, HOP_LENGTH, window)
    n_frames = 1 + (len(samples) - FFT_SIZE) // HOP_LENGTH
    n_bins = FFT_SIZE // 2 + 1
    assert stft.shape == (n_frames, n_bins)
    assert stft.dtype == np.complex64


def test_compute_magnitude_spectrum() -> None:
    """Verify magnitude spectrum is real and non-negative."""
    samples = _sine_wave(440.0, duration=0.5)
    window = hann_window(FFT_SIZE)
    stft = compute_stft(samples, FFT_SIZE, HOP_LENGTH, window)
    mag = compute_magnitude_spectrum(stft)
    assert mag.shape == stft.shape
    assert np.all(mag >= 0.0)


def test_compute_power_spectrum() -> None:
    """Verify power spectrum is real and non-negative."""
    samples = _sine_wave(440.0, duration=0.5)
    window = hann_window(FFT_SIZE)
    stft = compute_stft(samples, FFT_SIZE, HOP_LENGTH, window)
    mag = compute_magnitude_spectrum(stft)
    power = compute_power_spectrum(mag)
    assert power.shape == mag.shape
    assert np.all(power >= 0.0)


def test_dsp_config_constants() -> None:
    """Verify DSP config constants are correct."""
    assert FFT_SIZE == 4096
    assert HOP_LENGTH == 1024
    assert WINDOW == "hann"


# --- Sine wave tests ---


def test_sine_wave(tmp_path: Path) -> None:
    """Verify sine wave produces valid spectral features."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    centroid = result.feature_set.find_by_name("spectral.centroid")
    assert centroid is not None
    assert 200.0 < centroid.value < 800.0


# --- White noise tests ---


def test_white_noise(tmp_path: Path) -> None:
    """Verify white noise has high flatness."""
    samples = _white_noise(duration=1.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    flatness = result.feature_set.find_by_name("spectral.flatness")
    assert flatness is not None
    assert flatness.value > 0.3


# --- Silence tests ---


def test_silence(tmp_path: Path) -> None:
    """Verify all-zero input produces safe defaults."""
    samples = np.zeros(SAMPLE_RATE, dtype=np.float32)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    centroid = result.feature_set.find_by_name("spectral.centroid")
    assert centroid is not None
    assert centroid.value == 0.0
    flatness = result.feature_set.find_by_name("spectral.flatness")
    assert flatness is not None
    assert flatness.value == 0.0


# --- Impulse tests ---


def test_impulse(tmp_path: Path) -> None:
    """Verify single-sample impulse produces deterministic output."""
    samples = np.zeros(SAMPLE_RATE, dtype=np.float32)
    samples[0] = 1.0
    result = _analyze(samples, tmp_path)
    assert result.success is True
    assert len(result.feature_set) == 6


# --- Constant signal tests ---


def test_constant_signal(tmp_path: Path) -> None:
    """Verify constant DC signal has zero crossing rate = 0."""
    samples = np.full(SAMPLE_RATE, 0.5, dtype=np.float32)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    zcr = result.feature_set.find_by_name("spectral.zero_crossing_rate")
    assert zcr is not None
    assert zcr.value == pytest.approx(0.0, abs=1e-6)


# --- Mono / Stereo tests ---


def test_mono(tmp_path: Path) -> None:
    """Verify mono audio produces 6 features."""
    samples = _sine_wave(440.0, duration=0.5).reshape(-1, 1)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    assert len(result.feature_set) == 6


def test_stereo(tmp_path: Path) -> None:
    """Verify stereo audio converted to mono produces 6 features."""
    mono = _sine_wave(440.0, duration=0.5)
    samples = np.stack([mono, mono], axis=1)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    assert len(result.feature_set) == 6


def test_stereo_matches_mono(tmp_path: Path) -> None:
    """Verify stereo (L=R) produces same features as mono."""
    mono = _sine_wave(440.0, duration=0.5)
    result_mono = _analyze(mono.reshape(-1, 1), tmp_path)
    result_stereo = _analyze(np.stack([mono, mono], axis=1), tmp_path)

    for name in [
        "spectral.centroid",
        "spectral.bandwidth",
        "spectral.rolloff",
        "spectral.flatness",
        "spectral.zero_crossing_rate",
        "spectral.flux",
    ]:
        f_mono = result_mono.feature_set.find_by_name(name)
        f_stereo = result_stereo.feature_set.find_by_name(name)
        assert f_mono is not None
        assert f_stereo is not None
        assert f_mono.value == pytest.approx(f_stereo.value, abs=1e-4)


# --- Individual feature tests ---


def test_spectral_centroid(tmp_path: Path) -> None:
    """Verify spectral centroid is in valid range."""
    samples = _sine_wave(1000.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    centroid = result.feature_set.find_by_name("spectral.centroid")
    assert centroid is not None
    assert centroid.value >= 0.0
    assert centroid.value <= SAMPLE_RATE / 2


def test_spectral_bandwidth(tmp_path: Path) -> None:
    """Verify spectral bandwidth is non-negative."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    bw = result.feature_set.find_by_name("spectral.bandwidth")
    assert bw is not None
    assert bw.value >= 0.0


def test_spectral_rolloff(tmp_path: Path) -> None:
    """Verify spectral rolloff is in valid frequency range."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    rolloff = result.feature_set.find_by_name("spectral.rolloff")
    assert rolloff is not None
    assert rolloff.value >= 0.0
    assert rolloff.value <= SAMPLE_RATE / 2


def test_spectral_flatness(tmp_path: Path) -> None:
    """Verify spectral flatness is in [0, 1]."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    flatness = result.feature_set.find_by_name("spectral.flatness")
    assert flatness is not None
    assert 0.0 <= flatness.value <= 1.0


def test_zero_crossing_rate(tmp_path: Path) -> None:
    """Verify zero crossing rate is in [0, 1]."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    zcr = result.feature_set.find_by_name("spectral.zero_crossing_rate")
    assert zcr is not None
    assert 0.0 <= zcr.value <= 1.0


def test_spectral_flux(tmp_path: Path) -> None:
    """Verify spectral flux is non-negative."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    flux = result.feature_set.find_by_name("spectral.flux")
    assert flux is not None
    assert flux.value >= 0.0


# --- Determinism test ---


def test_deterministic(tmp_path: Path) -> None:
    """Verify same input produces identical output."""
    samples = _sine_wave(440.0, duration=1.0)
    result1 = _analyze(samples, tmp_path)
    result2 = _analyze(samples, tmp_path)

    f1 = {f.name: f.value for f in result1.feature_set}
    f2 = {f.name: f.value for f in result2.feature_set}
    assert f1.keys() == f2.keys()
    for key in f1:
        assert f1[key] == pytest.approx(f2[key], abs=1e-10)


# --- Error handling tests ---


def test_nan_raises(tmp_path: Path) -> None:
    """Verify NaN values raise AnalysisError."""
    samples = np.zeros(100, dtype=np.float32)
    samples[50] = np.nan
    with pytest.raises(AnalysisError, match="NaN"):
        _analyze(samples, tmp_path)


def test_inf_raises(tmp_path: Path) -> None:
    """Verify infinite values raise AnalysisError."""
    samples = np.zeros(100, dtype=np.float32)
    samples[50] = np.inf
    with pytest.raises(AnalysisError, match="infinite"):
        _analyze(samples, tmp_path)


def test_empty_raises(tmp_path: Path) -> None:
    """Verify empty audio raises AnalysisError."""
    samples = np.array([], dtype=np.float32).reshape(0, 1)
    with pytest.raises(AnalysisError, match="no samples"):
        _analyze(samples, tmp_path)


# --- Feature names and units tests ---


def test_feature_names(tmp_path: Path) -> None:
    """Verify all 6 expected feature names are present."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    expected = {
        "spectral.centroid",
        "spectral.bandwidth",
        "spectral.rolloff",
        "spectral.flatness",
        "spectral.zero_crossing_rate",
        "spectral.flux",
    }
    actual = {f.name for f in result.feature_set}
    assert expected == actual


def test_feature_units(tmp_path: Path) -> None:
    """Verify Hz features have correct unit."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    centroid = result.feature_set.find_by_name("spectral.centroid")
    assert centroid is not None
    assert centroid.unit == "Hz"
    bw = result.feature_set.find_by_name("spectral.bandwidth")
    assert bw is not None
    assert bw.unit == "Hz"
    rolloff = result.feature_set.find_by_name("spectral.rolloff")
    assert rolloff is not None
    assert rolloff.unit == "Hz"


# --- Pipeline integration tests ---


def test_pipeline_integration(tmp_path: Path) -> None:
    """Verify pipeline runs SpectralAnalyzer via default registry."""
    registry = build_default_registry()
    analyzers = registry.list_analyzers()
    names = [a.name for a in analyzers]
    assert "spectral" in names

    pipeline = AnalysisPipeline(registry)
    samples = _sine_wave(440.0, duration=0.5).reshape(-1, 1)
    context = _make_context(samples, tmp_path)
    results = pipeline.run(context)

    spectral_results = [r for r in results if r.analyzer_name == "spectral"]
    assert len(spectral_results) == 1
    assert spectral_results[0].success is True
    assert len(spectral_results[0].feature_set) == 6


def test_persistence_integration(tmp_path: Path) -> None:
    """Verify pipeline with session persists spectral features to DB."""
    session = init_test_db(tmp_path)
    try:
        project = ProjectRepository.get_default_project(session)
        track = TrackRepository.create_track(
            session,
            project_id=project.id,
            relative_path="test.wav",
            original_filename="test.wav",
            sha256="hash_spectral_test",
            file_size=1000,
        )

        context = _make_context(
            _sine_wave(440.0, duration=0.5).reshape(-1, 1), tmp_path
        )
        context.track.id = track.id

        registry = build_default_registry()
        pipeline = AnalysisPipeline(registry)
        results = pipeline.run(context, session=session)
        session.commit()

        spectral_results = [r for r in results if r.analyzer_name == "spectral"]
        assert spectral_results[0].success is True

        runs = AnalysisRepository.list_runs(session, track.id)
        spectral_runs = [r for r in runs if r.analyzer_name == "spectral"]
        assert len(spectral_runs) == 1
        assert spectral_runs[0].success is True

        features = FeatureRepository.get_track_features(session, track.id)
        spectral_features = [
            f for f in features if f.analyzer_run.analyzer_name == "spectral"
        ]
        assert len(spectral_features) == 6
    finally:
        session.close()
