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

from api.models import EsExamRequirement, EsExamRequirementStudent
from api.models.audit_trail import AuditTrail
from api.models.audit_trail_details import AuditTrailDetails
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


class ExamRequirementUpdateStudent(AdminApiBase):
    def validate_request(self, request):
        rules = {
            "id": "required|array|exists:api.EsExamRequirement,id",
            "student_id": "nullable|array|exists:api.EsStudent,id",
        }

        attribute = {
            "id": "ID",
            "student_id": __("attr.student"),
        }

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

        exam_requirement_ids = request.data.get("id") or []
        return (
            EsExamRequirement.objects
            .prefetch_related("esexamrequirementstudent_set")
            .filter(id__in=exam_requirement_ids)
        )

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

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

            exam_requirement_ids = request.data.get("id") or []
            student_ids = request.data.get("student_id") or []
            # Deduplicate while preserving order.
            student_ids = list(dict.fromkeys(student_ids))

            if exam_requirement_ids:
                for obj in edit_obj:
                    existing_student_ids = set(
                        EsExamRequirementStudent.objects.filter(
                            exam_requirement_id=obj.id,
                            student_id__isnull=False,
                        )
                        .values_list("student_id", flat=True)
                    )

                    if set(existing_student_ids) != set(student_ids):
                        old_new_data.setdefault("old_data", {}).setdefault(obj.id, {})["student_id"] = sorted(existing_student_ids)
                        old_new_data.setdefault("new_data", {}).setdefault(obj.id, {})["student_id"] = sorted(student_ids)

                        # Null out removed students then delete those null rows.
                        qs_to_null = EsExamRequirementStudent.objects.filter(
                            exam_requirement_id=obj.id,
                        )
                        if student_ids:
                            qs_to_null = qs_to_null.exclude(student_id__in=student_ids)
                        qs_to_null.update(student_id=None)
                        EsExamRequirementStudent.objects.filter(
                            exam_requirement_id=obj.id,
                            student_id__isnull=True,
                        ).delete()

                        to_add = [sid for sid in student_ids if sid not in existing_student_ids]
                        if to_add:
                            EsExamRequirementStudent.objects.bulk_create(
                                [
                                    EsExamRequirementStudent(
                                        exam_requirement_id=obj.id,
                                        student_id=sid,
                                        created_by=request.user.id,
                                        updated_by=request.user.id,
                                    )
                                    for sid in to_add
                                ]
                            )

                        if student_ids:
                            EsExamRequirementStudent.objects.filter(
                                exam_requirement_id=obj.id,
                                student_id__in=student_ids,
                            ).update(
                                updated_by=request.user.id,
                                updated_at=timezone.now(),
                            )


            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 = {
                    "code": ",".join(edit_obj.values_list("code", 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)
