"""Unit tests for the feature definition registry.

Tests cover:
- FeatureDefinition immutability and validation
- FeatureRegistry registration, unregistration, lookup, listing
- Duplicate identifier rejection
- Category filtering
- Default registry completeness
- Cross-validation with analyzer output
"""

from __future__ import annotations

from dataclasses import FrozenInstanceError
from pathlib import Path

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 as build_analyzer_registry
from analyzer.dynamics.dynamic_analyzer import DynamicAnalyzer
from analyzer.feature import FeatureSet
from analyzer.harmony.harmony_analyzer import HarmonyAnalyzer
from analyzer.pipeline import AnalysisPipeline
from analyzer.rhythm.rhythm_analyzer import RhythmAnalyzer
from analyzer.spectral.spectral_analyzer import SpectralAnalyzer
from config.settings import Settings
from decoder.decoded_audio import DecodedAudio
from features.default_registry import build_default_feature_registry
from features.definition import (
    DataType,
    FeatureCategory,
    FeatureDefinition,
    NormalizationStrategy,
)
from features.registry import FeatureRegistry

# --- FeatureDefinition tests ---


def test_feature_definition_is_frozen() -> None:
    """Verify FeatureDefinition is immutable."""
    fd = FeatureDefinition(
        identifier="test.feature",
        display_name="Test Feature",
        category=FeatureCategory.SIGNAL,
        description="A test feature.",
        unit=None,
        data_type=DataType.FLOAT,
        normalization=NormalizationStrategy.IDENTITY,
        version="1.0.0",
    )
    with pytest.raises(FrozenInstanceError):
        fd.identifier = "other.feature"  # type: ignore[misc]


def test_feature_definition_stores_all_fields() -> None:
    """Verify all fields are stored correctly."""
    fd = FeatureDefinition(
        identifier="spectral.centroid",
        display_name="Spectral Centroid",
        category=FeatureCategory.SPECTRAL,
        description="Weighted mean frequency.",
        unit="Hz",
        data_type=DataType.FLOAT,
        normalization=NormalizationStrategy.MINMAX,
        version="2.1.0",
    )
    assert fd.identifier == "spectral.centroid"
    assert fd.display_name == "Spectral Centroid"
    assert fd.category == FeatureCategory.SPECTRAL
    assert fd.description == "Weighted mean frequency."
    assert fd.unit == "Hz"
    assert fd.data_type == DataType.FLOAT
    assert fd.normalization == NormalizationStrategy.MINMAX
    assert fd.version == "2.1.0"


def test_feature_definition_empty_identifier_raises() -> None:
    """Verify empty identifier is rejected."""
    with pytest.raises(ValueError, match="identifier"):
        FeatureDefinition(
            identifier="",
            display_name="Test",
            category=FeatureCategory.SIGNAL,
            description="Test.",
            unit=None,
            data_type=DataType.FLOAT,
            normalization=NormalizationStrategy.IDENTITY,
            version="1.0.0",
        )


def test_feature_definition_whitespace_identifier_raises() -> None:
    """Verify whitespace-only identifier is rejected."""
    with pytest.raises(ValueError, match="identifier"):
        FeatureDefinition(
            identifier="   ",
            display_name="Test",
            category=FeatureCategory.SIGNAL,
            description="Test.",
            unit=None,
            data_type=DataType.FLOAT,
            normalization=NormalizationStrategy.IDENTITY,
            version="1.0.0",
        )


def test_feature_definition_empty_display_name_raises() -> None:
    """Verify empty display_name is rejected."""
    with pytest.raises(ValueError, match="display_name"):
        FeatureDefinition(
            identifier="test.feature",
            display_name="",
            category=FeatureCategory.SIGNAL,
            description="Test.",
            unit=None,
            data_type=DataType.FLOAT,
            normalization=NormalizationStrategy.IDENTITY,
            version="1.0.0",
        )


# --- FeatureRegistry tests ---


