"""
==================================================
Smart Attendance AI
Trainer Service
==================================================
Handles face registration (training).
==================================================
"""

import time
from typing import List

from config.settings import settings

from core.image_data import ImageData
from core.training_result import TrainingResult

from services.inference_service import InferenceService
from services.embedding_service import EmbeddingService
from services.similarity_service import SimilarityService

from utils.embedding import EmbeddingUtils
from utils.exceptions import (
    TooFewTrainingImagesException,
    TooManyTrainingImagesException,
    MultipleTrainingFacesException,
    NoEmbeddingGeneratedException
)


class TrainerService:

    # ==========================================
    # Train Student
    # ==========================================

    @classmethod
    def train(
        cls,
        student_id: str,
        images: List[ImageData]
    ) -> TrainingResult:
        """
        Register a student's face using
        multiple training images.
        """

        start_time = time.perf_counter()

        cls._validate_image_count(images)

        embeddings = cls._extract_embeddings(
            images
        )

        embeddings = cls._remove_duplicates(
            embeddings
        )

        average_embedding = cls._average_embedding(
            embeddings
        )

        processing_time_ms = int(
            (
                time.perf_counter()
                - start_time
            ) * 1000
        )

        return TrainingResult(

            student_id=student_id,

            images_received=len(images),

            images_processed=len(images),

            faces_detected=len(embeddings),

            embeddings_generated=len(embeddings),

            duplicates_removed=(
                len(images) - len(embeddings)
            ),

            average_embedding=average_embedding,

            embeddings=embeddings,

            processing_time_ms=processing_time_ms

        )

    # ==========================================
    # Validate Image Count
    # ==========================================

    @staticmethod
    def _validate_image_count(
        images: List[ImageData]
    ) -> None:

        count = len(images)

        if count < settings.MIN_TRAIN_IMAGES:

            raise TooFewTrainingImagesException()

        if count > settings.MAX_TRAIN_IMAGES:

            raise TooManyTrainingImagesException()

    # ==========================================
    # Extract Embeddings
    # ==========================================

    @classmethod
    def _extract_embeddings(
        cls,
        images: List[ImageData]
    ) -> List[List[float]]:
        """
        Extract one embedding from each
        training image.
        """

        embeddings: List[List[float]] = []

        for image in images:

            # Detect faces
            faces = InferenceService.infer(
                image
            )

            # No face
            if not faces:

                continue

            # Require exactly one face
            if (
                settings.REQUIRE_SINGLE_FACE
                and len(faces) != 1
            ):

                raise MultipleTrainingFacesException()

            # Generate embedding
            results = EmbeddingService.extract(
                faces
            )

            if not results:

                continue

            embedding = results[0].embedding

            # Validate embedding size
            if not EmbeddingUtils.validate_dimension(
                embedding
            ):

                continue

            embeddings.append(
                embedding
            )

        if not embeddings:

            raise NoEmbeddingGeneratedException()

        return embeddings

    # ==========================================
    # Remove Duplicate Embeddings
    # ==========================================

    @classmethod
    def _remove_duplicates(
        cls,
        embeddings: List[List[float]]
    ) -> List[List[float]]:

        if settings.ALLOW_DUPLICATE_IMAGES:

            return embeddings

        unique_embeddings: List[List[float]] = []

        for embedding in embeddings:

            duplicate = False

            for existing in unique_embeddings:

                similarity = (
                    SimilarityService.cosine_similarity(
                        embedding,
                        existing
                    )
                )

                if (
                    similarity >=
                    settings.DUPLICATE_FACE_THRESHOLD
                ):

                    duplicate = True

                    break

            if not duplicate:

                unique_embeddings.append(
                    embedding
                )

        return unique_embeddings

    # ==========================================
    # Average Embedding
    # ==========================================

    @staticmethod
    def _average_embedding(
        embeddings: List[List[float]]
    ) -> List[float]:
        """
        Compute the average embedding
        from all valid embeddings.
        """

        return EmbeddingUtils.average(
            embeddings
        )