from __future__ import annotations

from dataclasses import dataclass

from fastapi import HTTPException, status
from sqlalchemy.orm import Session, joinedload, selectinload

from app.models import User, UserCompletedAssessment, UserCompletedSession, UserRegisterStep, UserTrackTrace
from app.models_admin import MemberSessionProgress
from app.programme.models_programme import (
    ProgrammeAssessment,
    ProgrammeLevel,
    ProgrammeModule,
    ProgrammeSession,
    ProgrammeWeek,
    normalize_programme_level_slug,
    normalize_programme_module_kind,
    programme_level_slug_str,
    programme_module_kind_str,
)
from app.programme.services.programme_service import (
    _CompletionBatch,
    _ensure_assessment_done,
    _ensure_session_marker_done,
)
from app.schemas import (
    ProgrammeTrackResponse,
    TrackAssessmentOut,
    TrackCompleteRequest,
    TrackLevelOut,
    TrackModuleOut,
    TrackProgressStatus,
    TrackSessionOut,
    TrackWeekOut,
)
from app.services.user_register_step_service import (
    build_register_step_state,
    get_user_register_step,
)


@dataclass
class _TrackSessionRef:
    session: ProgrammeSession
    week: ProgrammeWeek
    module: ProgrammeModule
    level: ProgrammeLevel


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 exc:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="level slug is required",
        ) from exc


def _parse_module_kind(module_kind: str | None) -> str | None:
    if module_kind is None or not module_kind.strip():
        return None
    try:
        return normalize_programme_module_kind(module_kind)
    except ValueError as exc:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="module_kind is required",
        ) from exc


def _load_completed_sets(db: Session, user_id: int) -> tuple[set[int], set[int]]:
    completed_assessments = {
        row.assessment_id
        for row in db.query(UserCompletedAssessment).filter(UserCompletedAssessment.user_id == user_id).all()
    }
    completed_sessions = {
        row.session_id
        for row in db.query(UserCompletedSession).filter(UserCompletedSession.user_id == user_id).all()
    }
    progress_rows = (
        db.query(MemberSessionProgress)
        .filter(MemberSessionProgress.user_id == user_id, MemberSessionProgress.is_completed.is_(True))
        .all()
    )
    for row in progress_rows:
        if row.session_id is not None:
            completed_sessions.add(row.session_id)
    return completed_assessments, completed_sessions


def _session_is_completed(
    session: ProgrammeSession,
    *,
    completed_assessments: set[int],
    completed_sessions: set[int],
) -> bool:
    assessments = sorted(session.assessments, key=lambda item: item.sort_order)
    if assessments:
        return all(assessment.id in completed_assessments for assessment in assessments)
    return session.id in completed_sessions


def _ordered_session_refs(
    level: ProgrammeLevel,
    *,
    module_kind: str | None,
) -> list[_TrackSessionRef]:
    refs: list[_TrackSessionRef] = []
    modules = sorted(level.modules, key=lambda item: item.sort_order)
    if module_kind is not None:
        modules = [
            module
            for module in modules
            if programme_module_kind_str(module.kind) == module_kind
        ]

    for module in modules:
        for week in sorted(module.weeks, key=lambda item: item.week_index):
            for session in sorted(week.sessions, key=lambda item: item.session_number):
                refs.append(
                    _TrackSessionRef(session=session, week=week, module=module, level=level)
                )
    return refs


def _level_context_for_session(
    session: ProgrammeSession,
    programme_levels: list[ProgrammeLevel],
) -> tuple[ProgrammeLevel | None, ProgrammeModule | None]:
    week = session.week
    module = week.module if week else None
    if module is None:
        return None, None
    for level in programme_levels:
        for mod in level.modules:
            if mod.id == module.id:
                return level, mod
    return None, module


def _anchor_session_index(
    session_refs: list[_TrackSessionRef],
    anchor_session_id: int | None,
) -> int:
    if anchor_session_id is None:
        return 0
    for index, ref in enumerate(session_refs):
        if ref.session.id == anchor_session_id:
            return index
    return 0


def _registration_anchor_for_refs(
    session_refs: list[_TrackSessionRef],
    register_step: UserRegisterStep | None,
) -> int | None:
    """Only use the registration anchor when it belongs to this level/module view."""
    if register_step is None or register_step.session_id is None:
        return None
    if any(ref.session.id == register_step.session_id for ref in session_refs):
        return register_step.session_id
    return None


def _prior_levels_complete(
    levels: list[ProgrammeLevel],
    level_row: ProgrammeLevel,
    *,
    completed_assessments: set[int],
    completed_sessions: set[int],
) -> bool:
    for level in levels:
        if level.sort_order >= level_row.sort_order:
            return True
        refs = _ordered_session_refs(level, module_kind=None)
        if not refs:
            continue
        for ref in refs:
            if not _session_is_completed(
                ref.session,
                completed_assessments=completed_assessments,
                completed_sessions=completed_sessions,
            ):
                return False
    return True


def _level_is_accessible(
    db: Session,
    user_id: int,
    level_row: ProgrammeLevel,
    register_step: UserRegisterStep | None,
    programme_levels: list[ProgrammeLevel],
    *,
    completed_assessments: set[int],
    completed_sessions: set[int],
) -> bool:
    """Level is playable when every prior level's sessions are complete."""
    if not _prior_levels_complete(
        programme_levels,
        level_row,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
    ):
        return False

    # Prior levels are fully complete — unlock this level. Do not also require the
    # global open pointer to have moved forward; it may still reference the last
    # session of a finished level when the next level has no sessions yet.
    return True