def _make_definition(identifier: str = "test.feature") -> FeatureDefinition:
    """Create a minimal FeatureDefinition for testing."""
    return FeatureDefinition(
        identifier=identifier,
        display_name="Test Feature",
        category=FeatureCategory.SIGNAL,
        description="A test feature.",
        unit=None,
        data_type=DataType.FLOAT,
        normalization=NormalizationStrategy.IDENTITY,
        version="1.0.0",
    )


def test_registry_register_and_contains() -> None:
    """Verify registration and contains check."""
    registry = FeatureRegistry()
    fd = _make_definition("test.feature")
    registry.register(fd)
    assert registry.contains("test.feature")


def test_registry_contains_false_for_unregistered() -> None:
    """Verify contains returns False for unregistered identifier."""
    registry = FeatureRegistry()
    assert not registry.contains("nonexistent")


def test_registry_register_duplicate_raises() -> None:
    """Verify duplicate registration raises ValueError."""
    registry = FeatureRegistry()
    fd = _make_definition("test.feature")
    registry.register(fd)
    with pytest.raises(ValueError, match="already registered"):
        registry.register(_make_definition("test.feature"))


def test_registry_unregister() -> None:
    """Verify unregistration removes the definition."""
    registry = FeatureRegistry()
    registry.register(_make_definition("test.feature"))
    registry.unregister("test.feature")
    assert not registry.contains("test.feature")


def test_registry_unregister_missing_raises_keyerror() -> None:
    """Verify unregistering a missing identifier raises KeyError."""
    registry = FeatureRegistry()
    with pytest.raises(KeyError, match="not registered"):
        registry.unregister("nonexistent")


def test_registry_get() -> None:
    """Verify get returns the correct definition."""
    registry = FeatureRegistry()
    fd = _make_definition("test.feature")
    registry.register(fd)
    retrieved = registry.get("test.feature")
    assert retrieved is fd


def test_registry_get_missing_raises_keyerror() -> None:
    """Verify get raises KeyError for missing identifier."""
    registry = FeatureRegistry()
    with pytest.raises(KeyError, match="not registered"):
        registry.get("nonexistent")


def test_registry_list_returns_registration_order() -> None:
    """Verify list returns definitions in registration order."""
    registry = FeatureRegistry()
    fd1 = _make_definition("alpha.feature")
    fd2 = _make_definition("beta.feature")
    fd3 = _make_definition("gamma.feature")
    registry.register(fd1)
    registry.register(fd2)
    registry.register(fd3)
    defs = registry.list()
    assert len(defs) == 3
    assert defs[0] is fd1
    assert defs[1] is fd2
    assert defs[2] is fd3


def test_registry_list_empty() -> None:
    """Verify list on empty registry returns empty list."""
    registry = FeatureRegistry()
    assert registry.list() == []


def test_registry_count() -> None:
    """Verify count returns the correct number."""
    registry = FeatureRegistry()
    assert registry.count() == 0
    registry.register(_make_definition("a.feature"))
    registry.register(_make_definition("b.feature"))
    assert registry.count() == 2


def test_registry_list_by_category() -> None:
    """Verify list_by_category filters correctly."""
    registry = FeatureRegistry()
    signal_fd = FeatureDefinition(
        identifier="signal.test",
        display_name="Signal Test",
        category=FeatureCategory.SIGNAL,
        description="Signal feature.",
        unit=None,
        data_type=DataType.FLOAT,
        normalization=NormalizationStrategy.IDENTITY,
        version="1.0.0",
    )
    spectral_fd = FeatureDefinition(
        identifier="spectral.test",
        display_name="Spectral Test",
        category=FeatureCategory.SPECTRAL,
        description="Spectral feature.",
        unit="Hz",
        data_type=DataType.FLOAT,
        normalization=NormalizationStrategy.MINMAX,
        version="1.0.0",
    )
    registry.register(signal_fd)
    registry.register(spectral_fd)

    signal_defs = registry.list_by_category(FeatureCategory.SIGNAL)
    assert len(signal_defs) == 1
    assert signal_defs[0] is signal_fd

    spectral_defs = registry.list_by_category(FeatureCategory.SPECTRAL)
    assert len(spectral_defs) == 1
    assert spectral_defs[0] is spectral_fd

    rhythm_defs = registry.list_by_category(FeatureCategory.RHYTHM)
    assert len(rhythm_defs) == 0


