"""Tests for the RhythmAnalyzer and DSP rhythm 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.rhythm.rhythm_analyzer import RhythmAnalyzer
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, ONSET_HOP_SIZE, TEMPO_MAX_BPM, TEMPO_MIN_BPM
from dsp.rhythm import (
    compute_autocorrelation,
    compute_onset_envelope,
    compute_rhythm_density,
    detect_onsets,
    estimate_beat_period,
    estimate_tempo,
)
from dsp.spectrum import 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 RhythmAnalyzer on the given samples."""
    context = _make_context(samples, tmp_path)
    analyzer = RhythmAnalyzer()
    return analyzer.analyze(context)


def _sine_wave(
    freq: float,
    duration: float = 2.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 = 2.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)


def _click_track(
    bpm: float,
    duration: float = 3.0,
    sample_rate: int = SAMPLE_RATE,
    click_duration: float = 0.01,
) -> np.ndarray:
    """Generate a metronome click track at the given BPM."""
    n = int(sample_rate * duration)
    samples = np.zeros(n, dtype=np.float32)
    interval = int(sample_rate * 60.0 / bpm)
    click_len = int(sample_rate * click_duration)
    for pos in range(0, n - click_len, interval):
        t = np.arange(click_len) / sample_rate
        click = np.exp(-t * 500) * np.sin(2.0 * np.pi * 1000.0 * t)
        samples[pos : pos + click_len] = click.astype(np.float32)
    return samples


def _regular_pulse(
    interval_samples: int,
    duration: float = 3.0,
    sample_rate: int = SAMPLE_RATE,
) -> np.ndarray:
    """Generate regular pulses at fixed intervals."""
    n = int(sample_rate * duration)
    samples = np.zeros(n, dtype=np.float32)
    pulse_len = 50
    for pos in range(0, n - pulse_len, interval_samples):
        samples[pos : pos + pulse_len] = 1.0
    return samples


def _irregular_pulse(
    intervals: list[int],
    sample_rate: int = SAMPLE_RATE,
) -> np.ndarray:
    """Generate irregular pulses at varying intervals."""
    total = sum(intervals) + 100
    samples = np.zeros(total, dtype=np.float32)
    pos = 0
    pulse_len = 50
    for interval in intervals:
        samples[pos : pos + pulse_len] = 1.0
        pos += interval
    return samples


# --- DSP utility tests ---