def _module_row_for_kind(
    level_row: ProgrammeLevel,
    module_kind: str | None,
) -> ProgrammeModule | None:
    if module_kind is None:
        return None
    for mod in level_row.modules:
        if programme_module_kind_str(mod.kind) == module_kind:
            return mod
    return None


def _prior_modules_complete(
    level_row: ProgrammeLevel,
    module_row: ProgrammeModule,
    *,
    completed_assessments: set[int],
    completed_sessions: set[int],
) -> bool:
    modules = sorted(level_row.modules, key=lambda item: item.sort_order)
    for mod in modules:
        if mod.sort_order >= module_row.sort_order:
            return True
        refs = _ordered_session_refs(
            level_row,
            module_kind=programme_module_kind_str(mod.kind),
        )
        if not refs:
            continue
        for ref in refs:
            if not _session_is_completed(
                ref.session,
                completed_assessments=completed_assessments,
                completed_sessions=completed_sessions,
            ):
                return False
    return True


def _module_is_accessible(
    db: Session,
    user_id: int,
    level_row: ProgrammeLevel,
    module_row: ProgrammeModule,
    register_step: UserRegisterStep | None,
    programme_levels: list[ProgrammeLevel],
    *,
    completed_assessments: set[int],
    completed_sessions: set[int],
) -> bool:
    """Module is playable only when the level and every prior module in it are complete."""
    if not _level_is_accessible(
        db,
        user_id,
        level_row,
        register_step,
        programme_levels,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
    ):
        return False
    if not _prior_modules_complete(
        level_row,
        module_row,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
    ):
        return False

    # Prior modules in this level are fully complete — unlock this module. The global
    # open pointer can still sit on the last session of the previous module until the
    # user opens the next one; that must not keep module_unlocked false.
    return True


def _resolve_open_for_track_view(
    db: Session,
    user_id: int,
    level_row: ProgrammeLevel,
    *,
    module_kind: str | None,
    register_step: UserRegisterStep | None,
    completed_assessments: set[int],
    completed_sessions: set[int],
    programme_levels: list[ProgrammeLevel],
) -> tuple[int | None, int | None]:
    """
    Resolve the open session/assessment for a specific level (and optional module).
    Prior levels must be fully complete before anything in this level can open.
    """
    session_refs = _ordered_session_refs(level_row, module_kind=module_kind)
    if not _level_is_accessible(
        db,
        user_id,
        level_row,
        register_step,
        programme_levels,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
    ):
        return None, None

    if module_kind is not None:
        module_row = _module_row_for_kind(level_row, module_kind)
        if module_row is None or not _module_is_accessible(
            db,
            user_id,
            level_row,
            module_row,
            register_step,
            programme_levels,
            completed_assessments=completed_assessments,
            completed_sessions=completed_sessions,
        ):
            return None, None

    anchor_session_id = _registration_anchor_for_refs(session_refs, register_step)
    global_session, global_assessment = resolve_programme_open_position(
        db,
        user_id,
        register_step,
        programme_levels=programme_levels,
    )

    if global_session is not None:
        week = global_session.week
        module = week.module if week else None
        global_level = module.level if module else None
        if global_level is not None and global_level.sort_order > level_row.sort_order:
            return _resolve_open_targets(
                level_row,
                module_kind=module_kind,
                completed_assessments=completed_assessments,
                completed_sessions=completed_sessions,
                anchor_session_id=anchor_session_id,
            )
        if global_level is not None and global_level.id == level_row.id:
            module_kind_str = programme_module_kind_str(module.kind) if module else None
            if module_kind is not None and module is not None:
                filtered_module = _module_row_for_kind(level_row, module_kind)
                if filtered_module is not None and module.sort_order > filtered_module.sort_order:
                    return _resolve_open_targets(
                        level_row,
                        module_kind=module_kind,
                        completed_assessments=completed_assessments,
                        completed_sessions=completed_sessions,
                        anchor_session_id=anchor_session_id,
                    )
            if module_kind is None or module_kind == module_kind_str:
                return global_session.id, global_assessment.id if global_assessment else None

    return _resolve_open_targets(
        level_row,
        module_kind=module_kind,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
        anchor_session_id=anchor_session_id,
    )


def _is_intro_placeholder_session(session: ProgrammeSession, week: ProgrammeWeek) -> bool:
    """Intro-week rows seeded without real assessments (e.g. HB intro session 1)."""
    if week.week_index != 0:
        return False
    if session.assessments:
        return False
    title = (session.title or "").strip().lower()
    return title in {"", "session 1", "session 2", "session 3"} or title.startswith("session ")


