"""Same-time, same-location, and preceding-exam relation graphs and anchors."""

from collections import defaultdict, deque

from django.db.models import Q

from api.models import (
	EsExamActivity,
	EsExamRequirementPrecedingExamination,
	EsExamRequirementSameLocation,
	EsExamRequirementSameTime,
)
from api.services.exam_scheduler.catalog import _load_requirements
from api.services.exam_scheduler.common import _raise_validation_error
from api.services.exam_scheduler.slots import _duration_slots


def _build_connected_components(relation_model, requirement_ids):
	requirement_ids = list(dict.fromkeys(requirement_ids or []))
	parent = {requirement_id: requirement_id for requirement_id in requirement_ids}

	def find(value):
		while parent[value] != value:
			parent[value] = parent[parent[value]]
			value = parent[value]
		return value

	def union(left, right):
		left_root = find(left)
		right_root = find(right)
		if left_root != right_root:
			parent[right_root] = left_root

	relations = relation_model.objects.filter(
		exam_requirement_id__in=requirement_ids,
		exam_activity__exam_requirement_id__in=requirement_ids,
	).values_list("exam_requirement_id", "exam_activity__exam_requirement_id")

	for source_requirement_id, target_requirement_id in relations:
		if source_requirement_id in parent and target_requirement_id in parent:
			union(source_requirement_id, target_requirement_id)

	group_members = defaultdict(set)
	for group_id, exam_requirement_id, activity_requirement_id in relation_model.objects.filter(
		Q(exam_requirement_id__in=requirement_ids) | Q(exam_activity__exam_requirement_id__in=requirement_ids)
	).values_list("group_id", "exam_requirement_id", "exam_activity__exam_requirement_id"):
		if exam_requirement_id in parent:
			group_members[group_id].add(exam_requirement_id)
		if activity_requirement_id in parent:
			group_members[group_id].add(activity_requirement_id)
	for members in group_members.values():
		if not members:
			continue
		first_requirement_id = next(iter(members))
		for requirement_id in members:
			union(first_requirement_id, requirement_id)

	components = defaultdict(set)
	for requirement_id in requirement_ids:
		components[find(requirement_id)].add(requirement_id)

	component_list = [component for component in components.values() if component]
	component_list.sort(key=lambda item: (min(item), len(item)))
	component_map = {}
	for index, component in enumerate(component_list):
		for requirement_id in component:
			component_map[requirement_id] = index

	return component_list, component_map


def _scheduled_slot_window(requirement):
	if requirement is None:
		return None
	scheduled_activities = [
		activity
		for activity in requirement.esexamactivity_set.all()
		if activity.is_scheduled and activity.time_slot is not None
	]
	if not scheduled_activities:
		return None
	start_slots = {activity.time_slot for activity in scheduled_activities}
	if len(start_slots) != 1:
		return None
	start_slot = next(iter(start_slots))
	return start_slot, start_slot + _duration_slots(requirement)


def _preceding_adjacency():
	forward = defaultdict(set)
	reverse = defaultdict(set)
	for source_requirement_id, target_requirement_id in EsExamRequirementPrecedingExamination.objects.values_list(
		"exam_requirement_id",
		"exam_activity__exam_requirement_id",
	):
		if source_requirement_id is None or target_requirement_id is None:
			continue
		if source_requirement_id == target_requirement_id:
			continue
		forward[source_requirement_id].add(target_requirement_id)
		reverse[target_requirement_id].add(source_requirement_id)
	return forward, reverse


def _reachable_requirement_ids(adjacency, start_id):
	visited = {start_id}
	queue = deque([start_id])
	reachable = set()
	while queue:
		current_id = queue.popleft()
		for next_id in adjacency.get(current_id, set()):
			if next_id in visited:
				continue
			visited.add(next_id)
			reachable.add(next_id)
			queue.append(next_id)
	return reachable


def _transitive_selected_preceding_pairs(requirement_ids, forward=None):
	requirement_ids = set(requirement_ids or [])
	if not requirement_ids:
		return []
	if forward is None:
		forward, _reverse = _preceding_adjacency()
	pairs = []
	seen = set()
	for source_requirement_id in requirement_ids:
		for target_requirement_id in _reachable_requirement_ids(forward, source_requirement_id):
			if target_requirement_id not in requirement_ids or target_requirement_id == source_requirement_id:
				continue
			pair = (source_requirement_id, target_requirement_id)
			if pair in seen:
				continue
			seen.add(pair)
			pairs.append(pair)
	return pairs


