from __future__ import annotations

from sqlalchemy import func, text
from sqlalchemy.orm import Session, joinedload, selectinload

from app.models import User, UserCompletedAssessment, UserCompletedSession
from app.programme.models_programme import (
    CSV_LEVEL_SLUGS,
    LEVEL_DISPLAY_DEFAULTS,
    ProgrammeAssessment,
    ProgrammeLevel,
    ProgrammeModule,
    ProgrammeSession,
    ProgrammeWeek,
    normalize_programme_module_kind,
    parse_upp_assessment_csv,
    programme_level_slug_str,
    programme_module_kind_str,
    resolve_upp_csv_path,
)


def programme_schema_outdated(engine) -> bool:
    """
    True if programme tables exist but do not match current models (e.g. missing `kind`
    on programme_modules). `create_all()` never ALTERs old tables, so we must drop and recreate.
    """
    from sqlalchemy import inspect

    try:
        insp = inspect(engine)
        names = set(insp.get_table_names())
    except Exception:
        return False

    if "programme_modules" in names:
        cols = {c["name"] for c in insp.get_columns("programme_modules")}
        if "kind" not in cols:
            return True

    if "programme_sessions" in names:
        cols = {c["name"] for c in insp.get_columns("programme_sessions")}
        if "admin_checked" not in cols:
            return True
        try:
            checks = insp.get_check_constraints("programme_sessions")
        except Exception:
            checks = []
        for check in checks:
            name = (check.get("name") or "").lower()
            sqltext = (check.get("sqltext") or "").lower()
            if name == "ck_programme_session_number_range" or (
                "session_number" in sqltext and "<= 3" in sqltext
            ):
                return True

    return False


def sync_users_programme_columns(engine) -> None:
    """
    Add users.programme_session_id and users.programme_assessment_id when the table
    predates those ORM fields. create_all() does not ALTER existing tables.
    """
    from sqlalchemy import inspect, text

    url = str(engine.url).lower()
    try:
        insp = inspect(engine)
        if "users" not in insp.get_table_names():
            return
        cols = {c["name"] for c in insp.get_columns("users")}
    except Exception:
        return

    with engine.begin() as conn:
        if "programme_session_id" not in cols:
            if "postgresql" in url:
                conn.execute(text("ALTER TABLE users ADD COLUMN IF NOT EXISTS programme_session_id INTEGER"))
            else:
                try:
                    conn.execute(text("ALTER TABLE users ADD COLUMN programme_session_id INTEGER"))
                except Exception:
                    pass
        if "programme_assessment_id" not in cols:
            if "postgresql" in url:
                conn.execute(text("ALTER TABLE users ADD COLUMN IF NOT EXISTS programme_assessment_id INTEGER"))
            else:
                try:
                    conn.execute(text("ALTER TABLE users ADD COLUMN programme_assessment_id INTEGER"))
                except Exception:
                    pass

    if "postgresql" not in url:
        return

    try:
        insp = inspect(engine)
        names = set(insp.get_table_names())
    except Exception:
        return
    if "programme_sessions" not in names or "programme_assessments" not in names:
        return

    fk_session = """
        DO $$ BEGIN
            IF NOT EXISTS (
                SELECT 1 FROM pg_constraint c
                JOIN pg_class t ON c.conrelid = t.oid
                WHERE t.relname = 'users' AND c.conname = 'users_programme_session_id_fkey'
            ) THEN
                ALTER TABLE users ADD CONSTRAINT users_programme_session_id_fkey
                FOREIGN KEY (programme_session_id) REFERENCES programme_sessions(id) ON DELETE SET NULL;
            END IF;
        END $$;
    """
    fk_assessment = """
        DO $$ BEGIN
            IF NOT EXISTS (
                SELECT 1 FROM pg_constraint c
                JOIN pg_class t ON c.conrelid = t.oid
                WHERE t.relname = 'users' AND c.conname = 'users_programme_assessment_id_fkey'
            ) THEN
                ALTER TABLE users ADD CONSTRAINT users_programme_assessment_id_fkey
                FOREIGN KEY (programme_assessment_id) REFERENCES programme_assessments(id) ON DELETE SET NULL;
            END IF;
        END $$;
    """
    with engine.begin() as conn:
        conn.execute(text(fk_session))
        conn.execute(text(fk_assessment))