# --- Default registry tests ---


def test_default_registry_count() -> None:
    """Verify default registry has 55 definitions."""
    registry = build_default_feature_registry()
    assert registry.count() == 55


def test_default_registry_category_counts() -> None:
    """Verify each category has the expected number of definitions."""
    registry = build_default_feature_registry()
    assert len(registry.list_by_category(FeatureCategory.SIGNAL)) == 13
    assert len(registry.list_by_category(FeatureCategory.SPECTRAL)) == 6
    assert len(registry.list_by_category(FeatureCategory.DYNAMIC)) == 7
    assert len(registry.list_by_category(FeatureCategory.RHYTHM)) == 7
    assert len(registry.list_by_category(FeatureCategory.HARMONY)) == 22


def test_default_registry_all_identifiers_unique() -> None:
    """Verify all identifiers in the default registry are unique."""
    registry = build_default_feature_registry()
    identifiers = [d.identifier for d in registry.list()]
    assert len(identifiers) == len(set(identifiers))


def test_default_registry_identifiers_use_block_feature_format() -> None:
    """Verify all identifiers use block.feature format (contain a dot)."""
    registry = build_default_feature_registry()
    for d in registry.list():
        assert "." in d.identifier, f"Identifier '{d.identifier}' missing dot separator"


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


# --- Cross-validation with analyzer output ---


def _make_decoded_audio() -> DecodedAudio:
    """Create a DecodedAudio with a simple sine wave for testing."""
    sr = 44100
    t = np.linspace(0, 2.0, int(sr * 2.0), dtype=np.float32)
    samples = (0.5 * np.sin(2.0 * np.pi * 440.0 * t)).reshape(-1, 1)
    return DecodedAudio(
        samples=samples, sample_rate=sr, channels=1, duration=2.0, bit_depth=None
    )


def _make_context(tmp_path: Path) -> AnalysisContext:
    """Create an AnalysisContext for cross-validation testing."""
    from unittest.mock import MagicMock

    track = MagicMock()
    track.id = "test-track-id"
    track.relative_path = "test.wav"
    settings = Settings(
        project_name="Test",
        database_url=f"sqlite:///{tmp_path / 'test.db'}",
        log_level="DEBUG",
    )
    return AnalysisContext(
        track=track,
        decoded_audio=_make_decoded_audio(),
        settings=settings,
    )


def _get_all_feature_names(feature_set: FeatureSet) -> set[str]:
    """Extract all feature names from a FeatureSet."""
    return {f.name for f in feature_set}


@pytest.mark.parametrize(
    "analyzer_class,expected_count",
    [
        (BasicSignalAnalyzer, 9),
        (SpectralAnalyzer, 6),
        (DynamicAnalyzer, 7),
        (RhythmAnalyzer, 7),
        (HarmonyAnalyzer, 22),
    ],
)
def test_analyzer_features_exist_in_default_registry(
    analyzer_class: type, expected_count: int, tmp_path: Path
) -> None:
    """Verify every feature produced by each analyzer exists in the default registry."""
    registry = build_default_feature_registry()
    analyzer = analyzer_class()
    context = _make_context(tmp_path)
    result = analyzer.analyze(context)
    names = _get_all_feature_names(result.feature_set)
    assert len(names) == expected_count
    for name in names:
        assert registry.contains(name), (
            f"Feature '{name}' from {analyzer_class.__name__} "
            "not found in default registry"
        )


def test_pipeline_features_exist_in_default_registry(tmp_path: Path) -> None:
    """Verify all features from a full pipeline run exist in the registry."""
    feature_registry = build_default_feature_registry()
    analyzer_registry = build_analyzer_registry()
    pipeline = AnalysisPipeline(analyzer_registry)
    context = _make_context(tmp_path)
    results = pipeline.run(context)

    all_names: set[str] = set()
    for result in results:
        all_names |= _get_all_feature_names(result.feature_set)

    for name in all_names:
        assert feature_registry.contains(
            name
        ), f"Pipeline feature '{name}' not found in default registry"
