"""Direct schedule: pin selected activities to one time slot, packing same-location cohorts."""

from collections import defaultdict

from django.db import transaction
from rest_framework.exceptions import ValidationError

from api.models import (
	EsExamPeriodUnavailability,
	EsExamRequirementSameLocation,
	EsExamRequirementSameTime,
	EsSessionStart,
)
from api.services.exam_scheduler.catalog import (
	_build_location_catalog,
	_build_requirement_snapshots,
	_build_requirement_unavailability_map,
	_candidate_datetimes,
	_get_requirement_fixed_location,
	_get_requirement_fixed_slot,
	_load_external_snapshots,
	_load_requirements,
	_load_selected_activities_in_order,
	_requirement_real_size,
	_requirement_seat_usage,
	_requirement_student_group_tt_ids,
	_requirement_student_ids,
	_selected_activities_for_requirements,
)
from api.services.exam_scheduler.common import (
	_activity_label,
	_activity_needs_schedule,
	_apply_activity_changes,
	_build_response_payload,
	_normalize_ids,
	_skipped_activity_labels,
)
from api.services.exam_scheduler.constraints import (
	_choose_single_location,
	check_student_and_tt,
	check_time_constraints,
)
from api.services.exam_scheduler.occupancy import _same_location_sharing_for_slot
from api.services.exam_scheduler.relations import (
	_already_scheduled_preceding_bounds,
	_build_connected_components,
	_preceding_adjacency,
	_relation_group_ids_for_requirement,
	_seed_same_location_group_anchor_map,
	_seed_same_time_group_anchor_map,
	_set_relation_group_anchor_value,
	_transitive_selected_preceding_pairs,
)
from api.services.exam_scheduler.slots import _duration_slots, _existing_session_for_slot, _slot_to_datetime


