import calendar
import re
from collections import defaultdict
from datetime import datetime, time, timedelta
from pathlib import Path
from zipfile import ZipFile
import xml.etree.ElementTree as ET

from django.conf import settings
from django.core.management.base import BaseCommand, CommandError
from django.db import transaction
from django.db.models import Count

from api.models import (
    EsDepartment,
    EsExamActivity,
    EsExamPeriod,
    EsExamPeriodUnavailability,
    EsExamRequirement,
    EsExamRequirementSameTime,
    EsExamRequirementSameTimeGroup,
    EsExamRequirementStudent,
    EsInvigilator,
    EsInvigilatorRole,
    EsLocation,
    EsPos,
    EsSession,
    EsSessionInvigilator,
    EsSessionSeat,
    EsSessionStart,
    EsSessionStartDay,
    EsStudent,
    EsStudentGroup,
    EsStudentStudentGroup,
)
from api.services.student_groups import replace_student_group_rows
from api.views.admin.exam_requirement_relation_grouping import update_requirement_relations
from api.views.admin.session_generate import SessionGenerate


DEMO_RELATIVE_DIR = Path("demo") / "20260902"

PERIOD_MAIN_CODE = "DEMO-EOS1-2026"
PERIOD_RESIT_CODE = "DEMO-RESIT-EOS1-2026"
PERIOD_CODES = (PERIOD_MAIN_CODE, PERIOD_RESIT_CODE)

PROGRAMMES = {
    "BP125": "Bachelor of Pharmacy (Honours)",
    "DN125": "Bachelor of Science (Hons) Dietetics with Nutrition",
    "FSI125": "Bachelor of Science (Honours) in Food Science Innovation",
    "PC225": "Bachelor of Science (Hons) in Pharmaceutical Chemistry",
}

OUR_PROGRAMME_CODES = tuple(PROGRAMMES.keys())
ACTIVE = 1
SLOT_SECONDS = 1800
W_NS = "{http://schemas.openxmlformats.org/wordprocessingml/2006/main}"
WEEKDAYS = [0, 1, 2, 3, 4]

ROOM_RE = re.compile(r"\d+\.\d+\.\d+|LIT\s+TL\s+[\d.]+", re.I)
MODULE_CODE_RE = re.compile(r"^([A-Z]{2,4}\s*\d{4})\b")
PROGRAMME_CODE_RE = re.compile(r"\b(BP125|DN125|FSI125|PC225)\b", re.I)
SHARE_RE = re.compile(r"\*?\s*\(?\*?\s*Share[d]?\s+with\b", re.I)
PDF_PAPER_CORE_RE = re.compile(
    r"\(\s*(?P<code>[A-Z]{2,4}\s*\d{4})\s*\)\s+"
    r"(?P<day>Monday|Tuesday|Wednesday|Thursday|Friday|Saturday|Sunday)\s+"
    r"(?P<date>\d{1,2}\s+[A-Za-z]+,?\s+\d{4})\s+"
    r"(?P<start>\d{1,2}:\d{2}\s*[AP]M)\s+"
    r"(?P<end>\d{1,2}:\d{2}\s*[AP]M)",
    re.I,
)
PDF_REST_CUT_RE = re.compile(
    r"\*\s*subject\b|\bSemester\s+\d+\s+Coordinator\b|\bPre Board\b|\bBoard of Examiners\b",
    re.I,
)
PDF_TITLE_STOP = {
    "remarks", "location", "locations", "time", "end", "start", "dates", "days",
    "module", "modules", "shared", "share", "with", "pm", "am",
    "monday", "tuesday", "wednesday", "thursday", "friday", "saturday", "sunday",
    "january", "february", "march", "april", "may", "june", "july", "august",
    "september", "october", "november", "december",
    "examination", "semester", "coordinator", "director", "dean", "students",
    "total", "no", "imu", "university", "resit",
    "bp125", "dn125", "fsi125", "pc225", "bm125", "mb125", "s1",
}


def require_optional_import(module_name, pip_name=None):
    try:
        return __import__(module_name)
    except ImportError as exc:
        package = pip_name or module_name
        raise CommandError(
            f"Missing Python package '{module_name}'. Install it with: pip install {package}"
        ) from exc