def _postgres_column_udt(engine, table: str, column: str) -> str | None:
    from sqlalchemy import text

    with engine.connect() as conn:
        row = conn.execute(
            text(
                """
                SELECT udt_name
                FROM information_schema.columns
                WHERE table_schema = 'public'
                  AND table_name = :table_name
                  AND column_name = :column_name
                """
            ),
            {"table_name": table, "column_name": column},
        ).first()
    return row[0] if row else None


def _postgres_programme_enum_types(engine) -> list[str]:
    from sqlalchemy import text

    with engine.connect() as conn:
        rows = conn.execute(
            text(
                """
                SELECT t.typname
                FROM pg_catalog.pg_type t
                JOIN pg_catalog.pg_namespace n ON n.oid = t.typnamespace
                WHERE n.nspname = 'public'
                  AND t.typtype = 'e'
                  AND (
                    t.typname ILIKE '%programme%'
                    OR t.typname ILIKE '%modulekind%'
                    OR t.typname ILIKE '%levelslug%'
                  )
                """
            )
        ).fetchall()
    return [row[0] for row in rows]


def sync_programme_string_columns(engine) -> None:
    """
    Allow custom module kinds and level slugs in existing databases.

    Older deployments stored programme_modules.kind and programme_levels.slug as
    SQLAlchemy/PostgreSQL enums (foundations/developmental/creativity only).
    After admins add custom values like test_modules, reads fail with LookupError
    until columns are plain VARCHAR — this runs that migration on startup.
    """
    import logging

    from sqlalchemy import inspect, text

    logger = logging.getLogger(__name__)
    url = str(engine.url).lower()
    is_postgres = "postgresql" in url or "postgres" in url
    is_mysql = "mysql" in url or "mariadb" in url

    try:
        insp = inspect(engine)
        tables = set(insp.get_table_names())
    except Exception as exc:
        logger.warning("programme string-column sync skipped (inspect failed): %s", exc)
        return

    def _column_type(table: str, column: str) -> str | None:
        try:
            for col in insp.get_columns(table):
                if col["name"] == column:
                    return str(col.get("type", "")).lower()
        except Exception:
            return None
        return None

    def _needs_varchar_migration(table: str, column: str) -> bool:
        if is_postgres:
            udt = (_postgres_column_udt(engine, table, column) or "").lower()
            return udt not in ("varchar", "character varying", "text", "bpchar")
        col_type = _column_type(table, column) or ""
        return "enum" in col_type

    def _run(label: str, sql: str) -> None:
        try:
            with engine.begin() as conn:
                conn.execute(text(sql))
            logger.info("programme migration ok: %s", label)
        except Exception as exc:
            logger.warning("programme migration skipped (%s): %s", label, exc)

    if "programme_modules" in tables and _needs_varchar_migration("programme_modules", "kind"):
        if is_postgres:
            _run(
                "programme_modules.kind -> varchar",
                "ALTER TABLE programme_modules "
                "ALTER COLUMN kind TYPE VARCHAR(64) USING kind::text",
            )
        elif is_mysql:
            _run(
                "programme_modules.kind -> varchar",
                "ALTER TABLE programme_modules MODIFY kind VARCHAR(64) NOT NULL",
            )

    if "programme_levels" in tables and _needs_varchar_migration("programme_levels", "slug"):
        if is_postgres:
            _run(
                "programme_levels.slug -> varchar",
                "ALTER TABLE programme_levels "
                "ALTER COLUMN slug TYPE VARCHAR(64) USING slug::text",
            )
        elif is_mysql:
            _run(
                "programme_levels.slug -> varchar",
                "ALTER TABLE programme_levels MODIFY slug VARCHAR(64) NOT NULL",
            )

    if not is_postgres:
        return

    for enum_name in _postgres_programme_enum_types(engine):
        _run(f"drop type {enum_name}", f'DROP TYPE IF EXISTS "{enum_name}" CASCADE')

    for enum_name in (
        "programmemodulekind",
        "programmelevelslug",
        "programme_module_kind",
        "programme_level_slug",
    ):
        _run(f"drop type {enum_name}", f"DROP TYPE IF EXISTS {enum_name} CASCADE")


