from django.db.models import (
    Sum, Case, When, IntegerField, FloatField, F
)

# ======================================================
# WEIGHTS (OFFICIAL NECTA STYLE)
# ======================================================

SUB_WEIGHT_CASE = Case(
    When(grade="A", then=1),
    When(grade="B", then=2),
    When(grade="C", then=3),
    When(grade="D", then=4),
    When(grade="F", then=5),
    output_field=IntegerField(),
)

DIV_WEIGHT_CASE = Case(
    When(division="I", then=1),
    When(division="II", then=2),
    When(division="III", then=3),
    When(division="IV", then=4),
    When(division="0", then=5),
    output_field=IntegerField(),
)

SAT_DIVISIONS = ["I", "II", "III", "IV", "0"]

# ======================================================
# SUBJECT GPA (GENERIC – ANY LEVEL)
# ======================================================

def subject_gpa_queryset(base_qs):
    """
    Returns queryset with:
      - subject_code
      - weighted_sum
      - total_sum
      - gpa
    """
    return (
        base_qs
        .annotate(weight=SUB_WEIGHT_CASE)
        .values("subject_code")
        .annotate(
            weighted_sum=Sum(
                F("weight") * F("total"),
                output_field=FloatField()
            ),
            total_sum=Sum("total"),
            gpa=F("weighted_sum") / F("total_sum"),
        )
    )


# ======================================================
# DIVISION GPA (GENERIC – ANY LEVEL)
# ======================================================

def division_gpa_value(base_qs):
    """
    Returns a single GPA value (float)
    Excludes ABSENT (X)
    """
    agg = (
        base_qs
        .filter(division__in=SAT_DIVISIONS)
        .annotate(weight=DIV_WEIGHT_CASE)
        .aggregate(
            weighted_sum=Sum(
                F("weight") * F("total"),
                output_field=FloatField()
            ),
            total_sum=Sum("total"),
        )
    )

    if not agg["total_sum"]:
        return 0

    return round(agg["weighted_sum"] / agg["total_sum"], 9)


# ======================================================
# REGION GPA HELPERS
# ======================================================

def region_subject_gpa(year, exam_type_id, region_id):
    from .models import SchoolSubjectGradeSummary

    qs = SchoolSubjectGradeSummary.objects.filter(
        year=year,
        exam_type_id=exam_type_id,
        school__region_id=region_id,
    )
    return subject_gpa_queryset(qs)


def region_division_gpa(year, exam_type_id, region_id):
    from .models import SchoolDivisionSummary

    qs = SchoolDivisionSummary.objects.filter(
        year=year,
        exam_type_id=exam_type_id,
        school__region_id=region_id,
    )
    return division_gpa_value(qs)


# ======================================================
# DISTRICT GPA HELPERS
# ======================================================

def district_subject_gpa(year, exam_type_id, district_id):
    from .models import SchoolSubjectGradeSummary

    qs = SchoolSubjectGradeSummary.objects.filter(
        year=year,
        exam_type_id=exam_type_id,
        school__district_id=district_id,
    )
    return subject_gpa_queryset(qs)


def district_division_gpa(year, exam_type_id, district_id):
    from .models import SchoolDivisionSummary

    qs = SchoolDivisionSummary.objects.filter(
        year=year,
        exam_type_id=exam_type_id,
        school__district_id=district_id,
    )
    return division_gpa_value(qs)


# ======================================================
# WARD GPA HELPERS
# ======================================================

def ward_subject_gpa(year, exam_type_id, ward_id):
    from .models import SchoolSubjectGradeSummary

    qs = SchoolSubjectGradeSummary.objects.filter(
        year=year,
        exam_type_id=exam_type_id,
        school__ward_id=ward_id,
    )
    return subject_gpa_queryset(qs)


def ward_division_gpa(year, exam_type_id, ward_id):
    from .models import SchoolDivisionSummary

    qs = SchoolDivisionSummary.objects.filter(
        year=year,
        exam_type_id=exam_type_id,
        school__ward_id=ward_id,
    )
    return division_gpa_value(qs)


# ======================================================
# SCHOOL GPA HELPERS
# ======================================================

def school_subject_gpa(year, exam_type_id, school_id):
    from .models import SchoolSubjectGradeSummary

    qs = SchoolSubjectGradeSummary.objects.filter(
        year=year,
        exam_type_id=exam_type_id,
        school_id=school_id,
    )
    return subject_gpa_queryset(qs)