def _repair_intro_placeholder_sessions(
    db: Session,
    user_id: int,
    *,
    level_slug: str | None,
    module_kind: str | None,
    register_step: UserRegisterStep | None,
) -> None:
    level_slug_enum = _parse_level_slug(level_slug)
    module_kind_enum = _parse_module_kind(module_kind)
    level_row = _load_level(db, level_slug=level_slug_enum, register_step=register_step)
    completed_assessments, completed_sessions = _load_completed_sets(db, user_id)
    batch = _CompletionBatch(db, user_id)
    for ref in _ordered_session_refs(level_row, module_kind=module_kind_enum):
        if not _is_intro_placeholder_session(ref.session, ref.week):
            continue
        if _session_is_completed(
            ref.session,
            completed_assessments=completed_assessments,
            completed_sessions=completed_sessions,
        ):
            continue
        _ensure_session_marker_done(db, user_id, ref.session.id, batch)
    batch.flush()


def _resolve_open_targets(
    level: ProgrammeLevel,
    *,
    module_kind: str | None,
    completed_assessments: set[int],
    completed_sessions: set[int],
    anchor_session_id: int | None,
) -> tuple[int | None, int | None]:
    """
    First incomplete assessment (preferred) or session without assessments.
    Everything before the registration anchor is treated as already done.
    """
    session_refs = _ordered_session_refs(level, module_kind=module_kind)
    anchor_index = _anchor_session_index(session_refs, anchor_session_id)

    for index, ref in enumerate(session_refs):
        if index < anchor_index:
            continue

        session = ref.session
        assessments = sorted(session.assessments, key=lambda item: item.sort_order)
        if assessments:
            for assessment in assessments:
                if assessment.id not in completed_assessments:
                    return session.id, assessment.id
        elif not _session_is_completed(
            session,
            completed_assessments=completed_assessments,
            completed_sessions=completed_sessions,
        ):
            if _is_intro_placeholder_session(session, ref.week):
                continue
            return session.id, None

    return None, None


def _load_programme_levels_ordered(db: Session) -> list[ProgrammeLevel]:
    return (
        db.query(ProgrammeLevel)
        .options(
            selectinload(ProgrammeLevel.modules)
            .selectinload(ProgrammeModule.weeks)
            .selectinload(ProgrammeWeek.sessions)
            .selectinload(ProgrammeSession.assessments),
            selectinload(ProgrammeLevel.modules)
            .selectinload(ProgrammeModule.weeks)
            .joinedload(ProgrammeWeek.module)
            .joinedload(ProgrammeModule.level),
        )
        .order_by(ProgrammeLevel.sort_order)
        .all()
    )


def _ordered_session_refs_for_levels(levels: list[ProgrammeLevel]) -> list[_TrackSessionRef]:
    refs: list[_TrackSessionRef] = []
    for level in levels:
        refs.extend(_ordered_session_refs(level, module_kind=None))
    return refs


def resolve_programme_open_position(
    db: Session,
    user_id: int,
    register_step: UserRegisterStep | None,
    *,
    programme_levels: list[ProgrammeLevel] | None = None,
) -> tuple[ProgrammeSession | None, ProgrammeAssessment | None]:
    """
    First incomplete session/assessment from the member's registered level onward.
    When a module or level is fully complete, the position advances to the next one.
    """
    if register_step is None or not register_step.is_completed:
        return None, None

    levels = programme_levels if programme_levels is not None else _load_programme_levels_ordered(db)
    if not levels:
        return None, None

    session_refs = _ordered_session_refs_for_levels(levels)
    if not session_refs:
        return None, None

    completed_assessments, completed_sessions = _load_completed_sets(db, user_id)
    anchor_index = _anchor_session_index(session_refs, register_step.session_id)

    for index, ref in enumerate(session_refs):
        if index < anchor_index:
            continue

        session = ref.session
        assessments = sorted(session.assessments, key=lambda item: item.sort_order)
        if assessments:
            for assessment in assessments:
                if assessment.id not in completed_assessments:
                    return session, assessment
        elif not _session_is_completed(
            session,
            completed_assessments=completed_assessments,
            completed_sessions=completed_sessions,
        ):
            if _is_intro_placeholder_session(session, ref.week):
                continue
            return session, None

    # Every session in the programme is complete — no open step remains.
    return None, None


def _programme_summary_from_session(
    session: ProgrammeSession,
    assessment: ProgrammeAssessment | None,
    *,
    level: ProgrammeLevel | None = None,
    module: ProgrammeModule | None = None,
) -> dict:
    week = session.week
    mod = module or (week.module if week else None)
    level_row = level or (mod.level if mod else None)
    return {
        "level_slug": programme_level_slug_str(level_row.slug) if level_row else None,
        "level_display_name": level_row.display_name if level_row else None,
        "module_kind": programme_module_kind_str(mod.kind) if mod else None,
        "module_display_name": mod.display_name if mod else None,
        "week_index": week.week_index if week else None,
        "week_label": week.display_label if week else None,
        "session_number": session.session_number,
        "session_title": session.title,
        "session_admin_checked": session.admin_checked,
        "session_id": session.id,
        "assessment_id": assessment.id if assessment else None,
        "assessment_title": assessment.title if assessment else None,
    }


def resolve_user_programme_summary(
    db: Session,
    user: User,
    *,
    programme_levels: list[ProgrammeLevel] | None = None,
) -> dict | None:
    from app.programme.services.programme_service import load_user_programme_summary

    levels = programme_levels if programme_levels is not None else _load_programme_levels_ordered(db)
    register_step = get_user_register_step(db, user)
    session, assessment = resolve_programme_open_position(
        db, user.id, register_step, programme_levels=levels
    )
    if session is not None:
        level_row, module_row = _level_context_for_session(session, levels)
        return _programme_summary_from_session(
            session,
            assessment,
            level=level_row,
            module=module_row,
        )
    return load_user_programme_summary(user)