def recreate_programme_schema() -> None:
    """
    Drop all programme_* tables so `init_db()` / `create_all()` can recreate them from ORM.

    Use when Postgres already had partial/old programme tables (e.g. missing `kind` column).
    """
    from app.database import engine

    url = str(engine.url).lower()

    with engine.begin() as conn:
        try:
            conn.execute(
                text("UPDATE users SET programme_session_id = NULL, programme_assessment_id = NULL")
            )
        except Exception:
            pass

    if "postgresql" in url:
        drops = [
            "DROP TABLE IF EXISTS user_completed_assessments CASCADE",
            "DROP TABLE IF EXISTS user_completed_sessions CASCADE",
            "DROP TABLE IF EXISTS programme_assessments CASCADE",
            "DROP TABLE IF EXISTS programme_sessions CASCADE",
            "DROP TABLE IF EXISTS programme_weeks CASCADE",
            "DROP TABLE IF EXISTS programme_modules CASCADE",
            "DROP TABLE IF EXISTS programme_levels CASCADE",
        ]
        with engine.begin() as conn:
            for stmt in drops:
                conn.execute(text(stmt))
        return

    drops = [
        "DROP TABLE IF EXISTS user_completed_assessments",
        "DROP TABLE IF EXISTS user_completed_sessions",
        "DROP TABLE IF EXISTS programme_assessments",
        "DROP TABLE IF EXISTS programme_sessions",
        "DROP TABLE IF EXISTS programme_weeks",
        "DROP TABLE IF EXISTS programme_modules",
        "DROP TABLE IF EXISTS programme_levels",
    ]
    with engine.begin() as conn:
        conn.execute(text("PRAGMA foreign_keys=OFF"))
        for stmt in drops:
            conn.execute(text(stmt))
        conn.execute(text("PRAGMA foreign_keys=ON"))


def import_programmes_from_csv(db: Session, *, progress: bool = False) -> bool:
    """Insert programme tree from CSV; programme tables must exist and match current models."""
    return _seed_programme_tables_from_csv(db, progress=progress)


def seed_programmes_from_csv_if_empty(db: Session, *, progress: bool = False) -> bool:
    """
    If no programme_levels rows exist, import all UPP CSV files from the repo root.
    Returns True if seeding ran, False if data was already present.
    """
    if db.query(ProgrammeLevel).first() is not None:
        return False
    return _seed_programme_tables_from_csv(db, progress=progress)


def clear_programme_tables(db: Session) -> None:
    """Remove all programme hierarchy rows (users' programme FKs become NULL via ON DELETE SET NULL)."""
    db.query(UserCompletedAssessment).delete(synchronize_session=False)
    db.query(UserCompletedSession).delete(synchronize_session=False)
    db.query(ProgrammeAssessment).delete(synchronize_session=False)
    db.query(ProgrammeSession).delete(synchronize_session=False)
    db.query(ProgrammeWeek).delete(synchronize_session=False)
    db.query(ProgrammeModule).delete(synchronize_session=False)
    db.query(ProgrammeLevel).delete(synchronize_session=False)
    db.commit()


def reseed_programmes_from_csv(db: Session, *, progress: bool = False) -> bool:
    """Wipe programme tables and import again from CSV files. Returns True if any level was imported."""
    clear_programme_tables(db)
    return _seed_programme_tables_from_csv(db, progress=progress)


def _seed_programme_tables_from_csv(db: Session, *, progress: bool = False) -> bool:
    order = 0
    any_level = False
    for slug in CSV_LEVEL_SLUGS:
        path = resolve_upp_csv_path(slug)
        if path is None:
            continue

        any_level = True
        if progress:
            print(f"  → importing level {slug!r} …", flush=True)

        level = ProgrammeLevel(
            slug=slug,
            display_name=LEVEL_DISPLAY_DEFAULTS.get(slug, slug.replace("_", " ").title()),
            sort_order=order,
        )
        db.add(level)
        db.flush()

        modules_data = parse_upp_assessment_csv(path)
        mod_order = 0
        for kind, weeks_payload in modules_data:
            mod = ProgrammeModule(
                level_id=level.id,
                kind=programme_module_kind_str(kind),
                display_name=programme_module_kind_str(kind).replace("_", " ").title(),
                sort_order=mod_order,
            )
            db.add(mod)
            db.flush()
            mod_order += 1

            for wi, week in enumerate(weeks_payload):
                wk = ProgrammeWeek(
                    module_id=mod.id,
                    week_index=wi,
                    display_label=week["display_label"][:64],
                )
                db.add(wk)
                db.flush()
                sessions_for_week: list[tuple[ProgrammeSession, dict]] = []
                for sn, sdata in sorted(week["sessions"].items()):
                    sess = ProgrammeSession(
                        week_id=wk.id,
                        session_number=sn,
                        title=sdata["title"][:512],
                        admin_checked=False,
                    )
                    db.add(sess)
                    sessions_for_week.append((sess, sdata))
                db.flush()
                for sess, sdata in sessions_for_week:
                    for ai, atitle in enumerate(sdata["assessments"]):
                        db.add(
                            ProgrammeAssessment(
                                session_id=sess.id,
                                sort_order=ai,
                                title=atitle[:512],
                            )
                        )

        order += 1

    if not any_level:
        db.rollback()
        return False

    if progress:
        print("  → commit (sab rows DB me likh rahe hain) …", flush=True)

    try:
        db.commit()
    except Exception:
        db.rollback()
        raise
    return True


