import uuid
from datetime import datetime, timezone

from fastapi import HTTPException
from sqlalchemy import func, or_, select
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy.orm import selectinload

from app.core.actors import ActorType
from app.models.signalement import Signalement
from app.models.signalement_timeline import SignalementTimeline
from app.models.signalement_piece_jointe import SignalementPieceJointe
from app.models.note_interne import NoteInterne
from app.models.discussion import Discussion
from app.models.utilisateur import Utilisateur
from app.realtime import events as ws_events
from app.services.stats_service import StatsService
from app.schemas.signalement import (
    NoteInterneCreate,
    SignalementCreate,
)


def _gen_reference() -> str:
    from datetime import date

    today = date.today().strftime("%Y%m%d")
    short = str(uuid.uuid4()).replace("-", "")[:6].upper()
    return f"SIG-{today}-{short}"


def _serialize(
    s: Signalement,
    with_pieces: bool = False,
    has_unread_updates: bool = False,
    unread_messages_count: int = 0,
) -> dict:
    gestionnaire_nom = None
    if s.gestionnaire:
        gestionnaire_nom = f"{s.gestionnaire.prenom} {s.gestionnaire.nom}".strip()

    source_acteur = s.actor_type or ActorType.AUTH_USER.value
    acteur_id = s.actor_user_id if source_acteur == ActorType.AUTH_USER.value else s.actor_guest_id
    if not acteur_id:
        acteur_id = s.utilisateur_id

    data = {
        "id": str(s.id),
        "reference": s.reference,
        "titre": s.titre,
        "description": s.description,
        "type": s.type,
        "statut": s.statut,
        "estTraite": s.est_traite,
        "modeIdentification": s.mode_identification,
        "priorite": s.priorite,
        "dateIncident": s.date_incident,
        "lieuIncident": s.lieu_incident,
        "auteurPresume": s.auteur_presume,
        "aTemoins": s.a_temoins,
        "descriptionTemoins": s.description_temoins,
        "utilisateurId": str(s.utilisateur_id) if s.utilisateur_id else None,
        "sourceActeur": source_acteur,
        "acteurId": str(acteur_id) if acteur_id else None,
        "estFusionne": bool(
            s.actor_type == ActorType.AUTH_USER.value
            and s.actor_guest_id is None
            and s.actor_user_id is not None
        ),
        "gestionnaireId": str(s.gestionnaire_id) if s.gestionnaire_id else None,
        "gestionnaireNom": gestionnaire_nom,
        "bureauGenreId": s.bureau_genre_id,
        "dateCreation": s.date_creation,
        "dateModification": s.date_modification,
        "dateTraitement": s.date_traitement,
        "dateAssignation": s.date_assignation,
        "hasUnreadUpdates": has_unread_updates,
        "messagesNonLusCount": unread_messages_count,
    }
    if with_pieces:
        data["piecesJointes"] = [
            {
                "id": str(pj.id),
                "nomFichier": pj.nom_fichier,
                "typeFichier": pj.type_fichier,
                "tailleFichier": pj.taille_fichier,
                "url": pj.url,
                "filepath": pj.filepath,
                "fileServerId": pj.file_server_id,
                "dateUpload": pj.date_upload,
            }
            for pj in (s.pieces_jointes or [])
        ]
    return data