def _already_scheduled_preceding_bounds(requirement_ids, adjacency=None):
	requirement_ids = set(requirement_ids or [])
	min_start_by_requirement = {}
	max_end_by_requirement = {}
	if not requirement_ids:
		return min_start_by_requirement, max_end_by_requirement

	if adjacency is None:
		forward, reverse = _preceding_adjacency()
	else:
		forward, reverse = adjacency

	reachable_by_requirement = {}
	counterpart_ids = set()
	for requirement_id in requirement_ids:
		ancestors = _reachable_requirement_ids(reverse, requirement_id)
		descendants = _reachable_requirement_ids(forward, requirement_id)
		reachable_by_requirement[requirement_id] = (ancestors, descendants)
		counterpart_ids.update(ancestors)
		counterpart_ids.update(descendants)

	if not counterpart_ids:
		return min_start_by_requirement, max_end_by_requirement

	counterparts = {
		requirement.id: requirement
		for requirement in _load_requirements(list(counterpart_ids))
	}

	for requirement_id, (ancestors, descendants) in reachable_by_requirement.items():
		for ancestor_id in ancestors:
			window = _scheduled_slot_window(counterparts.get(ancestor_id))
			if window is None:
				continue
			_predecessor_start, predecessor_end = window
			current = min_start_by_requirement.get(requirement_id)
			min_start_by_requirement[requirement_id] = predecessor_end if current is None else max(current, predecessor_end)
		for descendant_id in descendants:
			window = _scheduled_slot_window(counterparts.get(descendant_id))
			if window is None:
				continue
			successor_start, _successor_end = window
			current = max_end_by_requirement.get(requirement_id)
			max_end_by_requirement[requirement_id] = successor_start if current is None else min(current, successor_start)

	return min_start_by_requirement, max_end_by_requirement


def _seed_relation_group_anchors_from_activities(relation_model, value_attr, requirement_ids, exclude_activity_ids=None):
	exclude_activity_ids = set(exclude_activity_ids or [])
	requirement_ids = set(requirement_ids or [])
	anchors = {}
	conflict_group_ids = set()
	if not requirement_ids:
		return anchors, conflict_group_ids

	group_ids = set(
		relation_model.objects.filter(
			Q(exam_requirement_id__in=requirement_ids) | Q(exam_activity__exam_requirement_id__in=requirement_ids)
		).values_list("group_id", flat=True)
	)
	if not group_ids:
		return anchors, conflict_group_ids

	group_requirement_ids = defaultdict(set)
	for group_id, exam_requirement_id, activity_requirement_id in relation_model.objects.filter(
		group_id__in=group_ids
	).values_list("group_id", "exam_requirement_id", "exam_activity__exam_requirement_id"):
		if exam_requirement_id is not None:
			group_requirement_ids[group_id].add(exam_requirement_id)
		if activity_requirement_id is not None:
			group_requirement_ids[group_id].add(activity_requirement_id)

	all_requirement_ids = set()
	for ids in group_requirement_ids.values():
		all_requirement_ids.update(ids)
	if not all_requirement_ids:
		return anchors, conflict_group_ids

	requirement_group_ids = defaultdict(set)
	for group_id, req_ids in group_requirement_ids.items():
		for req_id in req_ids:
			requirement_group_ids[req_id].add(group_id)

	activities = EsExamActivity.objects.filter(
		exam_requirement_id__in=all_requirement_ids,
		is_scheduled=True,
	)
	if exclude_activity_ids:
		activities = activities.exclude(id__in=exclude_activity_ids)

	for activity in activities:
		value = getattr(activity, value_attr)
		if value is None:
			continue
		for group_id in requirement_group_ids.get(activity.exam_requirement_id, ()):
			current_value = anchors.get(group_id)
			if current_value is None:
				anchors[group_id] = value
			elif current_value != value:
				conflict_group_ids.add(group_id)

	return anchors, conflict_group_ids


def _seed_same_time_group_anchor_map(requirement_ids, exclude_activity_ids=None):
	return _seed_relation_group_anchors_from_activities(
		EsExamRequirementSameTime,
		"time_slot",
		requirement_ids,
		exclude_activity_ids,
	)


def _seed_same_location_group_anchor_map(requirement_ids, exclude_activity_ids=None):
	return _seed_relation_group_anchors_from_activities(
		EsExamRequirementSameLocation,
		"location_id",
		requirement_ids,
		exclude_activity_ids,
	)


def _relation_group_ids_for_requirement(relation_model, requirement_id):
	return set(
		relation_model.objects.filter(
			Q(exam_requirement_id=requirement_id) | Q(exam_activity__exam_requirement_id=requirement_id)
		).values_list("group_id", flat=True)
	)


def _relation_group_anchor_value(relation_model, requirement_id, anchor_map):
	group_ids = _relation_group_ids_for_requirement(relation_model, requirement_id)
	values = {
		anchor_map[group_id]
		for group_id in group_ids
		if anchor_map.get(group_id) is not None
	}
	if len(values) > 1:
		_raise_validation_error("Exam activity relation group already has conflicting scheduled values.")
	return next(iter(values), None)


def _set_relation_group_anchor_value(relation_model, requirement_id, anchor_map, value):
	for group_id in _relation_group_ids_for_requirement(relation_model, requirement_id):
		anchor_map[group_id] = value