def get_level_by_slug(db: Session, slug: str) -> ProgrammeLevel | None:
    return db.query(ProgrammeLevel).filter(ProgrammeLevel.slug == slug).first()


def serialize_programme_catalog(
    db: Session,
    *,
    level_slug: str | None = None,
) -> dict:
    """Load the full programme hierarchy from SQL for frontend catalog screens."""
    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)

    levels = query.all()
    total_modules = 0
    total_weeks = 0
    total_sessions = 0
    total_assessments = 0
    levels_out: list[dict] = []

    for level in levels:
        modules_out: list[dict] = []
        level_weeks = 0
        level_sessions = 0
        level_assessments = 0

        for module in sorted(level.modules, key=lambda m: m.sort_order):
            weeks_out: list[dict] = []
            module_sessions = 0
            module_assessments = 0

            for week in sorted(module.weeks, key=lambda w: w.week_index):
                sessions_out: list[dict] = []
                week_assessments = 0

                for session in sorted(week.sessions, key=lambda s: s.session_number):
                    assessments = sorted(session.assessments, key=lambda a: a.sort_order)
                    assessments_out = [
                        {
                            "id": assessment.id,
                            "sort_order": assessment.sort_order,
                            "title": assessment.title,
                        }
                        for assessment in assessments
                    ]
                    week_assessments += len(assessments_out)
                    sessions_out.append(
                        {
                            "id": session.id,
                            "session_number": session.session_number,
                            "title": session.title,
                            "admin_checked": session.admin_checked,
                            "assessments_count": len(assessments_out),
                            "assessments": assessments_out,
                        }
                    )

                module_sessions += len(sessions_out)
                module_assessments += week_assessments
                weeks_out.append(
                    {
                        "id": week.id,
                        "week_index": week.week_index,
                        "display_label": week.display_label,
                        "sessions_count": len(sessions_out),
                        "assessments_count": week_assessments,
                        "sessions": sessions_out,
                    }
                )

            level_weeks += len(weeks_out)
            level_sessions += module_sessions
            level_assessments += module_assessments
            modules_out.append(
                {
                    "id": module.id,
                    "kind": programme_module_kind_str(module.kind),
                    "display_name": module.display_name,
                    "sort_order": module.sort_order,
                    "weeks_count": len(weeks_out),
                    "sessions_count": module_sessions,
                    "assessments_count": module_assessments,
                    "weeks": weeks_out,
                }
            )

        total_modules += len(modules_out)
        total_weeks += level_weeks
        total_sessions += level_sessions
        total_assessments += level_assessments
        levels_out.append(
            {
                "id": level.id,
                "slug": programme_level_slug_str(level.slug),
                "display_name": level.display_name,
                "sort_order": level.sort_order,
                "modules_count": len(modules_out),
                "weeks_count": level_weeks,
                "sessions_count": level_sessions,
                "assessments_count": level_assessments,
                "modules": modules_out,
            }
        )

    return {
        "total_levels": len(levels_out),
        "total_modules": total_modules,
        "total_weeks": total_weeks,
        "total_sessions": total_sessions,
        "total_assessments": total_assessments,
        "levels": levels_out,
    }


