"""Feature repository for AI MusiMuse.

This module defines :class:`FeatureRepository`, which manages
persistence and queries for :class:`~database.models.TrackFeature`
records.

Repositories are the only layer allowed to access SQLAlchemy.
No logging is performed inside repositories.
"""

from __future__ import annotations

from sqlalchemy import select
from sqlalchemy.orm import Session

from analyzer.feature import FeatureSet
from database.models import AnalyzerRun, TrackFeature, generate_uuid


class FeatureRepository:
    """Repository for :class:`TrackFeature` persistence and queries.

    Methods:
        save_feature_set: Persist a FeatureSet as TrackFeature rows.
        get_track_features: Get all features for a track.
        get_latest_features: Get features from the latest successful run
            per analyzer.
        get_feature: Get a single feature by track and name.
        has_analyzer_version: Check if an analyzer version has run.
    """

    @staticmethod
    def save_feature_set(
        session: Session,
        track_id: str,
        analyzer_run_id: str,
        feature_set: FeatureSet,
    ) -> int:
        """Persist a FeatureSet as TrackFeature rows.

        Each :class:`~analyzer.feature.Feature` is converted to one
        :class:`TrackFeature` row.  Values are serialized via
        ``str()`` for non-None, NULL for None.

        Args:
            session: An active SQLAlchemy session.
            track_id: UUID of the track.
            analyzer_run_id: UUID of the AnalyzerRun.
            feature_set: The FeatureSet to persist.

        Returns:
            The number of features saved.
        """
        count = 0
        for feature in feature_set:
            tf = TrackFeature(
                id=generate_uuid(),
                track_id=track_id,
                analyzer_run_id=analyzer_run_id,
                name=feature.name,
                value=str(feature.value) if feature.value is not None else None,
                unit=feature.unit,
            )
            session.add(tf)
            count += 1
        session.flush()
        return count

    @staticmethod
    def get_track_features(session: Session, track_id: str) -> list[TrackFeature]:
        """Get all features for a track.

        Args:
            session: An active SQLAlchemy session.
            track_id: UUID of the track.

        Returns:
            A list of :class:`TrackFeature` instances ordered by
            ``created_at`` ascending.
        """
        stmt = (
            select(TrackFeature)
            .where(TrackFeature.track_id == track_id)
            .order_by(TrackFeature.created_at)
        )
        return list(session.execute(stmt).scalars().all())

    @staticmethod
    def get_latest_features(session: Session, track_id: str) -> list[TrackFeature]:
        """Get features from the latest successful run per analyzer.

        For each analyzer that has run on this track, returns features
        from the most recent successful run only.

        Args:
            session: An active SQLAlchemy session.
            track_id: UUID of the track.

        Returns:
            A list of :class:`TrackFeature` instances from the latest
            successful run of each analyzer.
        """
        runs_stmt = (
            select(AnalyzerRun)
            .where(
                AnalyzerRun.track_id == track_id,
                AnalyzerRun.success.is_(True),
            )
            .order_by(
                AnalyzerRun.finished_at.desc(),
                AnalyzerRun.started_at.desc(),
            )
        )
        runs = list(session.execute(runs_stmt).scalars().all())

        seen_analyzers: set[str] = set()
        result: list[TrackFeature] = []
        for run in runs:
            if run.analyzer_name in seen_analyzers:
                continue
            seen_analyzers.add(run.analyzer_name)
            feat_stmt = (
                select(TrackFeature)
                .where(TrackFeature.analyzer_run_id == run.id)
                .order_by(TrackFeature.name)
            )
            result.extend(session.execute(feat_stmt).scalars().all())

        return result

    @staticmethod
    def get_feature(
        session: Session,
        track_id: str,
        name: str,
    ) -> TrackFeature | None:
        """Get a single feature by track and name.

        Returns the most recently created feature with the given name
        for the given track.

        Args:
            session: An active SQLAlchemy session.
            track_id: UUID of the track.
            name: Feature name.

        Returns:
            The most recent :class:`TrackFeature`, or ``None``.
        """
        stmt = (
            select(TrackFeature)
            .where(
                TrackFeature.track_id == track_id,
                TrackFeature.name == name,
            )
            .order_by(TrackFeature.created_at.desc())
        )
        return session.execute(stmt).scalars().first()

    @staticmethod
    def has_analyzer_version(
        session: Session,
        track_id: str,
        analyzer_name: str,
        analyzer_version: str,
    ) -> bool:
        """Check if an analyzer version has successfully run on a track.

        Args:
            session: An active SQLAlchemy session.
            track_id: UUID of the track.
            analyzer_name: Name of the analyzer.
            analyzer_version: Version of the analyzer.

        Returns:
            ``True`` if a successful run exists, ``False`` otherwise.
        """
        stmt = select(AnalyzerRun.id).where(
            AnalyzerRun.track_id == track_id,
            AnalyzerRun.analyzer_name == analyzer_name,
            AnalyzerRun.analyzer_version == analyzer_version,
            AnalyzerRun.success.is_(True),
        )
        return session.execute(stmt).first() is not None
