from collections import defaultdict
from datetime import timedelta

from django.db import transaction
from django.db.models import Prefetch
from django.utils import timezone
from rest_framework.exceptions import ValidationError

from api.models import EsExamActivity, EsInvigilator, EsInvigilatorSetting, EsSession, EsSessionInvigilator
from api.services.exam_scheduler import SLOT_MINUTES, _duration_slots, _slot_to_datetime
from api.services.invigilator_setting_v2 import _as_int, compute_rule_quantity, load_invigilator_setting
from api.services.tt_availability import tt_staff_conflict_at
from api.translation import __


def summarize_assigned_invigilators(invigilators):
	roles = {}
	seen = set()
	total = 0
	for invigilator in invigilators or []:
		if invigilator is None or invigilator.id in seen:
			continue
		seen.add(invigilator.id)
		total += 1
		role_ids = getattr(invigilator, "role_id_set", None)
		if role_ids is None:
			role_ids = {role.id for role in invigilator.role.all()}
		for role_id in role_ids:
			key = str(role_id)
			roles[key] = roles.get(key, 0) + 1
	return {
		"total": total,
		"roles": roles,
	}


def _duration_minutes(session):
	duration = session.duration
	if not duration:
		return 0
	return duration.hour * 60 + duration.minute


def session_end(session):
	return session.start_time + timedelta(minutes=_duration_minutes(session))


def _is_inclusive(rule):
	return rule.get("is_inclusive") == EsInvigilatorSetting.IS_INCLUSIVE_TO_CODE["true"]


def _consume_inclusive_credit(rule, quantity, credit):
	if rule.get("quantity_mode") != EsInvigilatorSetting.QUANTITY_MODE_TO_CODE["ratio"]:
		return quantity, credit
	used = min(quantity, credit)
	return quantity - used, credit - used


def sessions_overlap(left, right):
	if left is None or right is None or left.id == right.id:
		return False
	return left.start_time < session_end(right) and right.start_time < session_end(left)


def session_tt_window(session):
	activity = next((item for item in session.esexamactivity_set.all() if item.is_scheduled), None)
	session_start = session.start_time
	duration_minutes = _duration_minutes(session)
	if activity and activity.exam_requirement_id:
		session_start = _slot_to_datetime(
			activity.exam_requirement.exam_period,
			activity.time_slot,
		) or session.start_time
		if activity.exam_requirement:
			duration_minutes = _duration_slots(activity.exam_requirement) * SLOT_MINUTES
	return session_start, duration_minutes


def _role_id_set(invigilator):
	return {role.id for role in invigilator.role.all()}


def _as_session_list(sessions):
	if sessions is None:
		return []
	if isinstance(sessions, (list, tuple)):
		return list(sessions)
	return [sessions]


def _is_available(invigilator, sessions, local_ids, assigned_sessions):
	if invigilator.id in local_ids:
		return False
	for session in _as_session_list(sessions):
		for other in assigned_sessions.get(invigilator.id, []):
			if sessions_overlap(session, other):
				return False
		staff_tt_id = invigilator.staff.tt_id if invigilator.staff_id else None
		if staff_tt_id:
			session_start, duration_minutes = session_tt_window(session)
			if tt_staff_conflict_at([staff_tt_id], session_start, duration_minutes):
				return False
	return True


def _pick_people(candidates, quantity, sessions, local_ids, assigned_sessions, workload):
	chosen = []
	if quantity <= 0:
		return chosen
	ordered = sorted(candidates, key=lambda invigilator: (workload[invigilator.id], invigilator.id))
	for invigilator in ordered:
		if not _is_available(invigilator, sessions, local_ids, assigned_sessions):
			continue
		chosen.append(invigilator)
		local_ids.add(invigilator.id)
		if len(chosen) >= quantity:
			break
	return chosen


def _candidates_for_rule(invigilators, local_ids, role_ids):
	role_ids = {
		role_id for role_id in (role_ids or [])
		if role_id not in (None, "", EsInvigilatorSetting.ANY_ROLE_ID, str(EsInvigilatorSetting.ANY_ROLE_ID))
	}
	picked = []
	for invigilator in invigilators:
		if invigilator.id in local_ids:
			continue
		if role_ids and not (invigilator.role_id_set & role_ids):
			continue
		picked.append(invigilator)
	return picked


def _insufficient_error(session, missing):
	return ValidationError({
		"error": __(
			"validation.invigilator_assign_insufficient",
			name=session.name,
			count=missing,
		),
		"errors": {
			"id": __(
				"validation.invigilator_assign_insufficient",
				name=session.name,
				count=missing,
			),
		},
	})


def _slot_quantities_for_rules(rules, students):
	credit = 0
	quantities = []
	for rule in rules:
		quantity = compute_rule_quantity(rule, students)
		quantity, credit = _consume_inclusive_credit(rule, quantity, credit)
		quantities.append(quantity)
		if _is_inclusive(rule):
			credit += quantity
	return quantities


