"""Tests for feature storage and analysis persistence."""

from __future__ import annotations

from pathlib import Path
from unittest.mock import MagicMock

import numpy as np
import pytest
from sqlalchemy.exc import IntegrityError
from sqlalchemy.orm import Session

from analyzer.analyzer import Analyzer
from analyzer.context import AnalysisContext
from analyzer.feature import Feature, FeatureSet
from analyzer.pipeline import AnalysisPipeline
from analyzer.registry import AnalyzerRegistry
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 storage.analysis_repository import AnalysisRepository
from storage.feature_repository import FeatureRepository


def _create_track(session: Session, name: str = "song.wav") -> str:
    """Create a track in the database and return its ID."""
    project = ProjectRepository.get_default_project(session)
    track = TrackRepository.create_track(
        session,
        project_id=project.id,
        relative_path=name,
        original_filename=name,
        sha256=f"hash_{name}",
        file_size=1000,
    )
    return track.id


def _make_feature_set(
    analyzer_name: str = "test",
    version: str = "1.0.0",
) -> FeatureSet:
    """Create a FeatureSet with test features."""
    return FeatureSet(
        [
            Feature(
                name="signal.peak", value=0.8, analyzer=analyzer_name, version=version
            ),
            Feature(
                name="signal.rms", value=0.5, analyzer=analyzer_name, version=version
            ),
            Feature(
                name="signal.peak_db",
                value=1.94,
                unit="dBFS",
                analyzer=analyzer_name,
                version=version,
            ),
        ]
    )


def _make_context(track_id: str, tmp_path: Path) -> AnalysisContext:
    """Create an AnalysisContext for testing."""
    track = MagicMock()
    track.id = track_id
    track.relative_path = "test.wav"
    audio = DecodedAudio(
        samples=np.zeros((100, 1), dtype=np.float32),
        sample_rate=44100,
        channels=1,
        duration=0.00227,
        bit_depth=16,
    )
    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)


class _TestAnalyzer(Analyzer):
    """Simple test analyzer producing a fixed feature set."""

    name = "test"
    version = "1.0.0"

    def __init__(self, feature_set: FeatureSet | None = None) -> None:
        self._feature_set = feature_set or _make_feature_set()

    def analyze(self, context: AnalysisContext) -> AnalysisResult:
        return AnalysisResult(
            analyzer_name=self.name,
            analyzer_version=self.version,
            execution_time_ms=1.0,
            success=True,
            warnings=(),
            feature_set=self._feature_set,
        )


class _FailingAnalyzer(Analyzer):
    """Analyzer that always fails."""

    name = "failing"
    version = "1.0.0"

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


class _VersionedAnalyzer(Analyzer):
    """Analyzer with configurable version."""

    def __init__(self, version: str) -> None:
        self.version = version

    name = "test"

    def analyze(self, context: AnalysisContext) -> AnalysisResult:
        return AnalysisResult(
            analyzer_name=self.name,
            analyzer_version=self.version,
            execution_time_ms=1.0,
            success=True,
            warnings=(),
            feature_set=_make_feature_set(version=self.version),
        )


# --- AnalyzerRun tests ---


def test_analyzer_run_creation(tmp_path: Path) -> None:
    """Verify AnalyzerRun is created with correct fields."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        run = AnalysisRepository.create_run(session, track_id, "basic_signal", "1.0.0")
        assert run.id is not None
        assert run.track_id == track_id
        assert run.analyzer_name == "basic_signal"
        assert run.analyzer_version == "1.0.0"
        assert run.started_at is not None
        assert run.finished_at is None
        assert run.success is False
    finally:
        session.close()


def test_analyzer_run_completion(tmp_path: Path) -> None:
    """Verify finish_run sets completion fields."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        run = AnalysisRepository.create_run(session, track_id, "basic_signal", "1.0.0")
        AnalysisRepository.finish_run(
            session, run, execution_time_ms=42.5, success=True
        )
        assert run.finished_at is not None
        assert run.execution_time_ms == 42.5
        assert run.success is True
        assert run.warnings is None
    finally:
        session.close()


