import calendar
from collections import defaultdict
from datetime import datetime, timedelta

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

from api.models import EsExamActivity, EsExamRequirement, EsExamRequirementStudent, EsLocation, EsSession
from api.translation import __
from api.utils import get_exception_detail, log_critical_error
from api.validator import BaseValidator
from api.views.admin.base import AdminApiBase


def _duration_time(requirement):
	total_duration = timedelta()

	if requirement.writing_time:
		total_duration += timedelta(
			hours=requirement.writing_time.hour,
			minutes=requirement.writing_time.minute,
			seconds=requirement.writing_time.second,
		)
	if requirement.reading_time:
		total_duration += timedelta(
			hours=requirement.reading_time.hour,
			minutes=requirement.reading_time.minute,
			seconds=requirement.reading_time.second,
		)

	if total_duration == timedelta():
		total_duration = timedelta(hours=1)

	return total_duration


def _slot_to_datetime(time_slot):
	if time_slot is None:
		return None
	return datetime(1970, 1, 1) + timedelta(seconds=int(time_slot) * 1800)


def _time_to_timefield(duration_td):
	return (datetime.min + duration_td).time()


def _session_name(location, start_time):
	location_code = location.code if location else "ONLINE"
	return f"{location_code}: {start_time.strftime('%Y-%m-%d %H:%M:%S')}"


def _group_sort_key(item):
	time_slot, location_id = item[0]
	return (
		0 if time_slot is None else int(time_slot),
		1 if location_id is None else 0,
		-1 if location_id is None else int(location_id),
	)


def _serialize_exam_activity(activity):
	requirement = activity.exam_requirement
	return {
		"id": activity.id,
		"code": activity.code,
		"name": activity.name,
		"exam_requirement_id": activity.exam_requirement_id,
		# "exam_requirement_code": requirement.code,
		# "exam_requirement_name": requirement.name,
		"description": requirement.description,
		"fixed_start_date": requirement.fixed_start_date,
		"fixed_start_time": requirement.fixed_start_time,
		"is_scheduled": activity.is_scheduled,
		"students": requirement.planned_size,
		"real_size": EsExamRequirementStudent.objects.filter(exam_requirement_id=activity.exam_requirement_id, student_id__isnull=False).count()
	}


def _serialize_session(session, activities):
	return {
		"id": session.id,
		"name": session.name,
		"location": session.location.name if session.location else None,
		"description": session.description,
		"start_time": session.start_time,
		"duration": session.duration,
		"students_enrolled": session.students_enrolled,
		"invigilators_required": session.invigilators_required,
		"exam_activities": [_serialize_exam_activity(activity) for activity in activities],
	}


def _find_matching_session(start_time, location_id, exam_period_id):
	session_activity_queryset = (
		EsExamActivity.objects
		.select_related("exam_requirement")
		.filter(exam_requirement__exam_period_id=exam_period_id)
		.order_by("id")
	)
	candidates = list(
		EsSession.objects
		.select_related("location")
		.prefetch_related(Prefetch("esexamactivity_set", queryset=session_activity_queryset))
		.filter(
			start_time=start_time,
			location_id=location_id,
			description__isnull=True,
			invigilators_required=1,
			esexamactivity__exam_requirement__exam_period_id=exam_period_id,
		)
		.distinct()
		.order_by("id")
	)

	if not candidates:
		return None

	candidates.sort(key=lambda session: (-len(session.esexamactivity_set.all()), session.id))
	return candidates[0]