@transaction.atomic
def schedule_exam_requirements_direct(activity_ids, time_slot, user_id):
	activity_ids = _normalize_ids(activity_ids)
	time_slot = int(time_slot)
	if not activity_ids:
		return {
			"data": {"activities": [], "unscheduled_activities": [], "unscheduled_activities_text": "", "skipped_activities": []},
			"old_data": {},
			"new_data": {},
			"remark_codes": [],
		}

	selected_activities = _load_selected_activities_in_order(activity_ids)
	skipped_activity_labels = _skipped_activity_labels(selected_activities)
	selected_activities_to_schedule = [activity for activity in selected_activities if _activity_needs_schedule(activity)]
	conflict_activity_ids = set()
	conflict_activity_labels = []

	def add_conflict_activities(activities):
		for activity in activities:
			if activity.id in conflict_activity_ids:
				continue
			conflict_activity_ids.add(activity.id)
			conflict_activity_labels.append(_activity_label(activity))

	requirement_ids = list(dict.fromkeys(activity.exam_requirement_id for activity in selected_activities))
	requirements = _load_requirements(requirement_ids)
	requirement_by_id = {requirement.id: requirement for requirement in requirements}
	requirement_student_map = _build_requirement_snapshots(requirement_ids)
	selected_activity_ids = {activity.id for activity in selected_activities_to_schedule}
	location_catalog = _build_location_catalog()
	requirement_unavailability_map = _build_requirement_unavailability_map(requirement_ids)
	selected_activities_by_requirement = defaultdict(list)
	for activity in selected_activities_to_schedule:
		selected_activities_by_requirement[activity.exam_requirement_id].append(activity)

	preference_groups_by_period = defaultdict(list)
	for requirement in requirements:
		preference_groups_by_period[requirement.exam_period_id].append(requirement)

	changed_activities = []
	old_data = {}
	new_data = {}
	remark_codes = []

	for period_id, period_requirements in sorted(preference_groups_by_period.items()):
		exam_period = period_requirements[0].exam_period
		candidate_dt = _slot_to_datetime(exam_period, time_slot)
		period_selected_activities = [activity for activity in selected_activities_to_schedule if activity.exam_requirement.exam_period_id == period_id]
		if candidate_dt is None:
			add_conflict_activities(period_selected_activities)
			continue

		session_starts = list(
			EsSessionStart.objects
			.filter(exam_period_id=period_id, status=EsSessionStart.STATUS_TO_CODE["active"])
			.prefetch_related("essessionstartday_set")
			.order_by("start_time", "id")
		)
		if not session_starts:
			add_conflict_activities(period_selected_activities)
			continue

		unavailable_dates = set(
			EsExamPeriodUnavailability.objects
			.filter(exam_period_id=period_id)
			.values_list("unavailable_date", flat=True)
		)
		external = _load_external_snapshots(period_id, exclude_activity_ids=selected_activity_ids)
		activity_snapshots = list(external["activity_snapshots"])
		period_requirement_ids = [requirement.id for requirement in period_requirements]
		preceding_forward, preceding_reverse = _preceding_adjacency()
		preceding_min_start_by_requirement, preceding_max_end_by_requirement = _already_scheduled_preceding_bounds(
			period_requirement_ids,
			(preceding_forward, preceding_reverse),
		)
		same_time_group_anchor_by_group_id, same_time_group_conflict_ids = _seed_same_time_group_anchor_map(
			period_requirement_ids,
		)
		same_location_group_anchor_by_group_id, same_location_group_conflict_ids = _seed_same_location_group_anchor_map(
			period_requirement_ids,
		)
		same_time_components, same_time_component_map = _build_connected_components(EsExamRequirementSameTime, period_requirement_ids)
		same_location_components, same_location_component_map = _build_connected_components(EsExamRequirementSameLocation, period_requirement_ids)

		conflict_location_component_ids = set()
		same_location_anchor_by_component = {}
		same_time_component_conflict_ids = set()
		for component_id, requirement_set in enumerate(same_location_components):
			scheduled_locations = sorted(
				{
					activity.location_id
					for requirement_id in requirement_set
					for activity in requirement_by_id[requirement_id].esexamactivity_set.all()
					if activity.is_scheduled and activity.location_id is not None
				}
			)
			if len(scheduled_locations) > 1:
				conflict_location_component_ids.add(component_id)
			elif len(scheduled_locations) == 1:
				same_location_anchor_by_component[component_id] = scheduled_locations[0]

		precedence_conflict_component_ids = set()
		for source_requirement_id, target_requirement_id in _transitive_selected_preceding_pairs(period_requirement_ids, preceding_forward):
			if not selected_activities_by_requirement.get(source_requirement_id) or not selected_activities_by_requirement.get(target_requirement_id):
				continue
			source_component_id = same_time_component_map.get(source_requirement_id)
			target_component_id = same_time_component_map.get(target_requirement_id)
			if source_component_id is None or target_component_id is None:
				continue
			precedence_conflict_component_ids.add(source_component_id)
			precedence_conflict_component_ids.add(target_component_id)

		component_ids_in_order = sorted(
			range(len(same_time_components)),
			key=lambda component_id: (min(same_time_components[component_id]), len(same_time_components[component_id])),
		)

		for component_id in component_ids_in_order:
			requirement_set = same_time_components[component_id]
			component_requirement_ids = [requirement_id for requirement_id in requirement_set if selected_activities_by_requirement.get(requirement_id)]
			if not component_requirement_ids:
				continue

			component_activities = [
				activity
				for requirement_id in component_requirement_ids
				for activity in selected_activities_by_requirement[requirement_id]
			]

			def mark_same_time_component_conflict():
				same_time_component_conflict_ids.add(component_id)
				add_conflict_activities(_selected_activities_for_requirements(component_requirement_ids, selected_activities_by_requirement))

			if component_id in same_time_component_conflict_ids:
				mark_same_time_component_conflict()
				continue

			if any(_relation_group_ids_for_requirement(EsExamRequirementSameTime, requirement_id) & same_time_group_conflict_ids for requirement_id in component_requirement_ids):
				mark_same_time_component_conflict()
				continue
			if any(_relation_group_ids_for_requirement(EsExamRequirementSameLocation, requirement_id) & same_location_group_conflict_ids for requirement_id in component_requirement_ids):
				mark_same_time_component_conflict()
				continue

			if component_id in precedence_conflict_component_ids:
				mark_same_time_component_conflict()
				continue

			if any(same_location_component_map.get(requirement_id) in conflict_location_component_ids for requirement_id in component_requirement_ids):
				mark_same_time_component_conflict()
				continue

			component_conflict = False
			component_student_ids_by_requirement = {}
			component_duration_slots_by_requirement = {}
			component_group_tt_ids_by_requirement = {}
			requirement_same_time_group_ids = {}
			requirement_same_location_group_ids = {}

			for requirement_id in component_requirement_ids:
				requirement = requirement_by_id[requirement_id]
				requirement_same_time_group_ids[requirement_id] = _relation_group_ids_for_requirement(EsExamRequirementSameTime, requirement_id)
				requirement_same_location_group_ids[requirement_id] = _relation_group_ids_for_requirement(EsExamRequirementSameLocation, requirement_id)
				try:
					fixed_slot = _get_requirement_fixed_slot(requirement)
				except ValidationError:
					component_conflict = True
					break
				group_time_anchor_values = {
					same_time_group_anchor_by_group_id[group_id]
					for group_id in requirement_same_time_group_ids[requirement_id]
					if same_time_group_anchor_by_group_id.get(group_id) is not None
				}
				if len(group_time_anchor_values) > 1:
					component_conflict = True
					break
				group_time_anchor = next(iter(group_time_anchor_values), None)
				if group_time_anchor is not None and group_time_anchor != time_slot:
					component_conflict = True
					break
				group_location_anchor_values = {
					same_location_group_anchor_by_group_id[group_id]
					for group_id in requirement_same_location_group_ids[requirement_id]
					if same_location_group_anchor_by_group_id.get(group_id) is not None
				}
				if len(group_location_anchor_values) > 1:
					component_conflict = True
					break
				group_location_anchor = next(iter(group_location_anchor_values), None)

				if fixed_slot is not None and fixed_slot != time_slot:
					component_conflict = True
					break
				try:
					fixed_location = _get_requirement_fixed_location(requirement)
				except ValidationError:
					component_conflict = True
					break
				if group_location_anchor is not None and fixed_location is not None and group_location_anchor != fixed_location:
					component_conflict = True
					break

				requirement_duration_slots = _duration_slots(requirement)
				component_duration_slots_by_requirement[requirement_id] = requirement_duration_slots
				component_student_ids_by_requirement[requirement_id] = _requirement_student_ids(requirement)
				component_group_tt_ids_by_requirement[requirement_id] = _requirement_student_group_tt_ids(requirement)
				if not check_time_constraints(
					requirement,
					time_slot,
					candidate_dt,
					requirement_duration_slots,
					requirement_unavailability_map,
					legal_session_datetimes=set(_candidate_datetimes(requirement, exam_period, session_starts, unavailable_dates)),
					preceding_min_start=preceding_min_start_by_requirement.get(requirement_id),
					preceding_max_end=preceding_max_end_by_requirement.get(requirement_id),
				):
					component_conflict = True
					break

			if component_conflict:
				mark_same_time_component_conflict()
				continue

			for index, left_requirement_id in enumerate(component_requirement_ids):
				left_student_ids = component_student_ids_by_requirement[left_requirement_id]
				for right_requirement_id in component_requirement_ids[index + 1:]:
					if left_student_ids.intersection(component_student_ids_by_requirement[right_requirement_id]):
						component_conflict = True
						break
				if component_conflict:
					break

			if component_conflict:
				mark_same_time_component_conflict()
				continue

			for requirement_id in component_requirement_ids:
				if not check_student_and_tt(
					component_student_ids_by_requirement[requirement_id],
					component_group_tt_ids_by_requirement[requirement_id],
					time_slot,
					component_duration_slots_by_requirement[requirement_id],
					activity_snapshots,
					exam_period,
					exclude_requirement_ids={requirement_id},
				):
					component_conflict = True
					break

			if component_conflict:
				mark_same_time_component_conflict()
				continue

			location_group_map = defaultdict(list)
			for requirement_id in component_requirement_ids:
				location_group_map[same_location_component_map.get(requirement_id)].append(requirement_id)

			working_snapshots = list(activity_snapshots)
			component_snapshots = []
			group_assignments = {}
			location_groups = []
			for location_component_id, location_requirement_ids in location_group_map.items():
				group_requirements = [requirement_by_id[requirement_id] for requirement_id in location_requirement_ids]
				location_needed_requirements = [requirement for requirement in group_requirements if requirement.location_required]
				group_required_seats = sum(
					_requirement_seat_usage(requirement, _requirement_real_size(requirement))
					for requirement in location_needed_requirements
				)
				group_duration_slots = max(component_duration_slots_by_requirement[requirement_id] for requirement_id in location_requirement_ids)
				group_exclusive_use = any(requirement.exclusive_use for requirement in location_needed_requirements)
				group_minimum_split_size = max((requirement.minimum_split_size or 0 for requirement in location_needed_requirements), default=0)
				group_location_required = bool(location_needed_requirements)
				location_groups.append((location_component_id, location_requirement_ids, group_requirements, group_required_seats, group_duration_slots, group_exclusive_use, group_minimum_split_size, group_location_required))

			location_groups.sort(key=lambda item: (-item[3], min(item[1])))
			for location_component_id, location_requirement_ids, group_requirements, group_required_seats, group_duration_slots, group_exclusive_use, group_minimum_split_size, group_location_required in location_groups:
				if not group_location_required:
					group_assignments[location_component_id] = None
					continue
				if group_exclusive_use and sum(1 for requirement in group_requirements if requirement.location_required) > 1:
					component_conflict = True
					break
				group_location_anchor_values = set()
				for requirement_id in location_requirement_ids:
					for group_id in requirement_same_location_group_ids[requirement_id]:
						anchor_value = same_location_group_anchor_by_group_id.get(group_id)
						if anchor_value is not None:
							group_location_anchor_values.add(anchor_value)
					try:
						fixed_location = _get_requirement_fixed_location(requirement_by_id[requirement_id])
					except ValidationError:
						component_conflict = True
						break
					if fixed_location is not None:
						group_location_anchor_values.add(fixed_location)
				if component_conflict:
					break
				if len(group_location_anchor_values) > 1:
					component_conflict = True
					break
				group_location_anchor = next(iter(group_location_anchor_values), None)

				scheduled_locations = sorted(
					{
						activity.location_id
						for requirement_id in location_requirement_ids
						for activity in requirement_by_id[requirement_id].esexamactivity_set.all()
						if activity.is_scheduled and activity.location_id is not None
					}
				)
				if len(scheduled_locations) > 1:
					component_conflict = True
					break

				anchor_location_id = scheduled_locations[0] if len(scheduled_locations) == 1 else (group_location_anchor or same_location_anchor_by_component.get(location_component_id))
				available_locations = location_catalog
				if anchor_location_id is not None:
					available_locations = [location_candidate for location_candidate in location_catalog if location_candidate["location"].id == anchor_location_id]
					if not available_locations:
						component_conflict = True
						break

				sharing_conflict, sharing_exclude_ids, sharing_extra_seats, share_location_id = _same_location_sharing_for_slot(
					set(location_requirement_ids),
					working_snapshots,
					time_slot,
					group_duration_slots,
					anchor_location_id,
					bool(group_exclusive_use),
				)
				if sharing_conflict:
					component_conflict = True
					break
				if anchor_location_id is None and share_location_id is not None:
					available_locations = [
						location_candidate
						for location_candidate in location_catalog
						if location_candidate["location"].id == share_location_id
					]
					if not available_locations:
						component_conflict = True
						break

				chosen_location = _choose_single_location(
					group_required_seats + sharing_extra_seats,
					available_locations,
					working_snapshots,
					time_slot,
					group_duration_slots,
					group_exclusive_use,
					group_minimum_split_size,
					True,
					exclude_requirement_ids=set(location_requirement_ids) | sharing_exclude_ids,
					exam_period=exam_period,
				)
				if chosen_location is None:
					component_conflict = True
					break

				group_assignments[location_component_id] = chosen_location
				same_location_anchor_by_component[location_component_id] = chosen_location.id
				for requirement_id in location_requirement_ids:
					requirement = requirement_by_id[requirement_id]
					requirement_duration_slots = component_duration_slots_by_requirement[requirement_id]
					assigned_location = chosen_location if requirement.location_required else None
					component_snapshots.append(
						{
							"activity_id": selected_activities_by_requirement[requirement_id][0].id,
							"requirement_id": requirement_id,
							"start_slot": time_slot,
							"duration_slots": requirement_duration_slots,
							"location_id": assigned_location.id if assigned_location else None,
							"seat_usage": _requirement_seat_usage(requirement, _requirement_real_size(requirement)),
							"exclusive_use": bool(requirement.exclusive_use) if assigned_location else False,
							"location_is_partition": bool(assigned_location.is_partition) if assigned_location else False,
							"student_ids": set(component_student_ids_by_requirement[requirement_id]),
						}
					)
				working_snapshots = list(activity_snapshots) + list(component_snapshots)

			if component_conflict:
				mark_same_time_component_conflict()
				continue

			for requirement_id in component_requirement_ids:
				_set_relation_group_anchor_value(
					EsExamRequirementSameTime,
					requirement_id,
					same_time_group_anchor_by_group_id,
					time_slot,
				)
				location_component_id_for_requirement = same_location_component_map.get(requirement_id)
				chosen_location_for_requirement = group_assignments.get(location_component_id_for_requirement)
				if chosen_location_for_requirement is not None and requirement_by_id[requirement_id].location_required:
					_set_relation_group_anchor_value(
						EsExamRequirementSameLocation,
						requirement_id,
						same_location_group_anchor_by_group_id,
						chosen_location_for_requirement.id,
					)

			activity_snapshots.extend(component_snapshots)
			for location_component_id, location_requirement_ids, group_requirements, group_required_seats, group_duration_slots, group_exclusive_use, group_minimum_split_size, group_location_required in location_groups:
				location = group_assignments.get(location_component_id)
				for requirement_id in location_requirement_ids:
					requirement = requirement_by_id[requirement_id]
					requirement_activities = selected_activities_by_requirement[requirement_id]
					assigned_location = location if requirement.location_required else None
					remark_codes.extend(activity.code for activity in requirement_activities)
					for activity in requirement_activities:
						session = _existing_session_for_slot(exam_period, time_slot, assigned_location)
						old_value = {
							"time_slot": activity.time_slot,
							"location_id": activity.location_id,
							"session_id": activity.session_id,
							"is_scheduled": activity.is_scheduled,
						}
						new_value = {
							"time_slot": time_slot,
							"location_id": assigned_location.id if assigned_location else None,
							"session_id": session.id if session else None,
							"is_scheduled": True,
						}

						if old_value != new_value:
							old_data.setdefault(activity.id, old_value)
							new_data.setdefault(activity.id, new_value)
							activity.time_slot = time_slot
							activity.location = assigned_location
							activity.session = session
							activity.is_scheduled = True
							changed_activities.append(activity)

	_apply_activity_changes(changed_activities, user_id)

	payload_activities = _load_selected_activities_in_order(activity_ids)
	return {
		"data": {
			**_build_response_payload(payload_activities, requirement_student_map),
			"unscheduled_activities": conflict_activity_labels,
			"skipped_activities": skipped_activity_labels,
			# "unscheduled_activities_text": "\n".join(conflict_activity_labels),
		},
		"old_data": old_data,
		"new_data": new_data,
		"remark_codes": list(dict.fromkeys(remark_codes)),
	}
