"""Tests for the BasicSignalAnalyzer and default registry."""

from __future__ import annotations

from dataclasses import FrozenInstanceError
from pathlib import Path
from unittest.mock import MagicMock

import numpy as np
import pytest

from analyzer.basic.basic_signal_analyzer import BasicSignalAnalyzer
from analyzer.context import AnalysisContext
from analyzer.default_registry import build_default_registry
from analyzer.pipeline import AnalysisPipeline
from analyzer.registry import AnalyzerRegistry
from analyzer.result import AnalysisResult
from config.settings import Settings
from core.exceptions import AnalysisError
from decoder.decoded_audio import DecodedAudio


def _make_decoded_audio(
    samples: np.ndarray,
    sample_rate: int = 44100,
    duration: float | None = None,
    bit_depth: int | None = 16,
) -> DecodedAudio:
    """Create a DecodedAudio from a NumPy array.

    Args:
        samples: PCM samples with shape (n_frames, n_channels).
        sample_rate: Sample rate in Hz.
        duration: Duration in seconds (auto-calculated if None).
        bit_depth: Bit depth.

    Returns:
        A :class:`DecodedAudio` instance.
    """
    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.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 BasicSignalAnalyzer on the given samples."""
    context = _make_context(samples, tmp_path)
    analyzer = BasicSignalAnalyzer()
    return analyzer.analyze(context)


# --- Peak calculation tests ---


def test_peak_mono(tmp_path: Path) -> None:
    """Verify peak amplitude on mono signal."""
    samples = np.array([[0.5], [0.8], [-0.3], [0.1]], dtype=np.float32)
    result = _analyze(samples, tmp_path)
    peak = result.feature_set.find_by_name("signal.peak")
    assert peak is not None
    assert peak.value == pytest.approx(0.8, abs=1e-6)


def test_peak_stereo(tmp_path: Path) -> None:
    """Verify peak amplitude per channel on stereo."""
    samples = np.array(
        [[0.5, 0.2], [0.8, 0.9], [-0.3, -0.1], [0.1, 0.4]],
        dtype=np.float32,
    )
    result = _analyze(samples, tmp_path)
    peak_left = result.feature_set.find_by_name("signal.peak_left")
    peak_right = result.feature_set.find_by_name("signal.peak_right")
    assert peak_left is not None
    assert peak_right is not None
    assert peak_left.value == pytest.approx(0.8, abs=1e-6)
    assert peak_right.value == pytest.approx(0.9, abs=1e-6)


# --- RMS calculation tests ---


def test_rms_mono(tmp_path: Path) -> None:
    """Verify RMS on mono signal."""
    samples = np.array([[0.5], [-0.5], [0.5], [-0.5]], dtype=np.float32)
    result = _analyze(samples, tmp_path)
    rms = result.feature_set.find_by_name("signal.rms")
    assert rms is not None
    assert rms.value == pytest.approx(0.5, abs=1e-6)


def test_rms_stereo(tmp_path: Path) -> None:
    """Verify RMS per channel on stereo."""
    samples = np.array(
        [[1.0, 0.0], [0.0, 1.0], [1.0, 0.0], [0.0, 1.0]],
        dtype=np.float32,
    )
    result = _analyze(samples, tmp_path)
    rms_left = result.feature_set.find_by_name("signal.rms_left")
    rms_right = result.feature_set.find_by_name("signal.rms_right")
    assert rms_left is not None
    assert rms_right is not None
    expected = float(np.sqrt(0.5))
    assert rms_left.value == pytest.approx(expected, abs=1e-6)
    assert rms_right.value == pytest.approx(expected, abs=1e-6)


# --- Peak dB tests ---


def test_peak_db(tmp_path: Path) -> None:
    """Verify peak dBFS conversion."""
    samples = np.array([[0.5]], dtype=np.float32)
    result = _analyze(samples, tmp_path)
    peak_db = result.feature_set.find_by_name("signal.peak_db")
    assert peak_db is not None
    expected = 20.0 * np.log10(0.5)
    assert peak_db.value == pytest.approx(expected, abs=1e-4)


def test_peak_db_silence(tmp_path: Path) -> None:
    """Verify peak dBFS for zero signal is -inf."""
    samples = np.zeros((100, 1), dtype=np.float32)
    result = _analyze(samples, tmp_path)
    peak_db = result.feature_set.find_by_name("signal.peak_db")
    assert peak_db is not None
    assert peak_db.value == float("-inf")


# --- RMS dB tests ---


def test_rms_db(tmp_path: Path) -> None:
    """Verify RMS dBFS conversion."""
    samples = np.full((100, 1), 0.5, dtype=np.float32)
    result = _analyze(samples, tmp_path)
    rms_db = result.feature_set.find_by_name("signal.rms_db")
    assert rms_db is not None
    expected = 20.0 * np.log10(0.5)
    assert rms_db.value == pytest.approx(expected, abs=1e-4)


def test_rms_db_silence(tmp_path: Path) -> None:
    """Verify RMS dBFS for zero signal is -inf."""
    samples = np.zeros((100, 1), dtype=np.float32)
    result = _analyze(samples, tmp_path)
    rms_db = result.feature_set.find_by_name("signal.rms_db")
    assert rms_db is not None
    assert rms_db.value == float("-inf")


# --- DC Offset tests ---


def test_dc_offset(tmp_path: Path) -> None:
    """Verify DC offset is the mean of samples."""
    samples = np.array([[0.1], [0.2], [0.3], [0.4]], dtype=np.float32)
    result = _analyze(samples, tmp_path)
    dc = result.feature_set.find_by_name("signal.dc_offset")
    assert dc is not None
    assert dc.value == pytest.approx(0.25, abs=1e-6)


def test_dc_offset_zero(tmp_path: Path) -> None:
    """Verify DC offset is zero for symmetric signal."""
    samples = np.array([[0.5], [-0.5], [0.5], [-0.5]], dtype=np.float32)
    result = _analyze(samples, tmp_path)
    dc = result.feature_set.find_by_name("signal.dc_offset")
    assert dc is not None
    assert dc.value == pytest.approx(0.0, abs=1e-6)


# --- Silence Ratio tests ---


def test_silence_ratio(tmp_path: Path) -> None:
    """Verify silence ratio for all-silent signal."""
    samples = np.zeros((100, 1), dtype=np.float32)
    result = _analyze(samples, tmp_path)
    sr = result.feature_set.find_by_name("signal.silence_ratio")
    assert sr is not None
    assert sr.value == pytest.approx(1.0, abs=1e-6)


def test_silence_ratio_partial(tmp_path: Path) -> None:
    """Verify silence ratio for mixed signal."""
    samples = np.zeros((100, 1), dtype=np.float32)
    samples[50:] = 0.5
    result = _analyze(samples, tmp_path)
    sr = result.feature_set.find_by_name("signal.silence_ratio")
    assert sr is not None
    assert sr.value == pytest.approx(0.5, abs=1e-6)


def test_silence_ratio_none(tmp_path: Path) -> None:
    """Verify silence ratio is 0 for non-silent signal."""
    samples = np.full((100, 1), 0.5, dtype=np.float32)
    result = _analyze(samples, tmp_path)
    sr = result.feature_set.find_by_name("signal.silence_ratio")
    assert sr is not None
    assert sr.value == pytest.approx(0.0, abs=1e-6)


# --- Mono / Stereo feature name tests ---


def test_mono_features(tmp_path: Path) -> None:
    """Verify mono feature names are present."""
    samples = np.random.default_rng(42).standard_normal((100, 1)).astype(np.float32)
    result = _analyze(samples, tmp_path)
    assert result.feature_set.find_by_name("signal.peak") is not None
    assert result.feature_set.find_by_name("signal.rms") is not None
    assert result.feature_set.find_by_name("signal.peak_left") is None
    assert result.feature_set.find_by_name("signal.rms_left") is None


def test_stereo_features(tmp_path: Path) -> None:
    """Verify stereo feature names are present."""
    samples = np.random.default_rng(42).standard_normal((100, 2)).astype(np.float32)
    result = _analyze(samples, tmp_path)
    assert result.feature_set.find_by_name("signal.peak_left") is not None
    assert result.feature_set.find_by_name("signal.peak_right") is not None
    assert result.feature_set.find_by_name("signal.rms_left") is not None
    assert result.feature_set.find_by_name("signal.rms_right") is not None
    assert result.feature_set.find_by_name("signal.peak") is None
    assert result.feature_set.find_by_name("signal.rms") is None


# --- Special signal tests ---


def test_silence_input(tmp_path: Path) -> None:
    """Verify all-zero input produces safe values."""
    samples = np.zeros((100, 2), dtype=np.float32)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    assert result.feature_set.find_by_name("signal.peak_db").value == float("-inf")
    assert result.feature_set.find_by_name("signal.rms_db").value == float("-inf")
    assert result.feature_set.find_by_name(
        "signal.silence_ratio"
    ).value == pytest.approx(1.0)


def test_impulse_signal(tmp_path: Path) -> None:
    """Verify single-sample impulse."""
    samples = np.zeros((100, 1), dtype=np.float32)
    samples[50, 0] = 1.0
    result = _analyze(samples, tmp_path)
    peak = result.feature_set.find_by_name("signal.peak")
    assert peak.value == pytest.approx(1.0)
    peak_db = result.feature_set.find_by_name("signal.peak_db")
    assert peak_db.value == pytest.approx(0.0, abs=1e-4)


def test_constant_signal(tmp_path: Path) -> None:
    """Verify constant amplitude signal."""
    samples = np.full((100, 1), 0.5, dtype=np.float32)
    result = _analyze(samples, tmp_path)
    peak = result.feature_set.find_by_name("signal.peak")
    assert peak.value == pytest.approx(0.5)
    rms = result.feature_set.find_by_name("signal.rms")
    assert rms.value == pytest.approx(0.5)
    dc = result.feature_set.find_by_name("signal.dc_offset")
    assert dc.value == pytest.approx(0.5)


# --- Error handling tests ---


def test_empty_audio_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)


def test_nan_raises(tmp_path: Path) -> None:
    """Verify NaN values raise AnalysisError."""
    samples = np.array([[0.5], [np.nan], [0.3]], dtype=np.float32)
    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.array([[0.5], [np.inf], [0.3]], dtype=np.float32)
    with pytest.raises(AnalysisError, match="infinite"):
        _analyze(samples, tmp_path)


# --- Feature name and immutability tests ---


def test_expected_feature_names_mono(tmp_path: Path) -> None:
    """Verify all expected mono feature names are present."""
    samples = np.random.default_rng(42).standard_normal((100, 1)).astype(np.float32)
    result = _analyze(samples, tmp_path)
    expected = {
        "signal.duration_seconds",
        "signal.sample_rate",
        "signal.channels",
        "signal.peak",
        "signal.rms",
        "signal.peak_db",
        "signal.rms_db",
        "signal.dc_offset",
        "signal.silence_ratio",
    }
    actual = {f.name for f in result.feature_set}
    assert expected == actual


def test_expected_feature_names_stereo(tmp_path: Path) -> None:
    """Verify all expected stereo feature names are present."""
    samples = np.random.default_rng(42).standard_normal((100, 2)).astype(np.float32)
    result = _analyze(samples, tmp_path)
    expected = {
        "signal.duration_seconds",
        "signal.sample_rate",
        "signal.channels",
        "signal.peak_left",
        "signal.peak_right",
        "signal.peak_db",
        "signal.rms_left",
        "signal.rms_right",
        "signal.rms_db",
        "signal.dc_offset",
        "signal.silence_ratio",
    }
    actual = {f.name for f in result.feature_set}
    assert expected == actual


def test_feature_set_immutable(tmp_path: Path) -> None:
    """Verify result FeatureSet is frozen."""
    samples = np.zeros((10, 1), dtype=np.float32)
    result = _analyze(samples, tmp_path)
    with pytest.raises(FrozenInstanceError):
        result.feature_set.features = ()  # type: ignore[misc]


def test_result_success(tmp_path: Path) -> None:
    """Verify result has success=True."""
    samples = np.zeros((10, 1), dtype=np.float32)
    result = _analyze(samples, tmp_path)
    assert result.success is True
    assert result.analyzer_name == "basic_signal"
    assert result.analyzer_version == "1.0.0"


# --- Default registry tests ---


def test_default_registry() -> None:
    """Verify build_default_registry returns registry with expected analyzers."""
    registry = build_default_registry()
    analyzers = registry.list_analyzers()
    assert len(analyzers) == 5
    names = [a.name for a in analyzers]
    assert "basic_signal" in names
    assert "spectral" in names
    assert "dynamic" in names
    assert "rhythm" in names
    assert "harmony" in names
    basic = [a for a in analyzers if a.name == "basic_signal"][0]
    assert isinstance(basic, BasicSignalAnalyzer)


def test_default_registry_no_dummy() -> None:
    """Verify default registry does not contain DummyAnalyzer."""
    registry = build_default_registry()
    with pytest.raises(KeyError):
        registry.get("dummy")


def test_default_registry_returns_new_instance() -> None:
    """Verify each call returns a fresh registry."""
    r1 = build_default_registry()
    r2 = build_default_registry()
    assert r1 is not r2


# --- Pipeline execution tests ---


def test_pipeline_execution(tmp_path: Path) -> None:
    """Verify pipeline runs all default analyzers and produces results."""
    registry = build_default_registry()
    pipeline = AnalysisPipeline(registry)
    samples = np.random.default_rng(42).standard_normal((100, 2)).astype(np.float32)
    context = _make_context(samples, tmp_path)
    results = pipeline.run(context)

    assert len(results) == 5
    basic_results = [r for r in results if r.analyzer_name == "basic_signal"]
    assert len(basic_results) == 1
    assert basic_results[0].success is True
    assert len(basic_results[0].feature_set) == 11  # stereo


def test_pipeline_failure_isolation(tmp_path: Path) -> None:
    """Verify pipeline continues if analyzer raises AnalysisError."""

    class FailAnalyzer(BasicSignalAnalyzer):
        name = "fail_basic"
        version = "0.0.1"

        def analyze(self, context: AnalysisContext) -> AnalysisResult:
            raise AnalysisError("intentional")

    registry = AnalyzerRegistry()
    registry.register(FailAnalyzer())
    registry.register(BasicSignalAnalyzer())
    pipeline = AnalysisPipeline(registry)

    samples = np.zeros((10, 1), dtype=np.float32)
    context = _make_context(samples, tmp_path)
    results = pipeline.run(context)

    assert len(results) == 2
    assert results[0].success is False
    assert results[1].success is True


# --- Determinism test ---


def test_deterministic(tmp_path: Path) -> None:
    """Verify same input produces identical output."""
    samples = np.random.default_rng(42).standard_normal((100, 2)).astype(np.float32)
    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)
