from collections import defaultdict, deque

from rest_framework.exceptions import ValidationError

from api.models import EsExamActivity
from api.translation import __


def validate_relation_exam_activity_ids(data):
    """exam_activity_id may be empty or contain any number of distinct IDs."""
    return


def related_group_ids_for_requirement(relation_model, exam_requirement_id):
    return set(
        relation_model.objects.filter(exam_requirement_id=exam_requirement_id).values_list(
            "group_id", flat=True
        ).distinct()
    )


def related_exam_activity_ids_for_group_ids(relation_model, group_ids):
    if not group_ids:
        return set()
    return set(
        relation_model.objects.filter(group_id__in=group_ids).values_list(
            "exam_activity_id", flat=True
        )
    )


def owned_exam_activity_ids_for_requirement(exam_requirement_id):
    return set(
        EsExamActivity.objects.filter(exam_requirement_id=exam_requirement_id).values_list(
            "id", flat=True
        )
    )


def owned_exam_activity_ids_by_requirements(exam_requirement_ids):
    if not exam_requirement_ids:
        return set()
    return set(
        EsExamActivity.objects.filter(exam_requirement_id__in=exam_requirement_ids).values_list(
            "id", flat=True
        )
    )


def target_exam_activity_ids_for_relation_update(
    relation_model, exam_requirement_id, payload_activity_ids
):
    """
    Relation rows to write for one requirement.

    Linking via a foreign activity (not owned by this requirement) merges the payload
    with this requirement's owned activities already in the group, so e.g. req 48 with
    [61, 63] keeps owned activity 62 that was already grouped with 61.
    """
    payload = set(payload_activity_ids)
    owned_by_editor = owned_exam_activity_ids_for_requirement(exam_requirement_id)
    existing = related_exam_activity_ids_for_requirement(relation_model, exam_requirement_id)

    if payload - owned_by_editor:
        return payload | (owned_by_editor & existing)
    return payload


def direct_linked_exam_activity_ids_for_requirement(relation_model, exam_requirement_id):
    return set(
        relation_model.objects.filter(exam_requirement_id=exam_requirement_id).values_list(
            "exam_activity_id", flat=True
        )
    )


def direct_target_requirement_ids_for_requirement(relation_model, exam_requirement_id):
    return set(
        relation_model.objects.filter(exam_requirement_id=exam_requirement_id).values_list(
            "exam_activity__exam_requirement_id", flat=True
        ).distinct()
    )


def preceding_requirement_reachable_ids(relation_model, exam_requirement_id):
    relations = list(
        relation_model.objects.values_list(
            "exam_requirement_id",
            "exam_activity__exam_requirement_id",
        )
    )

    adjacency = defaultdict(set)
    for source_requirement_id, target_requirement_id in relations:
        adjacency[source_requirement_id].add(target_requirement_id)

    if exam_requirement_id not in adjacency:
        return set()

    visited = set()
    reachable_ids = set()
    queue = deque([exam_requirement_id])

    while queue:
        current_requirement_id = queue.popleft()
        if current_requirement_id in visited:
            continue

        visited.add(current_requirement_id)
        next_requirement_ids = adjacency[current_requirement_id] - visited
        reachable_ids.update(next_requirement_ids)
        queue.extend(next_requirement_ids)

    return reachable_ids


def validate_preceding_examination_relation_update(
    relation_model,
    source_requirement_ids,
    payload_exam_activity_ids,
):
    source_requirement_ids = set(source_requirement_ids or [])
    payload_exam_activity_ids = list(dict.fromkeys(payload_exam_activity_ids or []))

    if not source_requirement_ids or not payload_exam_activity_ids:
        return

    source_requirement_activity_ids = set(
        EsExamActivity.objects.filter(exam_requirement_id__in=source_requirement_ids).values_list(
            "id", flat=True
        )
    )
    own_activity_ids = set(payload_exam_activity_ids) & source_requirement_activity_ids
    if own_activity_ids:
        message = __("validation.preceding_examination_invalid")
        raise ValidationError({"error": message, "errors": {"exam_activity_id": message}})

    selected_requirement_ids = set(
        EsExamActivity.objects.filter(id__in=payload_exam_activity_ids).values_list(
            "exam_requirement_id", flat=True
        )
    ) - source_requirement_ids

    if not selected_requirement_ids:
        return

    direct_target_requirement_ids_by_source = {
        source_requirement_id: direct_target_requirement_ids_for_requirement(
            relation_model,
            source_requirement_id,
        )
        for source_requirement_id in source_requirement_ids
    }

    reachable_requirement_ids_by_target = {
        target_requirement_id: preceding_requirement_reachable_ids(
            relation_model,
            target_requirement_id,
        )
        for target_requirement_id in selected_requirement_ids
    }

    for source_requirement_id in source_requirement_ids:
        current_targets = direct_target_requirement_ids_by_source[source_requirement_id]
        for target_requirement_id in selected_requirement_ids:
            if target_requirement_id in current_targets:
                continue

            if source_requirement_id in reachable_requirement_ids_by_target[target_requirement_id]:
                message = __("validation.preceding_examination_invalid")
                raise ValidationError({"error": message, "errors": {"exam_activity_id": message}})


