from api.models import EsInvigilatorSetting, EsSessionInvigilator
from api.services.invigilator_setting_v2 import save_invigilator_setting
from api.tests.helpers import SchedulerTestCase, build_exam_world, post_signed

ASSIGN_URL = "/api/admin/session/assign-invigilator"


class SessionAssignInvigilatorTests(SchedulerTestCase):
	def _assigned_ids(self, session):
		return list(
			EsSessionInvigilator.objects.filter(session=session, invigilator_id__isnull=False)
			.order_by("invigilator_id")
			.values_list("invigilator_id", flat=True)
		)

	def _rule(self, *, role=None, role_ids=None, is_floating=0, is_inclusive=0, mode="fixed", quantity=1, students_count=None, priority=0, setting_id=None):
		if role_ids is None:
			role_ids = [role.id] if role is not None else []
		payload = {
			"role_ids": role_ids,
			"is_floating": is_floating,
			"is_inclusive": is_inclusive,
			"quantity_mode": EsInvigilatorSetting.QUANTITY_MODE_TO_CODE[mode],
			"quantity": quantity,
			"students_count": students_count,
			"priority": priority,
		}
		if setting_id is not None:
			payload["id"] = setting_id
		return payload

	def _people_with_role(self, world, role, count):
		people = []
		for _ in range(count):
			invigilator = world.add_invigilator()
			world.attach_invigilator_role(invigilator, role)
			people.append(invigilator)
		return people

	def test_example_1_shared_chief_and_per_room_ratio(self):
		world = build_exam_world()
		sessions = [world.add_session(students_enrolled=30) for _ in range(5)]
		world.mark_scheduled(world.activity, session=sessions[0])
		chief_role = world.add_invigilator_role(name="Chief")
		senior_role = world.add_invigilator_role(name="Senior")
		assistant_role = world.add_invigilator_role(name="Assistant")
		chiefs = self._people_with_role(world, chief_role, 1)
		seniors = self._people_with_role(world, senior_role, 5)
		assistants = self._people_with_role(world, assistant_role, 10)
		save_invigilator_setting({
			"configuration": [
				self._rule(role=chief_role, is_floating=1, quantity=1),
				self._rule(role=senior_role, mode="ratio", quantity=1, students_count=30),
				self._rule(role=assistant_role, mode="ratio", quantity=2, students_count=30),
			],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id for session in sessions]})
		self.assertEqual(response.status_code, 200)
		items = response.json()["data"]["list"]
		self.assertEqual(len(items), 5)
		chief_ids = {chiefs[0].id}
		for item in items:
			self.assertEqual(item["invigilators_required"], 4)
			self.assertEqual(item["total_assigned_invigilators"]["total"], 4)
			self.assertEqual(item["total_assigned_invigilators"]["roles"].get(str(chief_role.id)), 1)
			self.assertEqual(item["total_assigned_invigilators"]["roles"].get(str(senior_role.id)), 1)
			self.assertEqual(item["total_assigned_invigilators"]["roles"].get(str(assistant_role.id)), 2)
			self.assertEqual(set(item["invigilator_id"]) & chief_ids, chief_ids)
			self.assertEqual(set(item["floating_invigilator_id"]), chief_ids)
			self.assertFalse(set(item["floating_invigilator_id"]) - set(item["invigilator_id"]))
		shared_chief = set(items[0]["invigilator_id"]) & chief_ids
		for item in items[1:]:
			self.assertEqual(set(item["invigilator_id"]) & chief_ids, shared_chief)
		session_scope_ids = []
		for item in items:
			session_scope_ids.append(set(item["invigilator_id"]) - chief_ids)
		for index, left in enumerate(session_scope_ids):
			for right in session_scope_ids[index + 1:]:
				self.assertFalse(left & right)
		self.assertEqual(set().union(*session_scope_ids), {person.id for person in [*seniors, *assistants]})

	def test_example_2_floating_ratio_on_pooled_students(self):
		world = build_exam_world()
		first = world.add_session(students_enrolled=400)
		second = world.add_session(students_enrolled=600)
		world.mark_scheduled(world.activity, session=first)
		chief_role = world.add_invigilator_role(name="Chief")
		senior_role = world.add_invigilator_role(name="Senior")
		chiefs = self._people_with_role(world, chief_role, 5)
		seniors = self._people_with_role(world, senior_role, 2)
		save_invigilator_setting({
			"configuration": [
				self._rule(role=chief_role, is_floating=1, mode="ratio", quantity=1, students_count=200),
				self._rule(role=senior_role, quantity=1),
			],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [first.id, second.id]})
		self.assertEqual(response.status_code, 200)
		by_id = {item["session_id"]: item for item in response.json()["data"]["list"]}
		chief_ids = {person.id for person in chiefs}
		self.assertEqual(by_id[first.id]["invigilators_required"], 6)
		self.assertEqual(by_id[second.id]["invigilators_required"], 6)
		self.assertEqual(set(by_id[first.id]["invigilator_id"]) & chief_ids, chief_ids)
		self.assertEqual(set(by_id[second.id]["invigilator_id"]) & chief_ids, chief_ids)
		self.assertEqual(len(set(by_id[first.id]["invigilator_id"]) - chief_ids), 1)
		self.assertEqual(len(set(by_id[second.id]["invigilator_id"]) - chief_ids), 1)
		self.assertFalse(
			(set(by_id[first.id]["invigilator_id"]) - chief_ids)
			& (set(by_id[second.id]["invigilator_id"]) - chief_ids)
		)
		self.assertEqual(
			(set(by_id[first.id]["invigilator_id"]) - chief_ids)
			| (set(by_id[second.id]["invigilator_id"]) - chief_ids),
			{person.id for person in seniors},
		)

	def test_overlapping_floaters_are_shared(self):
		world = build_exam_world()
		first = world.add_session(students_enrolled=20)
		second = world.add_session(students_enrolled=20)
		chief_role = world.add_invigilator_role(name="Chief")
		chiefs = self._people_with_role(world, chief_role, 1)
		save_invigilator_setting({
			"configuration": [self._rule(role=chief_role, is_floating=1, quantity=1)],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [first.id, second.id]})
		self.assertEqual(response.status_code, 200)
		by_id = {item["session_id"]: item for item in response.json()["data"]["list"]}
		self.assertEqual(by_id[first.id]["invigilator_id"], [chiefs[0].id])
		self.assertEqual(by_id[second.id]["invigilator_id"], [chiefs[0].id])

	def test_session_scope_overlapping_sessions_do_not_share_invigilator(self):
		world = build_exam_world()
		first = world.add_session(students_enrolled=20)
		second = world.add_session(students_enrolled=20)
		world.mark_scheduled(world.activity, session=first)
		role = world.add_invigilator_role(name="Invigilator")
		invigilators = self._people_with_role(world, role, 2)
		save_invigilator_setting({
			"configuration": [self._rule(role=role, mode="ratio", quantity=1, students_count=20)],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [first.id, second.id]})
		self.assertEqual(response.status_code, 200)
		by_id = {item["session_id"]: item for item in response.json()["data"]["list"]}
		self.assertEqual(len(by_id[first.id]["invigilator_id"]), 1)
		self.assertEqual(len(by_id[second.id]["invigilator_id"]), 1)
		self.assertEqual(by_id[first.id]["floating_invigilator_id"], [])
		self.assertEqual(by_id[second.id]["floating_invigilator_id"], [])
		self.assertNotEqual(by_id[first.id]["invigilator_id"], by_id[second.id]["invigilator_id"])
		self.assertEqual(
			set(by_id[first.id]["invigilator_id"] + by_id[second.id]["invigilator_id"]),
			{invigilators[0].id, invigilators[1].id},
		)

	def test_multi_session_uses_each_student_count(self):
		world = build_exam_world()
		small = world.add_session(students_enrolled=20)
		large = world.add_session(
			start_time=world.datetime_for(1),
			students_enrolled=40,
		)
		world.mark_scheduled(world.activity, session=small)
		role = world.add_invigilator_role(name="Invigilator")
		self._people_with_role(world, role, 3)
		save_invigilator_setting({
			"configuration": [self._rule(role=role, mode="ratio", quantity=1, students_count=20)],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [small.id, large.id]})
		self.assertEqual(response.status_code, 200)
		by_id = {item["session_id"]: item for item in response.json()["data"]["list"]}
		self.assertEqual(by_id[small.id]["invigilators_required"], 1)
		self.assertEqual(by_id[large.id]["invigilators_required"], 2)
		self.assertEqual(len(by_id[small.id]["invigilator_id"]), 1)
		self.assertEqual(len(by_id[large.id]["invigilator_id"]), 2)

	def test_empty_rules_assign_zero_people(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=30)
		previous = world.add_invigilator()
		EsSessionInvigilator.objects.create(session=session, invigilator=previous)
		save_invigilator_setting({"configuration": []})

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id]})
		self.assertEqual(response.status_code, 200)
		item = response.json()["data"]["list"][0]
		self.assertEqual(item["invigilators_required"], 0)
		self.assertEqual(item["invigilator_id"], [])
		self.assertEqual(item["total_assigned_invigilators"], {"total": 0, "roles": {}})
		session.refresh_from_db()
		self.assertEqual(session.invigilators_required, 0)
		self.assertEqual(self._assigned_ids(session), [])

	def test_overlapping_shortage_does_not_write(self):
		world = build_exam_world()
		first = world.add_session(students_enrolled=20)
		second = world.add_session(students_enrolled=20)
		role = world.add_invigilator_role(name="Invigilator")
		self._people_with_role(world, role, 1)
		save_invigilator_setting({
			"configuration": [self._rule(role=role, mode="ratio", quantity=1, students_count=20)],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [first.id, second.id]})
		self.assertEqual(response.status_code, 400)
		self.assertFalse(EsSessionInvigilator.objects.filter(session__in=[first, second]).exists())
		first.refresh_from_db()
		second.refresh_from_db()
		self.assertEqual(first.invigilators_required, 1)
		self.assertEqual(second.invigilators_required, 1)

	def test_tt_busy_staff_is_skipped_and_fails_when_no_replacement(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=20)
		world.mark_scheduled(world.activity, session=session)
		role = world.add_invigilator_role(name="Invigilator")
		invigilator = world.add_invigilator(with_staff=True)
		world.attach_invigilator_role(invigilator, role)
		save_invigilator_setting({
			"configuration": [self._rule(role=role, mode="ratio", quantity=1, students_count=20)],
		})
		self.tt_staff_assign_conflict_at.return_value = True

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id]})
		self.assertEqual(response.status_code, 400)
		self.assertFalse(EsSessionInvigilator.objects.filter(session=session).exists())
		self.tt_staff_assign_conflict_at.assert_called()

	def test_insufficient_staff_returns_400_without_writes(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=40)
		role = world.add_invigilator_role(name="Invigilator")
		self._people_with_role(world, role, 1)
		save_invigilator_setting({
			"configuration": [self._rule(role=role, mode="ratio", quantity=1, students_count=20)],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id]})
		self.assertEqual(response.status_code, 400)
		self.assertIn("Not enough available invigilators", response.json()["error"])
		self.assertFalse(EsSessionInvigilator.objects.filter(session=session).exists())
		session.refresh_from_db()
		self.assertEqual(session.invigilators_required, 1)

	def test_replaces_previous_assignments(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=20)
		other = world.add_session(start_time=world.datetime_for(1), students_enrolled=20)
		role = world.add_invigilator_role(name="Invigilator")
		previous = world.add_invigilator()
		replacement = world.add_invigilator()
		world.attach_invigilator_role(previous, role)
		world.attach_invigilator_role(replacement, role)
		EsSessionInvigilator.objects.create(session=session, invigilator=previous)
		EsSessionInvigilator.objects.create(session=other, invigilator=previous)
		save_invigilator_setting({
			"configuration": [self._rule(role=role, mode="ratio", quantity=1, students_count=20)],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id]})
		self.assertEqual(response.status_code, 200)
		self.assertEqual(self._assigned_ids(session), [replacement.id])
		self.assertEqual(
			list(
				EsSessionInvigilator.objects.filter(session=other, invigilator_id__isnull=False)
				.values_list("invigilator_id", flat=True)
			),
			[previous.id],
		)

	def test_invalid_session_id_returns_400(self):
		response = post_signed(self.client, ASSIGN_URL, {"id": [999999999]})
		self.assertEqual(response.status_code, 400)

	def test_empty_role_ids_picks_any_active_invigilator(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=20)
		chief = world.add_invigilator_role(name="Chief")
		assistant = world.add_invigilator_role(name="Assistant")
		with_chief = self._people_with_role(world, chief, 1)[0]
		with_assistant = self._people_with_role(world, assistant, 1)[0]
		unassigned = world.add_invigilator()
		save_invigilator_setting({
			"configuration": [self._rule(role_ids=[], quantity=3)],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id]})
		self.assertEqual(response.status_code, 200)
		item = response.json()["data"]["list"][0]
		self.assertEqual(item["invigilators_required"], 3)
		self.assertEqual(set(item["invigilator_id"]), {with_chief.id, with_assistant.id, unassigned.id})

	def test_inclusive_senior_covers_ratio_at_50_students(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=50)
		senior_role = world.add_invigilator_role(name="Senior")
		extra_role = world.add_invigilator_role(name="Invigilator")
		seniors = self._people_with_role(world, senior_role, 1)
		self._people_with_role(world, extra_role, 1)
		save_invigilator_setting({
			"configuration": [
				self._rule(role=senior_role, is_inclusive=1, quantity=1),
				self._rule(role=extra_role, mode="ratio", quantity=1, students_count=50),
			],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id]})
		self.assertEqual(response.status_code, 200)
		item = response.json()["data"]["list"][0]
		self.assertEqual(item["invigilators_required"], 1)
		self.assertEqual(item["invigilator_id"], [seniors[0].id])
		self.assertEqual(item["total_assigned_invigilators"]["total"], 1)
		self.assertEqual(item["total_assigned_invigilators"]["roles"].get(str(senior_role.id)), 1)
		self.assertIsNone(item["total_assigned_invigilators"]["roles"].get(str(extra_role.id)))

	def test_inclusive_senior_needs_extra_at_100_students(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=100)
		senior_role = world.add_invigilator_role(name="Senior")
		extra_role = world.add_invigilator_role(name="Invigilator")
		seniors = self._people_with_role(world, senior_role, 1)
		extras = self._people_with_role(world, extra_role, 1)
		save_invigilator_setting({
			"configuration": [
				self._rule(role=senior_role, is_inclusive=1, quantity=1),
				self._rule(role=extra_role, mode="ratio", quantity=1, students_count=50),
			],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id]})
		self.assertEqual(response.status_code, 200)
		item = response.json()["data"]["list"][0]
		self.assertEqual(item["invigilators_required"], 2)
		self.assertEqual(set(item["invigilator_id"]), {seniors[0].id, extras[0].id})
		self.assertEqual(item["total_assigned_invigilators"]["total"], 2)
		self.assertEqual(item["total_assigned_invigilators"]["roles"].get(str(senior_role.id)), 1)
		self.assertEqual(item["total_assigned_invigilators"]["roles"].get(str(extra_role.id)), 1)

	def test_exclusive_senior_still_adds_at_50_students(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=50)
		senior_role = world.add_invigilator_role(name="Senior")
		extra_role = world.add_invigilator_role(name="Invigilator")
		seniors = self._people_with_role(world, senior_role, 1)
		extras = self._people_with_role(world, extra_role, 1)
		save_invigilator_setting({
			"configuration": [
				self._rule(role=senior_role, is_inclusive=0, quantity=1),
				self._rule(role=extra_role, mode="ratio", quantity=1, students_count=50),
			],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id]})
		self.assertEqual(response.status_code, 200)
		item = response.json()["data"]["list"][0]
		self.assertEqual(item["invigilators_required"], 2)
		self.assertEqual(set(item["invigilator_id"]), {seniors[0].id, extras[0].id})
		self.assertEqual(item["total_assigned_invigilators"]["total"], 2)

	def test_inclusive_credit_is_consumed_in_order(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=50)
		senior_role = world.add_invigilator_role(name="Senior")
		first_extra_role = world.add_invigilator_role(name="First Extra")
		second_extra_role = world.add_invigilator_role(name="Second Extra")
		seniors = self._people_with_role(world, senior_role, 1)
		first_extras = self._people_with_role(world, first_extra_role, 1)
		second_extras = self._people_with_role(world, second_extra_role, 1)
		save_invigilator_setting({
			"configuration": [
				self._rule(role=senior_role, is_inclusive=1, quantity=1),
				self._rule(role=first_extra_role, mode="ratio", quantity=1, students_count=50),
				self._rule(role=second_extra_role, mode="ratio", quantity=1, students_count=50),
			],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id]})
		self.assertEqual(response.status_code, 200)
		item = response.json()["data"]["list"][0]
		self.assertEqual(item["invigilators_required"], 2)
		self.assertEqual(set(item["invigilator_id"]), {seniors[0].id, second_extras[0].id})
		self.assertNotIn(first_extras[0].id, item["invigilator_id"])
		self.assertEqual(item["total_assigned_invigilators"]["roles"].get(str(senior_role.id)), 1)
		self.assertIsNone(item["total_assigned_invigilators"]["roles"].get(str(first_extra_role.id)))
		self.assertEqual(item["total_assigned_invigilators"]["roles"].get(str(second_extra_role.id)), 1)

	def test_higher_priority_rule_runs_before_list_order(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=50)
		senior_role = world.add_invigilator_role(name="Senior")
		extra_role = world.add_invigilator_role(name="Invigilator")
		seniors = self._people_with_role(world, senior_role, 1)
		extras = self._people_with_role(world, extra_role, 1)
		save_invigilator_setting({
			"configuration": [
				self._rule(role=senior_role, is_inclusive=1, quantity=1, priority=0),
				self._rule(role=extra_role, mode="ratio", quantity=1, students_count=50, priority=10),
			],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id]})
		self.assertEqual(response.status_code, 200)
		item = response.json()["data"]["list"][0]
		self.assertEqual(item["invigilators_required"], 2)
		self.assertEqual(set(item["invigilator_id"]), {seniors[0].id, extras[0].id})

	def test_same_priority_uses_lower_id_first(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=50)
		senior_role = world.add_invigilator_role(name="Senior")
		first_extra_role = world.add_invigilator_role(name="First Extra")
		second_extra_role = world.add_invigilator_role(name="Second Extra")
		seniors = self._people_with_role(world, senior_role, 1)
		first_extras = self._people_with_role(world, first_extra_role, 1)
		second_extras = self._people_with_role(world, second_extra_role, 1)
		created = save_invigilator_setting({
			"configuration": [
				self._rule(role=senior_role, is_inclusive=1, quantity=1),
				self._rule(role=first_extra_role, mode="ratio", quantity=1, students_count=50),
				self._rule(role=second_extra_role, mode="ratio", quantity=1, students_count=50),
			],
		})
		senior_id = created["configuration"][0]["id"]
		first_extra_id = created["configuration"][1]["id"]
		second_extra_id = created["configuration"][2]["id"]
		save_invigilator_setting({
			"configuration": [
				self._rule(role=second_extra_role, mode="ratio", quantity=1, students_count=50, setting_id=second_extra_id),
				self._rule(role=first_extra_role, mode="ratio", quantity=1, students_count=50, setting_id=first_extra_id),
				self._rule(role=senior_role, is_inclusive=1, quantity=1, setting_id=senior_id),
			],
		})

		response = post_signed(self.client, ASSIGN_URL, {"id": [session.id]})
		self.assertEqual(response.status_code, 200)
		item = response.json()["data"]["list"][0]
		self.assertEqual(item["invigilators_required"], 2)
		self.assertEqual(set(item["invigilator_id"]), {seniors[0].id, second_extras[0].id})
		self.assertNotIn(first_extras[0].id, item["invigilator_id"])

	def test_assigned_summary_counts_people_and_roles(self):
		world = build_exam_world()
		session = world.add_session(students_enrolled=20)
		world.mark_scheduled(world.activity, session=session)
		chief = world.add_invigilator_role(name="Chief")
		assistant = world.add_invigilator_role(name="Assistant")
		first = world.add_invigilator()
		second = world.add_invigilator()
		third = world.add_invigilator()
		world.attach_invigilator_role(first, chief)
		world.attach_invigilator_role(second, assistant)
		world.attach_invigilator_role(third, assistant)
		EsSessionInvigilator.objects.create(session=session, invigilator=first)
		EsSessionInvigilator.objects.create(session=session, invigilator=second)
		EsSessionInvigilator.objects.create(session=session, invigilator=third)

		listed = post_signed(
			self.client,
			"/api/admin/session/list",
			{"exam_period_id": world.period.id},
		)
		self.assertEqual(listed.status_code, 200)
		session_row = next(item for item in listed.json()["data"]["data"] if item["id"] == session.id)
		self.assertEqual(
			session_row["total_assigned_invigilators"],
			{
				"total": 3,
				"roles": {
					str(chief.id): 1,
					str(assistant.id): 2,
				},
			},
		)
