"""Load rooms, session starts, exam requirements, activities, and candidate datetimes."""

from collections import defaultdict
from datetime import datetime, timedelta

from django.db.models import Prefetch
from rest_framework.exceptions import ValidationError

from api.models import (
	EsExamActivity,
	EsExamRequirement,
	EsExamRequirementStudent,
	EsExamRequirementUnavailability,
	EsLocation,
	EsStudentStudentGroup,
)
from api.services.exam_scheduler.common import _activity_needs_schedule, _raise_validation_error
from api.services.exam_scheduler.slots import SLOT_MINUTES, _duration_slots, _slot_to_datetime


def _location_capacity(location):
	capacity = (location.row or 0) * (location.column or 0)
	capacity -= len(location.eslocationunavailableseat_set.all())
	return max(0, capacity)


def _build_location_catalog():
	locations = list(
		EsLocation.objects
		.filter(status=EsLocation.STATUS_TO_CODE["active"])
		.prefetch_related("eslocationunavailableseat_set")
		.order_by("id")
		.all()
	)
	catalog = []
	for location in locations:
		capacity = _location_capacity(location)
		if capacity <= 0:
			continue
		catalog.append(
			{
				"location": location,
				"capacity": capacity,
				"is_partition": bool(location.is_partition),
			}
		)

	catalog.sort(key=lambda item: (-item["capacity"], -int(item["is_partition"]), item["location"].id))
	return catalog


def _requirement_student_ids(requirement):
	return {
		relation.student_id
		for relation in requirement.esexamrequirementstudent_set.all()
		if relation.student_id is not None
	}


def _requirement_student_group_tt_ids(requirement):
	student_ids = _requirement_student_ids(requirement)
	if not student_ids:
		return set()

	return set(
		EsStudentStudentGroup.objects.filter(
			student_id__in=student_ids,
			student_group__tt_id__isnull=False,
		).values_list("student_group__tt_id", flat=True)
	)


def _requirement_real_size(requirement):
	return EsExamRequirementStudent.objects.filter(
		exam_requirement_id=requirement.id,
		student_id__isnull=False,
	).count()


def _requirement_seat_usage(requirement, real_size):
	return max(requirement.planned_size or 0, real_size)


def _iter_candidate_session_starts(requirement, exam_period, session_starts, unavailable_dates):
	candidates = []
	seen = set()

	for session_start in session_starts:
		if (session_start.start_time.hour * 3600 + session_start.start_time.minute * 60 + session_start.start_time.second) % (SLOT_MINUTES * 60) != 0:
			continue

		session_days = {
			relation.day
			for relation in session_start.essessionstartday_set.all()
		}

		current_date = exam_period.start_date
		while current_date <= exam_period.end_date:
			if current_date.weekday() not in session_days or current_date in unavailable_dates:
				current_date += timedelta(days=1)
				continue

			candidate_dt = datetime.combine(current_date, session_start.start_time)
			if requirement.fixed_start_date and candidate_dt.date() != requirement.fixed_start_date:
				current_date += timedelta(days=1)
				continue
			if requirement.fixed_start_time and candidate_dt.time() != requirement.fixed_start_time:
				current_date += timedelta(days=1)
				continue

			if candidate_dt not in seen:
				candidates.append((candidate_dt, session_start))
				seen.add(candidate_dt)

			current_date += timedelta(days=1)

	candidates.sort(key=lambda item: item[0])
	return candidates


def _candidate_datetimes(requirement, exam_period, session_starts, unavailable_dates):
	return [
		candidate_dt
		for candidate_dt, _session_start in _iter_candidate_session_starts(
			requirement, exam_period, session_starts, unavailable_dates
		)
	]


def _build_requirement_snapshots(requirement_ids):
	snapshots = defaultdict(set)
	relations = (
		EsExamRequirementStudent.objects
		.filter(exam_requirement_id__in=requirement_ids, student_id__isnull=False)
		.values_list("exam_requirement_id", "student_id")
	)
	for requirement_id, student_id in relations:
		snapshots[requirement_id].add(student_id)
	return snapshots