def sync_user_programme_pointer(db: Session, user: User) -> None:
    register_step = get_user_register_step(db, user)
    levels = _load_programme_levels_ordered(db)
    session, assessment = resolve_programme_open_position(
        db, user.id, register_step, programme_levels=levels
    )
    if session is None:
        return
    user.programme_session_id = session.id
    user.programme_assessment_id = assessment.id if assessment else None
    if register_step is None:
        return

    level_row, module_row = _level_context_for_session(session, levels)
    week = session.week
    if level_row is not None:
        register_step.level_id = level_row.id
        register_step.level_name = level_row.display_name
    if module_row is not None:
        register_step.module_id = module_row.id
        register_step.module_name = module_row.display_name
    if week is not None:
        register_step.week_id = week.id
        register_step.week_name = week.display_label
    register_step.session_id = session.id
    register_step.session_name = session.title


def _repair_stale_session_completions(
    db: Session,
    user_id: int,
    *,
    level_slug: str | None,
    module_kind: str | None,
    register_step: UserRegisterStep | None,
) -> None:
    """Mark sessions done when all assessments are already completed (fixes stuck open step)."""
    level_slug_enum = _parse_level_slug(level_slug)
    module_kind_enum = _parse_module_kind(module_kind)
    level_row = _load_level(db, level_slug=level_slug_enum, register_step=register_step)

    if module_kind_enum is None and register_step and register_step.module_id is not None:
        module_row = db.get(ProgrammeModule, register_step.module_id)
        if module_row is not None:
            module_kind_enum = programme_module_kind_str(module_row.kind)

    completed_assessments, _ = _load_completed_sets(db, user_id)
    batch = _CompletionBatch(db, user_id)
    for ref in _ordered_session_refs(level_row, module_kind=module_kind_enum):
        session = ref.session
        assessments = sorted(session.assessments, key=lambda item: item.sort_order)
        if not assessments:
            continue
        if all(assessment.id in completed_assessments for assessment in assessments):
            _ensure_session_marker_done(db, user_id, session.id, batch)
    batch.flush()


def _upsert_track_trace(
    db: Session,
    user_id: int,
    *,
    session: ProgrammeSession,
    week: ProgrammeWeek | None,
    assessment: ProgrammeAssessment | None = None,
    is_completed: bool = False,
) -> MemberSessionProgress:
    """Save progress to member_session_progress for dashboard / trace."""
    week_row = week or session.week
    row = (
        db.query(MemberSessionProgress)
        .filter(
            MemberSessionProgress.user_id == user_id,
            MemberSessionProgress.session_id == session.id,
        )
        .first()
    )
    if row is None:
        row = MemberSessionProgress(user_id=user_id)
        db.add(row)

    row.week_id = week_row.id if week_row else None
    row.week_name = week_row.display_label if week_row else None
    row.week_index = week_row.week_index if week_row else None
    row.session_id = session.id
    row.session_name = session.title
    row.session_number = session.session_number
    row.is_completed = is_completed
    return row


def _insert_user_track_trace(
    db: Session,
    user_id: int,
    *,
    session: ProgrammeSession,
    week: ProgrammeWeek | None,
    assessment: ProgrammeAssessment | None = None,
    action: str = "assessment_completed",
    session_fully_completed: bool = False,
) -> UserTrackTrace:
    """Append one row to user_track_traces when user clicks complete on Tracks page."""
    week_row = week or session.week
    module: ProgrammeModule | None = week_row.module if week_row else None
    level: ProgrammeLevel | None = module.level if module else None

    row = UserTrackTrace(
        user_id=user_id,
        level_id=level.id if level else None,
        level_slug=programme_level_slug_str(level.slug) if level else None,
        level_name=level.display_name if level else None,
        module_id=module.id if module else None,
        module_kind=programme_module_kind_str(module.kind) if module else None,
        module_name=module.display_name if module else None,
        week_id=week_row.id if week_row else None,
        week_name=week_row.display_label if week_row else None,
        week_index=week_row.week_index if week_row else None,
        session_id=session.id,
        session_name=session.title,
        session_number=session.session_number,
        assessment_id=assessment.id if assessment else None,
        assessment_title=assessment.title if assessment else None,
        action=action,
        session_fully_completed=session_fully_completed,
    )
    db.add(row)
    db.flush()
    return row


def list_user_track_traces(db: Session, user_id: int, *, limit: int = 100) -> list[UserTrackTrace]:
    return (
        db.query(UserTrackTrace)
        .filter(UserTrackTrace.user_id == user_id)
        .order_by(UserTrackTrace.created_at.desc(), UserTrackTrace.id.desc())
        .limit(limit)
        .all()
    )


def _assessment_status(
    assessment_id: int,
    session: ProgrammeSession,
    *,
    open_session_id: int | None,
    open_assessment_id: int | None,
    completed_assessments: set[int],
    completed_sessions: set[int],
) -> TrackProgressStatus:
    if assessment_id in completed_assessments:
        return TrackProgressStatus.completed
    if _session_is_completed(
        session,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
    ):
        return TrackProgressStatus.completed
    if open_assessment_id == assessment_id:
        return TrackProgressStatus.open
    if open_session_id == session.id and open_assessment_id is None:
        return TrackProgressStatus.open
    return TrackProgressStatus.locked