def _add_estimated_roles(roles, rule, quantity):
	if quantity <= 0:
		return
	role_ids = [
		role_id for role_id in (rule.get("role_ids") or [])
		if role_id not in (None, "", EsInvigilatorSetting.ANY_ROLE_ID, str(EsInvigilatorSetting.ANY_ROLE_ID))
	]
	if not role_ids:
		key = str(EsInvigilatorSetting.ANY_ROLE_ID)
		roles[key] = roles.get(key, 0) + quantity
		return
	for role_id in role_ids:
		key = str(role_id)
		roles[key] = roles.get(key, 0) + quantity


def _rules_by_priority(rows, is_floating):
	flag = EsInvigilatorSetting.IS_FLOATING_TO_CODE["true"]
	if is_floating:
		filtered = [rule for rule in rows if rule.get("is_floating") == flag]
	else:
		filtered = [rule for rule in rows if rule.get("is_floating") != flag]
	return sorted(
		filtered,
		key=lambda rule: (-_as_int(rule.get("priority"), 0), _as_int(rule.get("id"), 0)),
	)


def estimate_required_invigilators(sessions=None, setting=None):
	if setting is None:
		setting = load_invigilator_setting()
	rows = setting.get("configuration") or []
	floating_rules = _rules_by_priority(rows, is_floating=True)
	session_rules = _rules_by_priority(rows, is_floating=False)
	if sessions is None:
		sessions = list(EsSession.objects.order_by("start_time", "id"))
	else:
		sessions = list(sessions)

	total = 0
	roles = {}
	for cohort in _overlapping_cohorts(sessions):
		cohort_students = sum(max(session.students_enrolled or 0, 0) for session in cohort)
		for rule, quantity in zip(floating_rules, _slot_quantities_for_rules(floating_rules, cohort_students)):
			total += quantity
			_add_estimated_roles(roles, rule, quantity)
		for session in cohort:
			for rule, quantity in zip(session_rules, _slot_quantities_for_rules(session_rules, session.students_enrolled)):
				total += quantity
				_add_estimated_roles(roles, rule, quantity)
	return {
		"total": total,
		"roles": roles,
		"total_sessions": len(sessions),
	}


def _overlapping_cohorts(sessions):
	items = list(sessions)
	parent = list(range(len(items)))

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

	def union(left, right):
		root_left = find(left)
		root_right = find(right)
		if root_left != root_right:
			parent[root_right] = root_left

	for i, left in enumerate(items):
		for j in range(i + 1, len(items)):
			if sessions_overlap(left, items[j]):
				union(i, j)

	groups = defaultdict(list)
	for index, session in enumerate(items):
		groups[find(index)].append(session)
	return list(groups.values())


def _replace_session_invigilators(session, invigilator_ids, user_id, floating_invigilator_ids=None):
	existing_ids = set(
		EsSessionInvigilator.objects.filter(
			session_id=session.id,
			invigilator_id__isnull=False,
		).values_list("invigilator_id", flat=True)
	)
	floating_invigilator_ids = list(dict.fromkeys(floating_invigilator_ids or []))
	floating_set = set(floating_invigilator_ids)
	invigilator_ids = [
		invigilator_id
		for invigilator_id in dict.fromkeys(invigilator_ids)
		if invigilator_id not in floating_set
	]
	assigned_ids = list(dict.fromkeys(floating_invigilator_ids + invigilator_ids))

	qs_to_null = EsSessionInvigilator.objects.filter(session_id=session.id)
	if assigned_ids:
		qs_to_null = qs_to_null.exclude(invigilator_id__in=assigned_ids)
	qs_to_null.update(invigilator=None)
	EsSessionInvigilator.objects.filter(
		session_id=session.id,
		invigilator_id__isnull=True,
	).delete()

	for ids, is_floating in ((invigilator_ids, 0), (floating_invigilator_ids, 1)):
		if not ids:
			continue
		EsSessionInvigilator.objects.filter(
			session_id=session.id,
			invigilator_id__in=ids,
		).update(
			is_floating=is_floating,
			updated_by=user_id,
			updated_at=timezone.now(),
		)
		EsSessionInvigilator.bulk_insert(session.id, ids, user_id, is_floating=is_floating)

	return existing_ids, assigned_ids


