from fastapi import APIRouter, Depends, HTTPException, Query, status
from sqlalchemy.orm import Session, selectinload

from app.database import get_db
from app.dependencies import require_admin_token
from app.models import User
from app.programme.models_programme import (
    ProgrammeAssessment,
    ProgrammeLevel,
    ProgrammeModule,
    normalize_programme_level_slug,
    normalize_programme_module_kind,
    programme_level_slug_str,
    programme_module_kind_str,
    ProgrammeSession,
    ProgrammeWeek,
)
from app.programme.services.programme_service import serialize_programme_catalog
from app.schemas import (
    Message,
    ProgrammeAssessmentCreateRequest,
    ProgrammeAssessmentOut,
    ProgrammeAssessmentUpdateRequest,
    ProgrammeCatalogResponse,
    ProgrammeIdRequest,
    ProgrammeLevelCreateRequest,
    ProgrammeLevelOut,
    ProgrammeLevelTreeCreateRequest,
    ProgrammeLevelUpdateRequest,
    ProgrammeModuleCreateRequest,
    ProgrammeModuleOut,
    ProgrammeModuleUpdateRequest,
    ProgrammeSessionCreateRequest,
    ProgrammeSessionOut,
    ProgrammeSessionUpdateRequest,
    ProgrammeWeekCreateRequest,
    ProgrammeWeekOut,
    ProgrammeWeekUpdateRequest,
)

router = APIRouter(tags=["admin-programme"])


# ---------------------------------------------------------------------------
# Helpers
# ---------------------------------------------------------------------------

def _parse_level_slug(level: str | None) -> str | None:
    if level is None or not level.strip():
        return None
    try:
        return normalize_programme_level_slug(level)
    except ValueError as e:
        raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e


def _serialize_assessment(a: ProgrammeAssessment) -> dict:
    return {"id": a.id, "sort_order": a.sort_order, "title": a.title}


def _serialize_session(s: ProgrammeSession) -> dict:
    assessments = sorted(s.assessments, key=lambda x: x.sort_order)
    return {
        "id": s.id,
        "session_number": s.session_number,
        "title": s.title,
        "admin_checked": s.admin_checked,
        "assessments_count": len(assessments),
        "assessments": [_serialize_assessment(a) for a in assessments],
    }


def _serialize_week(w: ProgrammeWeek) -> dict:
    sessions = sorted(w.sessions, key=lambda x: x.session_number)
    total_assessments = sum(len(s.assessments) for s in sessions)
    return {
        "id": w.id,
        "week_index": w.week_index,
        "display_label": w.display_label,
        "sessions_count": len(sessions),
        "assessments_count": total_assessments,
        "sessions": [_serialize_session(s) for s in sessions],
    }


def _serialize_module(m: ProgrammeModule) -> dict:
    weeks = sorted(m.weeks, key=lambda x: x.week_index)
    total_sessions = sum(len(w.sessions) for w in weeks)
    total_assessments = sum(len(s.assessments) for w in weeks for s in w.sessions)
    return {
        "id": m.id,
        "kind": programme_module_kind_str(m.kind),
        "display_name": m.display_name,
        "sort_order": m.sort_order,
        "weeks_count": len(weeks),
        "sessions_count": total_sessions,
        "assessments_count": total_assessments,
        "weeks": [_serialize_week(w) for w in weeks],
    }


def _load_level_full(db: Session, level_id: int) -> ProgrammeLevel:
    level = (
        db.query(ProgrammeLevel)
        .options(
            selectinload(ProgrammeLevel.modules)
            .selectinload(ProgrammeModule.weeks)
            .selectinload(ProgrammeWeek.sessions)
            .selectinload(ProgrammeSession.assessments)
        )
        .filter(ProgrammeLevel.id == level_id)
        .first()
    )
    if not level:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Level not found")
    return level


def _level_out(level: ProgrammeLevel) -> dict:
    modules = sorted(level.modules, key=lambda x: x.sort_order)
    total_weeks = sum(len(m.weeks) for m in modules)
    total_sessions = sum(len(w.sessions) for m in modules for w in m.weeks)
    total_assessments = sum(len(s.assessments) for m in modules for w in m.weeks for s in w.sessions)
    return {
        "id": level.id,
        "slug": programme_level_slug_str(level.slug),
        "display_name": level.display_name,
        "sort_order": level.sort_order,
        "modules_count": len(modules),
        "weeks_count": total_weeks,
        "sessions_count": total_sessions,
        "assessments_count": total_assessments,
        "modules": [_serialize_module(m) for m in modules],
    }


