"""Tests for the DynamicAnalyzer and DSP dynamics 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.dynamics.dynamic_analyzer import DynamicAnalyzer
from analyzer.pipeline import AnalysisPipeline
from analyzer.result import AnalysisResult
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 RMS_HOP_SIZE, RMS_WINDOW_SIZE
from dsp.dynamics import (
    compute_clipping_ratio,
    compute_dynamic_range,
    compute_headroom,
    compute_rms_envelope,
)
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 DynamicAnalyzer on the given samples."""
    context = _make_context(samples, tmp_path)
    analyzer = DynamicAnalyzer()
    return analyzer.analyze(context)


def _sine_wave(
    freq: float,
    duration: float = 1.0,
    sample_rate: int = SAMPLE_RATE,
    amplitude: float = 0.5,
) -> np.ndarray:
    """Generate a sine wave at the given frequency."""
    t = np.arange(int(sample_rate * duration)) / sample_rate
    return (amplitude * 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_compute_rms_envelope() -> None:
    """Verify RMS envelope has correct shape."""
    samples = _sine_wave(440.0, duration=0.5)
    env = compute_rms_envelope(samples, RMS_WINDOW_SIZE, RMS_HOP_SIZE)
    n_frames = 1 + (len(samples) - RMS_WINDOW_SIZE) // RMS_HOP_SIZE
    assert len(env) == n_frames
    assert np.all(env >= 0.0)


def test_compute_rms_envelope_short_signal() -> None:
    """Verify RMS envelope handles signals shorter than window."""
    samples = np.array([0.1, 0.2, 0.3], dtype=np.float32)
    env = compute_rms_envelope(samples, RMS_WINDOW_SIZE, RMS_HOP_SIZE)
    assert len(env) == 1
    assert env[0] > 0.0


def test_compute_dynamic_range() -> None:
    """Verify dynamic range is non-negative."""
    env = np.array([0.1, 0.5, 0.2, 0.8, 0.3], dtype=np.float64)
    dr = compute_dynamic_range(env)
    assert dr >= 0.0


def test_compute_dynamic_range_silence() -> None:
    """Verify dynamic range is 0 for silent envelope."""
    env = np.zeros(10, dtype=np.float64)
    dr = compute_dynamic_range(env)
    assert dr == 0.0


def test_compute_headroom() -> None:
    """Verify headroom is non-negative."""
    assert compute_headroom(0.5) > 0.0
    assert compute_headroom(1.0) == 0.0
    assert compute_headroom(0.0) == float("inf")


def test_compute_clipping_ratio() -> None:
    """Verify clipping ratio is in [0, 1]."""
    samples = np.array([0.5, 1.0, 0.999, 0.3], dtype=np.float32)
    ratio = compute_clipping_ratio(samples)
    assert 0.0 <= ratio <= 1.0
    assert ratio == 0.5  # 2 of 4 samples >= 0.999


def test_compute_clipping_ratio_empty() -> None:
    """Verify clipping ratio is 0 for empty signal."""
    ratio = compute_clipping_ratio(np.array([], dtype=np.float32))
    assert ratio == 0.0


def test_dsp_config_constants() -> None:
    """Verify DSP config constants are correct."""
    assert RMS_WINDOW_SIZE == 2048
    assert RMS_HOP_SIZE == 512


# --- Sine wave test ---


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


# --- Silence test ---


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
    crest = result.feature_set.find_by_name("dynamic.crest_factor")
    assert crest is not None
    assert crest.value == 1.0
    crest_db = result.feature_set.find_by_name("dynamic.crest_factor_db")
    assert crest_db is not None
    assert crest_db.value == 0.0
    dr = result.feature_set.find_by_name("dynamic.range")
    assert dr is not None
    assert dr.value == 0.0
    var = result.feature_set.find_by_name("dynamic.rms_variability")
    assert var is not None
    assert var.value == 0.0


# --- Clipped signal test ---


def test_clipped_signal(tmp_path: Path) -> None:
    """Verify clipped signal has clipping ratio > 0."""
    samples = _sine_wave(440.0, duration=1.0, amplitude=1.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    clip = result.feature_set.find_by_name("dynamic.clipping_ratio")
    assert clip is not None
    assert clip.value > 0.0


# --- Impulse test ---


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) == 7


# --- White noise test ---


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


# --- Constant signal test ---


def test_constant_signal(tmp_path: Path) -> None:
    """Verify constant DC signal has low variability."""
    samples = np.full(SAMPLE_RATE, 0.5, dtype=np.float32)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    var = result.feature_set.find_by_name("dynamic.rms_variability")
    assert var is not None
    assert var.value == pytest.approx(0.0, abs=1e-6)


# --- Mono / Stereo tests ---


def test_mono(tmp_path: Path) -> None:
    """Verify mono audio produces 7 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) == 7


def test_stereo(tmp_path: Path) -> None:
    """Verify stereo audio converted to mono produces 7 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) == 7


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 [
        "dynamic.crest_factor",
        "dynamic.crest_factor_db",
        "dynamic.range",
        "dynamic.headroom",
        "dynamic.clipping_ratio",
        "dynamic.average_rms_db",
        "dynamic.rms_variability",
    ]:
        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-6)


