"""Greedy auto-schedule: first legal session start per requirement, with same-time rollback."""

from collections import defaultdict, deque

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_jobs,
	_build_requirement_snapshots,
	_build_requirement_unavailability_map,
	_get_requirement_fixed_location,
	_load_external_snapshots,
	_load_requirements,
	_load_selected_activities,
	_max_scheduled_end_slot_for_requirements,
	_scheduled_location_ids_for_requirements,
	_scheduled_time_slots_for_requirements,
	_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 evaluate_requirement_slot
from api.services.exam_scheduler.occupancy import (
	_requirement_overlaps_slot,
	_same_location_peer_requirement_ids,
)
from api.services.exam_scheduler.relations import (
	_already_scheduled_preceding_bounds,
	_build_connected_components,
	_preceding_adjacency,
	_relation_group_anchor_value,
	_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 _datetime_to_slot, _existing_session_for_slot, _slot_to_datetime


@transaction.atomic
def schedule_exam_requirements(activity_ids, user_id):
	activity_ids = _normalize_ids(activity_ids)
	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(activity_ids)
	skipped_activity_labels = _skipped_activity_labels(selected_activities)
	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}
	selected_activities_by_requirement = defaultdict(list)
	for activity in selected_activities:
		selected_activities_by_requirement[activity.exam_requirement_id].append(activity)

	conflict_activity_ids = set()
	conflict_activity_labels = []

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

	location_catalog = _build_location_catalog()
	requirement_student_map = _build_requirement_snapshots(requirement_ids)
	requirement_unavailability_map = _build_requirement_unavailability_map(requirement_ids)
	period_groups = defaultdict(list)
	for requirement in requirements:
		period_groups[requirement.exam_period_id].append(requirement)
	selected_activity_map = {activity.id: activity for activity in selected_activities}

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

	for period_id, period_requirements in period_groups.items():
		exam_period = period_requirements[0].exam_period
		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(
				[
					activity
					for requirement in period_requirements
					for activity in selected_activities_by_requirement.get(requirement.id, [])
				]
			)
			continue

		unavailable_dates = set(
			EsExamPeriodUnavailability.objects
			.filter(exam_period_id=period_id)
			.values_list("unavailable_date", flat=True)
		)
		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,
		)

		component_edges = defaultdict(set)
		component_predecessors = defaultdict(set)
		component_in_degree = defaultdict(int)
		internal_precedence_requirement_ids = set()
		for source_requirement_id, target_requirement_id in _transitive_selected_preceding_pairs(period_requirement_ids, preceding_forward):
			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
			if source_component_id == target_component_id:
				internal_precedence_requirement_ids.add(source_requirement_id)
				internal_precedence_requirement_ids.add(target_requirement_id)
				continue
			if target_component_id not in component_edges[source_component_id]:
				component_edges[source_component_id].add(target_component_id)
				component_predecessors[target_component_id].add(source_component_id)
				component_in_degree[target_component_id] += 1

		topological_components = []
		component_queue = deque(
			sorted(
				[
					component_id
					for component_id in range(len(same_time_components))
					if component_in_degree.get(component_id, 0) == 0
				],
				key=lambda component_id: min(same_time_components[component_id]),
			)
		)
		processed_components = set()
		while component_queue:
			component_id = component_queue.popleft()
			if component_id in processed_components:
				continue
			processed_components.add(component_id)
			topological_components.append(component_id)
			for next_component_id in sorted(component_edges.get(component_id, set()), key=lambda item: min(same_time_components[item])):
				component_in_degree[next_component_id] -= 1
				if component_in_degree[next_component_id] == 0:
					component_queue.append(next_component_id)

		cyclic_components = set(range(len(same_time_components))) - set(topological_components)
		component_order = {
			component_id: index
			for index, component_id in enumerate(topological_components)
		}
		same_time_anchor_by_component = {}
		same_location_anchor_by_component = {}
		same_time_component_conflict_ids = set()
		conflict_location_components = set()
		component_end_slot_by_component = {}
		component_changed_activity_ids_by_component = defaultdict(set)

		def _refresh_same_time_component_state(component_id):
			requirement_set = same_time_components[component_id]
			scheduled_slots = _scheduled_time_slots_for_requirements(
				requirement_set,
				requirement_by_id,
				selected_activities_by_requirement,
			)
			if len(scheduled_slots) > 1:
				return False
			if len(scheduled_slots) == 1:
				same_time_anchor_by_component[component_id] = scheduled_slots[0]
				component_end_slot_by_component[component_id] = _max_scheduled_end_slot_for_requirements(
					requirement_set,
					requirement_by_id,
					selected_activities_by_requirement,
				)
			else:
				same_time_anchor_by_component.pop(component_id, None)
				component_end_slot_by_component.pop(component_id, None)
			return True

		def _refresh_same_location_component_state(location_component_id):
			if location_component_id is None:
				return
			requirement_set = same_location_components[location_component_id]
			scheduled_locations = _scheduled_location_ids_for_requirements(
				requirement_set,
				requirement_by_id,
				selected_activities_by_requirement,
			)
			if len(scheduled_locations) > 1:
				conflict_location_components.add(location_component_id)
				same_location_anchor_by_component.pop(location_component_id, None)
			elif len(scheduled_locations) == 1:
				same_location_anchor_by_component[location_component_id] = scheduled_locations[0]
			else:
				same_location_anchor_by_component.pop(location_component_id, None)

		def _rollback_same_time_component(component_id):
			rolled_back_activity_ids = component_changed_activity_ids_by_component.get(component_id, set())
			if not rolled_back_activity_ids:
				return
			rolled_back_codes = set()
			for activity_id in list(rolled_back_activity_ids):
				activity = selected_activity_map.get(activity_id)
				if activity is None:
					continue
				old_value = old_data.get(activity_id)
				if old_value is None:
					continue
				activity.time_slot = old_value["time_slot"]
				activity.location_id = old_value["location_id"]
				activity.session_id = old_value["session_id"]
				activity.is_scheduled = old_value["is_scheduled"]
				if activity in changed_activities:
					changed_activities.remove(activity)
				new_data.pop(activity_id, None)
				old_data.pop(activity_id, None)
				rolled_back_codes.add(activity.code)
			activity_snapshots[:] = [snapshot for snapshot in activity_snapshots if snapshot["activity_id"] not in rolled_back_activity_ids]
			if rolled_back_codes:
				remark_codes[:] = [code for code in remark_codes if code not in rolled_back_codes]
			component_changed_activity_ids_by_component[component_id].clear()
			_refresh_same_time_component_state(component_id)
			for requirement_id in same_time_components[component_id]:
				_refresh_same_location_component_state(same_location_component_map.get(requirement_id))

		for component_id, requirement_set in enumerate(same_time_components):
			scheduled_slots = _scheduled_time_slots_for_requirements(
				requirement_set,
				requirement_by_id,
				selected_activities_by_requirement,
			)
			if len(scheduled_slots) > 1:
				add_conflict_activities(
					[
						activity
						for requirement_id in requirement_set
						for activity in selected_activities_by_requirement.get(requirement_id, [])
					]
				)
				cyclic_components.add(component_id)
			elif len(scheduled_slots) == 1:
				same_time_anchor_by_component[component_id] = scheduled_slots[0]
				component_end_slot_by_component[component_id] = _max_scheduled_end_slot_for_requirements(
					requirement_set,
					requirement_by_id,
					selected_activities_by_requirement,
				)

		for component_id, requirement_set in enumerate(same_location_components):
			scheduled_locations = _scheduled_location_ids_for_requirements(
				requirement_set,
				requirement_by_id,
				selected_activities_by_requirement,
			)
			if len(scheduled_locations) > 1:
				add_conflict_activities(
					[
						activity
						for requirement_id in requirement_set
						for activity in selected_activities_by_requirement.get(requirement_id, [])
					]
				)
				conflict_location_components.add(component_id)
			elif len(scheduled_locations) == 1:
				same_location_anchor_by_component[component_id] = scheduled_locations[0]
		activity_snapshots = list(_load_external_snapshots(period_id)["activity_snapshots"])
		jobs = _build_requirement_jobs(period_requirements, selected_activities_by_requirement, exam_period, session_starts, unavailable_dates)
		jobs.sort(
			key=lambda job: (
				component_order.get(same_time_component_map.get(job["requirement"].id, -1), len(component_order)),
				job["fixed_slot"] is None,
				len(job["candidates"]),
				-len(job["activities"]),
				-job["required_seats"],
				job["requirement"].id,
			),
		)
		for job in jobs:
			requirement = job["requirement"]
			activities = job["activities"]
			same_time_component_id = same_time_component_map.get(requirement.id)
			same_time_component_requirement_ids = same_time_components[same_time_component_id] if same_time_component_id is not None else {requirement.id}

			def mark_same_time_component_conflict():
				if same_time_component_id is not None:
					same_time_component_conflict_ids.add(same_time_component_id)
					_rollback_same_time_component(same_time_component_id)
				add_conflict_activities(_selected_activities_for_requirements(same_time_component_requirement_ids, selected_activities_by_requirement))

			if job.get("preconflict"):
				mark_same_time_component_conflict()
				continue
			if same_time_component_id in same_time_component_conflict_ids:
				mark_same_time_component_conflict()
				continue
			if any(group_id in same_time_group_conflict_ids for group_id in _relation_group_ids_for_requirement(EsExamRequirementSameTime, requirement.id)):
				mark_same_time_component_conflict()
				continue
			if any(group_id in same_location_group_conflict_ids for group_id in _relation_group_ids_for_requirement(EsExamRequirementSameLocation, requirement.id)):
				mark_same_time_component_conflict()
				continue

			try:
				group_time_anchor = _relation_group_anchor_value(EsExamRequirementSameTime, requirement.id, same_time_group_anchor_by_group_id)
				group_location_anchor = _relation_group_anchor_value(EsExamRequirementSameLocation, requirement.id, same_location_group_anchor_by_group_id)
				requirement_fixed_location = _get_requirement_fixed_location(requirement)
			except ValidationError:
				mark_same_time_component_conflict()
				continue
			component_anchor_slot = same_time_anchor_by_component.get(same_time_component_id)
			time_anchor_values = {
				value
				for value in (
					group_time_anchor,
					job["fixed_slot"],
					component_anchor_slot,
				)
				if value is not None
			}
			if len(time_anchor_values) > 1:
				mark_same_time_component_conflict()
				continue
			anchor_time_slot = next(iter(time_anchor_values), None)
			same_location_component_id = same_location_component_map.get(requirement.id)
			location_anchor_values = {
				value
				for value in (
					group_location_anchor,
					requirement_fixed_location,
					same_location_anchor_by_component.get(same_location_component_id),
				)
				if value is not None
			}
			if len(location_anchor_values) > 1:
				mark_same_time_component_conflict()
				continue
			anchor_location_id = next(iter(location_anchor_values), None)
			same_location_related_ids = _same_location_peer_requirement_ids({requirement.id})
			if same_location_component_id is not None:
				same_location_related_ids = same_location_related_ids | set(same_location_components[same_location_component_id])
			same_location_related_ids.discard(requirement.id)
			location_required_for_job = bool(requirement.location_required)
			candidate_slot = None
			assigned_locations = None
			predecessor_component_ids = component_predecessors.get(same_time_component_id, set())
			if predecessor_component_ids:
				if any(pred_component_id not in component_end_slot_by_component for pred_component_id in predecessor_component_ids):
					mark_same_time_component_conflict()
					continue
				component_min_start_slot = max(
					component_end_slot_by_component[pred_component_id]
					for pred_component_id in predecessor_component_ids
				)
			else:
				component_min_start_slot = None

			if same_time_component_id in cyclic_components or requirement.id in internal_precedence_requirement_ids:
				mark_same_time_component_conflict()
				continue

			if same_location_component_id in conflict_location_components:
				mark_same_time_component_conflict()
				continue

			if anchor_time_slot is not None:
				candidate_slots = [anchor_time_slot]
			else:
				candidate_slots = [
					_datetime_to_slot(exam_period, candidate_dt)
					for candidate_dt in job["candidates"]
				]

			same_location_sibling_ids = set()
			if same_location_component_id is not None:
				same_location_sibling_ids = set(same_location_components[same_location_component_id]) - {requirement.id}

			for candidate_slot_value in candidate_slots:
				if candidate_slot_value is None:
					continue
				candidate_dt = _slot_to_datetime(exam_period, candidate_slot_value)
				pending_same_time_location_ids = (
					set(same_time_component_requirement_ids) & (same_location_related_ids | same_location_sibling_ids)
				) - {requirement.id}
				pending_cohort_ids = [
					peer_id
					for peer_id in pending_same_time_location_ids
					if peer_id in selected_activities_by_requirement
					and peer_id in requirement_by_id
					and requirement_by_id[peer_id].location_required
					and not _requirement_overlaps_slot(peer_id, activity_snapshots, candidate_slot_value, job["duration_slots"])
				]
				assigned_locations = evaluate_requirement_slot(
					requirement,
					candidate_slot_value,
					candidate_dt,
					job["duration_slots"],
					exam_period,
					activity_snapshots,
					location_catalog,
					requirement_unavailability_map,
					job["student_ids"],
					job["group_tt_ids"],
					job["required_seats"],
					activity_count=len(activities),
					legal_session_datetimes=job["legal_session_datetimes"],
					preceding_min_start=preceding_min_start_by_requirement.get(requirement.id),
					preceding_max_end=preceding_max_end_by_requirement.get(requirement.id),
					min_start_slot=component_min_start_slot,
					anchor_location_id=anchor_location_id,
					location_required=location_required_for_job,
					include_location_sharing=location_required_for_job,
					same_location_sibling_ids=same_location_sibling_ids,
					pending_cohort_ids=pending_cohort_ids,
					requirement_by_id=requirement_by_id,
				)
				if assigned_locations is None:
					continue

				candidate_slot = candidate_slot_value
				break

			if candidate_slot is None or assigned_locations is None:
				mark_same_time_component_conflict()
				continue

			if component_anchor_slot is None and same_time_component_id is not None:
				same_time_anchor_by_component[same_time_component_id] = candidate_slot
				component_end_slot_by_component[same_time_component_id] = max(
					component_end_slot_by_component.get(same_time_component_id, 0),
					candidate_slot + job["duration_slots"],
				)
			elif same_time_component_id is not None:
				component_end_slot_by_component[same_time_component_id] = max(
					component_end_slot_by_component.get(same_time_component_id, 0),
					candidate_slot + job["duration_slots"],
				)

			if same_location_component_id is not None and same_location_anchor_by_component.get(same_location_component_id) is None:
				chosen_location_id = next(
					(
						location.id
						for location in assigned_locations
						if location is not None
					),
					None,
				)
				if chosen_location_id is not None:
					same_location_anchor_by_component[same_location_component_id] = chosen_location_id

			_set_relation_group_anchor_value(EsExamRequirementSameTime, requirement.id, same_time_group_anchor_by_group_id, candidate_slot)
			if any(location is not None for location in assigned_locations):
				chosen_location_id = next(location.id for location in assigned_locations if location is not None)
				_set_relation_group_anchor_value(EsExamRequirementSameLocation, requirement.id, same_location_group_anchor_by_group_id, chosen_location_id)

			remark_codes.extend(activity.code for activity in activities)

			for activity, location in zip(activities, assigned_locations):
				session = _existing_session_for_slot(exam_period, candidate_slot, 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": candidate_slot,
					"location_id": location.id if 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 = candidate_slot
					activity.location = location
					activity.session = session
					activity.is_scheduled = True
					changed_activities.append(activity)
					if same_time_component_id is not None:
						component_changed_activity_ids_by_component[same_time_component_id].add(activity.id)
				activity_snapshots.append(
					{
						"activity_id": activity.id,
						"requirement_id": requirement.id,
						"start_slot": candidate_slot,
						"duration_slots": job["duration_slots"],
						"location_id": location.id if location else None,
						"seat_usage": job["seat_usage"],
						"exclusive_use": bool(requirement.exclusive_use),
						"location_is_partition": bool(location.is_partition) if location else False,
						"student_ids": set(job["student_ids"]),
					}
				)

	_apply_activity_changes(changed_activities, user_id)

	payload_activities = _load_selected_activities(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)),
	}