def resolve_session(
    db: Session,
    *,
    level_slug: str,
    module_kind: str,
    week_index: int,
    session_number: int,
) -> ProgrammeSession | None:
    kind_str = normalize_programme_module_kind(module_kind)
    return (
        db.query(ProgrammeSession)
        .join(ProgrammeWeek, ProgrammeSession.week_id == ProgrammeWeek.id)
        .join(ProgrammeModule, ProgrammeWeek.module_id == ProgrammeModule.id)
        .join(ProgrammeLevel, ProgrammeModule.level_id == ProgrammeLevel.id)
        .filter(
            ProgrammeLevel.slug == level_slug,
            ProgrammeModule.kind == kind_str,
            ProgrammeWeek.week_index == week_index,
            ProgrammeSession.session_number == session_number,
        )
        .first()
    )


def session_assessment_count(db: Session, session_id: int) -> int:
    return (
        db.query(func.count(ProgrammeAssessment.id)).filter(ProgrammeAssessment.session_id == session_id).scalar()
        or 0
    )


def validate_user_programme_selection(
    db: Session,
    *,
    level_slug: str,
    module_kind: str,
    week_index: int,
    session_number: int,
    assessment_id: int | None,
) -> tuple[ProgrammeSession, ProgrammeAssessment | None]:
    if not (0 <= week_index <= 12):
        raise ValueError("week_index must be between 0 (intro) and 12")
    if session_number < 1:
        raise ValueError("session_number must be a positive integer")

    session = resolve_session(
        db,
        level_slug=level_slug,
        module_kind=module_kind,
        week_index=week_index,
        session_number=session_number,
    )
    if session is None:
        raise ValueError("No matching session for the given programme path")

    n = session_assessment_count(db, session.id)
    if n > 0:
        if assessment_id is None:
            raise ValueError(
                f"This session has {n} assessment(s). Send assessment_id from "
                f"GET /api/v1/admin/programme/catalog?level={level_slug}"
                f"&modules=true&weeks=true&sessions=true&assessments=true "
                f"(use the id inside that session's assessments list)."
            )
        a = db.get(ProgrammeAssessment, assessment_id)
        if a is None or a.session_id != session.id:
            raise ValueError(
                f"assessment_id {assessment_id} is not valid for "
                f"{level_slug}/{module_kind}/week_index={week_index}/session={session_number}. "
                f"Pick an id from that session in the programme catalog (not a random global id)."
            )
        return session, a

    if assessment_id is not None:
        raise ValueError(
            f"{level_slug}/{module_kind}/week_index={week_index}/session={session_number} "
            f"has no assessments in the database. Remove assessment_id from the body, or choose another "
            f"week/session that has assessments (see admin programme catalog)."
        )
    return session, None


class _CompletionBatch:
    """Batch programme completion writes (avoids one SELECT per assessment on register)."""

    def __init__(self, db: Session, user_id: int) -> None:
        self.db = db
        self.user_id = user_id
        self.done_assessments: set[int] = {
            row[0]
            for row in db.query(UserCompletedAssessment.assessment_id)
            .filter(UserCompletedAssessment.user_id == user_id)
            .all()
        }
        self.done_sessions: set[int] = {
            row[0]
            for row in db.query(UserCompletedSession.session_id)
            .filter(UserCompletedSession.user_id == user_id)
            .all()
        }
        self.new_assessments: list[int] = []
        self.new_sessions: list[int] = []

    def ensure_assessment(self, assessment_id: int) -> None:
        if assessment_id in self.done_assessments:
            return
        self.done_assessments.add(assessment_id)
        self.new_assessments.append(assessment_id)

    def ensure_session(self, session_id: int) -> None:
        if session_id in self.done_sessions:
            return
        self.done_sessions.add(session_id)
        self.new_sessions.append(session_id)

    def flush(self) -> None:
        if not self.new_assessments and not self.new_sessions:
            return
        bind = self.db.get_bind()
        if bind is not None and bind.dialect.name == "postgresql":
            from sqlalchemy.dialects.postgresql import insert as pg_insert

            if self.new_assessments:
                self.db.execute(
                    pg_insert(UserCompletedAssessment)
                    .values(
                        [
                            {"user_id": self.user_id, "assessment_id": assessment_id}
                            for assessment_id in self.new_assessments
                        ]
                    )
                    .on_conflict_do_nothing()
                )
            if self.new_sessions:
                self.db.execute(
                    pg_insert(UserCompletedSession)
                    .values(
                        [
                            {"user_id": self.user_id, "session_id": session_id}
                            for session_id in self.new_sessions
                        ]
                    )
                    .on_conflict_do_nothing()
                )
        else:
            for assessment_id in self.new_assessments:
                self.db.add(
                    UserCompletedAssessment(user_id=self.user_id, assessment_id=assessment_id)
                )
            for session_id in self.new_sessions:
                self.db.add(UserCompletedSession(user_id=self.user_id, session_id=session_id))
        self.new_assessments.clear()
        self.new_sessions.clear()