def _session_status(
    session: ProgrammeSession,
    *,
    open_session_id: int | None,
    open_assessment_id: int | None,
    completed_assessments: set[int],
    completed_sessions: set[int],
    prior_session_done: bool,
) -> TrackProgressStatus:
    if _session_is_completed(
        session,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
    ):
        return TrackProgressStatus.completed
    if not prior_session_done:
        return TrackProgressStatus.locked
    assessments = sorted(session.assessments, key=lambda item: item.sort_order)
    if assessments:
        if any(
            _assessment_status(
                assessment.id,
                session,
                open_session_id=open_session_id,
                open_assessment_id=open_assessment_id,
                completed_assessments=completed_assessments,
                completed_sessions=completed_sessions,
            )
            == TrackProgressStatus.open
            for assessment in assessments
        ):
            return TrackProgressStatus.open
        if all(assessment.id in completed_assessments for assessment in assessments):
            return TrackProgressStatus.completed
        return TrackProgressStatus.locked
    if open_session_id == session.id:
        return TrackProgressStatus.open
    return TrackProgressStatus.locked


def _week_status(
    week: ProgrammeWeek,
    session_statuses: list[TrackProgressStatus],
    *,
    prior_week_done: bool,
) -> TrackProgressStatus:
    if session_statuses and all(item == TrackProgressStatus.completed for item in session_statuses):
        return TrackProgressStatus.completed
    if not session_statuses:
        return (
            TrackProgressStatus.completed
            if prior_week_done
            else TrackProgressStatus.locked
        )
    if not prior_week_done:
        return TrackProgressStatus.locked
    if any(item == TrackProgressStatus.open for item in session_statuses):
        return TrackProgressStatus.open
    if any(item == TrackProgressStatus.completed for item in session_statuses):
        return TrackProgressStatus.open
    return TrackProgressStatus.locked


def _load_level(
    db: Session,
    *,
    level_slug: str | None,
    register_step: UserRegisterStep | None,
) -> ProgrammeLevel:
    query = (
        db.query(ProgrammeLevel)
        .options(
            selectinload(ProgrammeLevel.modules)
            .selectinload(ProgrammeModule.weeks)
            .selectinload(ProgrammeWeek.sessions)
            .selectinload(ProgrammeSession.assessments)
        )
        .order_by(ProgrammeLevel.sort_order)
    )
    if level_slug is not None:
        query = query.filter(ProgrammeLevel.slug == level_slug)
    elif register_step and register_step.level_id is not None:
        query = query.filter(ProgrammeLevel.id == register_step.level_id)

    level_row = query.first()
    if level_row is None:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Programme level not found")
    return level_row


def _sync_track_progress(
    db: Session,
    user: User,
    register_step: UserRegisterStep | None,
    *,
    level_slug: str | None,
    module_kind_enum: str | None,
) -> None:
    """Lightweight stale-session repair on track preview (reseed runs only via POST /track/reseed)."""
    _repair_stale_session_completions(
        db,
        user.id,
        level_slug=level_slug,
        module_kind=module_kind_enum,
        register_step=register_step,
    )
    _repair_intro_placeholder_sessions(
        db,
        user.id,
        level_slug=level_slug,
        module_kind=module_kind_enum,
        register_step=register_step,
    )
    if register_step is not None and register_step.is_completed:
        sync_user_programme_pointer(db, user)
    if db.new or db.dirty or db.deleted:
        db.commit()


def _peek_current_open(
    db: Session,
    user: User,
    register_step: UserRegisterStep | None,
    level_row: ProgrammeLevel,
    *,
    level_slug: str | None,
    module_kind_enum: str | None,
) -> dict:
    completed_assessments, completed_sessions = _load_completed_sets(db, user.id)
    programme_levels = _load_programme_levels_ordered(db)
    open_session_id, open_assessment_id = _resolve_open_for_track_view(
        db,
        user.id,
        level_row,
        module_kind=module_kind_enum,
        register_step=register_step,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
        programme_levels=programme_levels,
    )
    if open_session_id is None:
        return {}
    session_ref = next(
        (
            ref
            for ref in _ordered_session_refs(level_row, module_kind=module_kind_enum)
            if ref.session.id == open_session_id
        ),
        None,
    )
    if session_ref is None:
        return {}
    return {
        "type": "assessment" if open_assessment_id else "session",
        "assessment_id": open_assessment_id,
        "session_id": open_session_id,
        "week_id": session_ref.week.id,
        "module_id": session_ref.module.id,
        "module_kind": programme_module_kind_str(session_ref.module.kind),
        "level_id": level_row.id,
        "level_slug": programme_level_slug_str(level_row.slug),
    }