def test_analyzer_run_completion_with_warnings(tmp_path: Path) -> None:
    """Verify finish_run stores warnings."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        run = AnalysisRepository.create_run(session, track_id, "basic_signal", "1.0.0")
        AnalysisRepository.finish_run(
            session,
            run,
            execution_time_ms=10.0,
            success=True,
            warnings=("warn1", "warn2"),
        )
        assert run.warnings == "warn1\nwarn2"
    finally:
        session.close()


def test_analyzer_run_failure(tmp_path: Path) -> None:
    """Verify mark_failed sets failure fields."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        run = AnalysisRepository.create_run(session, track_id, "basic_signal", "1.0.0")
        AnalysisRepository.mark_failed(session, run, 5.0, "something went wrong")
        assert run.finished_at is not None
        assert run.success is False
        assert run.warnings == "something went wrong"
    finally:
        session.close()


# --- Feature persistence tests ---


def test_feature_persistence(tmp_path: Path) -> None:
    """Verify FeatureSet is persisted as TrackFeature rows."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        run = AnalysisRepository.create_run(session, track_id, "test", "1.0.0")
        fs = _make_feature_set()
        count = FeatureRepository.save_feature_set(session, track_id, run.id, fs)
        assert count == 3

        features = FeatureRepository.get_track_features(session, track_id)
        assert len(features) == 3
        names = {f.name for f in features}
        assert names == {"signal.peak", "signal.rms", "signal.peak_db"}
    finally:
        session.close()


def test_feature_lookup(tmp_path: Path) -> None:
    """Verify get_track_features returns saved features."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        run = AnalysisRepository.create_run(session, track_id, "test", "1.0.0")
        FeatureRepository.save_feature_set(
            session, track_id, run.id, _make_feature_set()
        )
        features = FeatureRepository.get_track_features(session, track_id)
        assert len(features) == 3
    finally:
        session.close()


def test_latest_feature_lookup(tmp_path: Path) -> None:
    """Verify get_latest_features returns most recent successful run."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)

        # First run (v1.0.0)
        run1 = AnalysisRepository.create_run(session, track_id, "test", "1.0.0")
        AnalysisRepository.finish_run(session, run1, 10.0, success=True)
        FeatureRepository.save_feature_set(
            session, track_id, run1.id, _make_feature_set()
        )

        import time as _time

        _time.sleep(0.05)

        # Second run (v1.1.0)
        run2 = AnalysisRepository.create_run(session, track_id, "test", "1.1.0")
        AnalysisRepository.finish_run(session, run2, 5.0, success=True)
        FeatureRepository.save_feature_set(
            session,
            track_id,
            run2.id,
            _make_feature_set(version="1.1.0"),
        )

        latest = FeatureRepository.get_latest_features(session, track_id)
        assert len(latest) == 3
        for f in latest:
            assert f.analyzer_run.analyzer_version == "1.1.0"
    finally:
        session.close()


def test_history_lookup(tmp_path: Path) -> None:
    """Verify list_runs returns all runs for a track."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        AnalysisRepository.create_run(session, track_id, "test", "1.0.0")
        AnalysisRepository.create_run(session, track_id, "test", "1.1.0")

        runs = AnalysisRepository.list_runs(session, track_id)
        assert len(runs) == 2
    finally:
        session.close()