def assign_invigilators_to_sessions(session_ids, user_id=None):
	session_ids = list(dict.fromkeys(session_ids or []))
	setting = load_invigilator_setting()
	rows = setting.get("configuration") or []
	floating_rules = _rules_by_priority(rows, is_floating=True)
	session_rules = _rules_by_priority(rows, is_floating=False)
	sessions = list(
		EsSession.objects.filter(id__in=session_ids)
		.prefetch_related(
			Prefetch(
				"esexamactivity_set",
				queryset=EsExamActivity.objects.filter(is_scheduled=True).select_related(
					"exam_requirement__exam_period",
				),
			)
		)
		.order_by("start_time", "id")
	)
	sessions_by_id = {session.id: session for session in sessions}

	invigilators = list(
		EsInvigilator.objects.filter(status=EsInvigilator.STATUS_TO_CODE["active"])
		.select_related("staff")
		.prefetch_related("role")
	)
	for invigilator in invigilators:
		invigilator.role_id_set = _role_id_set(invigilator)

	assigned_sessions = defaultdict(list)
	workload = defaultdict(int)
	existing_links = (
		EsSessionInvigilator.objects.filter(invigilator_id__isnull=False)
		.exclude(session_id__in=session_ids)
		.select_related("session")
	)
	for link in existing_links:
		if not link.session_id:
			continue
		assigned_sessions[link.invigilator_id].append(link.session)
		workload[link.invigilator_id] += 1

	picks = []
	for cohort in _overlapping_cohorts(sessions):
		cohort_students = sum(max(session.students_enrolled or 0, 0) for session in cohort)
		floating_ids = set()
		floating_people = []
		floating_credit = 0
		for rule in floating_rules:
			quantity = compute_rule_quantity(rule, cohort_students)
			quantity, floating_credit = _consume_inclusive_credit(rule, quantity, floating_credit)
			if quantity <= 0:
				continue
			candidates = _candidates_for_rule(invigilators, floating_ids, rule.get("role_ids"))
			picked = _pick_people(candidates, quantity, cohort, floating_ids, assigned_sessions, workload)
			if len(picked) < quantity:
				raise _insufficient_error(cohort[0], quantity - len(picked))
			floating_people.extend(picked)
			if _is_inclusive(rule):
				floating_credit += len(picked)

		for invigilator in floating_people:
			for session in cohort:
				assigned_sessions[invigilator.id].append(session)
				workload[invigilator.id] += 1

		for session in cohort:
			local_ids = {invigilator.id for invigilator in floating_people}
			chosen = list(floating_people)
			session_slots = 0
			session_credit = 0
			for rule in session_rules:
				quantity = compute_rule_quantity(rule, session.students_enrolled)
				quantity, session_credit = _consume_inclusive_credit(rule, quantity, session_credit)
				session_slots += quantity
				if quantity <= 0:
					continue
				candidates = _candidates_for_rule(invigilators, local_ids, rule.get("role_ids"))
				picked = _pick_people(candidates, quantity, session, local_ids, assigned_sessions, workload)
				if len(picked) < quantity:
					raise _insufficient_error(session, quantity - len(picked))
				chosen.extend(picked)
				if _is_inclusive(rule):
					session_credit += len(picked)
				for invigilator in picked:
					assigned_sessions[invigilator.id].append(session)
					workload[invigilator.id] += 1

			picks.append({
				"session": session,
				"required": session_slots + len(floating_people),
				"invigilators": chosen,
				"invigilator_ids": [invigilator.id for invigilator in chosen if invigilator not in floating_people],
				"floating_invigilator_ids": [invigilator.id for invigilator in floating_people],
			})

	old_data = {}
	new_data = {}
	result_list = []

	with transaction.atomic():
		for pick in picks:
			session = pick["session"]
			required = pick["required"]
			invigilator_ids = pick["invigilator_ids"]
			existing_ids, assigned_ids = _replace_session_invigilators(
				session,
				invigilator_ids,
				user_id,
				floating_invigilator_ids=pick["floating_invigilator_ids"],
			)

			if session.invigilators_required != required:
				old_data.setdefault(session.id, {})["invigilators_required"] = session.invigilators_required
				new_data.setdefault(session.id, {})["invigilators_required"] = required
				session.invigilators_required = required
				session.updated_by = user_id
				session.updated_at = timezone.now()
				session.save(update_fields=["invigilators_required", "updated_by", "updated_at"])

			if set(existing_ids) != set(assigned_ids):
				old_data.setdefault(session.id, {})["invigilator_id"] = sorted(existing_ids)
				new_data.setdefault(session.id, {})["invigilator_id"] = sorted(assigned_ids)

			result_list.append({
				"session_id": session.id,
				"students_enrolled": session.students_enrolled,
				"invigilators_required": required,
				"total_assigned_invigilators": summarize_assigned_invigilators(pick["invigilators"]),
				"invigilator_id": assigned_ids,
				"floating_invigilator_id": pick["floating_invigilator_ids"],
			})

	# Keep response order aligned with the requested ids where possible.
	ordered_result = []
	result_by_session = {item["session_id"]: item for item in result_list}
	for session_id in session_ids:
		if session_id in result_by_session:
			ordered_result.append(result_by_session[session_id])
	for item in result_list:
		if item not in ordered_result:
			ordered_result.append(item)

	return {
		"list": ordered_result,
		"old_data": old_data,
		"new_data": new_data,
		"session_names": [sessions_by_id[session_id].name for session_id in session_ids if session_id in sessions_by_id],
	}