def _pending_user_completed_session(db: Session, user_id: int, session_id: int) -> bool:
    for obj in db.new:
        if (
            isinstance(obj, UserCompletedSession)
            and obj.user_id == user_id
            and obj.session_id == session_id
        ):
            return True
    return False


def _ensure_assessment_done(
    db: Session,
    user_id: int,
    assessment_id: int,
    batch: _CompletionBatch | None = None,
) -> None:
    if batch is not None:
        batch.ensure_assessment(assessment_id)
        return
    exists = db.query(UserCompletedAssessment).filter_by(
        user_id=user_id, assessment_id=assessment_id
    ).first()
    if not exists:
        db.add(UserCompletedAssessment(user_id=user_id, assessment_id=assessment_id))


def _ensure_session_marker_done(
    db: Session,
    user_id: int,
    session_id: int,
    batch: _CompletionBatch | None = None,
) -> None:
    if batch is not None:
        batch.ensure_session(session_id)
        return
    if _pending_user_completed_session(db, user_id, session_id):
        return
    exists = db.query(UserCompletedSession).filter_by(
        user_id=user_id, session_id=session_id
    ).first()
    if not exists:
        db.add(UserCompletedSession(user_id=user_id, session_id=session_id))


def _complete_session_fully(
    db: Session,
    user_id: int,
    sess: ProgrammeSession,
    batch: _CompletionBatch | None = None,
) -> None:
    assessments = sorted(sess.assessments, key=lambda a: a.sort_order)
    if not assessments:
        _ensure_session_marker_done(db, user_id, sess.id, batch)
    else:
        for a in assessments:
            _ensure_assessment_done(db, user_id, a.id, batch)


def _complete_entire_week(
    db: Session,
    user_id: int,
    week: ProgrammeWeek,
    batch: _CompletionBatch | None = None,
) -> None:
    for sess in sorted(week.sessions, key=lambda s: s.session_number):
        _complete_session_fully(db, user_id, sess, batch)


def _complete_entire_module(
    db: Session,
    user_id: int,
    module: ProgrammeModule,
    batch: _CompletionBatch | None = None,
) -> None:
    for wk in sorted(module.weeks, key=lambda w: w.week_index):
        _complete_entire_week(db, user_id, wk, batch)


def _complete_entire_level(
    db: Session,
    user_id: int,
    level: ProgrammeLevel,
    batch: _CompletionBatch | None = None,
) -> None:
    for mod in sorted(level.modules, key=lambda m: m.sort_order):
        _complete_entire_module(db, user_id, mod, batch)


def _complete_partial_week(
    db: Session,
    user_id: int,
    week: ProgrammeWeek,
    target_sess: ProgrammeSession,
    target_assessment: ProgrammeAssessment | None,
    batch: _CompletionBatch | None = None,
) -> None:
    for sess in sorted(week.sessions, key=lambda s: s.session_number):
        if sess.session_number < target_sess.session_number:
            _complete_session_fully(db, user_id, sess, batch)
        elif sess.id == target_sess.id:
            assessments = sorted(sess.assessments, key=lambda a: a.sort_order)
            if assessments and target_assessment is not None:
                for a in assessments:
                    if a.sort_order < target_assessment.sort_order:
                        _ensure_assessment_done(db, user_id, a.id, batch)
            break
        else:
            break


def _complete_partial_module(
    db: Session,
    user_id: int,
    module: ProgrammeModule,
    target_week: ProgrammeWeek,
    target_sess: ProgrammeSession,
    target_assessment: ProgrammeAssessment | None,
    batch: _CompletionBatch | None = None,
) -> None:
    for wk in sorted(module.weeks, key=lambda w: w.week_index):
        if wk.week_index < target_week.week_index:
            _complete_entire_week(db, user_id, wk, batch)
        elif wk.id == target_week.id:
            _complete_partial_week(db, user_id, wk, target_sess, target_assessment, batch)
            break
        else:
            break