def get_programme_track(
    db: Session,
    user: User,
    *,
    level: str | None,
    module_kind: str | None,
    sync: bool = True,
) -> ProgrammeTrackResponse:
    register_step = get_user_register_step(db, user)
    registration_state = build_register_step_state(register_step)
    if not registration_state.is_completed:
        return ProgrammeTrackResponse(
            registration_completed=False,
            current_open=None,
            level=None,
        )

    level_slug = _parse_level_slug(level)
    module_kind_enum = _parse_module_kind(module_kind)
    level_row = _load_level(db, level_slug=level_slug, register_step=register_step)

    if module_kind_enum is None and register_step and register_step.module_id is not None:
        module_row = db.get(ProgrammeModule, register_step.module_id)
        if module_row is not None:
            module_kind_enum = programme_module_kind_str(module_row.kind)

    if sync:
        _sync_track_progress(
            db,
            user,
            register_step,
            level_slug=level_slug,
            module_kind_enum=module_kind_enum,
        )

    completed_assessments, completed_sessions = _load_completed_sets(db, user.id)
    programme_levels = _load_programme_levels_ordered(db)
    level_unlocked = _level_is_accessible(
        db,
        user.id,
        level_row,
        register_step,
        programme_levels,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
    )
    module_unlocked = True
    if module_kind_enum is not None:
        module_row = _module_row_for_kind(level_row, module_kind_enum)
        module_unlocked = module_row is not None and _module_is_accessible(
            db,
            user.id,
            level_row,
            module_row,
            register_step,
            programme_levels,
            completed_assessments=completed_assessments,
            completed_sessions=completed_sessions,
        )
    content_unlocked = level_unlocked and module_unlocked

    session_refs = _ordered_session_refs(level_row, module_kind=module_kind_enum)
    anchor_session_id = _registration_anchor_for_refs(session_refs, register_step)
    anchor_index = _anchor_session_index(session_refs, anchor_session_id)
    anchor_week_index: int | None = None
    if session_refs and anchor_session_id is not None and 0 <= anchor_index < len(session_refs):
        anchor_week_index = session_refs[anchor_index].week.week_index
    open_session_id, open_assessment_id = _resolve_open_for_track_view(
        db,
        user.id,
        level_row,
        module_kind=module_kind_enum,
        register_step=register_step,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
        programme_levels=programme_levels,
    )

    current_open: dict | None = None
    if open_session_id is not None:
        session_ref = next(
            (
                ref
                for ref in _ordered_session_refs(level_row, module_kind=module_kind_enum)
                if ref.session.id == open_session_id
            ),
            None,
        )
        if session_ref is not None:
            current_open = {
                "type": "assessment" if open_assessment_id else "session",
                "assessment_id": open_assessment_id,
                "session_id": open_session_id,
                "week_id": session_ref.week.id,
                "module_id": session_ref.module.id,
                "module_kind": programme_module_kind_str(session_ref.module.kind),
                "level_id": level_row.id,
                "level_slug": programme_level_slug_str(level_row.slug),
            }

    modules_out: list[TrackModuleOut] = []
    level_completed = 0
    level_total = 0
    global_session_index = 0

    modules = sorted(level_row.modules, key=lambda item: item.sort_order)
    if module_kind_enum is not None:
        filtered_modules = [
            module
            for module in modules
            if programme_module_kind_str(module.kind) == module_kind_enum
        ]
        # Keep full level data when the filter key does not match stored kinds.
        if filtered_modules:
            modules = filtered_modules

    for module in modules:
        weeks_out: list[TrackWeekOut] = []
        module_completed = 0
        module_total = 0
        prior_week_done = True

        for week in sorted(module.weeks, key=lambda item: item.week_index):
            sessions_out: list[TrackSessionOut] = []
            week_completed = 0
            session_statuses: list[TrackProgressStatus] = []
            prior_session_done = True

            for session in sorted(week.sessions, key=lambda item: item.session_number):
                assessments_out: list[TrackAssessmentOut] = []
                before_anchor = (
                    content_unlocked
                    and anchor_session_id is not None
                    and global_session_index < anchor_index
                )

                if before_anchor:
                    session_status = TrackProgressStatus.completed
                elif not content_unlocked:
                    session_status = TrackProgressStatus.locked
                else:
                    session_status = _session_status(
                        session,
                        open_session_id=open_session_id,
                        open_assessment_id=open_assessment_id,
                        completed_assessments=completed_assessments,
                        completed_sessions=completed_sessions,
                        prior_session_done=prior_session_done,
                    )
                session_statuses.append(session_status)

                for assessment in sorted(session.assessments, key=lambda item: item.sort_order):
                    if before_anchor:
                        assessment_status = TrackProgressStatus.completed
                    elif not content_unlocked:
                        assessment_status = TrackProgressStatus.locked
                    else:
                        assessment_status = _assessment_status(
                            assessment.id,
                            session,
                            open_session_id=open_session_id,
                            open_assessment_id=open_assessment_id,
                            completed_assessments=completed_assessments,
                            completed_sessions=completed_sessions,
                        )
                    assessments_out.append(
                        TrackAssessmentOut(
                            id=assessment.id,
                            sort_order=assessment.sort_order,
                            title=assessment.title,
                            status=assessment_status,
                        )
                    )

                if session_status == TrackProgressStatus.completed:
                    week_completed += 1
                    module_completed += 1
                    level_completed += 1
                module_total += 1
                level_total += 1
                prior_session_done = session_status == TrackProgressStatus.completed
                global_session_index += 1

                sessions_out.append(
                    TrackSessionOut(
                        id=session.id,
                        session_number=session.session_number,
                        title=session.title,
                        admin_checked=session.admin_checked,
                        status=session_status,
                        assessments_count=len(assessments_out),
                        assessments=assessments_out,
                    )
                )

            if not content_unlocked:
                week_status = TrackProgressStatus.locked
            elif not session_statuses and anchor_week_index is not None:
                if week.week_index < anchor_week_index:
                    week_status = TrackProgressStatus.completed
                elif week.week_index > anchor_week_index:
                    week_status = TrackProgressStatus.locked
                else:
                    week_status = TrackProgressStatus.open if prior_week_done else TrackProgressStatus.locked
            else:
                week_status = _week_status(week, session_statuses, prior_week_done=prior_week_done)
            prior_week_done = week_status == TrackProgressStatus.completed

            weeks_out.append(
                TrackWeekOut(
                    id=week.id,
                    week_index=week.week_index,
                    display_label=week.display_label,
                    status=week_status,
                    sessions_count=len(sessions_out),
                    completed_sessions_count=week_completed,
                    sessions=sessions_out,
                )
            )

        module_status = TrackProgressStatus.locked
        if weeks_out:
            if all(week.status == TrackProgressStatus.completed for week in weeks_out):
                module_status = TrackProgressStatus.completed
            elif any(week.status == TrackProgressStatus.open for week in weeks_out):
                module_status = TrackProgressStatus.open

        modules_out.append(
            TrackModuleOut(
                id=module.id,
                kind=programme_module_kind_str(module.kind),
                display_name=module.display_name,
                status=module_status,
                weeks=weeks_out,
                completed_sessions=module_completed,
                total_sessions=module_total,
            )
        )

    registration_selection: dict | None = None
    if register_step is not None:
        registration_selection = {
            "level_id": register_step.level_id,
            "level_name": register_step.level_name,
            "module_id": register_step.module_id,
            "module_name": register_step.module_name,
            "week_id": register_step.week_id,
            "week_name": register_step.week_name,
            "session_id": register_step.session_id,
            "session_name": register_step.session_name,
        }

    return ProgrammeTrackResponse(
        registration_completed=True,
        level_unlocked=level_unlocked,
        module_unlocked=module_unlocked,
        current_open=current_open,
        registration_selection=registration_selection,
        level=TrackLevelOut(
            id=level_row.id,
            slug=programme_level_slug_str(level_row.slug),
            display_name=level_row.display_name,
            modules=modules_out,
            completed_sessions=level_completed,
            total_sessions=level_total,
        ),
    )


