from django.utils import timezone
from rest_framework import status
from rest_framework.exceptions import ValidationError

from api.models import AuditTrail, AuditTrailDetails, EsExamActivity, EsInvigilator, EsSession, EsSessionInvigilator
from api.services.invigilator_assign import sessions_overlap
from api.services.tt_availability import tt_staff_conflict_at
from api.translation import __
from api.utils import get_exception_detail, get_ip, log_critical_error
from api.validator import BaseValidator
from api.views.admin.base import AdminApiBase
from api.services.exam_scheduler import SLOT_MINUTES, _duration_slots, _slot_to_datetime


class SessionUpdateInvigilator(AdminApiBase):
    def validate_request(self, request):
        rules = {
            "id": "required|array|exists:api.EsSession,id",
            "invigilator_id": "nullable|array|exists:api.EsInvigilator,id",
            "floating_invigilator_id": "nullable|array|exists:api.EsInvigilator,id",
        }

        attribute = {
            "id": "ID",
            "invigilator_id": __("attr.invigilator"),
            "floating_invigilator_id": __("attr.floating_invigilator_id"),
        }

        validator = BaseValidator(request.data, rules, attribute)
        error = validator.validate()
        if error:
            raise ValidationError(error)

        session_ids = request.data.get("id") or []
        return EsSession.objects.prefetch_related("essessioninvigilator_set").filter(id__in=session_ids)

    def post(self, request):
        try:
            edit_obj = self.validate_request(request)

            action = "session_update_invigilator"
            old_new_data = {
                "new_data": {},
                "old_data": {},
            }

            session_ids = request.data.get("id") or []
            invigilator_ids = list(dict.fromkeys(request.data.get("invigilator_id") or []))
            update_floating = "floating_invigilator_id" in request.data
            floating_invigilator_ids = list(dict.fromkeys(request.data.get("floating_invigilator_id") or []))
            overlap = sorted(set(invigilator_ids) & set(floating_invigilator_ids))
            if overlap:
                raise _overlap_error(overlap)

            if invigilator_ids:
                _reject_same_time_invigilators(list(edit_obj), invigilator_ids)

            if session_ids:
                for obj in edit_obj:
                    existing_by_flag = {
                        0: set(),
                        1: set(),
                    }
                    for invigilator_id, is_floating in EsSessionInvigilator.objects.filter(
                        session_id=obj.id,
                        invigilator_id__isnull=False,
                    ).values_list("invigilator_id", "is_floating"):
                        existing_by_flag[1 if is_floating else 0].add(invigilator_id)

                    existing_invigilator_ids = existing_by_flag[0]
                    existing_floating_ids = existing_by_flag[1]
                    effective_floating_ids = floating_invigilator_ids if update_floating else existing_floating_ids
                    kept_overlap = sorted(set(invigilator_ids) & set(effective_floating_ids))
                    if kept_overlap:
                        raise _overlap_error(kept_overlap)
                    session_changed = set(existing_invigilator_ids) != set(invigilator_ids)
                    floating_changed = update_floating and set(existing_floating_ids) != set(floating_invigilator_ids)
                    to_add = []
                    if session_changed:
                        to_add.extend(sid for sid in invigilator_ids if sid not in existing_invigilator_ids)
                    floating_to_add = []
                    if floating_changed:
                        floating_to_add.extend(
                            sid for sid in floating_invigilator_ids if sid not in existing_floating_ids
                        )

                    if to_add or floating_to_add:
                        activity = (
                            EsExamActivity.objects.filter(session_id=obj.id, is_scheduled=True)
                            .select_related("exam_requirement")
                            .first()
                        )
                        session_start = activity and _slot_to_datetime(
                            activity.exam_requirement.exam_period,
                            activity.time_slot,
                        ) or obj.start_time
                        duration_minutes = obj.duration.hour * 60 + obj.duration.minute
                        if activity and activity.exam_requirement:
                            duration_slots = _duration_slots(activity.exam_requirement)
                            duration_minutes = duration_slots * SLOT_MINUTES

                        busy_fields = {}
                        for field, ids in (
                            ("invigilator_id", to_add),
                            ("floating_invigilator_id", floating_to_add),
                        ):
                            if not ids:
                                continue
                            busy_invigilators = []
                            for invigilator in EsInvigilator.objects.filter(id__in=ids).select_related("staff"):
                                staff_tt_id = invigilator.staff.tt_id if invigilator.staff_id else None
                                if staff_tt_id and tt_staff_conflict_at([staff_tt_id], session_start, duration_minutes):
                                    busy_invigilators.append(invigilator.name or invigilator.code)
                            if busy_invigilators:
                                busy_fields[field] = __("validation.tt_staff_busy", values=", ".join(busy_invigilators))

                        if busy_fields:
                            raise ValidationError({
                                "error": next(iter(busy_fields.values())),
                                "errors": busy_fields,
                            })

                    if session_changed:
                        old_new_data.setdefault("old_data", {}).setdefault(obj.id, {})["invigilator_id"] = sorted(existing_invigilator_ids)
                        old_new_data.setdefault("new_data", {}).setdefault(obj.id, {})["invigilator_id"] = sorted(invigilator_ids)
                        _replace_session_invigilators(obj.id, invigilator_ids, request.user.id, is_floating=0)

                    if floating_changed:
                        old_new_data.setdefault("old_data", {}).setdefault(obj.id, {})["floating_invigilator_id"] = sorted(existing_floating_ids)
                        old_new_data.setdefault("new_data", {}).setdefault(obj.id, {})["floating_invigilator_id"] = sorted(floating_invigilator_ids)
                        _replace_session_invigilators(obj.id, floating_invigilator_ids, request.user.id, is_floating=1)

            if old_new_data["old_data"] or old_new_data["new_data"]:
                audit_trail = AuditTrail.objects.create(
                    user_id=request.user.id,
                    type=self.audit_type,
                    ip_address=get_ip(request),
                )

                remark_param = {
                    "name": ",".join(edit_obj.values_list("name", flat=True)),
                }

                AuditTrailDetails.custom_insert(
                    audit_trail=audit_trail,
                    action=action,
                    remark_param=remark_param,
                    new_data=old_new_data["new_data"],
                    old_data=old_new_data["old_data"],
                )

            response = {}

            return self.api_response(data=response)
        except ValidationError as e:
            first_message = e.detail["error"]
            errors = e.detail["errors"]
            return self.api_response(error=first_message, errors=errors, code=status.HTTP_400_BAD_REQUEST)
        except Exception as e:
            e_details = get_exception_detail(e)
            log_critical_error(user_id=None, descr=e_details["descr"], url=e_details["url"], trace=e_details["trace"])
            return self.api_response(error=__("message.internal_server_error"), code=status.HTTP_500_INTERNAL_SERVER_ERROR)