# --- Individual feature tests ---


def test_crest_factor(tmp_path: Path) -> None:
    """Verify crest factor >= 1.0."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    crest = result.feature_set.find_by_name("dynamic.crest_factor")
    assert crest is not None
    assert crest.value >= 1.0


def test_crest_factor_db(tmp_path: Path) -> None:
    """Verify crest factor dB >= 0.0."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    crest_db = result.feature_set.find_by_name("dynamic.crest_factor_db")
    assert crest_db is not None
    assert crest_db.value >= 0.0


def test_headroom(tmp_path: Path) -> None:
    """Verify headroom >= 0.0."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    headroom = result.feature_set.find_by_name("dynamic.headroom")
    assert headroom is not None
    assert headroom.value >= 0.0


def test_clipping_ratio(tmp_path: Path) -> None:
    """Verify clipping ratio is in [0, 1]."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    clip = result.feature_set.find_by_name("dynamic.clipping_ratio")
    assert clip is not None
    assert 0.0 <= clip.value <= 1.0


def test_dynamic_range(tmp_path: Path) -> None:
    """Verify dynamic range >= 0.0."""
    samples = _white_noise(duration=1.0)
    result = _analyze(samples, tmp_path)
    dr = result.feature_set.find_by_name("dynamic.range")
    assert dr is not None
    assert dr.value >= 0.0


def test_rms_variability(tmp_path: Path) -> None:
    """Verify RMS variability >= 0.0."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    var = result.feature_set.find_by_name("dynamic.rms_variability")
    assert var is not None
    assert var.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 7 expected feature names are present."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)
    expected = {
        "dynamic.crest_factor",
        "dynamic.crest_factor_db",
        "dynamic.range",
        "dynamic.headroom",
        "dynamic.clipping_ratio",
        "dynamic.average_rms_db",
        "dynamic.rms_variability",
    }
    actual = {f.name for f in result.feature_set}
    assert expected == actual


def test_feature_units(tmp_path: Path) -> None:
    """Verify correct units for each feature."""
    samples = _sine_wave(440.0, duration=1.0)
    result = _analyze(samples, tmp_path)

    crest = result.feature_set.find_by_name("dynamic.crest_factor")
    assert crest is not None
    assert crest.unit == "ratio"

    dr = result.feature_set.find_by_name("dynamic.range")
    assert dr is not None
    assert dr.unit == "dB"

    avg = result.feature_set.find_by_name("dynamic.average_rms_db")
    assert avg is not None
    assert avg.unit == "dBFS"

    var = result.feature_set.find_by_name("dynamic.rms_variability")
    assert var is not None
    assert var.unit == "dB"


# --- Pipeline integration tests ---


def test_pipeline_integration(tmp_path: Path) -> None:
    """Verify pipeline runs DynamicAnalyzer via default registry."""
    registry = build_default_registry()
    analyzers = registry.list_analyzers()
    names = [a.name for a in analyzers]
    assert "dynamic" 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)

    dynamic_results = [r for r in results if r.analyzer_name == "dynamic"]
    assert len(dynamic_results) == 1
    assert dynamic_results[0].success is True
    assert len(dynamic_results[0].feature_set) == 7


def test_persistence_integration(tmp_path: Path) -> None:
    """Verify pipeline with session persists dynamic 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_dynamic_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()

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

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

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