def group_ids_for_exam_activity_ids(relation_model, exam_activity_ids):
    if not exam_activity_ids:
        return set()
    return set(
        relation_model.objects.filter(exam_activity_id__in=exam_activity_ids).values_list(
            "group_id", flat=True
        ).distinct()
    )


def expanded_exam_activity_ids_for_requirement_update(exam_requirement_id, payload_exam_activity_ids):
    payload_exam_activity_ids = set(payload_exam_activity_ids or [])
    selected_requirement_ids = set(
        EsExamActivity.objects.filter(id__in=payload_exam_activity_ids).values_list(
            "exam_requirement_id", flat=True
        )
    )
    target_requirement_ids = selected_requirement_ids | {exam_requirement_id}
    return set(
        EsExamActivity.objects.filter(exam_requirement_id__in=target_requirement_ids).values_list(
            "id", flat=True
        )
    )


def expanded_exam_activity_ids_for_requirement_selection(exam_requirement_id, payload_exam_activity_ids):
    payload_exam_activity_ids = set(payload_exam_activity_ids or [])
    selected_requirement_ids = set(
        EsExamActivity.objects.filter(id__in=payload_exam_activity_ids)
        .exclude(exam_requirement_id=exam_requirement_id)
        .values_list("exam_requirement_id", flat=True)
    )
    if not selected_requirement_ids:
        return set()
    return set(
        EsExamActivity.objects.filter(exam_requirement_id__in=selected_requirement_ids).values_list(
            "id", flat=True
        )
    )


def removed_exam_activity_ids_for_relation_update(
    relation_model,
    exam_requirement_id,
    payload_activity_ids,
    component_requirement_ids,
):
    """
    Activities to drop from the component after an update.

    Replace mode (payload is only this requirement's owned activities):
    - Drop this requirement's own links not in the payload (e.g. req 49 [63, 64] drops 61).
    - Drop peer requirements' owned activities absent from the payload (e.g. req 48 [61, 62]
      unlinks 63 and 64) without touching the peer's other owned activities (62 stays on 48
      when req 49 returns to [63, 64]).

    Merge mode (payload includes a foreign activity): only drop non-owned activities in the
    merged view that are not in the payload.
    """
    payload = set(payload_activity_ids)
    existing = related_exam_activity_ids_for_requirement(relation_model, exam_requirement_id)
    direct_linked = direct_linked_exam_activity_ids_for_requirement(
        relation_model, exam_requirement_id
    )
    owned_by_editor = owned_exam_activity_ids_for_requirement(exam_requirement_id)
    foreign_in_payload = payload - owned_by_editor

    if not foreign_in_payload and payload <= owned_by_editor:
        owned_in_component = owned_exam_activity_ids_by_requirements(component_requirement_ids)
        owned_by_others = owned_in_component - owned_by_editor
        dropped_own_links = direct_linked - payload
        unlink_peer_owned = (existing - payload) & owned_by_others
        return dropped_own_links | unlink_peer_owned

    if foreign_in_payload:
        owned_in_component = owned_exam_activity_ids_by_requirements(component_requirement_ids)
        return (existing - payload) - owned_in_component

    return existing - payload


def related_exam_activity_ids_for_requirement(relation_model, exam_requirement_id):
    group_ids = related_group_ids_for_requirement(relation_model, exam_requirement_id)
    return related_exam_activity_ids_for_group_ids(relation_model, group_ids)


def related_foreign_exam_activity_ids_for_requirement(relation_model, exam_requirement_id):
    return related_exam_activity_ids_for_requirement(relation_model, exam_requirement_id) - owned_exam_activity_ids_for_requirement(exam_requirement_id)