def _build_requirement_unavailability_map(requirement_ids):
	unavailability_map = defaultdict(set)
	relations = (
		EsExamRequirementUnavailability.objects
		.filter(exam_requirement_id__in=requirement_ids)
		.values_list("exam_requirement_id", "time_slot")
	)
	for requirement_id, time_slot in relations:
		if time_slot is not None:
			unavailability_map[requirement_id].add(int(time_slot))
	return unavailability_map


def _load_selected_activities(activity_ids):
	activities = list(
		EsExamActivity.objects
		.select_related("exam_requirement", "exam_requirement__exam_period", "location")
		.prefetch_related("location__eslocationunavailableseat_set")
		.filter(id__in=activity_ids)
		.order_by("id")
	)
	if len(activities) != len(activity_ids):
		_raise_validation_error("One or more exam activities are invalid.")
	return activities


def _load_selected_activities_in_order(activity_ids):
	activities = _load_selected_activities(activity_ids)
	activity_map = {activity.id: activity for activity in activities}
	return [activity_map[activity_id] for activity_id in activity_ids]


def _selected_activities_for_requirements(requirement_ids, selected_activities_by_requirement):
	activities = []
	seen_activity_ids = set()
	for requirement_id in requirement_ids:
		for activity in selected_activities_by_requirement.get(requirement_id, []):
			if activity.id in seen_activity_ids:
				continue
			seen_activity_ids.add(activity.id)
			activities.append(activity)
	return activities


def _build_external_activity_snapshots(activities, student_map):
	snapshots = []
	for activity in activities:
		if activity.time_slot is None:
			continue
		requirement = activity.exam_requirement
		student_count = len(student_map.get(activity.exam_requirement_id, set()))
		snapshots.append(
			{
				"activity_id": activity.id,
				"requirement_id": activity.exam_requirement_id,
				"start_slot": activity.time_slot,
				"duration_slots": _duration_slots(requirement),
				"location_id": activity.location_id,
				"seat_usage": _requirement_seat_usage(requirement, student_count),
				"exclusive_use": bool(requirement.exclusive_use),
				"location_is_partition": bool(activity.location.is_partition) if activity.location else False,
				"student_ids": set(student_map.get(activity.exam_requirement_id, set())),
			}
		)

	return snapshots


def _load_requirements(requirement_ids):
	return list(
		EsExamRequirement.objects
		.select_related("exam_period")
		.prefetch_related(
			Prefetch(
				"esexamactivity_set",
				queryset=EsExamActivity.objects.select_related("location").prefetch_related("location__eslocationunavailableseat_set").order_by("id"),
			),
		)
		.filter(id__in=requirement_ids)
		.order_by("id")
	)


def _get_requirement_fixed_slot(requirement):
	scheduled_time_slots = {
		activity.time_slot
		for activity in requirement.esexamactivity_set.all()
		if activity.is_scheduled and activity.time_slot is not None
	}
	if len(scheduled_time_slots) > 1:
		_raise_validation_error(f"Exam requirement {requirement.code} already has conflicting scheduled activities.")
	return next(iter(scheduled_time_slots), None)


def _get_requirement_fixed_location(requirement):
	scheduled_location_ids = {
		activity.location_id
		for activity in requirement.esexamactivity_set.all()
		if activity.is_scheduled and activity.location_id is not None
	}
	if len(scheduled_location_ids) > 1:
		_raise_validation_error(f"Exam requirement {requirement.code} already has conflicting scheduled locations.")
	return next(iter(scheduled_location_ids), None)


def _scheduled_location_ids_for_requirements(requirement_ids, requirement_by_id, selected_activities_by_requirement):
	location_ids = set()
	for requirement_id in requirement_ids or []:
		requirement = requirement_by_id.get(requirement_id)
		if requirement is not None:
			location_ids.update(
				activity.location_id
				for activity in requirement.esexamactivity_set.all()
				if activity.is_scheduled and activity.location_id is not None
			)
		location_ids.update(
			activity.location_id
			for activity in selected_activities_by_requirement.get(requirement_id, [])
			if activity.is_scheduled and activity.location_id is not None
		)
	return sorted(location_ids)