def test_version_lookup(tmp_path: Path) -> None:
    """Verify find_run by track, name, version."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        AnalysisRepository.create_run(session, track_id, "test", "1.0.0")

        run = AnalysisRepository.find_run(session, track_id, "test", "1.0.0")
        assert run is not None
        assert run.analyzer_name == "test"

        missing = AnalysisRepository.find_run(session, track_id, "test", "2.0.0")
        assert missing is None
    finally:
        session.close()


def test_has_analyzer_version(tmp_path: Path) -> None:
    """Verify has_analyzer_version returns correct boolean."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        assert not FeatureRepository.has_analyzer_version(
            session, track_id, "test", "1.0.0"
        )

        run = AnalysisRepository.create_run(session, track_id, "test", "1.0.0")
        AnalysisRepository.finish_run(session, run, 10.0, success=True)
        FeatureRepository.save_feature_set(
            session, track_id, run.id, _make_feature_set()
        )

        assert FeatureRepository.has_analyzer_version(
            session, track_id, "test", "1.0.0"
        )
    finally:
        session.close()


def test_has_analyzer_version_failed_run(tmp_path: Path) -> None:
    """Verify has_analyzer_version is False for failed runs."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        run = AnalysisRepository.create_run(session, track_id, "test", "1.0.0")
        AnalysisRepository.mark_failed(session, run, 5.0, "error")

        assert not FeatureRepository.has_analyzer_version(
            session, track_id, "test", "1.0.0"
        )
    finally:
        session.close()


# --- Pipeline persistence tests ---


def test_pipeline_persistence(tmp_path: Path) -> None:
    """Verify pipeline with session saves results."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        context = _make_context(track_id, tmp_path)

        registry = AnalyzerRegistry()
        registry.register(_TestAnalyzer())
        pipeline = AnalysisPipeline(registry)
        results = pipeline.run(context, session=session)

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

        runs = AnalysisRepository.list_runs(session, track_id)
        assert len(runs) == 1
        assert runs[0].success is True

        features = FeatureRepository.get_track_features(session, track_id)
        assert len(features) == 3
    finally:
        session.close()


def test_pipeline_no_session(tmp_path: Path) -> None:
    """Verify pipeline without session works (backward compat)."""
    context = _make_context("fake-id", tmp_path)

    registry = AnalyzerRegistry()
    registry.register(_TestAnalyzer())
    pipeline = AnalysisPipeline(registry)
    results = pipeline.run(context)

    assert len(results) == 1
    assert results[0].success is True
    assert len(results[0].feature_set) == 3


def test_idempotent_execution(tmp_path: Path) -> None:
    """Verify pipeline skips when analyzer version already exists."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        context = _make_context(track_id, tmp_path)

        registry = AnalyzerRegistry()
        registry.register(_TestAnalyzer())
        pipeline = AnalysisPipeline(registry)

        # First run — should execute and persist
        results1 = pipeline.run(context, session=session)
        assert len(results1) == 1
        assert results1[0].success is True
        assert len(results1[0].feature_set) == 3

        # Second run — should skip (idempotent)
        results2 = pipeline.run(context, session=session)
        assert len(results2) == 1
        assert results2[0].success is True
        assert len(results2[0].feature_set) == 0  # skipped, empty feature set

        # Only one AnalyzerRun should exist
        runs = AnalysisRepository.list_runs(session, track_id)
        assert len(runs) == 1
    finally:
        session.close()


def test_pipeline_failure_persistence(tmp_path: Path) -> None:
    """Verify failed analyzer creates failed AnalyzerRun."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        context = _make_context(track_id, tmp_path)

        registry = AnalyzerRegistry()
        registry.register(_FailingAnalyzer())
        pipeline = AnalysisPipeline(registry)
        results = pipeline.run(context, session=session)

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

        runs = AnalysisRepository.list_runs(session, track_id)
        assert len(runs) == 1
        assert runs[0].success is False
        assert runs[0].warnings is not None

        # No features should be saved for a failed run
        features = FeatureRepository.get_track_features(session, track_id)
        assert len(features) == 0
    finally:
        session.close()


# --- Multiple version and track tests ---