def school_division_gpa(year, exam_type_id, school_id):
    from .models import SchoolDivisionSummary

    qs = SchoolDivisionSummary.objects.filter(
        year=year,
        exam_type_id=exam_type_id,
        school_id=school_id,
    )
    return division_gpa_value(qs)





def region_subject_gpa_value(year, exam_type_id, region_id):
    """
    Computes NECTA-correct SUBJECT GPA for a REGION.
    This is a GLOBAL weighted GPA across ALL subjects,
    NOT an average of subject GPAs.
    """

    from django.db.models import Sum, Case, When, IntegerField, FloatField, F
    from .models import SchoolSubjectGradeSummary

    weight_case = Case(
        When(grade="A", then=1),
        When(grade="B", then=2),
        When(grade="C", then=3),
        When(grade="D", then=4),
        When(grade="F", then=5),
        output_field=IntegerField(),
    )

    agg = (
        SchoolSubjectGradeSummary.objects
        .filter(
            year=year,
            exam_type_id=exam_type_id,
            school__region_id=region_id
        )
        .annotate(weight=weight_case)
        .aggregate(
            weighted_sum=Sum(
                F("weight") * F("total"),
                output_field=FloatField()
            ),
            total_sum=Sum("total"),
        )
    )


    qs = SchoolSubjectGradeSummary.objects.filter(
        year=year,
        exam_type_id=exam_type_id,
        school__region_id=region_id
    )

    # for g in ["A", "B", "C", "D", "F"]:
    #     total = qs.filter(grade=g).aggregate(t=Sum("total"))["t"] or 0
    #     print(f"Grade {g}: {total}")

    # grand_total = qs.aggregate(t=Sum("total"))["t"] or 0
    # print("TOTAL SUBJECT ENTRIES:", grand_total)


    if not agg["total_sum"]:
        return 0

    return round(agg["weighted_sum"] / agg["total_sum"], 9)




def district_subject_gpa_value(year, exam_type_id, district_id):
    """
    Computes NECTA-correct SUBJECT GPA for a DISTRICT.
    Global weighted GPA across ALL subjects.
    """
    from django.db.models import Sum, Case, When, IntegerField, FloatField, F
    from .models import SchoolSubjectGradeSummary

    weight_case = Case(
        When(grade="A", then=1),
        When(grade="B", then=2),
        When(grade="C", then=3),
        When(grade="D", then=4),
        When(grade="F", then=5),
        output_field=IntegerField(),
    )

    agg = (
        SchoolSubjectGradeSummary.objects
        .filter(
            year=year,
            exam_type_id=exam_type_id,
            school__district_id=district_id
        )
        .annotate(weight=weight_case)
        .aggregate(
            weighted_sum=Sum(
                F("weight") * F("total"),
                output_field=FloatField()
            ),
            total_sum=Sum("total"),
        )
    )

    if not agg["total_sum"]:
        return 0

    return round(agg["weighted_sum"] / agg["total_sum"], 9)




def ward_division_gpa(year, exam_type_id, ward_id):
    from django.db.models import Sum
    from .models import SchoolDivisionSummary

    DIV_WEIGHT = {"I":1,"II":2,"III":3,"IV":4,"0":5}

    qs = SchoolDivisionSummary.objects.filter(
        year=year,
        exam_type_id=exam_type_id,
        school__ward_id=ward_id,
        division__in=DIV_WEIGHT.keys()
    )

    ws = sum(DIV_WEIGHT[r.division] * r.total for r in qs)
    ts = sum(r.total for r in qs)

    return round(ws / ts, 9) if ts else 0


def ward_subject_gpa_value(year, exam_type_id, ward_id):
    from django.db.models import Sum, Case, When, IntegerField, FloatField, F
    from .models import SchoolSubjectGradeSummary

    weight = Case(
        When(grade="A", then=1),
        When(grade="B", then=2),
        When(grade="C", then=3),
        When(grade="D", then=4),
        When(grade="F", then=5),
        output_field=IntegerField(),
    )

    agg = (
        SchoolSubjectGradeSummary.objects
        .filter(
            year=year,
            exam_type_id=exam_type_id,
            school__ward_id=ward_id
        )
        .annotate(w=weight)
        .aggregate(
            ws=Sum(F("w") * F("total"), output_field=FloatField()),
            ts=Sum("total")
        )
    )

    return round(agg["ws"] / agg["ts"], 9) if agg["ts"] else 0
