from collections import defaultdict

from django.contrib.postgres.aggregates import ArrayAgg
from django.db.models import F
from rest_framework import serializers, status

from api.models import TtStaffSuitability, TtDepartment, TtZone, TtAvailability, TtStartPreference, TtUsagePreference, \
    TtSuitability, TtTag, TtSetting, TtModule, TtWeekPattern, TtWeek, TtAcademicTerm, TtPos, TtPosModuleGroup, \
    TtPosModuleGroupModule, TtPathway, TtPathwayPosModuleGroupModule, TtStudentSet
from api.views.admin.base import AdminApiBase
from api.validator import BaseValidator
from api.utils import decrypt_aes_128_cbc, encrypt_aes_128_cbc, log_critical_error, get_exception_detail, \
    calculate_preset_pattern_time
from rest_framework.exceptions import ValidationError
from api.translation import __


class AllocationRules(AdminApiBase):
    def validate_request(self,request):
        rules = {
            "academic_term_id": "required|exists:api.TtAcademicTerm,id",
            "pos_ids": "nullable|array|exists:api.TtPos,id",
            "allocation_by": "required",
        }
        # if need rename field, can put here
        attribute = {
            # "email": __("attr.email"),
        }
        # validate id first, cause need get pos by id, if wrong id direct return error
        validator = BaseValidator(request.data,rules,attribute)
        error = validator.validate()
        if error:
            raise ValidationError(error)
        # custom error after basic validation

        if error:
            raise ValidationError(error)

    def post(self, request):
        # Validate input
        try:
            self.api_log_skip_outgoing_data=True
            self.validate_request(request)

            academic_term_id = request.data["academic_term_id"]
            input_pos_ids = request.data.get("pos_ids")
            allocation_by = request.data["allocation_by"]

            pathway_data = (
                TtPathway.objects
                .filter(
                    academic_term_id=academic_term_id,
                    status=TtPathway.STATUS_TO_CODE['active']
                )
            )
            if input_pos_ids:
                pathway_data = pathway_data.filter(pos_id__in=input_pos_ids)

            pathway_data = (
                pathway_data.annotate(
                    module=ArrayAgg(
                        "ttpathwayposmodulegroupmodule__pos_module_group_module__module__name",
                        distinct=True
                    ),
                    module_ids=ArrayAgg(
                        "ttpathwayposmodulegroupmodule__pos_module_group_module__module_id",
                        distinct=True
                    ),
                    pos_name=F("pos__name"),
                )
                .values("id", "pos_id", "pos_name", "planned_size", "module_ids", "module")
            )

            module_ids = {module_id for row in pathway_data for module_id in (row["module_ids"] or [])}
            pos_ids = list(pathway_data.values_list("pos_id", flat=True))
            # calculate real size for the pathway
            student_sets = {
                s["id"]: s for s in (
                    TtStudentSet.objects
                    .filter(pos_id__in=pos_ids, academic_term_id=academic_term_id)
                    .annotate(
                        module_ids=ArrayAgg("module__id", distinct=True)
                    )
                    .values("id", "planned_size", "module_ids","pos_id")
                )
            }

            for p in pathway_data:
                total_size = sum(
                    s["planned_size"]
                    for s in student_sets.values()
                    if s["pos_id"] == p["pos_id"] # must include pos_id, cause pathway_string_2 is module id, possible will same as other pathway, if no pos, will have calculation problem when select multiple pos
                    and set(s["module_ids"] or []) == set(p["module_ids"] or [])
                )
                p['real_size'] = total_size

            rules = {}
            rules["allocation_method"] = {
                "spread": __("attr.allocation_method_name.spread"),
                "clump": __("attr.allocation_method_name.clump"),
            }
            rules["module"] = list(TtModule.objects.filter(id__in=module_ids).values("id","name"))

            response = {
                "list": pathway_data,
                "rules": rules
            }

            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)