def sync_relation_rows_to_group(relation_model, group_id, exam_activity_ids):
    if not exam_activity_ids:
        return

    activity_owner_map = dict(
        EsExamActivity.objects.filter(id__in=exam_activity_ids).values_list(
            "id", "exam_requirement_id"
        )
    )

    for exam_activity_id in exam_activity_ids:
        exam_requirement_id = activity_owner_map.get(exam_activity_id)
        if exam_requirement_id is None:
            continue

        rows = list(
            relation_model.objects.filter(
                exam_activity_id=exam_activity_id,
            ).order_by("id")
        )

        if rows:
            keeper = next(
                (
                    row
                    for row in rows
                    if row.exam_requirement_id == exam_requirement_id
                ),
                rows[0],
            )
            if keeper.group_id != group_id or keeper.exam_requirement_id != exam_requirement_id:
                keeper.group_id = group_id
                keeper.exam_requirement_id = exam_requirement_id
                keeper.save(update_fields=["group_id", "exam_requirement_id"])
            for row in rows:
                if row.id != keeper.id:
                    row.delete()
        else:
            relation_model.objects.create(
                group_id=group_id,
                exam_requirement_id=exam_requirement_id,
                exam_activity_id=exam_activity_id,
            )


def propagate_removed_exam_activities(
    relation_model,
    removed_exam_activity_ids,
    edited_requirement_ids,
    component_requirement_ids,
    prior_reqs_by_activity,
):
    """
    After relation rows for edited requirements are rewritten, drop removed activities.

    If an activity was tied to only one requirement in the pre-update component, remove
    all of its rows in that component (e.g. peer's only activity leaves the shared group).

    If multiple requirements shared that activity, remove only rows for edited
    requirements so peers keep their links.
    """
    if not removed_exam_activity_ids or not component_requirement_ids:
        return
    edited_set = set(edited_requirement_ids)
    for activity_id in removed_exam_activity_ids:
        prior_reqs = prior_reqs_by_activity.get(activity_id, set())
        if len(prior_reqs) <= 1:
            relation_model.objects.filter(
                exam_requirement_id__in=component_requirement_ids,
                exam_activity_id=activity_id,
            ).delete()
        else:
            relation_model.objects.filter(
                exam_requirement_id__in=edited_set,
                exam_activity_id=activity_id,
            ).delete()


def regroup_exam_requirement_relations(relation_model, group_model):
    relations = list(
        relation_model.objects.select_related("exam_activity").order_by("id")
    )

    if not relations:
        group_model.objects.all().delete()
        return

    used_group_ids = set()
    relations_by_group_activity = defaultdict(list)

    for relation in relations:
        relations_by_group_activity[(relation.group_id, relation.exam_activity_id)].append(relation)

    for grouped_relations in relations_by_group_activity.values():
        keeper = next(
            (
                relation
                for relation in grouped_relations
                if relation.exam_requirement_id == relation.exam_activity.exam_requirement_id
            ),
            grouped_relations[0],
        )
        expected_exam_requirement_id = keeper.exam_activity.exam_requirement_id
        if keeper.exam_requirement_id != expected_exam_requirement_id:
            keeper.exam_requirement_id = expected_exam_requirement_id
            keeper.save(update_fields=["exam_requirement_id"])

        used_group_ids.add(keeper.group_id)
        for relation in grouped_relations:
            if relation.id != keeper.id:
                relation.delete()

    group_model.objects.exclude(id__in=used_group_ids).delete()


def update_requirement_relations(relation_model, group_model, exam_requirement_id, payload_exam_activity_ids):
    payload_exam_activity_ids = list(dict.fromkeys(payload_exam_activity_ids or []))
    current_related_exam_activity_ids = related_exam_activity_ids_for_requirement(
        relation_model,
        exam_requirement_id,
    )
    target_exam_activity_ids = expanded_exam_activity_ids_for_requirement_update(
        exam_requirement_id,
        payload_exam_activity_ids,
    )

    if current_related_exam_activity_ids == target_exam_activity_ids:
        current_exam_activity_ids = sorted(current_related_exam_activity_ids)
        return {
            "changed": False,
            "old_exam_activity_ids": current_exam_activity_ids,
            "new_exam_activity_ids": current_exam_activity_ids,
        }

    current_group_ids = related_group_ids_for_requirement(relation_model, exam_requirement_id)

    if current_group_ids:
        relation_model.objects.filter(group_id__in=current_group_ids).exclude(
            exam_activity_id__in=target_exam_activity_ids,
        ).delete()

    target_group_ids = group_ids_for_exam_activity_ids(relation_model, target_exam_activity_ids)

    if target_group_ids:
        target_group_id = min(target_group_ids)
        relation_model.objects.filter(group_id__in=target_group_ids).update(
            group_id=target_group_id,
        )
    else:
        target_group_id = group_model.objects.create().id

    sync_relation_rows_to_group(relation_model, target_group_id, target_exam_activity_ids)

    regroup_exam_requirement_relations(relation_model, group_model)

    updated_related_exam_activity_ids = related_exam_activity_ids_for_requirement(
        relation_model,
        exam_requirement_id,
    )

    return {
        "changed": True,
        "old_exam_activity_ids": sorted(current_related_exam_activity_ids),
        "new_exam_activity_ids": sorted(updated_related_exam_activity_ids),
    }