def test_compute_onset_envelope() -> None:
    """Verify onset envelope has correct shape."""
    samples = _click_track(120, duration=1.0)
    window = hann_window(FFT_SIZE)
    env = compute_onset_envelope(samples, FFT_SIZE, ONSET_HOP_SIZE, window)
    n_frames = max(1, 1 + (len(samples) - FFT_SIZE) // ONSET_HOP_SIZE)
    assert len(env) == max(0, n_frames - 1)
    assert np.all(env >= 0.0)


def test_compute_autocorrelation() -> None:
    """Verify autocorrelation is normalized and starts at 1.0."""
    signal = _click_track(120, duration=1.0)
    acf = compute_autocorrelation(signal)
    assert len(acf) == len(signal)
    assert acf[0] == pytest.approx(1.0, abs=1e-6)


def test_estimate_tempo_click_track() -> None:
    """Verify tempo estimation on a 120 BPM click track."""
    samples = _click_track(120, duration=3.0)
    window = hann_window(FFT_SIZE)
    env = compute_onset_envelope(samples, FFT_SIZE, ONSET_HOP_SIZE, window)
    tempo, strength = estimate_tempo(
        env,
        SAMPLE_RATE,
        ONSET_HOP_SIZE,
        TEMPO_MIN_BPM,
        TEMPO_MAX_BPM,
    )
    assert tempo > 0.0
    assert 80.0 < tempo < 180.0
    assert 0.0 <= strength <= 1.0


def test_estimate_tempo_silence() -> None:
    """Verify tempo estimation returns 0 for silence."""
    env = np.zeros(100, dtype=np.float64)
    tempo, strength = estimate_tempo(
        env,
        SAMPLE_RATE,
        ONSET_HOP_SIZE,
        TEMPO_MIN_BPM,
        TEMPO_MAX_BPM,
    )
    assert tempo == 0.0
    assert strength == 0.0


def test_estimate_beat_period() -> None:
    """Verify beat period computation."""
    assert estimate_beat_period(120.0) == pytest.approx(0.5, abs=1e-6)
    assert estimate_beat_period(0.0) == 0.0


def test_compute_rhythm_density() -> None:
    """Verify rhythm density is non-negative."""
    samples = _click_track(120, duration=2.0)
    window = hann_window(FFT_SIZE)
    env = compute_onset_envelope(samples, FFT_SIZE, ONSET_HOP_SIZE, window)
    density = compute_rhythm_density(env, SAMPLE_RATE, ONSET_HOP_SIZE)
    assert density >= 0.0


def test_compute_rhythm_density_silence() -> None:
    """Verify rhythm density is 0 for silence."""
    env = np.zeros(100, dtype=np.float64)
    density = compute_rhythm_density(env, SAMPLE_RATE, ONSET_HOP_SIZE)
    assert density == 0.0


def test_detect_onsets() -> None:
    """Verify onset detection returns valid values."""
    samples = _click_track(120, duration=2.0)
    window = hann_window(FFT_SIZE)
    env = compute_onset_envelope(samples, FFT_SIZE, ONSET_HOP_SIZE, window)
    count, rate, regularity = detect_onsets(env, SAMPLE_RATE, ONSET_HOP_SIZE)
    assert count >= 0
    assert rate >= 0.0
    assert 0.0 <= regularity <= 1.0


def test_detect_onsets_silence() -> None:
    """Verify onset detection on silence."""
    env = np.zeros(100, dtype=np.float64)
    count, rate, regularity = detect_onsets(env, SAMPLE_RATE, ONSET_HOP_SIZE)
    assert count == 0
    assert rate == 0.0
    assert regularity == 0.0


def test_dsp_config_constants() -> None:
    """Verify DSP config constants are correct."""
    assert ONSET_HOP_SIZE == 512
    assert TEMPO_MIN_BPM == 40
    assert TEMPO_MAX_BPM == 220


# --- Silence test ---


def test_silence(tmp_path: Path) -> None:
    """Verify all-zero input produces safe defaults."""
    samples = np.zeros(SAMPLE_RATE * 2, dtype=np.float32)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    tempo = result.feature_set.find_by_name("rhythm.tempo")
    assert tempo is not None
    assert tempo.value == 0.0
    strength = result.feature_set.find_by_name("rhythm.beat_strength")
    assert strength is not None
    assert strength.value == 0.0
    regularity = result.feature_set.find_by_name("rhythm.regularity")
    assert regularity is not None
    assert regularity.value == 0.0


# --- Metronome click track test ---


def test_metronome_click_track(tmp_path: Path) -> None:
    """Verify click track at 120 BPM produces tempo near 120."""
    samples = _click_track(120, duration=3.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    tempo = result.feature_set.find_by_name("rhythm.tempo")
    assert tempo is not None
    assert tempo.value > 0.0
    assert 80.0 < tempo.value < 180.0


# --- Regular pulse test ---


def test_regular_pulse(tmp_path: Path) -> None:
    """Verify regular pulse has high regularity."""
    interval = SAMPLE_RATE // 4  # 4 pulses per second
    samples = _regular_pulse(interval, duration=3.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    regularity = result.feature_set.find_by_name("rhythm.regularity")
    assert regularity is not None
    assert regularity.value > 0.3


# --- Irregular pulse test ---


def test_irregular_pulse(tmp_path: Path) -> None:
    """Verify irregular pulse has lower regularity than regular."""
    intervals = [
        SAMPLE_RATE // 2,
        SAMPLE_RATE // 8,
        SAMPLE_RATE // 3,
        SAMPLE_RATE // 6,
        SAMPLE_RATE // 4,
        SAMPLE_RATE // 10,
    ]
    samples = _irregular_pulse(intervals)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    regularity = result.feature_set.find_by_name("rhythm.regularity")
    assert regularity is not None
    assert regularity.value >= 0.0
    assert regularity.value <= 1.0


# --- White noise test ---


def test_white_noise(tmp_path: Path) -> None:
    """Verify white noise produces deterministic output."""
    samples = _white_noise(duration=2.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    assert len(result.feature_set) == 7


# --- Sine wave test ---


def test_sine_wave(tmp_path: Path) -> None:
    """Verify sine wave produces deterministic output."""
    samples = _sine_wave(440.0, duration=2.0)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    assert len(result.feature_set) == 7


# --- Mono / Stereo tests ---


def test_mono(tmp_path: Path) -> None:
    """Verify mono audio produces 7 features."""
    samples = _click_track(120, duration=1.0).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 = _click_track(120, duration=1.0)
    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 = _click_track(120, duration=1.0)
    result_mono = _analyze(mono.reshape(-1, 1), tmp_path)
    result_stereo = _analyze(np.stack([mono, mono], axis=1), tmp_path)

    for name in [
        "rhythm.tempo",
        "rhythm.beat_period",
        "rhythm.density",
        "rhythm.beat_strength",
        "rhythm.onset_count",
        "rhythm.onset_rate",
        "rhythm.regularity",
    ]:
        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_tempo_estimation(tmp_path: Path) -> None:
    """Verify tempo is in valid range or 0."""
    samples = _click_track(120, duration=3.0)
    result = _analyze(samples, tmp_path)
    tempo = result.feature_set.find_by_name("rhythm.tempo")
    assert tempo is not None
    assert tempo.value == 0.0 or TEMPO_MIN_BPM <= tempo.value <= TEMPO_MAX_BPM


def test_onset_detection(tmp_path: Path) -> None:
    """Verify onset count is non-negative."""
    samples = _click_track(120, duration=2.0)
    result = _analyze(samples, tmp_path)
    count = result.feature_set.find_by_name("rhythm.onset_count")
    assert count is not None
    assert count.value >= 0


def test_beat_strength(tmp_path: Path) -> None:
    """Verify beat strength in [0, 1]."""
    samples = _click_track(120, duration=2.0)
    result = _analyze(samples, tmp_path)
    strength = result.feature_set.find_by_name("rhythm.beat_strength")
    assert strength is not None
    assert 0.0 <= strength.value <= 1.0


def test_onset_count(tmp_path: Path) -> None:
    """Verify onset count is non-negative integer."""
    samples = _click_track(120, duration=2.0)
    result = _analyze(samples, tmp_path)
    count = result.feature_set.find_by_name("rhythm.onset_count")
    assert count is not None
    assert count.value >= 0
    assert count.value == int(count.value)


def test_onset_rate(tmp_path: Path) -> None:
    """Verify onset rate >= 0."""
    samples = _click_track(120, duration=2.0)
    result = _analyze(samples, tmp_path)
    rate = result.feature_set.find_by_name("rhythm.onset_rate")
    assert rate is not None
    assert rate.value >= 0.0


def test_rhythm_density(tmp_path: Path) -> None:
    """Verify rhythm density >= 0."""
    samples = _click_track(120, duration=2.0)
    result = _analyze(samples, tmp_path)
    density = result.feature_set.find_by_name("rhythm.density")
    assert density is not None
    assert density.value >= 0.0


def test_rhythm_regularity(tmp_path: Path) -> None:
    """Verify rhythm regularity in [0, 1]."""
    samples = _click_track(120, duration=2.0)
    result = _analyze(samples, tmp_path)
    regularity = result.feature_set.find_by_name("rhythm.regularity")
    assert regularity is not None
    assert 0.0 <= regularity.value <= 1.0


# --- Determinism test ---


def test_deterministic(tmp_path: Path) -> None:
    """Verify same input produces identical output."""
    samples = _click_track(120, duration=2.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 = _click_track(120, duration=2.0)
    result = _analyze(samples, tmp_path)
    expected = {
        "rhythm.tempo",
        "rhythm.beat_period",
        "rhythm.density",
        "rhythm.beat_strength",
        "rhythm.onset_count",
        "rhythm.onset_rate",
        "rhythm.regularity",
    }
    actual = {f.name for f in result.feature_set}
    assert expected == actual


def test_feature_units(tmp_path: Path) -> None:
    """Verify correct units for tempo and beat period."""
    samples = _click_track(120, duration=2.0)
    result = _analyze(samples, tmp_path)

    tempo = result.feature_set.find_by_name("rhythm.tempo")
    assert tempo is not None
    assert tempo.unit == "BPM"

    beat_period = result.feature_set.find_by_name("rhythm.beat_period")
    assert beat_period is not None
    assert beat_period.unit == "seconds"


# --- Pipeline integration tests ---


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

    pipeline = AnalysisPipeline(registry)
    samples = _click_track(120, duration=1.0).reshape(-1, 1)
    context = _make_context(samples, tmp_path)
    results = pipeline.run(context)

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


def test_persistence_integration(tmp_path: Path) -> None:
    """Verify pipeline with session persists rhythm 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_rhythm_test",
            file_size=1000,
        )

        context = _make_context(
            _click_track(120, duration=1.0).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()

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

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

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