import json
import math

from django.db.models import Count, Q
from rest_framework import status
from rest_framework.exceptions import ValidationError
from django.db import connection
from api.models import TtModule, TtActivityTemplate, TtActivity, TtPos, TtSetting
from api.models.tt_location import TtLocation
from api.models.tt_staff import TtStaff
from api.models.tt_student_set import 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
from api.views.medical.base import MedicalApiBase


class ScheduleListActivity(MedicalApiBase):
    def validate_request(self,request):
        rules = {
            "academic_term_id": "required|exists:api.TtAcademicTerm,id",
            "module_id": "nullable",# only can filter by module, if no pass in means is get all module activity
        }
        # 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)
            setting_params = {
                "violation_key",
            }
            settings = TtSetting.get_multiple_setting(setting_params)
            violation_key = json.loads(settings["violation_key"])
            # data = []
            data = self.get_data(request)

            response = {
                "data": data,
                "violation_key":violation_key
            }
            # log_critical_error(user_id=None,descr=len(connection.queries),url="testing")

            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_matching_jta_parent_ids(self, academic_term_id, filter_q):
        return TtActivity.objects.filter(
            academic_term_id=academic_term_id,
            is_jta=1,
            jta_parent_id__isnull=False,
        ).filter(filter_q).values_list("jta_parent_id", flat=True)

    def get_jta_variant_ids(self, academic_term_id, variant_parent_ids):
        return TtActivity.objects.filter(
            academic_term_id=academic_term_id,
            is_variant=1,
            variant_parent_id__in=variant_parent_ids,
        ).values_list("id", flat=True)


    def get_data(self, request):
        activities = (
            TtActivity.objects
            .filter(academic_term_id=request.data.get('academic_term_id'))
            .exclude(is_jta=1, jta_parent_id__isnull=False)
            .annotate(# using for do filter
                week_pattern_week_count=Count("week_pattern__week", distinct=True),
                week_count=Count("week", distinct=True),
            )
            .filter(# only 1 week activity will count as medical activity
                (Q(week_pattern_id__isnull=False) & Q(week_pattern_week_count=1))
                |
                Q(week_count=1)
            )
            .select_related(
                "activity_template", "department", "academic_term", "module",
                "week_pattern", "activity_type", "zone", "availability",
                "start_preference", "usage_preference"
            )
            .prefetch_related(
                "week", "week_pattern__week", "tag__tag", "staff", "staff_preset",
                "staff_suitability", "location", "location_preset",
                "location_suitability", "student_set", "jta_children__module"
            )
            .order_by(self.list_default_order_column)
        )
        module_id = request.data.get("module_id")
        academic_term_id = request.data.get("academic_term_id")

        filter_pattern = Q()
        if module_id:
            filter_pattern = Q(module_id=module_id)
        matching_jta_parent_ids = self.get_matching_jta_parent_ids(academic_term_id, filter_pattern)
        # variant parent id in matching_jta_parent_ids also get it to display
        jta_variant_ids = self.get_jta_variant_ids(academic_term_id,matching_jta_parent_ids) if matching_jta_parent_ids else []

        activities = activities.filter(
            filter_pattern |
            Q(id__in=matching_jta_parent_ids) |
            Q(id__in=jta_variant_ids)
        ).distinct()

        return_data = []
        for val in activities:
            activity_data = self.get_activity_data(val)
            return_data.append(activity_data)
        return list(return_data)

    def get_activity_data(self,val):
        # Basic Data
        data = {
            "id": val.id,
            "code": val.code,
            "name": val.name,
            "desc": val.desc,
            "activity_template_name": val.activity_template.name if val.activity_template else None,
            "activity_template_id": val.activity_template_id,
            "department": val.department.name if val.department else None,
            "department_id": val.department_id,
            "sequence_number": val.sequence_number,
            "academic_term": val.academic_term.name if val.academic_term else None,
            "academic_term_id": val.academic_term_id,
            "activity_type": val.activity_type.name if val.activity_type else None,
            "activity_type_color": val.activity_type.color if val.activity_type else None,
            "activity_type_id": val.activity_type_id,
            "zone": val.zone.name if val.zone else None,
            "zone_id": val.zone_id,
            "module_id": val.module_id,
            "duration": val.duration,
            "slot_required": val.slot_required,
            "is_jta": val.is_jta,
            "jta_parent_id": val.jta_parent_id,
            "is_variant": val.is_variant,
            "variant_parent_id": val.variant_parent_id,
            "scheduled_start_time": val.scheduled_start_time,
            "scheduled_day": __("attr.days_name." + str(val.scheduled_day)) if val.scheduled_day is not None else None,
            "suggested_day": __("attr.days_name." + str(val.suggested_day)) if val.suggested_day is not None else None,
            "scheduled_day_id": val.scheduled_day,
            "suggested_day_id": val.suggested_day,
            "suggested_time": val.suggested_time,
            "suggested_time_slot": val.suggested_time_slot,
            "staff_requirement_type": __("attr.staff_requirement_type_name." + str(val.staff_requirement_type)),
            "staff_requirement": val.staff_requirement,
            "location_requirement_type": __("attr.location_requirement_type_name." + str(val.location_requirement_type)),
            "location_requirement": val.location_requirement,
            "planned_size": val.planned_size,
            "real_size": sum((student.planned_size or 0) for student in val.student_set.all()),
            "module_size": val.module.planned_size if val.module else 0,
            "scheduled": val.scheduled,
        }

        if val.jta_children.all():
            data["module_name"] = ",".join(child.module.name for child in val.jta_children.all() if child.module)
            data["module_code"] = ",".join(child.module.code for child in val.jta_children.all() if child.module)
        else:
            data["module_name"] = val.module.name if val.module else None
            data["module_code"] = val.module.code if val.module else None


        # Week Pattern Logic
        # weeks = val.week_pattern.week.all() if val.week_pattern else val.week.all()
        # data["week_pattern"] = [{"id": w.id, "week": w.week, "start_date": w.start_date} for w in weeks]
        # data["week_pattern_id"] = val.week_pattern.id if val.week_pattern else None
        # data["week_pattern_name"] = val.week_pattern.name if val.week_pattern else "-"

        # Preference Patterns
        for pref in ['availability', 'start_preference', 'usage_preference']:
            obj = getattr(val, pref)
            data[f"{pref}_id"] = obj.id if obj else None
            data[f"{pref}_name"] = obj.name if obj else "-"
            data[f"{pref}_pattern"] = obj.pattern if obj else getattr(val, f"{pref}_pattern")

        # Many-to-Many Relationships
        m2m_fields = ['staff', 'staff_preset', 'staff_suitability', 'location', 'location_preset',
                      'location_suitability', 'student_set']
        for field in m2m_fields:
            if field in ["student_set"]:
                data[field] = [{"id": o.id, "name": o.name, "code": o.code, "planned_size": o.planned_size} for o in getattr(val, field).all()]
            else:
                data[field] = [{"id": o.id, "name": o.name, "code": o.code} for o in getattr(val, field).all()]

        # Tags (Special case)
        data["tag"] = [{"id": t.tag.id, "name": t.tag.name} for t in val.tag.all()]

        # Extra Fields
        for i in range(1, 11):
            field_name = f"extra_data_{i}"
            data[field_name] = getattr(val, field_name, None)

        return data