import json

from rest_framework import status
from rest_framework.exceptions import ValidationError

from django.core.serializers.json import DjangoJSONEncoder
from django.db import transaction
from api.models import EsExamPeriod
from api.models import EsCache
from api.models import EsSetting
from api.models import EsExamRequirement
from api.models import EsStudent
from api.models import EsInvigilator
from api.models import EsSessionInvigilator
from api.models.es_department import EsDepartment
from api.models.es_exam_activity import EsExamActivity
from api.models.es_invigilator_role import EsInvigilatorRole
from api.models.es_location import EsLocation
from api.models import EsSessionStart

from django.db.models import Count, Q, F, IntegerField, ExpressionWrapper, Sum

from api.translation import __
from api.utils import log_critical_error, get_exception_detail
from api.validator import BaseValidator
from api.views.admin.base import AdminApiBase
from django.utils import timezone
from datetime import datetime, timedelta

class DashboardInfo(AdminApiBase):
	default_cache_expiry_minutes = 5

	def validate_request(self, request):
		rules = {
			"id": "nullable|integer|exists:api.EsExamPeriod,id",
		}
		attribute = {}
		# validate id first, cause need get object by id, if wrong id direct return error
		validator = BaseValidator(request.data, rules, attribute)
		error = validator.validate()
		if error:
			raise ValidationError(error)

		# custom error after basic validation
		if error:
			raise ValidationError(error)

	def get_section_cache_expiry_minutes(self, sections):
		setting_params = {
			section: f"dashboard_info_exp:{section}"
			for section in sections
		}
		setting_values = EsSetting.get_multiple_setting(
			setting_params.values()
		)

		expiry_minutes = {}
		for section, param in setting_params.items():
			try:
				value = int(setting_values.get(param))
			except (TypeError, ValueError):
				value = self.default_cache_expiry_minutes

			if value <= 0:
				value = self.default_cache_expiry_minutes

			expiry_minutes[section] = value

		return expiry_minutes

	def get_cached_sections(self, cache_key_prefix, section_builders):
		now = timezone.now()
		section_expiry_minutes = self.get_section_cache_expiry_minutes(
			section_builders.keys()
		)
		section_keys = {
			section: f"{cache_key_prefix}:{section}"
			for section in section_builders
		}

		# Load every section cache in one query to keep cache hits inexpensive.
		cache_entries = {
			entry.cache_key: entry
			for entry in EsCache.objects.filter(
				cache_key__in=section_keys.values()
			)
		}

		sections = {}
		cache_updates = {}
		section_updated_datetimes = []
		for section, builder in section_builders.items():
			cache_key = section_keys[section]
			cache_entry = cache_entries.get(cache_key)
			cache_cutoff = now - timedelta(
				minutes=section_expiry_minutes[section]
			)

			if (
				cache_entry
				and cache_entry.data is not None
				and cache_entry.updated_at >= cache_cutoff
			):
				sections[section] = cache_entry.data
				section_updated_datetimes.append(cache_entry.updated_at)
				# print(f"Cache hit for section '{section}', using cached data.")
				continue

			# print(f"Cache miss for section '{section}', building fresh data...")
			# Only an absent or expired section runs its dashboard calculation.
			section_data = builder()
			sections[section] = section_data
			cache_updates[cache_key] = json.loads(
				json.dumps(section_data, cls=DjangoJSONEncoder)
			)

		# Keep each section as an independent, transaction-safe cache row.
		if cache_updates:
			with transaction.atomic():
				for cache_key, cache_data in cache_updates.items():
					cache_entry, _ = EsCache.objects.update_or_create(
						cache_key=cache_key,
						defaults={"data": cache_data},
					)
					section_updated_datetimes.append(cache_entry.updated_at)

		# Represent the oldest section currently included in the dashboard.
		last_update_datetime = min(section_updated_datetimes, default=None)
		return sections, last_update_datetime

	def post(self, request):
		# Validate input
		try:
			self.api_log_skip_outgoing_data = True
			self.validate_request(request)

			exam_period_id = request.data.get("id")
			obj = None
			if exam_period_id not in (None, ""):
				obj = (
					EsExamPeriod.objects
					.prefetch_related("esexamperiodunavailability_set")
					.filter(id=exam_period_id)
					.first()
				)

			info = None
			if obj:
				cache_key_prefix = f"dashboard_info:exam_period:{obj.id}"
				sections, last_update_datetime = self.get_cached_sections(
					cache_key_prefix,
					{
						"risk_summary": lambda: self.get_risk_summary_info(obj),
						"kpis": lambda: self.get_kpis_info(obj),
						"upcomingExams": lambda: self.get_upcoming_exams(obj),
						"examsPerDay": lambda: self.get_exams_per_day(obj),
						"slotHeatmap": lambda: self.get_activity_count_by_session_start(obj),
						"utilization": lambda: self.get_utilization_info(obj),
						"invigilatorWorkload": lambda: self.get_invigilator_workload_info(obj),
						"invigilatorsByRole": lambda: self.get_invigilators_by_role(obj),
						"studentsByDepartment": lambda: self.get_students_by_department(obj),
					},
				)

				info = {
					"id": obj.id,
					"code": obj.code,
					"name": obj.name,
					"start_date": obj.start_date,
					"end_date": obj.end_date,
					"start_time": obj.start_time,
					"end_time": obj.end_time,
					"status": obj.status,
					"status_text": __("attr.exam_period_status." + str(obj.status)),
					"last_update_datetime": last_update_datetime,
				}

				info['risk_summary'] = sections["risk_summary"]
				info["kpis"] = sections["kpis"]
				info["risk_summary"]["unscheduled"] = info["kpis"]["unscheduled"]["total"]
				info["readiness"] = {
							"scheduled": info["kpis"]["scheduled"]["total"],
							"unscheduled": info["kpis"]["unscheduled"]["total"]
						}

				info["upcomingExams"] = sections["upcomingExams"]
				info["examsPerDay"] = sections["examsPerDay"]
				info["slotHeatmap"] = sections["slotHeatmap"]
				info["utilization"] = sections["utilization"]
				info["invigilatorWorkload"] = sections["invigilatorWorkload"]
				info["invigilatorsByRole"] = sections["invigilatorsByRole"]
				info["studentsByDepartment"] = sections["studentsByDepartment"]

			data = {}
			data["exam_periods"] = [
				{
					"id": exam_period.id,
					"code": exam_period.code,
					"name": exam_period.name,
				}
				for exam_period in EsExamPeriod.objects.filter(status=EsExamPeriod.STATUS_TO_CODE["active"]).order_by("name").all()
			]

			response = {
				"info": info,
				"rules": data,
			}

			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)

	def get_risk_summary_info(self, exam_period):
		days_left = self.get_days_left(exam_period)

		return {
			"daysToStart": days_left["daysToStart"],
			"daysToEnd": days_left["daysToEnd"],
			# "daysLeft": days_left['daysToStart'],

			"unscheduled": 0,
			# "conflicts": -999,
			# "roomsOverCapacity": -999,
			# "accommodationsUnhandled": -999,
		}
	
	def get_days_left(self, exam_period):
		today = timezone.now().date()
		days_to_start = max((exam_period.start_date - today).days, 0)
		days_to_end = max((exam_period.end_date - today).days, 0)
		
		return {
			"daysToStart": days_to_start,
			"daysToEnd": days_to_end,
		}
	
	def get_kpis_info(self, exam_period):
		total_requirements = (
			EsExamRequirement.objects
			.filter(exam_period = exam_period)
			.count()
		)
		
		activity_summary = (
			EsExamActivity.objects
			.filter(exam_requirement__exam_period=exam_period)
			.aggregate(
				total=Count("id"),
				scheduled=Count("id", filter=Q(is_scheduled=True)),
			)
		)

		students_summary = (
			EsStudent.objects
			.filter(status = EsStudent.STATUS_TO_CODE["active"])
			.aggregate(
				total = Count("id"),
				extraProvision = Count("id", filter = Q(need_extra_provision=True)),
			)
		)

		total_invigilators = (
			EsInvigilator.objects
			.filter(status = EsInvigilator.STATUS_TO_CODE["active"])
			.count()
		)

		total_available_seats = (
			EsLocation.objects
			.annotate(
				available_seats=ExpressionWrapper(
					F("row") * F("column") - Count("eslocationunavailableseat", distinct=True),
					output_field=IntegerField(),
				)
			)
			.aggregate(total=Sum("available_seats"))
		)["total"] or 0

		return {
			"requirements": {
				"total": total_requirements
			},
			"scheduled": {
				"value": activity_summary["scheduled"],
				"total": activity_summary["total"],
				"percent": round((activity_summary["scheduled"] / activity_summary["total"]) * 100, 2) if activity_summary["total"] > 0 else 0
			},
			"unscheduled": {
				"total": activity_summary["total"] - activity_summary["scheduled"]
			},
			"students": {
				"total": students_summary["total"],
				"extraProvision": students_summary["extraProvision"],
			},
			"invigilators": {
				"total": total_invigilators,
			},
			"seats": {
				"total": total_available_seats,
			},
			"daysToEnd": {
				"value": self.get_days_left(exam_period)["daysToEnd"],
			},
		}
	
	def get_upcoming_exams(self, exam_period):
		now = timezone.now()
		epoch = datetime(1970, 1, 1)

		current_slot = int((now - epoch).total_seconds() // 1800)

		# If currently between slots, move to the next slot.
		if self.slot_to_datetime(current_slot) < now:
			current_slot += 1

		activities = (
			EsExamActivity.objects
			.select_related("location", "exam_requirement")
			.annotate(
				real_size=Count(
					"exam_requirement__esexamrequirementstudent",
					filter=Q(
						exam_requirement__esexamrequirementstudent__student__isnull=False
					),
					distinct=True,
				)
			)
			.filter(
				exam_requirement__exam_period=exam_period,
				is_scheduled=True,
				time_slot__isnull=False,
				time_slot__gte=current_slot,
			)
			.order_by("time_slot")[:6]
		)

		return [
			{
				"id": activity.id,
				"code": activity.code,
				"name": activity.name,
				"date": (
					f"{self.slot_to_datetime(activity.time_slot):%a}, "
					f"{self.slot_to_datetime(activity.time_slot).day} "
					f"{self.slot_to_datetime(activity.time_slot):%b}"
				),
				"time": self.slot_to_datetime(
					activity.time_slot
				).strftime("%H:%M"),
				"location": (
					activity.location.name
					if activity.location
					else None
				),
				"students": {
					"planned_size":
						activity.exam_requirement.planned_size or 0,
					"real_size": activity.real_size or 0,
				},
			}
			for activity in activities
		]
	
	def get_exams_per_day(self, exam_period):
		student_count_by_date = {}

		# Include every date in the exam period, even days with no exams.
		current_date = exam_period.start_date

		while current_date <= exam_period.end_date:
			student_count_by_date[current_date] = 0
			current_date += timedelta(days=1)

		activities = (
			EsExamActivity.objects
			.annotate(
				real_size=Count(
					"exam_requirement__esexamrequirementstudent",
					filter=Q(
						exam_requirement__esexamrequirementstudent__student__isnull=False
					),
					distinct=True,
				)
			)
			.filter(
				exam_requirement__exam_period=exam_period,
				is_scheduled=True,
				time_slot__isnull=False,
			)
			.order_by("time_slot")
		)

		for activity in activities:
			scheduled_at = self.slot_to_datetime(activity.time_slot)

			if scheduled_at is None:
				continue

			exam_date = scheduled_at.date()

			# Ignore activities accidentally scheduled outside the period.
			if exam_date not in student_count_by_date:
				continue

			student_count_by_date[exam_date] += activity.real_size or 0

		return {
			"dates": [
				f"{date:%b} {date.day}"
				for date in student_count_by_date
			],
			"students": [
				student_count_by_date[date]
				for date in student_count_by_date
			],
		}
	
	def get_activity_count_by_session_start(self, exam_period):
		session_starts = (
			EsSessionStart.objects
			.filter(
				exam_period=exam_period,
				status=EsSessionStart.STATUS_TO_CODE["active"],
			)
			.order_by("start_time", "id")
		)

		# Unique session times, preserving their configured order.
		slots = []

		for session_start in session_starts:
			slot = session_start.start_time.strftime("%H:%M")

			if slot not in slots:
				slots.append(slot)

		slot_set = set(slots)

		# Every calendar date within the exam period.
		dates = []
		current_date = exam_period.start_date

		while current_date <= exam_period.end_date:
			dates.append(current_date)
			current_date += timedelta(days=1)

		date_set = set(dates)

		# Monday=0 through Sunday=6.
		weekday_names = [
			"Mon",
			"Tue",
			"Wed",
			"Thu",
			"Fri",
			"Sat",
			"Sun",
		]

		by_date_counts = {}
		by_weekday_counts = {}

		activity_time_slots = (
			EsExamActivity.objects
			.filter(
				exam_requirement__exam_period=exam_period,
				is_scheduled=True,
				time_slot__isnull=False,
			)
			.values_list("time_slot", flat=True)
		)

		for time_slot in activity_time_slots:
			scheduled_at = self.slot_to_datetime(time_slot)

			if scheduled_at is None:
				continue

			scheduled_date = scheduled_at.date()
			scheduled_time = scheduled_at.strftime("%H:%M")

			if scheduled_date not in date_set:
				continue

			if scheduled_time not in slot_set:
				continue

			# Count by exact date.
			date_key = (
				scheduled_time,
				scheduled_date,
			)
			by_date_counts[date_key] = (
				by_date_counts.get(date_key, 0) + 1
			)

			# Count by weekday.
			weekday_key = (
				scheduled_time,
				scheduled_date.weekday(),
			)
			by_weekday_counts[weekday_key] = (
				by_weekday_counts.get(weekday_key, 0) + 1
			)

		return {
			"by_date": {
				"slots": slots,
				"days": [
					f"{date:%a} ({date:%b} {date.day})"
					for date in dates
				],
				"matrix": [
					[
						by_date_counts.get((slot, date), 0)
						for date in dates
					]
					for slot in slots
				],
			},
			"by_days": {
				"slots": slots,
				"days": weekday_names,
				"matrix": [
					[
						by_weekday_counts.get(
							(slot, weekday),
							0,
						)
						for weekday in range(7)
					]
					for slot in slots
				],
			},
		}

	def get_utilization_info(self, exam_period):
		unavailable_dates = set(
			exam_period
			.esexamperiodunavailability_set
			.values_list("unavailable_date", flat=True)
		)

		session_starts = list(
			EsSessionStart.objects
			.filter(
				exam_period=exam_period,
				status=EsSessionStart.STATUS_TO_CODE["active"],
			)
			.prefetch_related("essessionstartday_set")
			.order_by("start_time", "id")
		)

		# Cache the available weekdays for each session start.
		session_start_weekdays = {
			session_start.id: {
				item.day
				for item in session_start.essessionstartday_set.all()
			}
			for session_start in session_starts
		}

		# Build all unique date/session opportunities.
		valid_opportunities = set()
		current_date = exam_period.start_date

		while current_date <= exam_period.end_date:
			if current_date not in unavailable_dates:
				for session_start in session_starts:
					available_weekdays = session_start_weekdays[
						session_start.id
					]

					if current_date.weekday() in available_weekdays:
						valid_opportunities.add((
							current_date,
							session_start.start_time.strftime("%H:%M"),
						))

			current_date += timedelta(days=1)

		available_session_count = len(valid_opportunities)

		# Calculate usable seats directly in the locations query.
		locations = (
			EsLocation.objects
			.filter(
				status=EsLocation.STATUS_TO_CODE["active"]
			)
			.annotate(
				total_seats=ExpressionWrapper(
					F("row") * F("column")
					- Count(
						"eslocationunavailableseat",
						distinct=True,
					),
					output_field=IntegerField(),
				)
			)
			.order_by("name")
		)

		location_usage = {
			location.id: 0
			for location in locations
		}

		activities = (
			EsExamActivity.objects
			.select_related("exam_requirement")
			.annotate(
				real_size=Count(
					"exam_requirement__esexamrequirementstudent",
					filter=Q(
						exam_requirement__esexamrequirementstudent__student__isnull=False
					),
					distinct=True,
				)
			)
			.filter(
				exam_requirement__exam_period=exam_period,
				is_scheduled=True,
				time_slot__isnull=False,
				location__isnull=False,
				location__status=EsLocation.STATUS_TO_CODE["active"],
			)
		)

		for activity in activities:
			scheduled_at = self.slot_to_datetime(
				activity.time_slot
			)

			if scheduled_at is None:
				continue

			opportunity = (
				scheduled_at.date(),
				scheduled_at.strftime("%H:%M"),
			)

			# Only count valid sessions within the exam period.
			if opportunity not in valid_opportunities:
				continue

			if activity.location_id not in location_usage:
				continue

			location_usage[activity.location_id] += (
				activity.real_size or 0
			)

		utilization = []

		for location in locations:
			capacity_per_session = max(
				location.total_seats or 0,
				0,
			)

			# Total available seats across the entire exam period.
			total = (
				capacity_per_session
				* available_session_count
			)

			used = location_usage[location.id]

			percent = (
				round((used / total) * 100, 2)
				if total > 0
				else 0
			)

			utilization.append({
				"location": location.name,
				"used": used,
				"total": total,
				"percent": percent,
			})

		# Sort by descending percent, then ascending location name.
		utilization.sort(
			key=lambda item: (
				-item["percent"],
				item["location"],
			)
		)

		return utilization[:10]
	
	def get_invigilator_workload_info(self, exam_period):
		#list top 10 invigilators with the most sessions in the given exam period
		workload = (
			EsSessionInvigilator.objects
			.filter(
				invigilator__status=EsInvigilator.STATUS_TO_CODE["active"],
				session__esexamactivity__exam_requirement__exam_period_id=exam_period.id,
				session__esexamactivity__is_scheduled=True,
			)
			.values("invigilator_id", "invigilator__name")
			.annotate(session_count=Count("session_id", distinct=True))
			.order_by("-session_count", "invigilator__name")[:10]
		)

		return {
			"names": [item["invigilator__name"] for item in workload],
			"sessions": [item["session_count"] for item in workload],
		}

	def get_invigilators_by_role(self, exam_period):
		invigilators_by_role = (
			EsInvigilatorRole.objects
			.filter(
				status=EsInvigilatorRole.STATUS_TO_CODE["active"]
			)
			.values(
				"id",
				"code",
				"name",
				"color",
			)
			.annotate(
				invigilator_count=Count(
					"esinvigilatorinvigilatorrole",
					filter=Q(
						esinvigilatorinvigilatorrole__invigilator__status=1
					),
					distinct=True,
				)
			)
			.order_by("name")
		)
		return [
			{
				"role": item["name"],
				"count": item["invigilator_count"],
				"color": item["color"],
			}
			for item in invigilators_by_role
		]

	def get_department_by_students(self, exam_period):
		students_by_department = (
			EsStudent.objects
			.filter(status=EsStudent.STATUS_TO_CODE["active"])
			.values("department__id", "department__name")
			.annotate(student_count=Count("id"))
			.order_by("department__name")
		)

		return [
			{
				"department": item["department__name"],
				"count": item["student_count"],
			}
			for item in students_by_department
		]
	
	def get_students_by_department(self, exam_period):
		students_by_department = (
			EsDepartment.objects
			.filter(status=EsDepartment.STATUS_TO_CODE["active"])
			.values("id", "name")
			.annotate(student_count=Count("esstudent"))
			.order_by("name")
		)

		return [
			{
				"department": item["name"],
				"count": item["student_count"],
			}
			for item in students_by_department
		]
	
	@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)