def test_multiple_analyzer_versions(tmp_path: Path) -> None:
    """Verify two versions of same analyzer coexist."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        context = _make_context(track_id, tmp_path)

        # Run v1.0.0
        registry1 = AnalyzerRegistry()
        registry1.register(_VersionedAnalyzer("1.0.0"))
        pipeline1 = AnalysisPipeline(registry1)
        pipeline1.run(context, session=session)

        # Run v1.1.0
        registry2 = AnalyzerRegistry()
        registry2.register(_VersionedAnalyzer("1.1.0"))
        pipeline2 = AnalysisPipeline(registry2)
        pipeline2.run(context, session=session)

        runs = AnalysisRepository.list_runs(session, track_id)
        assert len(runs) == 2
        versions = {r.analyzer_version for r in runs}
        assert versions == {"1.0.0", "1.1.0"}
    finally:
        session.close()


def test_multiple_tracks(tmp_path: Path) -> None:
    """Verify features for multiple tracks don't cross-contaminate."""
    session = init_test_db(tmp_path)
    try:
        track_id_1 = _create_track(session, "song1.wav")
        track_id_2 = _create_track(session, "song2.wav")

        context1 = _make_context(track_id_1, tmp_path)
        context2 = _make_context(track_id_2, tmp_path)

        registry = AnalyzerRegistry()
        registry.register(_TestAnalyzer())
        pipeline = AnalysisPipeline(registry)

        pipeline.run(context1, session=session)
        pipeline.run(context2, session=session)

        features1 = FeatureRepository.get_track_features(session, track_id_1)
        features2 = FeatureRepository.get_track_features(session, track_id_2)
        assert len(features1) == 3
        assert len(features2) == 3

        runs1 = AnalysisRepository.list_runs(session, track_id_1)
        runs2 = AnalysisRepository.list_runs(session, track_id_2)
        assert len(runs1) == 1
        assert len(runs2) == 1
        assert runs1[0].track_id == track_id_1
        assert runs2[0].track_id == track_id_2
    finally:
        session.close()


# --- Database constraint tests ---


def test_unique_constraint(tmp_path: Path) -> None:
    """Verify duplicate (track_id, analyzer_name, version) raises."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)
        AnalysisRepository.create_run(session, track_id, "test", "1.0.0")

        with pytest.raises(IntegrityError):
            AnalysisRepository.create_run(session, track_id, "test", "1.0.0")
        session.rollback()
    finally:
        session.close()


def test_feature_never_overwritten(tmp_path: Path) -> None:
    """Verify re-saving creates new rows, doesn't update existing."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)

        # First save
        run1 = AnalysisRepository.create_run(session, track_id, "test", "1.0.0")
        AnalysisRepository.finish_run(session, run1, 10.0, success=True)
        FeatureRepository.save_feature_set(
            session, track_id, run1.id, _make_feature_set()
        )

        # Second save with different version
        run2 = AnalysisRepository.create_run(session, track_id, "test", "1.1.0")
        AnalysisRepository.finish_run(session, run2, 5.0, success=True)
        FeatureRepository.save_feature_set(
            session, track_id, run2.id, _make_feature_set(version="1.1.0")
        )

        # Both sets of features should exist (history preserved)
        all_features = FeatureRepository.get_track_features(session, track_id)
        assert len(all_features) == 6  # 3 from each run
    finally:
        session.close()


def test_get_feature_by_name(tmp_path: Path) -> None:
    """Verify get_feature returns most recent feature by name."""
    session = init_test_db(tmp_path)
    try:
        track_id = _create_track(session)

        run1 = AnalysisRepository.create_run(session, track_id, "test", "1.0.0")
        AnalysisRepository.finish_run(session, run1, 10.0, success=True)
        FeatureRepository.save_feature_set(
            session, track_id, run1.id, _make_feature_set()
        )

        feature = FeatureRepository.get_feature(session, track_id, "signal.peak")
        assert feature is not None
        assert feature.name == "signal.peak"
        assert feature.value == "0.8"
    finally:
        session.close()
