from collections import defaultdict
from contextlib import suppress

from django.db import transaction
from django.db.models import Count, Q
from rest_framework import status

from api.helper.resource_map_helper import helper_recalculate_and_get_redis_resource_map
from api.models import TtPosModuleGroupModule, TtSetting, TtStudentAcademicTerm, TtStudentPos, TtStudentPosModule, \
    TtPathway, TtStudentPathway, TtStudentSetStudent, TtStudentResourceMap, TtWeek, TtStudentSet
from api.models.audit_trail import AuditTrail
from api.models.audit_trail_details import AuditTrailDetails
from api.models.tt_pos import TtPos
from api.models.tt_student import TtStudent
from api.views.admin.base import AdminApiBase
from api.validator import BaseValidator
from api.utils import (
    bulk_sync_to_redis,
    generate_kafka_data_update,
    get_exception_detail,
    get_ip,
    get_multiple_old_new_data,
    log_critical_error,
)
from rest_framework.exceptions import ValidationError
from api.translation import __
from backend.kafka import send_request
from backend.redis_client import redis_client


class StudentUpdate(AdminApiBase):
    def validate_request(self, request, settings):
        rules = {
            "id": "required|array|exists:api.TtStudent,id",
        }
        validator = BaseValidator(request.data, rules, {})
        error = validator.validate()
        if error:
            raise ValidationError(error)

        edit_students = (
            TtStudent.objects
            .select_related("department")
            .prefetch_related(
                "academic_term",
                "pos",
                "pos_module__module",
                "pos_module__pos_module_group__module_group",
                "pos_module__pos_module_group__pos",
                "tag__tag",
            )
            .filter(id__in=request.data.get("id"))
        )

        rules = {
            "name": "nullable",
            "desc": "nullable",
            "email": "nullable",
            "department_id": "nullable|exists:api.TtDepartment,id",
            "academic_term_ids": "nullable|array|exists:api.TtAcademicTerm,id",
            "pos_ids": "nullable|array|exists:api.TtPos,id",
            "pos_module_ids": "nullable|array|exists:api.TtPosModuleGroupModule,id",
            "tag": "nullable|array|exists:api.TtTag,id",
        }

        if request.data.get("code") and edit_students.count() == 1:
            student = edit_students.first()
            if request.data.get("code") != student.code:
                rules["code"] = "nullable|not_exists:api.TtStudent,code"
        
        validator = BaseValidator(request.data, rules, {})
        error = validator.validate()
        if error:
            raise ValidationError(error)

        pos_ids = request.data.get("pos_ids")
        pos_module_ids = request.data.get("pos_module_ids")

        if pos_module_ids:
            allowed_pos_module_ids = set(
                TtPosModuleGroupModule.objects.filter(
                    pos_module_group__pos_id__in=(pos_ids or []),
                    id__in=pos_module_ids,
                ).values_list("id", flat=True)
            )
            
            invalid_pos_module_ids = set(pos_module_ids) - allowed_pos_module_ids
            pos_module_names = ", ".join(
                TtPosModuleGroupModule.objects
                .filter(id__in=invalid_pos_module_ids)
                .values_list("module__name", flat=True)
            )
            pos_ids_name = ", ".join(TtPos.objects.filter(id__in=(pos_ids or [])).values_list("name", flat=True))
            if invalid_pos_module_ids:
                raise ValidationError({
                    "error": __('validation.module_not_in_pos', module_name=pos_module_names,pos_name=pos_ids_name),
                    "errors": {"pos_module_ids": [__('validation.module_not_in_pos', module_name=pos_module_names,pos_name=pos_ids_name)]}
                })

        return edit_students

    def post(self, request):
        try:
            settings = TtSetting.get_multiple_setting({"slot_per_week"})
            slot_per_week = int(settings["slot_per_week"])
            edit_students = self.validate_request(request, settings)
            action = "student_update"
            student_ids = request.data.get('id')
            pos_ids = request.data.get("pos_ids")
            academic_term_ids = request.data.get("academic_term_ids")
            redis_data = {"update": {"student": []}}
            affected_student_set_ids = set()# used to update redis data
            batch_size = 500
            redis_student_set_table = "student_set"
            update_fields = {}
            ignore_fields = [
                "id",
                "academic_term_ids",
                "pos_ids",
                "pos_module_ids",
                "tag",
                "timestamp",
                "signature",
            ]
            if edit_students.count() > 1 or not request.data.get("code"):
                ignore_fields.append("code")

            valid_fields = {f.column for f in TtStudent._meta.fields}
            for key, val in request.data.items():
                if key not in ignore_fields and key in valid_fields:
                    update_fields[key] = val

            old_new_data = get_multiple_old_new_data(
                edit_students,
                request.data.items(),
                ignore_fields,
                valid_fields,
            )
            if update_fields:
                update_fields["updated_by"] = request.user.id
                edit_students.update(**update_fields)

            all_relation_idata = {
                "academic_term": [],
                "pos": [],
                "pos_module": []
            }

            student_ids_need_check_pathway = set()
            # bulk update have different handle for relation
            if edit_students.count() > 1:
                for edit_student in edit_students:
                    for field in ["academic_term_ids", "pos_ids", "pos_module_ids"]:
                        if field in request.data:
                            match(field):
                                case "academic_term_ids":
                                    relation_key = "academic_term_id"
                                case "pos_ids":
                                    relation_key = "pos_id"
                                case "pos_module_ids":
                                    relation_key = "pos_module_group_module_id"
                            input_ids = request.data.get(field) or []
                            relation_name = field.removesuffix("_ids")

                            old_ids = [obj.id for obj in getattr(edit_student, relation_name).all()]

                            old_set = set(old_ids)
                            input_set = set(input_ids)
                            if not input_set.issubset(old_set):
                                new_ids = old_ids + [i_id for i_id in input_ids if i_id not in old_set]
                                old_new_data["old_data"].setdefault(edit_student.id, {})[field] = old_ids
                                old_new_data["new_data"].setdefault(edit_student.id, {})[field] = new_ids
                                insert_ids = [i_id for i_id in input_ids if i_id not in old_set]
                                if insert_ids:
                                    for insert_id in insert_ids:
                                        all_relation_idata[relation_name].append({
                                            "student_id": edit_student.id,
                                            relation_key: insert_id
                                        })

                relation_map = {
                    "academic_term": TtStudentAcademicTerm,
                    "pos": TtStudentPos,
                    "pos_module": TtStudentPosModule,
                }

                # Perform bulk insert per model
                batch_size = 500
                with transaction.atomic():
                    for insert_table_name, insert_data in all_relation_idata.items():
                        if insert_data:
                            table = relation_map[insert_table_name]
                            insert_objects = [table(**data) for data in insert_data]

                            table.objects.bulk_create(
                                insert_objects,
                                batch_size=batch_size,
                                ignore_conflicts=True
                            )

                # extra handle when got update pos_ids, need check got pos in same term or not, if got, remove old
                if "pos_ids" in request.data:
                    # specific the term FE pass in and exclude the pos_ids pass in by FE, if got pos_id found means those pos_id with student relation is need to remove
                    duplicate_term_pos_ids = list(
                        TtStudentPos.objects.filter(
                            student_id__in=student_ids,
                            pos__academic_term_id__in=academic_term_ids
                        )
                        .exclude(pos_id__in=pos_ids)
                        .values_list("pos_id", flat=True)
                        .distinct()
                    )
                    # since already exclude pos_ids pass in by FE, so result in duplicate_term_pos_ids was all need to delete
                    if duplicate_term_pos_ids:
                        with transaction.atomic():
                            # must get the id,student_id first, cause need get the deleted ids for remove from audit trail
                            deleted_student_pos_modules = list(
                                TtStudentPosModule.objects.filter(
                                    student_id__in=student_ids,
                                    pos_module_group_module__pos_module_group__pos_id__in=duplicate_term_pos_ids
                                ).values("id", "student_id")
                            )
                            if deleted_student_pos_modules:
                                deleted_module_ids = []
                                for deleted_student_pos_module in deleted_student_pos_modules:
                                    student_id, deleted_id = deleted_student_pos_module["student_id"], deleted_student_pos_module["id"]
                                    deleted_module_ids.append(deleted_id)
                                    # if got any error like missing keys, None values, or missing value, will ignore
                                    with suppress(KeyError, TypeError, ValueError):
                                        old_new_data["new_data"][student_id]["pos_module"].remove(deleted_id)
                                TtStudentPosModule.objects.filter(id__in=deleted_module_ids).delete()

                            # get the student_pos needed delete record
                            deleted_student_pos_records = list(
                                TtStudentPos.objects.filter(
                                    student_id__in=student_ids,
                                    pos_id__in=duplicate_term_pos_ids
                                ).values("id", "student_id", "pos_id")
                            )

                            if deleted_student_pos_records:
                                pos_delete_ids = []
                                # Update audit log specifically per student
                                for deleted_student_pos_record in deleted_student_pos_records:
                                    pos_delete_ids.append(deleted_student_pos_record["id"])
                                    student_id, pos_id = deleted_student_pos_record["student_id"], deleted_student_pos_record["pos_id"]
                                    with suppress(KeyError, TypeError, ValueError):
                                        old_new_data["new_data"][student_id]["pos"].remove(pos_id)

                                TtStudentPos.objects.filter(id__in=pos_delete_ids).delete()

                if "pos_module_ids" in request.data:
                    for student_id in request.data.get("id"):
                        student_ids_need_check_pathway.add(student_id)

            else:# if only 1 student, can used back previous handle
                if "academic_term_ids" in request.data:
                    new_ids = request.data.get("academic_term_ids", [])
                    old_new_data = self.update_m2m_field_bulk(edit_students, "academic_term", new_ids, old_new_data)
                    student_ids_need_check_pathway.add(student_ids[0])

                if "pos_ids" in request.data:
                    new_ids = request.data.get("pos_ids", [])
                    old_new_data = self.update_m2m_field_bulk(edit_students, "pos", new_ids, old_new_data)
                    student_ids_need_check_pathway.add(student_ids[0])

                if "pos_module_ids" in request.data:
                    new_ids = request.data.get("pos_module_ids", [])
                    old_new_data = self.update_m2m_field_bulk(edit_students, "pos_module", new_ids, old_new_data)
                    student_ids_need_check_pathway.add(student_ids[0])

            if "tag" in request.data:
                new_ids = request.data.get("tag", [])
                old_new_data = self.update_m2m_field_bulk(edit_students, "tag", new_ids, old_new_data, False, True)

            # will get all user in this list to check need add/remove any pathway or not
            student_ids_need_check_pathway = list(student_ids_need_check_pathway)
            if student_ids_need_check_pathway:
                student_pathway_idata = []
                student_pathway_delete_data = []

                students_check_pathway = TtStudent.objects.filter(
                    id__in=student_ids_need_check_pathway
                ).prefetch_related(
                    "pathway",
                    "pos_module__pos_module_group",
                    "student_set__module"
                )
                for student_check_pathway in students_check_pathway:
                    # check got which pathway need delete
                    current_pathways = list(student_check_pathway.pathway.all())
                    existing_pathway_ids = [p for p in current_pathways]

                    pos_module_group_module = list(student_check_pathway.pos_module.all())
                    student_sets = list(student_check_pathway.student_set.all())

                    student_set_module_map = {}
                    for student_set in student_sets:
                        student_set_module_map[student_set.id] = {ssm.id for ssm in student_set.module.all()}
                    # group pos with pos_module_id by pos_id as key
                    pos_module_group_module_map = {}
                    # pathway string 2 used to check need remove student from student set or not, cause student set only can check by using pathway string 2
                    # even is different pos, this pathway string 2 can include all, later will check student set pathway_string 2 in student pathway string 2 or not
                    student_module_set = set()
                    for val in pos_module_group_module:
                        student_module_set.add(val.module_id)
                        pos_module_group_module_map.setdefault(val.pos_module_group.pos_id, []).append(val.id)
                    student_module_set = sorted(student_module_set)

                    pathway_string_list = []
                    # using latest student_pos_module to get latest pathway_string, will used to check which pathway need delete/insert
                    for key, val in pos_module_group_module_map.items():
                        pathway_string = ";".join(map(str, sorted(val)))
                        pathway_string_list.append(pathway_string)

                    # current_pathways is existing student pathway, not yet added for new, so if current pathway not in pathway_string, means need to delete
                    for current_pathway in current_pathways:
                        if current_pathway.pathway_string not in pathway_string_list:
                            student_pathway_delete_data.append({
                                "student_id": student_check_pathway.id,
                                "pathway_id": current_pathway.id,
                            })

                    pathways = TtPathway.objects.filter(pathway_string__in=pathway_string_list)
                    # if pathway id not found in existing_pathway_ids, means is need insert
                    for pathway in pathways:
                        if pathway.id not in existing_pathway_ids:
                            student_pathway_idata.append({
                                "student_id": student_check_pathway.id,
                                "pathway_id": pathway.id,
                            })

                    student_set_student_delete_data = []
                    # check student module id set got include each student set module id or not
                    for student_set_id, module_ids_set in student_set_module_map.items():
                        # if student set module not fully include student module id set, need remove student from the student set
                        if not module_ids_set.issubset(student_module_set):
                            student_set_student_delete_data.append({
                                "student_id": student_check_pathway.id,
                                "student_set_id": student_set_id,
                            })

                    if student_pathway_idata:
                        student_pathway_objects = [TtStudentPathway(**data) for data in student_pathway_idata]
                        with transaction.atomic():
                            for i in range(0, len(student_pathway_objects), batch_size):
                                TtStudentPathway.objects.bulk_create(
                                    student_pathway_objects[i:i + batch_size],
                                    batch_size=batch_size,
                                    ignore_conflicts=True
                                )
                    if student_pathway_delete_data:
                        student_pathway_delete_query = Q()
                        for pair in student_pathway_delete_data:
                            student_pathway_delete_query |= Q(
                                student_id=pair["student_id"],
                                pathway_id=pair["pathway_id"]
                            )
                        if student_pathway_delete_query:
                            TtStudentPathway.objects.filter(student_pathway_delete_query).delete()

                    if student_set_student_delete_data:
                        student_set_student_delete_query = Q()
                        recalculate_student_ids = set()
                        student_set_ids = set()
                        for pair in student_set_student_delete_data:
                            recalculate_student_ids.add(pair["student_id"])
                            student_set_ids.add(pair["student_set_id"])
                            affected_student_set_ids.add(pair["student_set_id"])
                            student_set_student_delete_query |= Q(
                                student_id=pair["student_id"],
                                student_set_id=pair["student_set_id"]
                            )
                        if student_set_student_delete_query:
                            TtStudentSetStudent.objects.filter(student_set_student_delete_query).delete()
                            # when got remove relation between student set and student, only need recalculate the resource map for related student
                            # need get all student set having which week, then student will follow the week id to update resource map also
                            recalculate_week_ids = TtStudentSet.objects.filter(
                                id__in=student_set_ids
                            ).values_list("academic_term__week__id", flat=True).distinct()
                            if recalculate_week_ids and recalculate_student_ids:
                                redis_data = helper_recalculate_and_get_redis_resource_map(
                                    redis_data=redis_data,
                                    student_ids=list(recalculate_student_ids),
                                    week_ids=list(recalculate_week_ids),
                                    slot_per_week=slot_per_week,
                                )
            if "academic_term_ids" in request.data and student_ids:
                affected_student_ids = list(student_ids)
                default_resource_map_pattern = "0" * int(settings["slot_per_week"])

                existing_weeks = defaultdict(set)
                for student_id, week_id in (
                    TtStudentResourceMap.objects
                    .filter(student_id__in=affected_student_ids)
                    .values_list("student_id", "week_id")
                ):
                    existing_weeks[student_id].add(week_id)

                term_rows = list(
                    TtStudentAcademicTerm.objects
                    .filter(student_id__in=affected_student_ids)
                    .values_list("student_id", "academic_term_id")
                )
                relevant_term_ids = {academic_term_id for _, academic_term_id in term_rows}
                term_week_map = defaultdict(set)
                for academic_term_id, week_id in (
                    TtWeek.objects
                    .filter(ttacademicterm__id__in=relevant_term_ids)
                    .values_list("ttacademicterm__id", "id")
                ):
                    term_week_map[academic_term_id].add(week_id)

                resource_map_rows = []
                for student_id, _ in term_rows:
                    desired_week_ids = set()
                    for _, academic_term_id in [row for row in term_rows if row[0] == student_id]:
                        desired_week_ids.update(term_week_map.get(academic_term_id, set()))

                    missing_week_ids = sorted(desired_week_ids - existing_weeks.get(student_id, set()))
                    for week_id in missing_week_ids:
                        resource_map_rows.append(
                            TtStudentResourceMap(
                                student_id=student_id,
                                week_id=week_id,
                                pattern=default_resource_map_pattern,
                            )
                        )

                if resource_map_rows:
                    TtStudentResourceMap.objects.bulk_create(
                        resource_map_rows,
                        batch_size=500,
                        ignore_conflicts=True,
                    )

            # Read the final database state after all relation and pathway changes.
            # Do not reuse edit_students: its prefetched relations may be stale.
            redis_students = (
                TtStudent.objects.filter(id__in=student_ids)
                .prefetch_related(
                    "academic_term",
                    "pos",
                    "pos_module__pos_module_group",
                    "student_set",
                    "ttstudentresourcemap_set__week",
                )
            )
            for student in redis_students:
                resource_map = {
                    int(resource.week.week) - 1: resource.pattern
                    for resource in student.ttstudentresourcemap_set.all()
                }
                redis_data["update"]["student"].append({
                    "id": student.id,
                    "code": student.code,
                    "name": student.name,
                    "email": student.email,
                    "department_id": student.department_id,
                    "academic_term_ids": sorted(term.id for term in student.academic_term.all()),
                    "pos_ids": sorted(pos.id for pos in student.pos.all()),
                    "pos_module": [
                        {"module_id": module.module_id, "pos_id": module.pos_module_group.pos_id}
                        for module in sorted(student.pos_module.all(), key=lambda module: module.id)
                    ],
                    "student_set_ids": sorted(student_set.id for student_set in student.student_set.all()),
                    "resource_map": dict(sorted(resource_map.items())),
                })
            # update student set redis for student_ids
            if affected_student_set_ids:
                student_sets = TtStudentSet.objects.filter(id__in=affected_student_set_ids).prefetch_related("student")
                for student_set in student_sets:
                    ss_student_ids = list({s.id for s in student_set.student.all()})
                    redis_data["update"].setdefault(redis_student_set_table, []).append({"id":student_set.id,"student_ids": ss_student_ids,})

            bulk_sync_to_redis(redis_client, redis_data)

            # kafka push
            if old_new_data["old_data"] or old_new_data["new_data"]:
                method = action
                kafka_topic = self.kafka_config["MICROSERVICES_TT_TOPIC"]
                kafka_students = []
                if kafka_topic:
                    new_data = old_new_data.get("new_data")
                    relation_fields = ["tag"]
                    relation_data = generate_kafka_data_update(new_data, relation_fields)
                    update_fields.pop("updated_by", None)

                    all_pos_module_ids = set()
                    # get all pos_module then using 1 query to get all needed info, to avoid have N+1 query in loop
                    for student in edit_students:
                        student_pos_modules = old_new_data.get("new_data", {}).get(student.id, {}).get("pos_module") or []
                        all_pos_module_ids.update(student_pos_modules)

                    pos_module_lookup = {}
                    if all_pos_module_ids:
                        records = (
                            TtPosModuleGroupModule.objects
                            .filter(id__in=all_pos_module_ids)
                            .select_related("pos_module_group", "module")
                        )
                        pos_module_lookup = {pm.id: pm for pm in records}

                    for student in edit_students:
                        data = {"id": student.id, **update_fields, **relation_data}

                        student_new_data = old_new_data.get("new_data", {}).get(student.id, {})
                        if student_new_data.get("academic_term"):
                            data["academic_term"] = student_new_data["academic_term"]
                        if student_new_data.get("pos"):
                            data["pos"] = student_new_data["pos"]

                        pos_module_ids = student_new_data.get("pos_module")
                        if pos_module_ids:
                            data["pos_module"] = []
                            for pm_id in pos_module_ids:
                                pos_module = pos_module_lookup.get(pm_id)
                                if pos_module:
                                    data["pos_module"].append({
                                        "pos_module_group_module_id": pos_module.id,
                                        "pos_id": pos_module.pos_module_group.pos_id,
                                        "module_group_id": pos_module.pos_module_group.module_group_id,
                                        "module_id": pos_module.module.id,
                                    })
                        
                        kafka_students.append(data)
                    kafka_request_data = {
                        "session_id": request.user.name,
                        "student": kafka_students,
                    }
                    send_request(kafka_topic, kafka_request_data, None, method)

                audit_trail = AuditTrail.objects.create(
                    user_id=request.user.id,
                    type=self.audit_type,
                    ip_address=get_ip(request),
                )
                name = ",".join(edit_students.values_list("name", flat=True))
                remark_param = {"name": name}
                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)