class SignalementService:
    def __init__(self, db: AsyncSession) -> None:
        self.db = db

    async def creer(
        self,
        actor_type: str,
        actor_id: uuid.UUID,
        data: SignalementCreate,
        actor_name: str | None = None,
    ) -> dict:
        user_id = actor_id if actor_type == ActorType.AUTH_USER.value else None
        signalement = Signalement(
            reference=_gen_reference(),
            titre=data.titre,
            description=data.description,
            type=data.type,
            statut="NOUVEAU",
            est_traite=False,
            mode_identification=data.modeIdentification,
            date_incident=data.dateIncident,
            lieu_incident=data.lieuIncident,
            auteur_presume=data.auteurPresume,
            a_temoins=data.aTemoins,
            description_temoins=data.descriptionTemoins,
            actor_type=actor_type,
            actor_user_id=user_id,
            actor_guest_id=(actor_id if actor_type == ActorType.GUEST.value else None),
            utilisateur_id=user_id,
            bureau_genre_id=data.bureauGenreId,
        )
        self.db.add(signalement)
        await self.db.flush()

        timeline_event = SignalementTimeline(
            signalement_id=signalement.id,
            type_evenement="CREATION",
            description="Signalement créé",
            actor_type=actor_type,
            actor_user_id=user_id,
            actor_guest_id=(actor_id if actor_type == ActorType.GUEST.value else None),
            effectue_par_id=user_id,
            effectue_par_nom=actor_name,
        )
        self.db.add(timeline_event)

        if data.piecesJointes:
            for pj in data.piecesJointes:
                piece_jointe = SignalementPieceJointe(
                    signalement_id=signalement.id,
                    file_server_id=pj.get("fileId") or pj.get("file_server_id"),
                    filepath=pj.get("filepath"),
                    nom_fichier=pj.get("nomFichier") or pj.get("nom") or pj.get("fileName"),
                    type_fichier=pj.get("typeFichier") or pj.get("type") or pj.get("fileType"),
                    taille_fichier=pj.get("tailleFichier") or pj.get("taille") or pj.get("fileSize") or 0,
                    url=pj.get("url"),
                    date_upload=datetime.now(timezone.utc),
                )
                self.db.add(piece_jointe)

        await self.db.commit()
        await self.db.refresh(signalement, ["pieces_jointes"])

        result = _serialize(signalement, with_pieces=True)
        await ws_events.broadcast_signalement_created(str(signalement.id), result)
        return result

    async def get_by_id(self, signalement_id: str | uuid.UUID, current_user: Utilisateur) -> dict:
        s = await self._get_or_404(signalement_id)
        self._check_access(s, current_user)
        await self.db.refresh(s, ["pieces_jointes"])
        return _serialize(s, with_pieces=True)

    async def get_by_id_for_actor(
        self,
        signalement_id: str | uuid.UUID,
        actor_type: str,
        actor_id: uuid.UUID,
        current_user: Utilisateur | None = None,
    ) -> dict:
        s = await self._get_or_404(signalement_id)
        self._check_access_actor(s, actor_type=actor_type, actor_id=actor_id, current_user=current_user)
        await self.db.refresh(s, ["pieces_jointes"])
        return _serialize(s, with_pieces=True)

    async def get_mes_signalements(
        self, utilisateur_id: uuid.UUID, page: int = 0, size: int = 20
    ) -> tuple[list[dict], int]:
        q = (
            select(Signalement)
            .options(
                selectinload(Signalement.gestionnaire),
                selectinload(Signalement.pieces_jointes),
            )
            .where(
                or_(
                    Signalement.utilisateur_id == utilisateur_id,
                    Signalement.actor_user_id == utilisateur_id,
                )
            )
        )
        total = (await self.db.execute(select(func.count()).select_from(q.subquery()))).scalar_one()
        rows = (
            await self.db.execute(
                q.order_by(Signalement.date_creation.desc()).offset(page * size).limit(size)
            )
        ).scalars().all()

        unread_by_signalement = await self._get_unread_counts_for_user(
            signalement_ids=[r.id for r in rows],
            utilisateur_id=utilisateur_id,
        )
        return [
            _serialize(
                s,
                with_pieces=True,
                has_unread_updates=unread_by_signalement.get(str(s.id), 0) > 0,
                unread_messages_count=unread_by_signalement.get(str(s.id), 0),
            )
            for s in rows
        ], total

    async def get_guest_signalements(
        self,
        guest_id: uuid.UUID,
        page: int = 0,
        size: int = 20,
    ) -> tuple[list[dict], int]:
        q = (
            select(Signalement)
            .options(
                selectinload(Signalement.gestionnaire),
                selectinload(Signalement.pieces_jointes),
            )
            .where(Signalement.actor_guest_id == guest_id)
        )
        total = (await self.db.execute(select(func.count()).select_from(q.subquery()))).scalar_one()
        rows = (
            await self.db.execute(
                q.order_by(Signalement.date_creation.desc()).offset(page * size).limit(size)
            )
        ).scalars().all()

        unread_by_signalement = await self._get_unread_counts_for_guest(
            signalement_ids=[r.id for r in rows],
            guest_id=guest_id,
        )
        return [
            _serialize(
                s,
                with_pieces=True,
                has_unread_updates=unread_by_signalement.get(str(s.id), 0) > 0,
                unread_messages_count=unread_by_signalement.get(str(s.id), 0),
            )
            for s in rows
        ], total

    async def get_non_traites(
        self, page: int = 0, size: int = 20, statut: str | None = None
    ) -> tuple[list[dict], int]:
        q = (
            select(Signalement)
            .options(
                selectinload(Signalement.gestionnaire),
                selectinload(Signalement.pieces_jointes),
            )
            .where(
                Signalement.statut.notin_(["TRAITE", "CLOTURE"])
                if statut is None
                else Signalement.statut == statut
            )
        )
        total = (await self.db.execute(select(func.count()).select_from(q.subquery()))).scalar_one()
        rows = (
            await self.db.execute(
                q.order_by(Signalement.date_creation.desc()).offset(page * size).limit(size)
            )
        ).scalars().all()

        unread_by_signalement = await self._get_unread_counts_for_staff(
            signalement_ids=[r.id for r in rows],
        )
        return [
            _serialize(
                s,
                with_pieces=True,
                has_unread_updates=unread_by_signalement.get(str(s.id), 0) > 0,
                unread_messages_count=unread_by_signalement.get(str(s.id), 0),
            )
            for s in rows
        ], total

    async def update_statut(
        self, signalement_id: str | uuid.UUID, nouveau_statut: str, current_user: Utilisateur
    ) -> dict:
        s = await self._get_or_404(signalement_id)

        if (
            str(s.gestionnaire_id) != str(current_user.id)
            and current_user.type_utilisateur not in ("ADMIN_SYSTEME",)
        ):
            raise HTTPException(status_code=403, detail="Permission insuffisante")

        ancien_statut = s.statut
        s.statut = nouveau_statut
        s.est_traite = nouveau_statut in ("TRAITE", "CLOTURE")
        if nouveau_statut in ("TRAITE", "CLOTURE"):
            s.date_traitement = datetime.now(timezone.utc)

        timeline_event = SignalementTimeline(
            signalement_id=s.id,
            type_evenement="CHANGEMENT_STATUT",
            description=f"Statut changé de {ancien_statut} vers {nouveau_statut}",
            effectue_par_id=current_user.id,
            effectue_par_nom=f"{current_user.prenom} {current_user.nom}",
            valeur_precedente=ancien_statut,
            nouvelle_valeur=nouveau_statut,
        )
        self.db.add(timeline_event)
        await self.db.commit()
        await self.db.refresh(s, ["pieces_jointes", "gestionnaire"])

        result = _serialize(s, with_pieces=True)
        await ws_events.broadcast_statut_updated(str(s.id), result)
        dashboard = await StatsService(self.db).get_dashboard()
        await ws_events.broadcast_stats_updated(
            {
                "signalementsEnAttente": dashboard.get("signalementsEnAttente", 0),
                "evenementsAVenir": dashboard.get("evenementsAVenir", 0),
                "equipeAlertes": dashboard.get("equipeAlertes", 0),
                "moderationEnAttente": dashboard.get("moderationEnAttente", 0),
            }
        )
        return result

    async def assigner(
        self, signalement_id: str | uuid.UUID, gestionnaire_id: str | uuid.UUID | None, current_user: Utilisateur
    ) -> dict:
        signalement_uuid = signalement_id if isinstance(signalement_id, uuid.UUID) else uuid.UUID(signalement_id)

        # Atomic check using select for update
        stmt = select(Signalement).where(Signalement.id == signalement_uuid).with_for_update()
        s = (await self.db.execute(stmt)).scalar_one_or_none()
        if not s:
            raise HTTPException(status_code=404, detail="Signalement non trouvé")

        # Double assignment check
        if s.gestionnaire_id is not None:
            if current_user.type_utilisateur != "ADMIN_SYSTEME":
                if str(s.gestionnaire_id) != str(current_user.id):
                    raise HTTPException(
                        status_code=400,
                        detail="Ce dossier est déjà pris en charge par un autre gestionnaire"
                    )

        if gestionnaire_id is None or str(gestionnaire_id) == "None" or str(gestionnaire_id) == "":
            old_manager_name = ""
            if s.gestionnaire:
                old_manager_name = f"{s.gestionnaire.prenom} {s.gestionnaire.nom}".strip()
            s.gestionnaire_id = None
            s.date_assignation = None
            s.statut = "NOUVEAU"
            gest_name_desc = f"Désassigné (précédemment assigné à {old_manager_name})" if old_manager_name else "Désassigné"
        else:
            gest_uuid = gestionnaire_id if isinstance(gestionnaire_id, uuid.UUID) else uuid.UUID(gestionnaire_id)
            gestionnaire = (
                await self.db.execute(
                    select(Utilisateur).where(Utilisateur.id == gest_uuid)
                )
            ).scalar_one_or_none()
            if not gestionnaire:
                raise HTTPException(status_code=404, detail="Gestionnaire introuvable")

            s.gestionnaire_id = gestionnaire.id
            s.date_assignation = datetime.now(timezone.utc)
            if s.statut == "NOUVEAU":
                s.statut = "ASSIGNE"
            gest_name_desc = f"Assigné à {gestionnaire.prenom} {gestionnaire.nom}"

        timeline_event = SignalementTimeline(
            signalement_id=s.id,
            type_evenement="ASSIGNATION",
            description=gest_name_desc,
            effectue_par_id=current_user.id,
            effectue_par_nom=f"{current_user.prenom} {current_user.nom}",
        )
        self.db.add(timeline_event)
        await self.db.commit()

        # Reload with selectinload to prevent MissingGreenlet lazy-loading exceptions on serialized properties
        stmt = select(Signalement).where(Signalement.id == signalement_uuid).options(
            selectinload(Signalement.gestionnaire),
            selectinload(Signalement.pieces_jointes)
        )
        s = (await self.db.execute(stmt)).scalar_one()

        result = _serialize(s, with_pieces=True)
        await ws_events.broadcast_signalement_assigned(str(s.id), result)
        return result

    async def ajouter_note(
        self, signalement_id: str | uuid.UUID, data: NoteInterneCreate, current_user: Utilisateur
    ) -> dict:
        s = await self._get_or_404(signalement_id)

        note = NoteInterne(
            signalement_id=s.id,
            actor_type=ActorType.AUTH_USER.value,
            actor_user_id=current_user.id,
            actor_guest_id=None,
            auteur_id=current_user.id,
            contenu=data.contenu,
        )
        self.db.add(note)
        await self.db.commit()
        await self.db.refresh(note)

        result = {
            "id": str(note.id),
            "auteurId": str(note.auteur_id) if note.auteur_id else None,
            "sourceActeur": note.actor_type,
            "acteurId": str(note.actor_user_id) if note.actor_user_id else None,
            "contenu": note.contenu,
            "dateCreation": note.date_creation,
        }
        await ws_events.broadcast_note_added(str(s.id), result)
        return result

    async def get_timeline(self, signalement_id: str | uuid.UUID, current_user: Utilisateur) -> list[dict]:
        s = await self._get_or_404(signalement_id)
        self._check_access(s, current_user)

        rows = (
            await self.db.execute(
                select(SignalementTimeline)
                .where(SignalementTimeline.signalement_id == s.id)
                .order_by(SignalementTimeline.timestamp.asc())
            )
        ).scalars().all()

        return [
            {
                "id": str(r.id),
                "typeEvenement": r.type_evenement,
                "description": r.description,
                "timestamp": r.timestamp,
                "sourceActeur": r.actor_type,
                "acteurId": str(r.actor_user_id or r.actor_guest_id)
                if (r.actor_user_id or r.actor_guest_id)
                else None,
                "effectueParNom": r.effectue_par_nom,
                "valeurPrecedente": r.valeur_precedente,
                "nouvelleValeur": r.nouvelle_valeur,
            }
            for r in rows
        ]

    async def get_timeline_for_actor(
        self,
        signalement_id: str | uuid.UUID,
        actor_type: str,
        actor_id: uuid.UUID,
        current_user: Utilisateur | None = None,
    ) -> list[dict]:
        s = await self._get_or_404(signalement_id)
        self._check_access_actor(s, actor_type=actor_type, actor_id=actor_id, current_user=current_user)

        rows = (
            await self.db.execute(
                select(SignalementTimeline)
                .where(SignalementTimeline.signalement_id == s.id)
                .order_by(SignalementTimeline.timestamp.asc())
            )
        ).scalars().all()

        return [
            {
                "id": str(r.id),
                "typeEvenement": r.type_evenement,
                "description": r.description,
                "timestamp": r.timestamp,
                "sourceActeur": r.actor_type,
                "acteurId": str(r.actor_user_id or r.actor_guest_id)
                if (r.actor_user_id or r.actor_guest_id)
                else None,
                "effectueParNom": r.effectue_par_nom,
                "valeurPrecedente": r.valeur_precedente,
                "nouvelleValeur": r.nouvelle_valeur,
            }
            for r in rows
        ]

    def _check_access(self, s: Signalement, user: Utilisateur) -> None:
        if user.type_utilisateur in ("ADMIN_SYSTEME",):
            return
        if str(s.actor_user_id or s.utilisateur_id) == str(user.id):
            return
        if str(s.gestionnaire_id) == str(user.id):
            return
        if user.type_utilisateur == "GESTIONNAIRE" and s.gestionnaire_id is None:
            return
        raise HTTPException(status_code=403, detail="Accès refusé")

    def _check_access_actor(
        self,
        s: Signalement,
        actor_type: str,
        actor_id: uuid.UUID,
        current_user: Utilisateur | None = None,
    ) -> None:
        if actor_type == ActorType.AUTH_USER.value:
            if current_user is None:
                raise HTTPException(status_code=401, detail="AUTH_REQUIRED")
            self._check_access(s, current_user)
            return

        if str(s.actor_guest_id) == str(actor_id):
            return
        raise HTTPException(status_code=403, detail="GUEST_ACCESS_DENIED")

    async def _get_or_404(self, signalement_id: str | uuid.UUID) -> Signalement:
        signalement_uuid = signalement_id if isinstance(signalement_id, uuid.UUID) else uuid.UUID(signalement_id)
        s = (
            await self.db.execute(
                select(Signalement)
                .options(selectinload(Signalement.gestionnaire))
                .where(Signalement.id == signalement_uuid)
            )
        ).scalar_one_or_none()
        if not s:
            raise HTTPException(status_code=404, detail="Signalement non trouvé")
        return s

    async def _get_unread_counts_for_user(
        self,
        signalement_ids: list[uuid.UUID],
        utilisateur_id: uuid.UUID,
    ) -> dict[str, int]:
        if not signalement_ids:
            return {}

        rows = (
            await self.db.execute(
                select(Discussion.signalement_id, Discussion.non_lus_utilisateur).where(
                    or_(
                        Discussion.utilisateur_id == utilisateur_id,
                        Discussion.actor_user_id == utilisateur_id,
                    ),
                    Discussion.signalement_id.is_not(None),
                    Discussion.signalement_id.in_(signalement_ids),
                )
            )
        ).all()
        return {
            str(signalement_id): int(non_lus or 0)
            for signalement_id, non_lus in rows
            if signalement_id is not None
        }

    async def _get_unread_counts_for_guest(
        self,
        signalement_ids: list[uuid.UUID],
        guest_id: uuid.UUID,
    ) -> dict[str, int]:
        if not signalement_ids:
            return {}

        rows = (
            await self.db.execute(
                select(Discussion.signalement_id, Discussion.non_lus_utilisateur).where(
                    Discussion.actor_guest_id == guest_id,
                    Discussion.signalement_id.is_not(None),
                    Discussion.signalement_id.in_(signalement_ids),
                )
            )
        ).all()
        return {
            str(signalement_id): int(non_lus or 0)
            for signalement_id, non_lus in rows
            if signalement_id is not None
        }

    async def _get_unread_counts_for_staff(
        self,
        signalement_ids: list[uuid.UUID],
    ) -> dict[str, int]:
        if not signalement_ids:
            return {}

        rows = (
            await self.db.execute(
                select(Discussion.signalement_id, Discussion.non_lus_gestionnaire).where(
                    Discussion.signalement_id.is_not(None),
                    Discussion.signalement_id.in_(signalement_ids),
                )
            )
        ).all()
        return {
            str(signalement_id): int(non_lus or 0)
            for signalement_id, non_lus in rows
            if signalement_id is not None
        }