def _scheduled_time_slots_for_requirements(requirement_ids, requirement_by_id, selected_activities_by_requirement):
	time_slots = set()
	for requirement_id in requirement_ids or []:
		requirement = requirement_by_id.get(requirement_id)
		if requirement is not None:
			time_slots.update(
				activity.time_slot
				for activity in requirement.esexamactivity_set.all()
				if activity.is_scheduled and activity.time_slot is not None
			)
		time_slots.update(
			activity.time_slot
			for activity in selected_activities_by_requirement.get(requirement_id, [])
			if activity.is_scheduled and activity.time_slot is not None
		)
	return sorted(time_slots)


def _max_scheduled_end_slot_for_requirements(requirement_ids, requirement_by_id, selected_activities_by_requirement):
	end_slots = []
	for requirement_id in requirement_ids or []:
		requirement = requirement_by_id.get(requirement_id)
		duration = _duration_slots(requirement) if requirement is not None else None
		if requirement is not None:
			end_slots.extend(
				activity.time_slot + duration
				for activity in requirement.esexamactivity_set.all()
				if activity.is_scheduled and activity.time_slot is not None
			)
		for activity in selected_activities_by_requirement.get(requirement_id, []):
			if not activity.is_scheduled or activity.time_slot is None:
				continue
			activity_duration = duration
			if activity_duration is None:
				activity_requirement = requirement_by_id.get(activity.exam_requirement_id)
				if activity_requirement is None:
					continue
				activity_duration = _duration_slots(activity_requirement)
			end_slots.append(activity.time_slot + activity_duration)
	return max(end_slots) if end_slots else None


def _build_requirement_jobs(requirements, selected_activities_by_requirement, exam_period, session_starts, unavailable_dates):
	jobs = []
	for requirement in requirements:
		selected_activities = selected_activities_by_requirement.get(requirement.id, [])
		if not selected_activities:
			continue

		activities_to_schedule = [activity for activity in selected_activities if _activity_needs_schedule(activity)]
		if not activities_to_schedule:
			continue

		student_ids = _requirement_student_ids(requirement)
		required_seats = _requirement_seat_usage(requirement, len(student_ids))
		preconflict = False
		try:
			fixed_slot = _get_requirement_fixed_slot(requirement)
		except ValidationError:
			preconflict = True
			fixed_slot = None
			legal_candidates = []
			candidates = []
		else:
			legal_candidates = _candidate_datetimes(requirement, exam_period, session_starts, unavailable_dates)
			if fixed_slot is None:
				candidates = legal_candidates
			else:
				candidates = [
					_slot_to_datetime(exam_period, fixed_slot),
				]

		jobs.append(
			{
				"requirement": requirement,
				"activities": activities_to_schedule,
				"selected_activities": selected_activities,
				"candidates": candidates,
				"legal_session_datetimes": set(legal_candidates),
				"duration_slots": _duration_slots(requirement),
				"required_seats": required_seats,
				"seat_usage": required_seats,
				"student_ids": student_ids,
				"group_tt_ids": _requirement_student_group_tt_ids(requirement),
				"fixed_slot": fixed_slot,
				"preconflict": preconflict,
			}
		)

	jobs.sort(key=lambda job: (job["fixed_slot"] is None, len(job["candidates"]), -len(job["activities"]), -job["required_seats"], job["requirement"].id))
	return jobs


def _load_external_snapshots(period_id, exclude_activity_ids=None):
	exclude_activity_ids = set(exclude_activity_ids or [])
	external_activities = list(
		EsExamActivity.objects
		.select_related("exam_requirement", "exam_requirement__exam_period", "location")
		.prefetch_related("location__eslocationunavailableseat_set")
		.filter(exam_requirement__exam_period_id=period_id, is_scheduled=True)
		.exclude(id__in=exclude_activity_ids)
		.order_by("id")
	)

	student_map = _build_requirement_snapshots([activity.exam_requirement_id for activity in external_activities])
	return {
		"activities": external_activities,
		"activity_snapshots": _build_external_activity_snapshots(external_activities, student_map),
		"student_map": student_map,
	}