def _maybe_complete_session_after_assessment(
    db: Session,
    user_id: int,
    session_id: int,
    completed_assessments: set[int],
) -> None:
    if not session_id:
        return
    session = (
        db.query(ProgrammeSession)
        .options(selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeSession.id == session_id)
        .first()
    )
    if session is None:
        return
    assessments = sorted(session.assessments, key=lambda item: item.sort_order)
    if assessments and all(assessment.id in completed_assessments for assessment in assessments):
        _ensure_session_marker_done(db, user_id, session.id)


def complete_track_item(db: Session, user: User, body: TrackCompleteRequest) -> ProgrammeTrackResponse:
    register_step = get_user_register_step(db, user)
    if register_step is None or not register_step.is_completed:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="Complete registration steps before updating track progress.",
        )

    level_filter = body.level
    module_filter = body.module_kind
    if level_filter is None and register_step.level_id is not None:
        level_row = db.get(ProgrammeLevel, register_step.level_id)
        if level_row is not None:
            level_filter = programme_level_slug_str(level_row.slug)
    if module_filter is None and register_step.module_id is not None:
        module_row = db.get(ProgrammeModule, register_step.module_id)
        if module_row is not None:
            module_filter = programme_module_kind_str(module_row.kind)

    if body.assessment_id is None and body.session_id is None:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="Send assessment_id or session_id.",
        )
    if body.assessment_id is not None and body.session_id is not None:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="Send only one of assessment_id or session_id.",
        )

    level_slug_enum = _parse_level_slug(level_filter)
    module_kind_enum = _parse_module_kind(module_filter)
    if module_kind_enum is None and register_step.module_id is not None:
        module_row = db.get(ProgrammeModule, register_step.module_id)
        if module_row is not None:
            module_kind_enum = programme_module_kind_str(module_row.kind)

    level_row = _load_level(db, level_slug=level_slug_enum, register_step=register_step)
    programme_levels = _load_programme_levels_ordered(db)
    completed_assessments, completed_sessions = _load_completed_sets(db, user.id)
    if not _level_is_accessible(
        db,
        user.id,
        level_row,
        register_step,
        programme_levels,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
    ):
        raise HTTPException(
            status_code=status.HTTP_403_FORBIDDEN,
            detail="Complete all sessions in the previous level before updating this level.",
        )

    target_module_row: ProgrammeModule | None = None
    if body.session_id is not None:
        session_row = (
            db.query(ProgrammeSession)
            .options(joinedload(ProgrammeSession.week).joinedload(ProgrammeWeek.module))
            .filter(ProgrammeSession.id == body.session_id)
            .first()
        )
        if session_row is not None and session_row.week is not None:
            target_module_row = session_row.week.module
    elif body.assessment_id is not None:
        assessment_row = db.get(ProgrammeAssessment, body.assessment_id)
        if assessment_row is not None:
            session_row = (
                db.query(ProgrammeSession)
                .options(joinedload(ProgrammeSession.week).joinedload(ProgrammeWeek.module))
                .filter(ProgrammeSession.id == assessment_row.session_id)
                .first()
            )
            if session_row is not None and session_row.week is not None:
                target_module_row = session_row.week.module
    if target_module_row is not None and not _module_is_accessible(
        db,
        user.id,
        level_row,
        target_module_row,
        register_step,
        programme_levels,
        completed_assessments=completed_assessments,
        completed_sessions=completed_sessions,
    ):
        raise HTTPException(
            status_code=status.HTTP_403_FORBIDDEN,
            detail="Complete all sessions in the previous module before updating this module.",
        )

    _sync_track_progress(
        db,
        user,
        register_step,
        level_slug=level_slug_enum,
        module_kind_enum=module_kind_enum,
    )

    current_open = _peek_current_open(
        db,
        user,
        register_step,
        level_row,
        level_slug=level_slug_enum,
        module_kind_enum=module_kind_enum,
    )

    if body.assessment_id is not None:
        completed_assessments, _ = _load_completed_sets(db, user.id)
        if body.assessment_id in completed_assessments:
            return get_programme_track(
                db, user, level=level_filter, module_kind=module_filter, sync=False
            )

        expected_id = current_open.get("assessment_id")
        if expected_id != body.assessment_id:
            _repair_stale_session_completions(
                db,
                user.id,
                level_slug=level_filter,
                module_kind=module_filter,
                register_step=register_step,
            )
            db.flush()
            current_open = _peek_current_open(
                db,
                user,
                register_step,
                level_row,
                level_slug=level_slug_enum,
                module_kind_enum=module_kind_enum,
            )
            expected_id = current_open.get("assessment_id")
            if expected_id != body.assessment_id:
                raise HTTPException(
                    status_code=status.HTTP_400_BAD_REQUEST,
                    detail="This assessment is not the current open step.",
                )
        _ensure_assessment_done(db, user.id, body.assessment_id)
        db.flush()
        completed_assessments, _ = _load_completed_sets(db, user.id)
        assessment = db.get(ProgrammeAssessment, body.assessment_id)
        session = (
            db.query(ProgrammeSession)
            .options(
                joinedload(ProgrammeSession.week)
                .joinedload(ProgrammeWeek.module)
                .joinedload(ProgrammeModule.level),
                joinedload(ProgrammeSession.assessments),
            )
            .filter(ProgrammeSession.id == (assessment.session_id if assessment else 0))
            .first()
        )
        if session is not None:
            assessments = sorted(session.assessments, key=lambda item: item.sort_order)
            all_done = bool(
                assessments
                and all(item.id in completed_assessments for item in assessments)
            )
            _upsert_track_trace(
                db,
                user.id,
                session=session,
                week=session.week,
                assessment=assessment,
                is_completed=all_done,
            )
            _insert_user_track_trace(
                db,
                user.id,
                session=session,
                week=session.week,
                assessment=assessment,
                action="assessment_completed",
                session_fully_completed=all_done,
            )
        _maybe_complete_session_after_assessment(
            db, user.id, current_open.get("session_id") or 0, completed_assessments
        )
    else:
        session = (
            db.query(ProgrammeSession)
            .options(
                selectinload(ProgrammeSession.assessments),
                joinedload(ProgrammeSession.week)
                .joinedload(ProgrammeWeek.module)
                .joinedload(ProgrammeModule.level),
            )
            .filter(ProgrammeSession.id == body.session_id)
            .first()
        )
        if session is None:
            raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Session not found")

        assessments = sorted(session.assessments, key=lambda item: item.sort_order)
        completed_assessments, _ = _load_completed_sets(db, user.id)
        if assessments:
            if not all(assessment.id in completed_assessments for assessment in assessments):
                raise HTTPException(
                    status_code=status.HTTP_400_BAD_REQUEST,
                    detail="Complete all assessments in this session first.",
                )
        elif current_open.get("session_id") != body.session_id:
            raise HTTPException(
                status_code=status.HTTP_400_BAD_REQUEST,
                detail="This session is not the current open step.",
            )

        _ensure_session_marker_done(db, user.id, body.session_id)
        _upsert_track_trace(
            db,
            user.id,
            session=session,
            week=session.week,
            is_completed=True,
        )
        _insert_user_track_trace(
            db,
            user.id,
            session=session,
            week=session.week,
            action="session_completed",
            session_fully_completed=True,
        )
    sync_user_programme_pointer(db, user)
    db.commit()

    try:
        from app.routers.admin.users import invalidate_list_users_cache

        invalidate_list_users_cache()
    except Exception:
        pass

    from app.services.admin_notification_service import notify_admin_track_submit

    if body.assessment_id is not None and assessment is not None:
        session_label = (
            (session.title if session and session.title else None)
            or (f"Session {session.session_number}" if session and session.session_number else None)
            or f"Assessment {body.assessment_id}"
        )
        notify_admin_track_submit(
            db, user, step_label=session_label, reference_id=body.assessment_id
        )
    elif body.session_id is not None and session is not None:
        session_label = (
            session.title.strip()
            if session.title
            else (f"Session {session.session_number}" if session.session_number else f"Session {body.session_id}")
        )
        notify_admin_track_submit(db, user, step_label=session_label, reference_id=body.session_id)

    return get_programme_track(
        db, user, level=level_filter, module_kind=module_filter, sync=False
    )