def cell_text(value):
    if value is None:
        return ""
    if isinstance(value, float) and value.is_integer():
        return str(int(value))
    return str(value).strip()


def normalize_space(value):
    return re.sub(r"\s+", " ", (value or "").replace("\u202f", " ")).strip()


def normalize_module_code(value):
    return re.sub(r"\s+", "", (value or "").upper())


def demo_dir():
    return Path(settings.BASE_DIR) / DEMO_RELATIVE_DIR


def snap_time(value):
    total = value.hour * 60 + value.minute
    snapped = ((total + 15) // 30) * 30
    if snapped >= 24 * 60:
        snapped = 24 * 60 - 30
    return time(snapped // 60, snapped % 60)


def datetime_to_slot(value):
    epoch_seconds = calendar.timegm(value.timetuple())
    if epoch_seconds < 0 or epoch_seconds % SLOT_SECONDS != 0:
        return None
    return int(epoch_seconds // SLOT_SECONDS)


def parse_clock(value):
    text = normalize_space(value).upper().replace(".", "")
    for fmt in ("%I:%M %p", "%I:%M%p", "%H:%M"):
        try:
            return datetime.strptime(text, fmt).time()
        except ValueError:
            continue
    raise CommandError(f"Unrecognised time: {value!r}")


def parse_date(value):
    text = normalize_space(value).replace(",", "")
    for fmt in ("%d/%m/%Y", "%d/%m/%y", "%d %B %Y", "%d %b %Y"):
        try:
            return datetime.strptime(text, fmt).date()
        except ValueError:
            continue
    raise CommandError(f"Unrecognised date: {value!r}")


def duration_time(start, end):
    start_dt = datetime.combine(datetime.min, start)
    end_dt = datetime.combine(datetime.min, end)
    if end_dt <= start_dt:
        end_dt += timedelta(days=1)
    delta = end_dt - start_dt
    if delta <= timedelta():
        delta = timedelta(hours=1)
    seconds = int(delta.total_seconds())
    hours, rem = divmod(seconds, 3600)
    minutes, secs = divmod(rem, 60)
    return time(hours, minutes, secs)


def expand_venue_shorthand(text):
    def replacer(match):
        previous = match.group(1)
        suffix = match.group(2)
        building = previous.rsplit(".", 1)[0]
        return f"{previous} {building}.{suffix}"

    return re.sub(r"(\d+\.\d+\.\d+)\s*(?:and|&)\s*(\d{2})\b", replacer, text, flags=re.I)


def parse_ci_codes(raw):
    text = normalize_space(raw).upper()
    if not text:
        return []
    parts = re.split(r"[\s,;/&]+", text)
    return list(dict.fromkeys(part for part in parts if re.fullmatch(r"[A-Z]{2,5}", part)))


def parse_venues(raw):
    text = expand_venue_shorthand(normalize_space(raw))
    if not text:
        return []
    rooms = [normalize_space(item) for item in ROOM_RE.findall(text)]
    if rooms:
        return list(dict.fromkeys(rooms))
    parts = [part.strip() for part in re.split(r"\s*(?:,|&| and )\s*", text, flags=re.I) if part.strip()]
    return list(dict.fromkeys(parts))


def split_venue_and_remarks(rest):
    text = normalize_space(rest)
    if not text:
        return "", ""
    share_match = SHARE_RE.search(text)
    if share_match:
        return text[: share_match.start()].strip(" -"), text[share_match.start():].strip(" -()")
    trailing = re.search(
        r"^(?P<venue>.*?)(?P<remarks>(?:\s+(?:[A-Z]{2,4}\d{3}(?:\s*S\d)?(?:,\s*)?))+)$",
        text,
        re.I,
    )
    if trailing and trailing.group("remarks").strip():
        return trailing.group("venue").strip(" -"), trailing.group("remarks").strip(" -")
    return text, ""


def strip_share_from_name(name):
    text = normalize_space(name)
    share_match = SHARE_RE.search(text)
    if share_match:
        text = text[: share_match.start()]
    return text.strip(" -()*")


def parse_module_cell(raw):
    text = normalize_space(raw)
    remarks = ""
    share_match = SHARE_RE.search(text)
    if share_match:
        remarks = text[share_match.start():].strip(" -()*")
        text = text[: share_match.start()].strip(" -()*")
    code_match = MODULE_CODE_RE.match(text)
    if not code_match:
        return None, strip_share_from_name(text), remarks
    code = normalize_module_code(code_match.group(1))
    title = text[code_match.end():].strip(" -")
    return code, title or code, remarks


def find_header_index(rows):
    for index, row in enumerate(rows):
        joined = " ".join(cell_text(cell) for cell in row).lower()
        if "student id" in joined and "full name" in joined:
            return index
    return None


def header_map(row):
    mapping = {}
    for index, cell in enumerate(row):
        key = normalize_space(cell_text(cell)).lower()
        if key:
            mapping[key] = index
    return mapping


def parse_student_rows(rows, programme):
    header_index = find_header_index(rows)
    if header_index is None:
        raise CommandError(f"Could not find student header row for {programme}")
    columns = header_map(rows[header_index])
    id_col = columns.get("student id")
    name_col = columns.get("full name")
    email_col = next((index for key, index in columns.items() if "email" in key), None)
    if id_col is None or name_col is None:
        raise CommandError(f"Student list for {programme} is missing Student ID or Full Name")

    students = []
    for row in rows[header_index + 1:]:
        code = cell_text(row[id_col] if id_col < len(row) else "")
        name = cell_text(row[name_col] if name_col < len(row) else "")
        if not code or not name:
            continue
        email = cell_text(row[email_col] if email_col is not None and email_col < len(row) else "") or None
        students.append({
            "code": code,
            "name": name,
            "email": email,
            "programme": programme,
        })
    return students


def load_xlsx_rows(path):
    openpyxl = require_optional_import("openpyxl")
    workbook = openpyxl.load_workbook(path, data_only=True)
    sheet = workbook.active
    return [list(row) for row in sheet.iter_rows(values_only=True)]


def load_xls_rows(path):
    xlrd = require_optional_import("xlrd")
    book = xlrd.open_workbook(path)
    sheet = book.sheet_by_index(0)
    return [[sheet.cell_value(row_idx, col_idx) for col_idx in range(sheet.ncols)] for row_idx in range(sheet.nrows)]


def parse_student_file(path):
    match = re.search(r"StudentList-([A-Z0-9]+)\.", path.name, re.I)
    if not match:
        raise CommandError(f"Cannot determine programme from {path.name}")
    programme = match.group(1).upper()
    if path.suffix.lower() == ".xlsx":
        rows = load_xlsx_rows(path)
    elif path.suffix.lower() == ".xls":
        rows = load_xls_rows(path)
    else:
        raise CommandError(f"Unsupported student list: {path.name}")
    return parse_student_rows(rows, programme)


def word_cell_text(cell):
    return normalize_space("".join((node.text or "") for node in cell.iter(f"{W_NS}t")))


def iter_docx_tables(path):
    with ZipFile(path) as archive:
        root = ET.fromstring(archive.read("word/document.xml"))
    for table in root.iter(f"{W_NS}tbl"):
        rows = []
        for table_row in table.findall(f"{W_NS}tr"):
            rows.append([word_cell_text(cell) for cell in table_row.findall(f"{W_NS}tc")])
        if rows:
            yield rows


def is_exam_table(rows):
    if not rows:
        return False
    header = " ".join(rows[0]).lower()
    return "module" in header and "start time" in header and "date" in header


def table_column_map(header):
    mapping = {}
    for index, cell in enumerate(header):
        key = normalize_space(cell).lower()
        if key:
            mapping[key] = index
    return mapping


def paper_from_row(row, columns, programme, sitting):
    module_idx = next((columns[key] for key in columns if "module" in key), 0)
    date_idx = next((columns[key] for key in columns if "date" in key), None)
    start_idx = next((columns[key] for key in columns if "start" in key), None)
    end_idx = next((columns[key] for key in columns if "end" in key), None)
    venue_idx = next((columns[key] for key in columns if key in ("venue", "location") or "venue" in key or "location" in key), None)
    ci_idx = next((columns[key] for key in columns if key in ("ci", "chief invigilator") or key.startswith("ci ")), None)
    if date_idx is None or start_idx is None or end_idx is None:
        return None

    def col(index):
        if index is None or index >= len(row):
            return ""
        return row[index]

    module_raw = col(module_idx)
    if not module_raw or module_raw.lower().startswith("module"):
        return None
    module_code, title, name_remarks = parse_module_cell(module_raw)
    if not module_code:
        return None
    venue_raw, venue_remarks = split_venue_and_remarks(col(venue_idx))
    remarks = normalize_space(" ".join(part for part in (name_remarks, venue_remarks) if part))
    start = parse_clock(col(start_idx))
    end = parse_clock(col(end_idx))
    exam_date = parse_date(col(date_idx))
    return {
        "programme": programme,
        "sitting": sitting,
        "module_code": module_code,
        "name": title or module_code,
        "exam_date": exam_date,
        "start_time": start,
        "end_time": end,
        "venues": parse_venues(venue_raw),
        "remarks": remarks,
        "ci_codes": parse_ci_codes(col(ci_idx)),
    }


def parse_docx_timetable(path, programme):
    exam_tables = [rows for rows in iter_docx_tables(path) if is_exam_table(rows)]
    if len(exam_tables) < 2:
        raise CommandError(f"{path.name} does not contain main and resit exam tables")
    papers = []
    for sitting, rows in (("main", exam_tables[0]), ("resit", exam_tables[1])):
        columns = table_column_map(rows[0])
        for row in rows[1:]:
            paper = paper_from_row(row, columns, programme, sitting)
            if paper:
                papers.append(paper)
    return papers


def extract_pdf_module_name(prefix):
    collected = []
    for word in reversed(normalize_space(prefix).split()):
        cleaned = re.sub(r"[^A-Za-z0-9/]+", "", word).lower()
        if not cleaned:
            continue
        if cleaned in PDF_TITLE_STOP and cleaned not in {"of", "in", "and", "the"}:
            break
        if re.fullmatch(r"\d+", cleaned) or ROOM_RE.fullmatch(word) or re.fullmatch(r"\d{1,2}:\d{2}", word):
            break
        collected.append(word.strip("()"))
        if len(collected) >= 8:
            break
    collected.reverse()
    while collected and collected[0].lower() in {"of", "in", "and", "the"}:
        collected.pop(0)
    return normalize_space(" ".join(collected))


def parse_pdf_page_papers(text, programme, sitting):
    collapsed = normalize_space(text)
    matches = list(PDF_PAPER_CORE_RE.finditer(collapsed))
    papers = []
    for index, match in enumerate(matches):
        name = extract_pdf_module_name(collapsed[: match.start()])
        next_start = matches[index + 1].start() if index + 1 < len(matches) else len(collapsed)
        rest = collapsed[match.end(): next_start].strip()
        if index + 1 < len(matches):
            next_name = extract_pdf_module_name(collapsed[: matches[index + 1].start()])
            if next_name and rest.lower().endswith(next_name.lower()):
                rest = rest[: -len(next_name)].strip()
        rest = PDF_REST_CUT_RE.split(rest, maxsplit=1)[0].strip()
        venue_raw, remarks = split_venue_and_remarks(rest)
        module_code = normalize_module_code(match.group("code"))
        papers.append({
            "programme": programme,
            "sitting": sitting,
            "module_code": module_code,
            "name": name or module_code,
            "exam_date": parse_date(match.group("date")),
            "start_time": parse_clock(match.group("start")),
            "end_time": parse_clock(match.group("end")),
            "venues": parse_venues(venue_raw),
            "remarks": remarks,
            "ci_codes": [],
        })
    return papers


def parse_pdf_timetable(path, programme):
    pypdf = require_optional_import("pypdf")
    reader = pypdf.PdfReader(str(path))
    papers = []
    for page in reader.pages:
        text = page.extract_text() or ""
        sitting = "resit" if re.search(r"\bresit\b", text, re.I) else "main"
        papers.extend(parse_pdf_page_papers(text, programme, sitting))
    if not papers:
        raise CommandError(f"No exam papers parsed from {path.name}")
    return papers


def parse_timetable_file(path):
    match = re.search(r"Examination-Timetable\s+([A-Z0-9]+)\.", path.name, re.I)
    if not match:
        raise CommandError(f"Cannot determine programme from {path.name}")
    programme = match.group(1).upper()
    suffix = path.suffix.lower()
    if suffix == ".docx":
        return parse_docx_timetable(path, programme)
    if suffix == ".pdf":
        return parse_pdf_timetable(path, programme)
    raise CommandError(f"Unsupported timetable: {path.name}")


def requirement_code(paper):
    suffix = "-R" if paper["sitting"] == "resit" else ""
    return f"{paper['programme']}-{paper['module_code']}{suffix}"


def get_or_create_by_code(model, code, defaults, update_existing=True):
    obj = model.objects.filter(code=code).order_by("id").first()
    if obj is None:
        return model.objects.create(code=code, **defaults), True
    if update_existing:
        dirty = False
        for key, value in defaults.items():
            if getattr(obj, key) != value:
                setattr(obj, key, value)
                dirty = True
        if dirty:
            obj.save()
    return obj, False


class Command(BaseCommand):
    help = "Import demo/20260902 student lists and exam timetables for demo use."

    def add_arguments(self, parser):
        parser.add_argument(
            "--reset",
            action="store_true",
            help="Delete previously imported demo exam periods, sessions, and students, then re-import.",
        )
        parser.add_argument(
            "--dry-run",
            action="store_true",
            help="Parse files and print a summary without writing to the database.",
        )

    def handle(self, *args, **options):
        source_dir = demo_dir()
        if not source_dir.is_dir():
            raise CommandError(f"Demo folder not found: {source_dir}")

        students_by_programme, papers = self.parse_sources(source_dir)
        self.write_parse_summary(students_by_programme, papers)
        if options["dry_run"]:
            self.stdout.write(self.style.WARNING("Dry run only — no database changes."))
            return

        with transaction.atomic():
            if options["reset"]:
                self.reset_demo_data(students_by_programme)
            self.import_demo(students_by_programme, papers)

        self.stdout.write(self.style.SUCCESS("Demo import complete."))

    def parse_sources(self, source_dir):
        students_by_programme = defaultdict(list)
        for path in sorted(source_dir.glob("StudentList-*")):
            for student in parse_student_file(path):
                students_by_programme[student["programme"]].append(student)

        papers = []
        for path in sorted(source_dir.glob("Examination-Timetable*")):
            papers.extend(parse_timetable_file(path))
        if not papers:
            raise CommandError("No timetable papers were parsed.")
        return students_by_programme, papers

    def write_parse_summary(self, students_by_programme, papers):
        self.stdout.write("Students (by student-list file code):")
        for programme in OUR_PROGRAMME_CODES:
            paper_count = sum(1 for paper in papers if paper["programme"] == programme)
            self.stdout.write(
                f"  {programme}: {len(students_by_programme.get(programme, []))} students "
                f"-> {paper_count} modules"
            )
        self.stdout.write(f"Papers: {len(papers)}")
        for sitting in ("main", "resit"):
            sitting_papers = [paper for paper in papers if paper["sitting"] == sitting]
            self.stdout.write(f"  {sitting}: {len(sitting_papers)}")
        ci_codes = sorted({code for paper in papers for code in paper.get("ci_codes") or []})
        self.stdout.write(f"Chief invigilators (CI): {', '.join(ci_codes) if ci_codes else '(none parsed)'}")

    def reset_demo_data(self, students_by_programme):
        periods = list(EsExamPeriod.objects.filter(code__in=PERIOD_CODES))
        period_ids = [period.id for period in periods]
        session_ids = list(
            EsExamActivity.objects.filter(exam_requirement__exam_period_id__in=period_ids)
            .exclude(session_id=None)
            .values_list("session_id", flat=True)
            .distinct()
        )
        if session_ids:
            EsSessionSeat.objects.filter(session_id__in=session_ids).delete()
            EsSessionInvigilator.objects.filter(session_id__in=session_ids).delete()
            EsExamActivity.objects.filter(session_id__in=session_ids).update(session=None)
            EsSession.objects.filter(id__in=session_ids).delete()

        if period_ids:
            EsSessionStartDay.objects.filter(session_start__exam_period_id__in=period_ids).delete()
            EsSessionStart.objects.filter(exam_period_id__in=period_ids).delete()
            EsExamPeriodUnavailability.objects.filter(exam_period_id__in=period_ids).delete()
            EsExamPeriod.objects.filter(id__in=period_ids).delete()
            EsExamRequirementSameTimeGroup.objects.annotate(
                n=Count("esexamrequirementsametime")
            ).filter(n=0).delete()

        student_codes = [
            student["code"]
            for students in students_by_programme.values()
            for student in students
        ]
        if student_codes:
            demo_students = EsStudent.objects.filter(
                code__in=student_codes,
                name__regex=r"^(BP125|DN125|FSI125|PC225)_",
            )
            demo_ids = list(demo_students.values_list("id", flat=True))
            EsStudentStudentGroup.objects.filter(student_id__in=demo_ids).delete()
            EsExamRequirementStudent.objects.filter(student_id__in=demo_ids).delete()
            demo_students.delete()

        self.stdout.write(self.style.WARNING("Reset previous demo exam periods, sessions, and students."))

    def import_demo(self, students_by_programme, papers):
        department, _ = get_or_create_by_code(
            EsDepartment,
            "IMU",
            {"name": "IMU University", "status": ACTIVE},
        )
        programmes = {}
        groups = {}
        for code, name in PROGRAMMES.items():
            pos, _ = get_or_create_by_code(
                EsPos,
                code,
                {"name": name, "desc": None, "department": department, "status": ACTIVE},
            )
            group, _ = get_or_create_by_code(
                EsStudentGroup,
                f"{code}-S1",
                {
                    "name": f"{name} Semester 1",
                    "desc": None,
                    "department": department,
                    "status": ACTIVE,
                },
            )
            programmes[code] = pos
            groups[code] = group

        students_by_id = {}
        for programme, students in students_by_programme.items():
            pos = programmes[programme]
            group = groups[programme]
            for row in students:
                student, _ = get_or_create_by_code(
                    EsStudent,
                    row["code"],
                    {
                        "name": row["name"],
                        "email": row["email"],
                        "phone": None,
                        "status": ACTIVE,
                        "need_extra_provision": False,
                        "department": department,
                        "enrolled_programme": pos,
                    },
                )
                replace_student_group_rows(student.id, [group.id])
                students_by_id.setdefault(programme, []).append(student)

        self.stdout.write(f"Students upserted: {sum(len(items) for items in students_by_id.values())}")

        venues = []
        for paper in papers:
            venues.extend(paper["venues"])
        locations = {}
        for venue in dict.fromkeys(venues):
            location, created = get_or_create_by_code(
                EsLocation,
                venue,
                {
                    "name": venue,
                    "seat_order": EsLocation.SEAT_ORDER_TO_CODE["column_first"],
                    "row": 10,
                    "column": 10,
                    "status": ACTIVE,
                },
                update_existing=False,
            )
            if created:
                self.stdout.write(f"Created location {venue}")
            locations[venue] = location

        periods = {
            "main": get_or_create_by_code(
                EsExamPeriod,
                PERIOD_MAIN_CODE,
                {
                    "name": "End-of-Semester 1 2026",
                    "start_date": datetime(2026, 1, 26).date(),
                    "end_date": datetime(2026, 2, 7).date(),
                    "status": ACTIVE,
                },
            )[0],
            "resit": get_or_create_by_code(
                EsExamPeriod,
                PERIOD_RESIT_CODE,
                {
                    "name": "Resit End-of-Semester 1 2026",
                    "start_date": datetime(2026, 2, 23).date(),
                    "end_date": datetime(2026, 2, 28).date(),
                    "status": ACTIVE,
                },
            )[0],
        }

        start_times_by_sitting = defaultdict(set)
        for paper in papers:
            start_times_by_sitting[paper["sitting"]].add(snap_time(paper["start_time"]))
        for sitting, times in start_times_by_sitting.items():
            period = periods[sitting]
            for start in sorted(times):
                code = f"{period.code}-{start.strftime('%H%M')}"
                session_start, _ = get_or_create_by_code(
                    EsSessionStart,
                    code,
                    {
                        "exam_period": period,
                        "start_time": start,
                        "status": ACTIVE,
                    },
                )
                EsSessionStartDay.bulk_insert(session_start.id, WEEKDAYS)

        requirements = {}
        for paper in papers:
            code = requirement_code(paper)
            planned_size = len(students_by_programme.get(paper["programme"], []))
            requirement, _ = get_or_create_by_code(
                EsExamRequirement,
                code,
                {
                    "exam_period": periods[paper["sitting"]],
                    "name": paper["name"][:150],
                    "description": paper["remarks"] or None,
                    "planned_size": planned_size,
                    "writing_time": duration_time(paper["start_time"], paper["end_time"]),
                    "reading_time": time(0, 0, 0),
                    "fixed_start_date": paper["exam_date"],
                    "fixed_start_time": snap_time(paper["start_time"]),
                    "location_required": True,
                    "exclusive_use": False,
                },
            )
            paper["requirement"] = requirement
            requirements[code] = requirement
            self.sync_activities(paper, requirement, locations)
            assigned = self.assign_students(
                requirement,
                students_by_id.get(paper["programme"], []),
            )
            self.stdout.write(
                f"  {code}: assigned {assigned} students from StudentList-{paper['programme']}"
            )

        self.stdout.write(f"Exam requirements upserted: {len(requirements)}")
        self.apply_same_time_groups(papers)
        self.schedule_activities(papers, locations, periods)
        self.assign_session_invigilators(papers, department)
        self.stdout.write("Demo data is ready.")

    def sync_activities(self, paper, requirement, locations):
        venues = paper["venues"] or [None]
        existing = {
            activity.code: activity
            for activity in EsExamActivity.objects.filter(exam_requirement=requirement).order_by("id")
        }
        kept_codes = []
        for index, venue in enumerate(venues, start=1):
            activity_code = f"{requirement.code}-{index:02d}"
            kept_codes.append(activity_code)
            defaults = {
                "exam_requirement": requirement,
                "name": requirement.name,
                "is_scheduled": False,
                "location": locations.get(venue) if venue else None,
                "time_slot": None,
                "session": None,
            }
            activity = existing.get(activity_code)
            if activity is None:
                EsExamActivity.objects.create(code=activity_code, **defaults)
            else:
                activity.name = requirement.name
                activity.location = defaults["location"]
                activity.save(update_fields=["name", "location", "updated_at"])

        extras = EsExamActivity.objects.filter(exam_requirement=requirement).exclude(code__in=kept_codes)
        extra_session_ids = list(extras.exclude(session_id=None).values_list("session_id", flat=True))
        extras.delete()
        if extra_session_ids:
            empty_sessions = EsSession.objects.filter(id__in=extra_session_ids).annotate(
                n=Count("esexamactivity")
            ).filter(n=0)
            empty_ids = list(empty_sessions.values_list("id", flat=True))
            if empty_ids:
                EsSessionSeat.objects.filter(session_id__in=empty_ids).delete()
                EsSessionInvigilator.objects.filter(session_id__in=empty_ids).delete()
                empty_sessions.delete()

    def assign_students(self, requirement, students):
        student_ids = [student.id for student in students]
        existing = set(
            EsExamRequirementStudent.objects.filter(exam_requirement=requirement)
            .exclude(student_id=None)
            .values_list("student_id", flat=True)
        )
        wanted = set(student_ids)
        extra_ids = existing - wanted
        if extra_ids:
            EsExamRequirementStudent.objects.filter(
                exam_requirement=requirement,
                student_id__in=extra_ids,
            ).delete()
        to_create = [
            EsExamRequirementStudent(exam_requirement=requirement, student=student)
            for student in students
            if student.id not in existing
        ]
        if to_create:
            EsExamRequirementStudent.objects.bulk_create(to_create, ignore_conflicts=True)
        return len(wanted)

    def assign_session_invigilators(self, papers, department):
        role, _ = get_or_create_by_code(
            EsInvigilatorRole,
            "CI",
            {"name": "Chief Invigilator", "desc": "Chief invigilator from examination timetable", "status": ACTIVE},
        )
        invigilators = {}
        for paper in papers:
            for ci_code in paper.get("ci_codes") or []:
                invigilator, created = get_or_create_by_code(
                    EsInvigilator,
                    ci_code,
                    {
                        "name": ci_code,
                        "status": ACTIVE,
                        "department": department,
                    },
                    update_existing=False,
                )
                invigilators[ci_code] = invigilator
                if created:
                    self.stdout.write(f"Created invigilator {ci_code}")
                invigilator.role.add(role)

        assigned_sessions = 0
        for paper in papers:
            ci_codes = paper.get("ci_codes") or []
            if not ci_codes:
                continue
            invigilator_ids = [invigilators[code].id for code in ci_codes if code in invigilators]
            if not invigilator_ids:
                continue
            activities = (
                EsExamActivity.objects
                .filter(exam_requirement=paper["requirement"], session_id__isnull=False)
                .select_related("session")
            )
            for activity in activities:
                session = activity.session
                existing = set(
                    EsSessionInvigilator.objects.filter(
                        session_id=session.id,
                        invigilator_id__isnull=False,
                    ).values_list("invigilator_id", flat=True)
                )
                to_add = [invigilator_id for invigilator_id in invigilator_ids if invigilator_id not in existing]
                if to_add:
                    EsSessionInvigilator.bulk_insert(session.id, to_add)
                total = len(existing | set(invigilator_ids))
                if session.invigilators_required != total:
                    session.invigilators_required = total
                    session.save(update_fields=["invigilators_required", "updated_at"])
                assigned_sessions += 1
        self.stdout.write(f"Sessions with CI assigned: {assigned_sessions}")

    def apply_same_time_groups(self, papers):
        by_key = defaultdict(list)
        for paper in papers:
            key = (paper["sitting"], paper["exam_date"], paper["start_time"])
            by_key[key].append(paper)

        parent = {}

        def find(value):
            parent.setdefault(value, value)
            while parent[value] != value:
                parent[value] = parent[parent[value]]
                value = parent[value]
            return value

        def union(left, right):
            left_root, right_root = find(left), find(right)
            if left_root != right_root:
                parent[right_root] = left_root

        for group_papers in by_key.values():
            for paper in group_papers:
                mentioned = {
                    code.upper()
                    for code in PROGRAMME_CODE_RE.findall(paper["remarks"] or "")
                    if code.upper() in OUR_PROGRAMME_CODES and code.upper() != paper["programme"]
                }
                if not mentioned:
                    continue
                for other in group_papers:
                    if other["programme"] in mentioned:
                        union(requirement_code(paper), requirement_code(other))

        groups = defaultdict(set)
        for code in parent:
            groups[find(code)].add(code)

        grouped = 0
        for members in groups.values():
            if len(members) < 2:
                continue
            requirement_ids = list(
                EsExamRequirement.objects.filter(code__in=members).values_list("id", flat=True)
            )
            activity_ids = list(
                EsExamActivity.objects.filter(exam_requirement_id__in=requirement_ids).values_list("id", flat=True)
            )
            for requirement_id in requirement_ids:
                update_requirement_relations(
                    EsExamRequirementSameTime,
                    EsExamRequirementSameTimeGroup,
                    requirement_id,
                    activity_ids,
                )
            grouped += 1
        self.stdout.write(f"Same-time groups: {grouped}")

    def schedule_activities(self, papers, locations, periods):
        activities = EsExamActivity.objects.filter(
            exam_requirement__exam_period__code__in=PERIOD_CODES
        )
        session_ids = list(activities.exclude(session_id=None).values_list("session_id", flat=True).distinct())
        activities.update(session=None)
        if session_ids:
            empty_sessions = EsSession.objects.filter(id__in=session_ids).annotate(
                n=Count("esexamactivity")
            ).filter(n=0)
            empty_ids = list(empty_sessions.values_list("id", flat=True))
            if empty_ids:
                EsSessionSeat.objects.filter(session_id__in=empty_ids).delete()
                EsSessionInvigilator.objects.filter(session_id__in=empty_ids).delete()
                empty_sessions.delete()

        scheduled = 0
        for paper in papers:
            snapped_start = datetime.combine(paper["exam_date"], snap_time(paper["start_time"]))
            time_slot = datetime_to_slot(snapped_start)
            if time_slot is None:
                raise CommandError(f"Could not convert {snapped_start} to a 30-minute time slot")
            venues = paper["venues"] or [None]
            for index, venue in enumerate(venues, start=1):
                activity_code = f"{paper['requirement'].code}-{index:02d}"
                updated = EsExamActivity.objects.filter(code=activity_code).update(
                    is_scheduled=True,
                    time_slot=time_slot,
                    location=locations.get(venue) if venue else None,
                    session=None,
                )
                scheduled += updated
        self.stdout.write(f"Activities scheduled: {scheduled}")

        generator = SessionGenerate()
        for sitting, period in periods.items():
            created = generator.generate_sessions(None, period.id)
            self.stdout.write(f"Sessions generated for {sitting}: {len(created)}")