# ---------------------------------------------------------------------------
# Catalogue (read)
# ---------------------------------------------------------------------------

@router.get(
    "/programme/catalog",
    response_model=ProgrammeCatalogResponse,
    summary="Get full programme hierarchy",
)
def get_programme_catalog(
    level: str | None = Query(None, description="Optional level filter: beginner, hb, intermediate, hi, advanced"),
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    return serialize_programme_catalog(db, level_slug=_parse_level_slug(level))


# ---------------------------------------------------------------------------
# Levels
# ---------------------------------------------------------------------------

@router.get(
    "/programme/levels",
    response_model=list[ProgrammeLevelOut],
    summary="List all programme levels",
)
def list_levels(
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    levels = (
        db.query(ProgrammeLevel)
        .options(
            selectinload(ProgrammeLevel.modules)
            .selectinload(ProgrammeModule.weeks)
            .selectinload(ProgrammeWeek.sessions)
            .selectinload(ProgrammeSession.assessments)
        )
        .order_by(ProgrammeLevel.sort_order)
        .all()
    )
    return [ProgrammeLevelOut.model_validate(_level_out(lv)) for lv in levels]


@router.get(
    "/programme/levels/{level_id}",
    response_model=ProgrammeLevelOut,
    summary="Get one level with full tree",
)
def get_level(
    level_id: int,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    level = _load_level_full(db, level_id)
    return ProgrammeLevelOut.model_validate(_level_out(level))


@router.post(
    "/programme/levels",
    response_model=ProgrammeLevelOut,
    status_code=status.HTTP_201_CREATED,
    summary="Create a new programme level",
)
def create_level(
    body: ProgrammeLevelCreateRequest,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    try:
        slug_str = normalize_programme_level_slug(body.slug)
    except ValueError as e:
        raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
    if db.query(ProgrammeLevel).filter(ProgrammeLevel.slug == slug_str).first():
        raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Level with this slug already exists")
    level = ProgrammeLevel(slug=slug_str, display_name=body.display_name.strip(), sort_order=body.sort_order)
    db.add(level)
    db.commit()
    level = _load_level_full(db, level.id)
    return ProgrammeLevelOut.model_validate(_level_out(level))


@router.post(
    "/programme/levels/full",
    response_model=ProgrammeLevelOut,
    status_code=status.HTTP_201_CREATED,
    summary="Create a new programme level with full nested tree (modules/weeks/sessions/assessments)",
)
def create_level_full(
    body: ProgrammeLevelTreeCreateRequest,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    """
    Admin UI helper: create a full custom level in one request.
    """
    try:
        slug_str = normalize_programme_level_slug(body.slug)
    except ValueError as e:
        raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e

    if db.query(ProgrammeLevel).filter(ProgrammeLevel.slug == slug_str).first():
        raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Level with this slug already exists")

    # Create everything inside one transaction so partial trees aren't saved.
    try:
        level = ProgrammeLevel(
            slug=slug_str,
            display_name=body.display_name.strip(),
            sort_order=body.sort_order,
        )
        db.add(level)
        db.flush()  # allocate level.id

        for mod_in in body.modules:
            try:
                kind_str = normalize_programme_module_kind(mod_in.kind)
            except ValueError as e:
                raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e

            mod = ProgrammeModule(
                level_id=level.id,
                kind=kind_str,
                display_name=mod_in.display_name.strip(),
                sort_order=mod_in.sort_order,
            )
            db.add(mod)
            db.flush()  # allocate mod.id

            # prevent duplicates in the same module payload
            seen_week_indexes: set[int] = set()
            for week_in in mod_in.weeks:
                if week_in.week_index in seen_week_indexes:
                    raise HTTPException(
                        status_code=status.HTTP_400_BAD_REQUEST,
                        detail=f"Duplicate week_index {week_in.week_index} in module {kind_str}",
                    )
                seen_week_indexes.add(week_in.week_index)

                week = ProgrammeWeek(
                    module_id=mod.id,
                    week_index=week_in.week_index,
                    display_label=week_in.display_label.strip(),
                )
                db.add(week)
                db.flush()  # allocate week.id

                seen_session_numbers: set[int] = set()
                for sess_in in week_in.sessions:
                    if sess_in.session_number in seen_session_numbers:
                        raise HTTPException(
                            status_code=status.HTTP_400_BAD_REQUEST,
                            detail=f"Duplicate session_number {sess_in.session_number} in week_index {week_in.week_index}",
                        )
                    seen_session_numbers.add(sess_in.session_number)

                    sess = ProgrammeSession(
                        week_id=week.id,
                        session_number=sess_in.session_number,
                        title=sess_in.title.strip(),
                        admin_checked=bool(sess_in.admin_checked),
                    )
                    db.add(sess)
                    db.flush()  # allocate session.id

                    seen_sort_orders: set[int] = set()
                    for a_in in sess_in.assessments:
                        if a_in.sort_order in seen_sort_orders:
                            raise HTTPException(
                                status_code=status.HTTP_400_BAD_REQUEST,
                                detail=f"Duplicate assessment sort_order {a_in.sort_order} in session_number {sess_in.session_number}",
                            )
                        seen_sort_orders.add(a_in.sort_order)
                        db.add(
                            ProgrammeAssessment(
                                session_id=sess.id,
                                sort_order=a_in.sort_order,
                                title=a_in.title.strip(),
                            )
                        )

        db.commit()
    except HTTPException:
        db.rollback()
        raise
    except Exception as e:
        db.rollback()
        raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e

    level = _load_level_full(db, level.id)
    return ProgrammeLevelOut.model_validate(_level_out(level))


@router.put(
    "/programme/levels/{level_id}",
    response_model=ProgrammeLevelOut,
    summary="Update a programme level display name / sort order",
)
def update_level(
    level_id: int,
    body: ProgrammeLevelUpdateRequest,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    level = db.get(ProgrammeLevel, level_id)
    if not level:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Level not found")
    if body.display_name is not None:
        level.display_name = body.display_name.strip()
    if body.sort_order is not None:
        level.sort_order = body.sort_order
    db.commit()
    level = _load_level_full(db, level.id)
    return ProgrammeLevelOut.model_validate(_level_out(level))


@router.delete(
    "/programme/levels/{level_id}",
    response_model=Message,
    summary="Delete a programme level (cascades to modules/weeks/sessions/assessments)",
)
def delete_level(
    level_id: int,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    level = db.get(ProgrammeLevel, level_id)
    if not level:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Level not found")
    db.delete(level)
    db.commit()
    return Message(message="Level deleted")


# ---------------------------------------------------------------------------
# Modules
# ---------------------------------------------------------------------------

@router.get(
    "/programme/levels/{level_id}/modules",
    response_model=list[ProgrammeModuleOut],
    summary="List modules for a level",
)
def list_modules(
    level_id: int,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    level = db.get(ProgrammeLevel, level_id)
    if not level:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Level not found")
    modules = (
        db.query(ProgrammeModule)
        .options(
            selectinload(ProgrammeModule.weeks)
            .selectinload(ProgrammeWeek.sessions)
            .selectinload(ProgrammeSession.assessments)
        )
        .filter(ProgrammeModule.level_id == level_id)
        .order_by(ProgrammeModule.sort_order)
        .all()
    )
    return [ProgrammeModuleOut.model_validate(_serialize_module(m)) for m in modules]


@router.post(
    "/programme/modules",
    response_model=ProgrammeModuleOut,
    status_code=status.HTTP_201_CREATED,
    summary="Create a module inside a level",
)
def create_module(
    body: ProgrammeModuleCreateRequest,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    if not db.get(ProgrammeLevel, body.level_id):
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Level not found")
    try:
        kind_str = normalize_programme_module_kind(body.kind)
    except ValueError as e:
        raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(e)) from e
    existing = (
        db.query(ProgrammeModule)
        .filter(ProgrammeModule.level_id == body.level_id, ProgrammeModule.kind == kind_str)
        .first()
    )
    if existing:
        raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Module with this type already exists in this level")
    mod = ProgrammeModule(
        level_id=body.level_id,
        kind=kind_str,
        display_name=body.display_name.strip(),
        sort_order=body.sort_order,
    )
    db.add(mod)
    db.commit()
    db.refresh(mod)
    mod = (
        db.query(ProgrammeModule)
        .options(selectinload(ProgrammeModule.weeks).selectinload(ProgrammeWeek.sessions).selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeModule.id == mod.id)
        .first()
    )
    return ProgrammeModuleOut.model_validate(_serialize_module(mod))


@router.put(
    "/programme/modules/{module_id}",
    response_model=ProgrammeModuleOut,
    summary="Update a module display name / sort order",
)
def update_module(
    module_id: int,
    body: ProgrammeModuleUpdateRequest,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    mod = db.get(ProgrammeModule, module_id)
    if not mod:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Module not found")
    if body.display_name is not None:
        mod.display_name = body.display_name.strip()
    if body.sort_order is not None:
        mod.sort_order = body.sort_order
    db.commit()
    mod = (
        db.query(ProgrammeModule)
        .options(selectinload(ProgrammeModule.weeks).selectinload(ProgrammeWeek.sessions).selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeModule.id == module_id)
        .first()
    )
    return ProgrammeModuleOut.model_validate(_serialize_module(mod))


@router.delete(
    "/programme/modules/{module_id}",
    response_model=Message,
    summary="Delete a module (cascades weeks → sessions → assessments)",
)
def delete_module(
    module_id: int,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    mod = (
        db.query(ProgrammeModule)
        .options(
            selectinload(ProgrammeModule.weeks)
            .selectinload(ProgrammeWeek.sessions)
            .selectinload(ProgrammeSession.assessments)
        )
        .filter(ProgrammeModule.id == module_id)
        .first()
    )
    if not mod:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Module not found")
    db.delete(mod)
    db.commit()
    return Message(message="Module deleted")


# ---------------------------------------------------------------------------
# Weeks
# ---------------------------------------------------------------------------

@router.get(
    "/programme/modules/{module_id}/weeks",
    response_model=list[ProgrammeWeekOut],
    summary="List weeks for a module",
)
def list_weeks(
    module_id: int,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    if not db.get(ProgrammeModule, module_id):
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Module not found")
    weeks = (
        db.query(ProgrammeWeek)
        .options(selectinload(ProgrammeWeek.sessions).selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeWeek.module_id == module_id)
        .order_by(ProgrammeWeek.week_index)
        .all()
    )
    return [ProgrammeWeekOut.model_validate(_serialize_week(w)) for w in weeks]


@router.post(
    "/programme/weeks",
    response_model=ProgrammeWeekOut,
    status_code=status.HTTP_201_CREATED,
    summary="Create a week inside a module",
)
def create_week(
    body: ProgrammeWeekCreateRequest,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    if not db.get(ProgrammeModule, body.module_id):
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Module not found")
    existing = (
        db.query(ProgrammeWeek)
        .filter(ProgrammeWeek.module_id == body.module_id, ProgrammeWeek.week_index == body.week_index)
        .first()
    )
    if existing:
        raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Week index already exists in this module")
    week = ProgrammeWeek(module_id=body.module_id, week_index=body.week_index, display_label=body.display_label.strip())
    db.add(week)
    db.commit()
    db.refresh(week)
    week = (
        db.query(ProgrammeWeek)
        .options(selectinload(ProgrammeWeek.sessions).selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeWeek.id == week.id)
        .first()
    )
    return ProgrammeWeekOut.model_validate(_serialize_week(week))


@router.put(
    "/programme/weeks/{week_id}",
    response_model=ProgrammeWeekOut,
    summary="Update a week label or index",
)
def update_week(
    week_id: int,
    body: ProgrammeWeekUpdateRequest,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    week = db.get(ProgrammeWeek, week_id)
    if not week:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Week not found")
    if body.display_label is not None:
        week.display_label = body.display_label.strip()
    if body.week_index is not None:
        week.week_index = body.week_index
    db.commit()
    week = (
        db.query(ProgrammeWeek)
        .options(selectinload(ProgrammeWeek.sessions).selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeWeek.id == week_id)
        .first()
    )
    return ProgrammeWeekOut.model_validate(_serialize_week(week))


@router.delete(
    "/programme/weeks/{week_id}",
    response_model=Message,
    summary="Delete a week (cascades sessions and assessments)",
)
def delete_week(
    week_id: int,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    week = (
        db.query(ProgrammeWeek)
        .options(selectinload(ProgrammeWeek.sessions).selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeWeek.id == week_id)
        .first()
    )
    if not week:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Week not found")
    db.delete(week)
    db.commit()
    return Message(message="Week deleted")


# ---------------------------------------------------------------------------
# Sessions
# ---------------------------------------------------------------------------

@router.get(
    "/programme/weeks/{week_id}/sessions",
    response_model=list[ProgrammeSessionOut],
    summary="List sessions for a week",
)
def list_sessions(
    week_id: int,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    if not db.get(ProgrammeWeek, week_id):
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Week not found")
    sessions = (
        db.query(ProgrammeSession)
        .options(selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeSession.week_id == week_id)
        .order_by(ProgrammeSession.session_number)
        .all()
    )
    return [ProgrammeSessionOut.model_validate(_serialize_session(s)) for s in sessions]


@router.post(
    "/programme/sessions",
    response_model=ProgrammeSessionOut,
    status_code=status.HTTP_201_CREATED,
    summary="Create a session inside a week",
)
def create_session(
    body: ProgrammeSessionCreateRequest,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    if not db.get(ProgrammeWeek, body.week_id):
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Week not found")
    existing = (
        db.query(ProgrammeSession)
        .filter(ProgrammeSession.week_id == body.week_id, ProgrammeSession.session_number == body.session_number)
        .first()
    )
    if existing:
        raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Session number already exists in this week")
    sess = ProgrammeSession(
        week_id=body.week_id,
        session_number=body.session_number,
        title=body.title.strip(),
        admin_checked=body.admin_checked,
    )
    db.add(sess)
    db.commit()
    db.refresh(sess)
    sess = (
        db.query(ProgrammeSession)
        .options(selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeSession.id == sess.id)
        .first()
    )
    return ProgrammeSessionOut.model_validate(_serialize_session(sess))


@router.put(
    "/programme/sessions/{session_id}",
    response_model=ProgrammeSessionOut,
    summary="Update a session title or number",
)
def update_session(
    session_id: int,
    body: ProgrammeSessionUpdateRequest,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    sess = db.get(ProgrammeSession, session_id)
    if not sess:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Session not found")
    if body.title is not None:
        sess.title = body.title.strip()
    if body.session_number is not None:
        sess.session_number = body.session_number
    if body.admin_checked is not None:
        sess.admin_checked = body.admin_checked
    db.commit()
    sess = (
        db.query(ProgrammeSession)
        .options(selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeSession.id == session_id)
        .first()
    )
    return ProgrammeSessionOut.model_validate(_serialize_session(sess))


@router.delete(
    "/programme/sessions/{session_id}",
    response_model=Message,
    summary="Delete a session (cascades assessments)",
)
def delete_session(
    session_id: int,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    sess = (
        db.query(ProgrammeSession)
        .options(selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeSession.id == session_id)
        .first()
    )
    if not sess:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Session not found")
    db.delete(sess)
    db.commit()
    return Message(message="Session deleted")


# ---------------------------------------------------------------------------
# Assessments
# ---------------------------------------------------------------------------

@router.get(
    "/programme/sessions/{session_id}/assessments",
    response_model=list[ProgrammeAssessmentOut],
    summary="List assessment tasks for a session",
)
def list_assessments(
    session_id: int,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    if not db.get(ProgrammeSession, session_id):
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Session not found")
    assessments = (
        db.query(ProgrammeAssessment)
        .filter(ProgrammeAssessment.session_id == session_id)
        .order_by(ProgrammeAssessment.sort_order)
        .all()
    )
    return [ProgrammeAssessmentOut.model_validate(_serialize_assessment(a)) for a in assessments]


@router.post(
    "/programme/assessments",
    response_model=ProgrammeAssessmentOut,
    status_code=status.HTTP_201_CREATED,
    summary="Create an assessment task inside a session",
)
def create_assessment(
    body: ProgrammeAssessmentCreateRequest,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    if not db.get(ProgrammeSession, body.session_id):
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Session not found")
    existing = (
        db.query(ProgrammeAssessment)
        .filter(ProgrammeAssessment.session_id == body.session_id, ProgrammeAssessment.sort_order == body.sort_order)
        .first()
    )
    if existing:
        raise HTTPException(status_code=status.HTTP_409_CONFLICT, detail="Assessment with this sort_order already exists in this session")
    assessment = ProgrammeAssessment(
        session_id=body.session_id,
        title=body.title.strip(),
        sort_order=body.sort_order,
    )
    db.add(assessment)
    db.commit()
    db.refresh(assessment)
    return ProgrammeAssessmentOut.model_validate(_serialize_assessment(assessment))


@router.put(
    "/programme/assessments/{assessment_id}",
    response_model=ProgrammeAssessmentOut,
    summary="Update an assessment task title or sort order",
)
def update_assessment(
    assessment_id: int,
    body: ProgrammeAssessmentUpdateRequest,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    assessment = db.get(ProgrammeAssessment, assessment_id)
    if not assessment:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Assessment not found")
    if body.title is not None:
        assessment.title = body.title.strip()
    if body.sort_order is not None:
        assessment.sort_order = body.sort_order
    db.commit()
    db.refresh(assessment)
    return ProgrammeAssessmentOut.model_validate(_serialize_assessment(assessment))


@router.delete(
    "/programme/assessments/{assessment_id}",
    response_model=Message,
    summary="Delete an assessment task",
)
def delete_assessment(
    assessment_id: int,
    db: Session = Depends(get_db),
    _: User = Depends(require_admin_token),
):
    assessment = db.get(ProgrammeAssessment, assessment_id)
    if not assessment:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Assessment not found")
    db.delete(assessment)
    db.commit()
    return Message(message="Assessment deleted")
