import math
from datetime import datetime, timedelta

from rest_framework import status
from rest_framework.exceptions import ValidationError

from django.db.models import F, ExpressionWrapper, IntegerField, Value
from django.db.models.functions import ExtractHour, ExtractMinute, ExtractSecond, Coalesce, Cast

from api.models import EsSuitability
from api.models import EsExamActivity, EsExamRequirementStudent
from api.translation import __
from api.utils import log_critical_error, get_exception_detail
from api.views.admin.base import AdminApiBase


class ExamActivityList(AdminApiBase):
	list_default_order_column = "id"

	@staticmethod
	def build_duration(writing_time, reading_time):
		total_duration = timedelta()

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

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

		return (datetime.min + total_duration).time()

	@staticmethod
	def build_duration_slots(writing_time, reading_time):
		duration = ExamActivityList.build_duration(writing_time, reading_time)
		total_seconds = duration.hour * 3600 + duration.minute * 60 + duration.second
		return max(1, math.ceil(total_seconds / (30 * 60)))

	@staticmethod
	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 post(self, request):
		try:
			self.api_log_skip_outgoing_data = True
			data = self.get_data(request)
			return self.api_response(data=data)
		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)

	def get_data(self, request):
		params = request.data
		exam_period_id = params.get("exam_period_id")

		query = (
			EsExamActivity.objects
			.select_related("exam_requirement")
			.prefetch_related("invigilator_suitability", "location_suitability")
			.annotate(
				_duration_sec=Cast(
					Coalesce(ExtractHour("exam_requirement__writing_time"), Value(0)) * 3600
					+ Coalesce(ExtractMinute("exam_requirement__writing_time"), Value(0)) * 60
					+ Coalesce(ExtractSecond("exam_requirement__writing_time"), Value(0))
					+ Coalesce(ExtractHour("exam_requirement__reading_time"), Value(0)) * 3600
					+ Coalesce(ExtractMinute("exam_requirement__reading_time"), Value(0)) * 60
					+ Coalesce(ExtractSecond("exam_requirement__reading_time"), Value(0)),
					output_field=IntegerField(),
				),
			)
			.order_by("id")
		)
		query = query.annotate(
			_duration_slots=ExpressionWrapper(
				Cast(
					(F("_duration_sec") + 1799) / 1800,
					output_field=IntegerField(),
				),
				output_field=IntegerField(),
			),
		)
		query = query.annotate(
			_scheduled_end_epoch=ExpressionWrapper(
				(F("time_slot") + F("_duration_slots")) * 1800,
				output_field=IntegerField(),
			),
		)

		if exam_period_id is not None:
			query = query.filter(exam_requirement__exam_period_id=exam_period_id)

		objs = self.exam_activity_listing_filter(request, query)
		pagination = self.exam_activity_listing_pagination(request, objs)

		return_data = []

		for val in pagination["paginated_data"]:
			data = {}
			data["id"] = val.id
			data["location_id"] = val.location_id
			data["location"] = val.location.name if val.location else None
			# data["session_id"] = val.session_id
			# data["session"] = val.session.name if val.session else None
			data["code"] = val.code
			data["name"] = val.name
			data["time_slot"] = val.time_slot
			data["is_scheduled"] = val.is_scheduled
			data["invigilator_suitability"] = [
				{"id": suitability.id, "name": suitability.name}
				for suitability in val.invigilator_suitability.all()
			]
			data["location_suitability"] = [
				{"id": suitability.id, "name": suitability.name}
				for suitability in val.location_suitability.all()
			]
			duration = self.build_duration(val.exam_requirement.writing_time, val.exam_requirement.reading_time)
			data["duration"] = duration
			duration_slots = self.build_duration_slots(val.exam_requirement.writing_time, val.exam_requirement.reading_time)
			scheduled_start_at = self.slot_to_datetime(val.time_slot)
			scheduled_end_at = scheduled_start_at + timedelta(minutes=duration_slots * 30) if scheduled_start_at else None
			data["scheduled_start_at"] = scheduled_start_at
			data["scheduled_end_at"] = scheduled_end_at

			data["exam_requirement_id"] = val.exam_requirement_id
			data["exam_requirement_code"] = val.exam_requirement.code
			data["exam_requirement_name"] = val.exam_requirement.name
			data["exam_requirements"] = val.exam_requirement.name
			data["exam_period_id"] = val.exam_requirement.exam_period_id
			data["exam_period"] = val.exam_requirement.exam_period.name if val.exam_requirement.exam_period else None
			data["planned_size"] = val.exam_requirement.planned_size
			data["writing_time"] = val.exam_requirement.writing_time
			data["fixed_start_date"] = val.exam_requirement.fixed_start_date
			data["reading_time"] = val.exam_requirement.reading_time
			data["fixed_start_time"] = val.exam_requirement.fixed_start_time
			data["location_required"] = val.exam_requirement.location_required
			data["exclusive_use"] = val.exam_requirement.exclusive_use
			data["minimum_split_size"] = val.exam_requirement.minimum_split_size
			data["earliest_start"] = val.exam_requirement.earliest_start
			data["latest_end"] = val.exam_requirement.latest_end
			
			data["real_size"] = EsExamRequirementStudent.objects.filter(exam_requirement_id=val.exam_requirement_id, student_id__isnull=False).count()


			return_data.append(data)

		return {
			"data": return_data,
			"total": pagination["total"],
			"page": pagination["page"],
			"per_page": pagination["per_page"],
		}

	def exam_activity_listing_filter(self, request, model):
		filters = request.data.get("filter")
		if not isinstance(filters, dict) or not filters:
			return model

		for key, val in filters.items():
			if val in (None, "", [], {}):
				continue

			match key:
				case "id_from":
					try:
						model = model.filter(id__gte=int(val))
					except (TypeError, ValueError):
						continue
				case "id_to":
					try:
						model = model.filter(id__lte=int(val))
					except (TypeError, ValueError):
						continue
				case "code" | "name":
					model = model.filter(**{f"{key}__icontains": val})
				case "is_scheduled":
					if isinstance(val, str):
						normalized = val.strip().lower()
						if normalized in ("true", "1", "yes", "y"):
							model = model.filter(is_scheduled=True)
						elif normalized in ("false", "0", "no", "n"):
							model = model.filter(is_scheduled=False)
					else:
						model = model.filter(is_scheduled=bool(val))
				case "location":
					model = model.filter(location__name__icontains=val)
				case "location_id":
					if isinstance(val, (list, tuple, set)):
						model = model.filter(location_id__in=list(val))
					else:
						model = model.filter(location_id=val)
				case "exam_requirement":
					model = model.filter(
						exam_requirement__code__icontains=val
					) | model.filter(
						exam_requirement__name__icontains=val
					)
				case "exam_requirement_id":
					if isinstance(val, (list, tuple, set)):
						model = model.filter(exam_requirement_id__in=list(val))
					else:
						model = model.filter(exam_requirement_id=val)
				case "writing_time":
					model = model.filter(exam_requirement__writing_time__icontains=val)
				case "writing_time_from":
					try:
						model = model.filter(exam_requirement__writing_time__gte=datetime.strptime(val, "%H:%M:%S").time())
					except (TypeError, ValueError):
						try:
							model = model.filter(exam_requirement__writing_time__gte=datetime.strptime(val, "%H:%M").time())
						except (TypeError, ValueError):
							continue
				case "writing_time_to":
					try:
						model = model.filter(exam_requirement__writing_time__lte=datetime.strptime(val, "%H:%M:%S").time())
					except (TypeError, ValueError):
						try:
							model = model.filter(exam_requirement__writing_time__lte=datetime.strptime(val, "%H:%M").time())
						except (TypeError, ValueError):
							continue
				case "reading_time":
					model = model.filter(exam_requirement__reading_time__icontains=val)
				case "reading_time_from":
					try:
						model = model.filter(exam_requirement__reading_time__gte=datetime.strptime(val, "%H:%M:%S").time())
					except (TypeError, ValueError):
						try:
							model = model.filter(exam_requirement__reading_time__gte=datetime.strptime(val, "%H:%M").time())
						except (TypeError, ValueError):
							continue
				case "reading_time_to":
					try:
						model = model.filter(exam_requirement__reading_time__lte=datetime.strptime(val, "%H:%M:%S").time())
					except (TypeError, ValueError):
						try:
							model = model.filter(exam_requirement__reading_time__lte=datetime.strptime(val, "%H:%M").time())
						except (TypeError, ValueError):
							continue
				case "planned_size_min":
					try:
						model = model.filter(exam_requirement__planned_size__gte=int(val))
					except (TypeError, ValueError):
						continue
				case "planned_size_max":
					try:
						model = model.filter(exam_requirement__planned_size__lte=int(val))
					except (TypeError, ValueError):
						continue
				case "duration_from":
					try:
						t = datetime.strptime(val, "%H:%M:%S").time()
						secs = t.hour * 3600 + t.minute * 60 + t.second
					except (TypeError, ValueError):
						try:
							t = datetime.strptime(val, "%H:%M").time()
							secs = t.hour * 3600 + t.minute * 60
						except (TypeError, ValueError):
							continue
					model = model.filter(_duration_sec__gte=secs)
				case "duration_to":
					try:
						t = datetime.strptime(val, "%H:%M:%S").time()
						secs = t.hour * 3600 + t.minute * 60 + t.second
					except (TypeError, ValueError):
						try:
							t = datetime.strptime(val, "%H:%M").time()
							secs = t.hour * 3600 + t.minute * 60
						except (TypeError, ValueError):
							continue
					model = model.filter(_duration_sec__lte=secs)
				case "scheduled_start_at_from":
					try:
						dt = datetime.fromisoformat(val)
						if dt.tzinfo is None:
							target_slot = math.ceil((dt - datetime(1970, 1, 1)).total_seconds() / 1800)
						else:
							target_slot = math.ceil(dt.timestamp() / 1800)
					except (TypeError, ValueError):
						continue
					model = model.filter(time_slot__gte=target_slot)
				case "scheduled_start_at_to":
					try:
						dt = datetime.fromisoformat(val)
						if dt.tzinfo is None:
							target_slot = int((dt - datetime(1970, 1, 1)).total_seconds() / 1800)
						else:
							target_slot = int(dt.timestamp() / 1800)
					except (TypeError, ValueError):
						continue
					model = model.filter(time_slot__lte=target_slot)
				case "scheduled_end_at_from":
					try:
						dt = datetime.fromisoformat(val)
						if dt.tzinfo is None:
							target_epoch = int((dt - datetime(1970, 1, 1)).total_seconds())
						else:
							target_epoch = int(dt.timestamp())
					except (TypeError, ValueError):
						continue
					model = model.filter(_scheduled_end_epoch__gte=target_epoch)
				case "scheduled_end_at_to":
					try:
						dt = datetime.fromisoformat(val)
						if dt.tzinfo is None:
							target_epoch = int((dt - datetime(1970, 1, 1)).total_seconds())
						else:
							target_epoch = int(dt.timestamp())
					except (TypeError, ValueError):
						continue
					model = model.filter(_scheduled_end_epoch__lte=target_epoch)
				case _:
					continue
		return model

	def exam_activity_listing_pagination(self, request, model):
		page, per_page = self.parse_pagination(request)
		sort_by = request.data.get("sort_by")
		order_by = request.data.get("order_by")

		total = model.count()

		sort_field_map = {
			"duration": "_duration_sec",
			"location": "location__name",
			"scheduled_start_at": "time_slot",
			"scheduled_end_at": "_scheduled_end_epoch",
			"exam_requirement": "exam_requirement__name",
			"planned_size": "exam_requirement__planned_size",
			"writing_time": "exam_requirement__writing_time",
			"reading_time": "exam_requirement__reading_time",
		}

		if sort_by:
			sort_column = sort_field_map.get(sort_by, sort_by)
			if str(order_by).lower() == "desc":
				ordering_string = f"-{sort_column}"
			else:
				ordering_string = sort_column
			model = model.order_by(ordering_string)
		else:
			model = model.order_by(self.list_default_order_column)

		if per_page != -1:
			start_index = (page - 1) * per_page
			end_index = start_index + per_page
			paginated_data = model[start_index:end_index]
		else:
			paginated_data = model

		return {
			"total": total,
			"paginated_data": paginated_data,
			"page": page,
			"per_page": per_page,
		}