def _complete_partial_level(
    db: Session,
    user_id: int,
    level: ProgrammeLevel,
    target_module: ProgrammeModule,
    target_week: ProgrammeWeek,
    target_sess: ProgrammeSession,
    target_assessment: ProgrammeAssessment | None,
    batch: _CompletionBatch | None = None,
) -> None:
    for mod in sorted(level.modules, key=lambda m: m.sort_order):
        if mod.sort_order < target_module.sort_order:
            _complete_entire_module(db, user_id, mod, batch)
        elif mod.id == target_module.id:
            _complete_partial_module(
                db, user_id, mod, target_week, target_sess, target_assessment, batch
            )
            break
        else:
            break


def seed_prior_programme_completions(
    db: Session,
    user_id: int,
    target_session: ProgrammeSession,
    target_assessment: ProgrammeAssessment | None,
) -> None:
    """
    Mark everything strictly before the user's chosen position as completed:
    all prior levels fully; within selected level, prior modules/weeks/sessions/assessments.
    """
    sid = target_session.id
    sess_row = (
        db.query(ProgrammeSession)
        .options(
            joinedload(ProgrammeSession.assessments),
            joinedload(ProgrammeSession.week)
            .joinedload(ProgrammeWeek.module)
            .joinedload(ProgrammeModule.level),
            joinedload(ProgrammeSession.week).joinedload(ProgrammeWeek.sessions).joinedload(ProgrammeSession.assessments),
        )
        .filter(ProgrammeSession.id == sid)
        .first()
    )
    if sess_row is None:
        return

    tw = sess_row.week
    tm = tw.module
    tl = tm.level

    levels = (
        db.query(ProgrammeLevel)
        .options(
            joinedload(ProgrammeLevel.modules)
            .joinedload(ProgrammeModule.weeks)
            .joinedload(ProgrammeWeek.sessions)
            .joinedload(ProgrammeSession.assessments)
        )
        .order_by(ProgrammeLevel.sort_order)
        .all()
    )

    batch = _CompletionBatch(db, user_id)
    for level in levels:
        if level.sort_order < tl.sort_order:
            _complete_entire_level(db, user_id, level, batch)
        elif level.id == tl.id:
            _complete_partial_level(
                db, user_id, level, tm, tw, sess_row, target_assessment, batch
            )
            break
    batch.flush()


def load_user_programme_summary(user: User) -> dict | None:
    """Build a dict for UserPublic.programme from a User ORM instance (relationships optional)."""
    if user.programme_session_id is None:
        return None
    sess = user.programme_session
    if sess is None:
        return {"session_id": user.programme_session_id, "assessment_id": user.programme_assessment_id}
    week = sess.week
    mod = week.module if week else None
    level = mod.level if mod else None
    ass = user.programme_assessment
    return {
        "level_slug": programme_level_slug_str(level.slug) if level else None,
        "level_display_name": level.display_name if level 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": sess.session_number,
        "session_title": sess.title,
        "session_admin_checked": sess.admin_checked,
        "session_id": sess.id,
        "assessment_id": ass.id if ass else None,
        "assessment_title": ass.title if ass else None,
    }


def user_programme_load_options():
    return (
        joinedload(User.programme_session)
        .joinedload(ProgrammeSession.week)
        .joinedload(ProgrammeWeek.module)
        .joinedload(ProgrammeModule.level),
        joinedload(User.programme_assessment),
    )


def serialize_user_public(user: User, db: Session) -> "UserPublic":
    from app.schemas import RoleEnum, UserProgrammeProgressPublic, UserPublic

    raw = load_user_programme_summary(user)
    if raw:
        ca = (
            db.query(func.count(UserCompletedAssessment.user_id))
            .filter(UserCompletedAssessment.user_id == user.id)
            .scalar()
            or 0
        )
        cs = (
            db.query(func.count(UserCompletedSession.user_id))
            .filter(UserCompletedSession.user_id == user.id)
            .scalar()
            or 0
        )
        raw["prior_completed_assessment_count"] = int(ca)
        raw["prior_completed_session_markers_count"] = int(cs)
    prog = UserProgrammeProgressPublic.model_validate(raw) if raw else None
    return UserPublic(
        id=user.id,
        email=user.email,
        full_name=user.full_name,
        role=RoleEnum(user.role.value),
        is_active=user.is_active,
        created_at=user.created_at,
        programme=prog,
    )
