from django.contrib.postgres.aggregates import ArrayAgg, StringAgg
from django.db.models import Sum, Count, Q
from rest_framework import status
from rest_framework.exceptions import ValidationError

from api.models import TtStaff, TtTagRelation, TtLocation, TtModule, TtPos, TtPathway, TtStudentSet
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

class PosShowPathway(AdminApiBase):
    def validate_request(self,request):
        rules = {
            "pos_id": "required|exists:api.TtPos,id",
        }
        # 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)

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

            data = self.get_data(request)
            response = {
                "data": data
            }
            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 get_data(self,request):
        pathway = (
            TtPathway.objects
            .filter(pos_id=request.data.get("pos_id"))
            .select_related("pos")
            .annotate(
                module_name=ArrayAgg(
                    "ttpathwayposmodulegroupmodule__pos_module_group_module__module__name",
                    order_by="ttpathwayposmodulegroupmodule__pos_module_group_module__module__name"
                ),
                module_ids=ArrayAgg(
                    "ttpathwayposmodulegroupmodule__pos_module_group_module__module__id",
                    distinct=True,
                    order_by="ttpathwayposmodulegroupmodule__pos_module_group_module__module__id"
                ),
            )
            .order_by("pos__name")
        )

        student_set = (
            TtStudentSet.objects.filter(pos_id=request.data.get("pos_id"))
            .prefetch_related("module")
        )
        # having a set key by pathway_string and value is the student_set planned size, so in pathway_part, will used pathway_string_2 to find
        student_set_size = {}
        for s in student_set:
            # generate pathway string as the size key
            module_ids = sorted(s.module.values_list("id", flat=True))
            student_set_pathway_string =";".join(str(mid) for mid in module_ids)
            student_set_size[student_set_pathway_string] = student_set_size.get(student_set_pathway_string, 0) + s.planned_size

        return_data=[]

        for val in pathway:
            data={}
            data["id"] = val.id
            data['pos_name'] = val.pos.name if val.pos else None
            data['module'] = val.module_name
            data['planned_size'] = val.planned_size
            # if pathway_string_2 is key in student_set_size, means this set of data is pathway real_size
            data['real_size'] = student_set_size.get(val.pathway_string_2,0)


            # module_ids = [int(x) for x in val.pathway_string_2.split(";")]
            # real_size = TtStudentSet.objects.filter(ttstudentsetactivity__activity__module_id__in=val.module_ids).sum("planned_size")
            # data['real_size'] = real_size

            # module_ids = val.module_ids
            # module_count = len(module_ids)
            # # got module only calculate real size
            # if module_count > 0:
            #     # student set activity the activity module_id must exact match the module_ids and the activity count must match with module count only can count as real size
            #     real_size = (
            #         TtStudentSet.objects
            #         .filter(ttstudentsetactivity__activity__module_id__in=module_ids)
            #         .annotate(
            #             matched_count=Count(
            #                 "ttstudentsetactivity__activity__module_id",
            #                 distinct=True,
            #                 filter=Q(ttstudentsetactivity__activity__module_id__in=module_ids)
            #             )
            #         )
            #         .filter(matched_count=module_count)
            #         .aggregate(total=Sum("planned_size"))["total"]
            #     )
            #     data['real_size'] = real_size if real_size else 0

            return_data.append(data)

        return return_data