def _reject_same_time_invigilators(sessions, invigilator_ids):
    session_ids = [session.id for session in sessions]
    for index, left in enumerate(sessions):
        for right in sessions[index + 1:]:
            if sessions_overlap(left, right):
                raise _session_overlap_error(invigilator_ids)

    other_links = (
        EsSessionInvigilator.objects.filter(
            invigilator_id__in=invigilator_ids,
            session_id__isnull=False,
        )
        .exclude(session_id__in=session_ids)
        .select_related("session")
    )
    busy_ids = set()
    for link in other_links:
        if any(sessions_overlap(session, link.session) for session in sessions):
            busy_ids.add(link.invigilator_id)
    if busy_ids:
        raise _session_overlap_error(sorted(busy_ids))


def _session_overlap_error(invigilator_ids):
    names = EsInvigilator.objects.filter(id__in=invigilator_ids).values_list("name", flat=True)
    message = __("validation.invigilator_session_overlap", values=", ".join(names))
    return ValidationError({
        "error": message,
        "errors": {
            "invigilator_id": message,
        },
    })


def _overlap_error(invigilator_ids):
    message = __(
        "validation.not_in",
        field=__("attr.floating_invigilator_id"),
        values=", ".join(str(invigilator_id) for invigilator_id in invigilator_ids),
    )
    return ValidationError({
        "error": message,
        "errors": {
            "floating_invigilator_id": message,
        },
    })


def _replace_session_invigilators(session_id, invigilator_ids, user_id, is_floating):
    qs_to_null = EsSessionInvigilator.objects.filter(session_id=session_id, is_floating=is_floating)
    if invigilator_ids:
        qs_to_null = qs_to_null.exclude(invigilator_id__in=invigilator_ids)
    qs_to_null.update(invigilator=None)
    EsSessionInvigilator.objects.filter(
        session_id=session_id,
        invigilator_id__isnull=True,
    ).delete()

    if invigilator_ids:
        EsSessionInvigilator.objects.filter(
            session_id=session_id,
            invigilator_id__in=invigilator_ids,
        ).exclude(is_floating=is_floating).update(
            is_floating=is_floating,
            updated_by=user_id,
            updated_at=timezone.now(),
        )
        EsSessionInvigilator.bulk_insert(session_id, invigilator_ids, user_id, is_floating=is_floating)
        EsSessionInvigilator.objects.filter(
            session_id=session_id,
            invigilator_id__in=invigilator_ids,
            is_floating=is_floating,
        ).update(
            updated_by=user_id,
            updated_at=timezone.now(),
        )