from rest_framework import status
from rest_framework.exceptions import ValidationError
from datetime import datetime

from api.models import EsExamActivity, EsExamRequirementStudent, EsSessionInvigilator
from api.models.es_session import EsSession
from api.services.invigilator_assign import summarize_assigned_invigilators
from api.translation import __
from api.utils import log_critical_error, get_exception_detail
from api.validator import BaseValidator
from api.views.admin.base import AdminApiBase
from django.db.models import Prefetch


class SessionList(AdminApiBase):
    list_default_order_column = "start_time"

    def validate_request(self, request):
        rules = {
            "exam_period_id": "required|exists:api.EsExamPeriod,id",
        }

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

    def post(self, request):
        try:
            self.validate_request(request)
            self.api_log_skip_outgoing_data = True
            data = self.get_data(request)
            return self.api_response(data=data)
        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 get_data(self, request):
        exam_period_id = request.data["exam_period_id"]
        exam_activity_queryset = (
            EsExamActivity.objects
            .select_related("exam_requirement")
            .filter(exam_requirement__exam_period_id=exam_period_id)
        )
        session_invigilator_queryset = (
            EsSessionInvigilator.objects
            .filter(invigilator_id__isnull=False)
            .select_related("invigilator")
            .prefetch_related("invigilator__role")
        )
        objs = (
            EsSession.objects
            .select_related("location")
            .prefetch_related(
                Prefetch("esexamactivity_set", queryset=exam_activity_queryset),
                Prefetch("essessioninvigilator_set", queryset=session_invigilator_queryset),
            )
            .filter(esexamactivity__exam_requirement__exam_period_id=exam_period_id)
        )
        objs = self.session_listing_filter(request, objs)
        objs = objs.distinct()
        pagination = self.session_listing_pagination(request, objs)

        return_data = []

        for val in pagination["paginated_data"]:
            data = {
                "id": val.id,
                "name": val.name,
                "location": val.location.name if val.location else None,
                "description": val.description,
                "start_time": val.start_time,
                "duration": val.duration,
                "students_enrolled": val.students_enrolled,
                "invigilators_required": val.invigilators_required,
                "total_assigned_invigilators": summarize_assigned_invigilators(
                    [
                        relation.invigilator
                        for relation in val.essessioninvigilator_set.all()
                    ]
                ),
                "exam_activities": [
                    {
                        "id": activity.id,
                        "code": activity.code,
                        "name": activity.name,
                        "exam_requirement_id": activity.exam_requirement_id,
                        # "exam_requirement_code": activity.exam_requirement.code,
                        # "exam_requirement_name": activity.exam_requirement.name,
                        "description": activity.exam_requirement.description,
                        "fixed_start_date": activity.exam_requirement.fixed_start_date,
                        "fixed_start_time": activity.exam_requirement.fixed_start_time,
                        "is_scheduled": activity.is_scheduled,
                        "students": activity.exam_requirement.planned_size,
                        "real_size": EsExamRequirementStudent.objects.filter(exam_requirement_id=activity.exam_requirement_id, student_id__isnull=False).count()
                    }
                    for activity in val.esexamactivity_set.all()
                ],
            }
            return_data.append(data)

        return {
            "data": return_data,
            "total": pagination["total"],
            "page": pagination["page"],
            "per_page": pagination["per_page"],
        }

    def session_listing_filter(self, request, model):
        filters = request.data.get("filter")
        if not isinstance(filters, dict) or not filters:
            return model

        for key, val in filters.items():
            if val in (None, "", [], {}):
                continue

            match key:
                case "name":
                    model = model.filter(name__icontains=val)
                case "location":
                    model = model.filter(location__name__icontains=val)
                case "location_id":
                    if isinstance(val, (list, tuple, set)):
                        model = model.filter(location_id__in=list(val))
                    else:
                        model = model.filter(location_id=val)
                case "start_time_from":
                    try:
                        model = model.filter(start_time__gte=datetime.fromisoformat(val))
                    except (TypeError, ValueError):
                        continue
                case "start_time_to":
                    try:
                        model = model.filter(start_time__lte=datetime.fromisoformat(val))
                    except (TypeError, ValueError):
                        continue
                case "duration_from":
                    try:
                        model = model.filter(duration__gte=datetime.strptime(val, "%H:%M:%S").time())
                    except (TypeError, ValueError):
                        try:
                            model = model.filter(duration__gte=datetime.strptime(val, "%H:%M").time())
                        except (TypeError, ValueError):
                            continue
                case "duration_to":
                    try:
                        model = model.filter(duration__lte=datetime.strptime(val, "%H:%M:%S").time())
                    except (TypeError, ValueError):
                        try:
                            model = model.filter(duration__lte=datetime.strptime(val, "%H:%M").time())
                        except (TypeError, ValueError):
                            continue
                case "students_enrolled_min":
                    try:
                        model = model.filter(students_enrolled__gte=int(val))
                    except (TypeError, ValueError):
                        continue
                case "students_enrolled_max":
                    try:
                        model = model.filter(students_enrolled__lte=int(val))
                    except (TypeError, ValueError):
                        continue
                case "invigilators_required_min":
                    try:
                        model = model.filter(invigilators_required__gte=int(val))
                    except (TypeError, ValueError):
                        continue
                case "invigilators_required_max":
                    try:
                        model = model.filter(invigilators_required__lte=int(val))
                    except (TypeError, ValueError):
                        continue
                case _:
                    continue

        return model

    def session_listing_pagination(self, request, model):
        page, per_page = self.parse_pagination(request)
        sort_by = request.data.get("sort_by")
        order_by = request.data.get("order_by")

        total = model.count()

        if sort_by:
            if str(order_by).lower() == "desc":
                ordering_string = f"-{sort_by}"
            else:
                ordering_string = sort_by
            model = model.order_by(ordering_string)
        else:
            model = model.order_by(self.list_default_order_column, "location_id", "name")

        if per_page != -1:
            start_index = (page - 1) * per_page
            end_index = start_index + per_page
            paginated_data = model[start_index:end_index]
        else:
            paginated_data = model

        return {
            "total": total,
            "paginated_data": paginated_data,
            "page": page,
            "per_page": per_page,
        }