class SessionGenerate(AdminApiBase):
	def validate_request(self, request):
		rules = {
			"exam_period_id": "required|exists:api.EsExamPeriod,id",
		}

		validator = BaseValidator(request.data, rules, {})
		error = validator.validate()
		if error:
			raise ValidationError(error)

	def post(self, request):
		try:
			self.validate_request(request)
			self.api_log_skip_outgoing_data = True
			generated_sessions = self.generate_sessions(request.user.id, request.data["exam_period_id"])
			response = {
				"data": generated_sessions,
			}
			return self.api_response(data=response)
		except ValidationError as e:
			first_message = e.detail["error"]
			errors = e.detail["errors"]
			return self.api_response(error=first_message, errors=errors, code=status.HTTP_400_BAD_REQUEST)
		except Exception as e:
			e_details = get_exception_detail(e)
			log_critical_error(user_id=None, descr=e_details["descr"], url=e_details["url"], trace=e_details["trace"])
			return self.api_response(error=__("message.internal_server_error"), code=status.HTTP_500_INTERNAL_SERVER_ERROR)

	@transaction.atomic
	def generate_sessions(self, user_id, exam_period_id):
		activities = list(
			EsExamActivity.objects
			.select_related("exam_requirement", "location", "session")
			.filter(
				is_scheduled=True,
				time_slot__isnull=False,
				session__isnull=True,
				exam_requirement__exam_period_id=exam_period_id,
			)
			.order_by("time_slot", "location_id", "id")
		)

		if not activities:
			return []

		groups = defaultdict(list)
		for activity in activities:
			groups[(activity.time_slot, activity.location_id)].append(activity)

		generated_sessions = []
		activities_to_update = []
		now = timezone.now()

		for (time_slot, location_id), group_activities in sorted(groups.items(), key=_group_sort_key):
			location = group_activities[0].location
			start_time = _slot_to_datetime(time_slot)
			if start_time is None:
				continue

			requirements_by_id = {}
			for activity in group_activities:
				requirements_by_id[activity.exam_requirement_id] = activity.exam_requirement

			requirements = list(requirements_by_id.values())
			if not requirements:
				continue

			duration_td = max((_duration_time(requirement) for requirement in requirements), default=timedelta(hours=1))
			session_name = _session_name(location, start_time)

			duration = _time_to_timefield(duration_td)
			session = _find_matching_session(start_time, location_id, exam_period_id)
			if session is None:
				#students_enrolled = sum(requirement.planned_size or 0 for requirement in requirements)
				students_enrolled = sum(EsExamRequirementStudent.objects.filter(exam_requirement_id=requirement.id, student_id__isnull=False).count() for requirement in requirements)
				session = EsSession.objects.create(
					name=session_name,
					location=location,
					description=None,
					start_time=start_time,
					duration=duration,
					students_enrolled=students_enrolled,
					invigilators_required=1,
					created_by=user_id,
				)
				session_activities = list(group_activities)
			else:
				existing_activities = list(session.esexamactivity_set.all())
				session_activity_map = {activity.id: activity for activity in existing_activities}
				for activity in group_activities:
					session_activity_map[activity.id] = activity
				session_activities = sorted(session_activity_map.values(), key=lambda activity: activity.id)

				requirements_by_id = {activity.exam_requirement_id: activity.exam_requirement for activity in session_activities}
				combined_requirements = list(requirements_by_id.values())
				#students_enrolled = sum(requirement.planned_size or 0 for requirement in combined_requirements)
				students_enrolled = sum(EsExamRequirementStudent.objects.filter(exam_requirement_id=requirement.id, student_id__isnull=False).count() for requirement in combined_requirements)
				duration = _time_to_timefield(max((_duration_time(requirement) for requirement in combined_requirements), default=timedelta(hours=1)))

				updated_fields = []
				if session.name != session_name:
					session.name = session_name
					updated_fields.append("name")
				if session.location_id != location_id:
					session.location = location
					updated_fields.append("location")
				if session.description is not None:
					session.description = None
					updated_fields.append("description")
				if session.start_time != start_time:
					session.start_time = start_time
					updated_fields.append("start_time")
				if session.duration != duration:
					session.duration = duration
					updated_fields.append("duration")
				if session.students_enrolled != students_enrolled:
					session.students_enrolled = students_enrolled
					updated_fields.append("students_enrolled")
				if session.invigilators_required != 1:
					session.invigilators_required = 1
					updated_fields.append("invigilators_required")
				if updated_fields:
					session.updated_by = user_id
					session.updated_at = now
					updated_fields.extend(["updated_by", "updated_at"])
					session.save(update_fields=updated_fields)

			for activity in group_activities:
				activity.session = session
				activity.updated_by = user_id
				activity.updated_at = now
				activities_to_update.append(activity)

			generated_sessions.append(_serialize_session(session, session_activities))

		if activities_to_update:
			EsExamActivity.objects.bulk_update(
				activities_to_update,
				["session", "updated_by", "updated_at"],
			)

		return generated_sessions