from __future__ import annotations

from collections import defaultdict
from dataclasses import dataclass, field

from sqlalchemy import func, or_
from sqlalchemy.orm import Session, selectinload

from app.models import User, UserCompletedAssessment, UserCompletedSession, UserRegisterStep
from app.models_admin import MemberSessionProgress
from app.programme.models_programme import (
    ProgrammeAssessment,
    ProgrammeLevel,
    ProgrammeModule,
    ProgrammeSession,
    ProgrammeWeek,
    programme_level_slug_str,
    programme_module_kind_str,
)
from app.services.admin_service import calculate_member_points, get_member_points, media_public_url
from app.services.user_register_step_service import get_user_register_step


@dataclass
class _DashboardSessionRow:
    session_id: int
    session_number: int
    session_title: str
    week_id: int
    week_index: int
    week_label: str
    module_id: int
    module_kind: str
    module_name: str
    module_sort_order: int
    assessments: list[tuple[int, str]] = field(default_factory=list)


def _resolve_current_level(
    db: Session,
    user: User,
    *,
    register_step: UserRegisterStep | None = None,
) -> ProgrammeLevel | None:
    """Pick the level shown on the user's home/dashboard screen."""
    step = register_step or db.query(UserRegisterStep).filter(UserRegisterStep.user_id == user.id).first()
    if step:
        if step.level_id is not None:
            level = db.get(ProgrammeLevel, step.level_id)
            if level is not None:
                return level
        if step.level_name:
            raw = step.level_name.strip()
            level = (
                db.query(ProgrammeLevel)
                .filter(
                    or_(
                        func.lower(ProgrammeLevel.display_name) == raw.lower(),
                        func.lower(ProgrammeLevel.slug) == raw.lower(),
                    )
                )
                .first()
            )
            if level is not None:
                return level

    if user.programme_session_id is not None:
        session = (
            db.query(ProgrammeSession)
            .join(ProgrammeWeek, ProgrammeSession.week_id == ProgrammeWeek.id)
            .join(ProgrammeModule, ProgrammeWeek.module_id == ProgrammeModule.id)
            .join(ProgrammeLevel, ProgrammeModule.level_id == ProgrammeLevel.id)
            .options(selectinload(ProgrammeSession.week).selectinload(ProgrammeWeek.module).selectinload(ProgrammeModule.level))
            .filter(ProgrammeSession.id == user.programme_session_id)
            .first()
        )
        if session and session.week and session.week.module:
            return session.week.module.level

    return db.query(ProgrammeLevel).order_by(ProgrammeLevel.sort_order).first()


def _is_session_completed(
    session: ProgrammeSession,
    week: ProgrammeWeek,
    *,
    completed_session_ids: set[int],
    completed_assessment_ids: set[int],
    completed_week_session_keys: set[tuple[int, int]],
    completed_index_session_keys: set[tuple[int, int]],
) -> bool:
    assessments = sorted(session.assessments, key=lambda item: item.sort_order)
    if assessments:
        return all(assessment.id in completed_assessment_ids for assessment in assessments)
    return (
        session.id in completed_session_ids
        or (week.id, session.session_number) in completed_week_session_keys
        or (week.week_index, session.session_number) in completed_index_session_keys
    )


def _is_flat_session_completed(
    row: _DashboardSessionRow,
    *,
    completed_session_ids: set[int],
    completed_assessment_ids: set[int],
    completed_week_session_keys: set[tuple[int, int]],
    completed_index_session_keys: set[tuple[int, int]],
) -> bool:
    if row.assessments:
        return all(assessment_id in completed_assessment_ids for assessment_id, _ in row.assessments)
    return (
        row.session_id in completed_session_ids
        or (row.week_id, row.session_number) in completed_week_session_keys
        or (row.week_index, row.session_number) in completed_index_session_keys
    )


def _load_completion_sets(
    db: Session, user_id: int
) -> tuple[set[int], set[int], set[tuple[int, int]], set[tuple[int, int]]]:
    completed_assessment_ids = {
        row[0]
        for row in db.query(UserCompletedAssessment.assessment_id)
        .filter(UserCompletedAssessment.user_id == user_id)
        .all()
    }
    completed_track_session_ids = {
        row[0]
        for row in db.query(UserCompletedSession.session_id)
        .filter(UserCompletedSession.user_id == user_id)
        .all()
    }
    completed_progress = (
        db.query(
            MemberSessionProgress.session_id,
            MemberSessionProgress.week_id,
            MemberSessionProgress.week_index,
            MemberSessionProgress.session_number,
        )
        .filter(
            MemberSessionProgress.user_id == user_id,
            MemberSessionProgress.is_completed.is_(True),
        )
        .all()
    )
    completed_session_ids = completed_track_session_ids | {
        session_id for session_id, _, _, _ in completed_progress if session_id is not None
    }
    completed_week_session_keys = {
        (week_id, session_number)
        for _, week_id, _, session_number in completed_progress
        if week_id is not None and session_number is not None
    }
    completed_index_session_keys = {
        (week_index, session_number)
        for session_id, week_id, week_index, session_number in completed_progress
        if session_id is None
        and week_id is None
        and week_index is not None
        and session_number is not None
    }
    return (
        completed_assessment_ids,
        completed_session_ids,
        completed_week_session_keys,
        completed_index_session_keys,
    )


def _load_level_sessions_flat(
    db: Session,
    level_id: int,
    *,
    module_id: int | None = None,
) -> list[_DashboardSessionRow]:
    """Flat session rows for one level — avoids loading the full ORM programme tree."""
    query = (
        db.query(
            ProgrammeSession.id,
            ProgrammeSession.session_number,
            ProgrammeSession.title,
            ProgrammeWeek.id,
            ProgrammeWeek.week_index,
            ProgrammeWeek.display_label,
            ProgrammeModule.id,
            ProgrammeModule.kind,
            ProgrammeModule.display_name,
            ProgrammeModule.sort_order,
        )
        .join(ProgrammeWeek, ProgrammeSession.week_id == ProgrammeWeek.id)
        .join(ProgrammeModule, ProgrammeWeek.module_id == ProgrammeModule.id)
        .filter(ProgrammeModule.level_id == level_id)
        .order_by(ProgrammeModule.sort_order, ProgrammeWeek.week_index, ProgrammeSession.session_number)
    )
    if module_id is not None:
        query = query.filter(ProgrammeModule.id == module_id)

    raw_rows = query.all()
    if not raw_rows:
        return []

    session_ids = [row[0] for row in raw_rows]
    assessments_by_session: dict[int, list[tuple[int, str]]] = defaultdict(list)
    if session_ids:
        for session_id, assessment_id, assessment_title in (
            db.query(
                ProgrammeAssessment.session_id,
                ProgrammeAssessment.id,
                ProgrammeAssessment.title,
            )
            .filter(ProgrammeAssessment.session_id.in_(session_ids))
            .order_by(ProgrammeAssessment.sort_order)
            .all()
        ):
            assessments_by_session[session_id].append((assessment_id, assessment_title))

    return [
        _DashboardSessionRow(
            session_id=session_id,
            session_number=session_number,
            session_title=title or f"Session {session_number}",
            week_id=week_id,
            week_index=week_index,
            week_label=week_label,
            module_id=module_id,
            module_kind=programme_module_kind_str(module_kind),
            module_name=module_name,
            module_sort_order=module_sort_order,
            assessments=assessments_by_session.get(session_id, []),
        )
        for session_id, session_number, title, week_id, week_index, week_label, module_id, module_kind, module_name, module_sort_order in raw_rows
    ]


def _compute_module_progress(
    session_rows: list[_DashboardSessionRow],
    *,
    completed_session_ids: set[int],
    completed_assessment_ids: set[int],
    completed_week_session_keys: set[tuple[int, int]],
    completed_index_session_keys: set[tuple[int, int]],
) -> tuple[list[dict], int, int]:
    modules_out: list[dict] = []
    total_sessions = 0
    completed_sessions = 0
    weeks_by_module: dict[int, set[int]] = defaultdict(set)

    by_module: dict[int, list[_DashboardSessionRow]] = defaultdict(list)
    for row in session_rows:
        by_module[row.module_id].append(row)
        weeks_by_module[row.module_id].add(row.week_index)

    for module_id in sorted(by_module, key=lambda mid: by_module[mid][0].module_sort_order):
        rows = by_module[module_id]
        module_total = 0
        module_completed = 0
        sample = rows[0]
        for row in rows:
            module_total += 1
            if _is_flat_session_completed(
                row,
                completed_session_ids=completed_session_ids,
                completed_assessment_ids=completed_assessment_ids,
                completed_week_session_keys=completed_week_session_keys,
                completed_index_session_keys=completed_index_session_keys,
            ):
                module_completed += 1

        total_sessions += module_total
        completed_sessions += module_completed
        modules_out.append(
            {
                "module_id": module_id,
                "module_kind": sample.module_kind,
                "module_name": sample.module_name,
                "completed_sessions": module_completed,
                "total_sessions": module_total,
                "progress_percent": round((module_completed / module_total) * 100) if module_total else 0,
                "total_weeks": len(weeks_by_module[module_id]),
            }
        )

    return modules_out, total_sessions, completed_sessions


def _build_current_step_from_sessions(
    session_rows: list[_DashboardSessionRow],
    register_step: UserRegisterStep | None,
    level: ProgrammeLevel | None,
    *,
    completed_assessment_ids: set[int],
    completed_session_ids: set[int],
) -> dict | None:
    if level is None or register_step is None or not register_step.is_completed:
        return None
    if not session_rows:
        return None

    ordered = sorted(
        session_rows,
        key=lambda row: (row.module_sort_order, row.week_index, row.session_number),
    )

    anchor_index = 0
    if register_step.session_id is not None:
        for index, row in enumerate(ordered):
            if row.session_id == register_step.session_id:
                anchor_index = index
                break

    for index, row in enumerate(ordered):
        if index < anchor_index:
            continue

        for assessment_id, assessment_title in row.assessments:
            if assessment_id not in completed_assessment_ids:
                return {
                    "week_id": row.week_id,
                    "week_label": row.week_label,
                    "session_id": row.session_id,
                    "session_title": row.session_title,
                    "session_number": row.session_number,
                    "assessment_id": assessment_id,
                    "assessment_title": assessment_title,
                    "module_kind": row.module_kind,
                    "level_slug": programme_level_slug_str(level.slug),
                    "label": " · ".join(
                        [row.week_label, f"Session {row.session_number}", assessment_title]
                    ),
                }

        if row.assessments:
            continue
        if row.session_id in completed_session_ids:
            continue

        return {
            "week_id": row.week_id,
            "week_label": row.week_label,
            "session_id": row.session_id,
            "session_title": row.session_title,
            "session_number": row.session_number,
            "assessment_id": None,
            "assessment_title": None,
            "module_kind": row.module_kind,
            "level_slug": programme_level_slug_str(level.slug),
            "label": " · ".join([row.week_label, f"Session {row.session_number}"]),
        }

    return None


def _build_track_current_step(track) -> dict | None:
    current_open = track.current_open or {}
    if not current_open.get("session_id"):
        return None

    label_parts: list[str] = []
    week_label = None
    session_title = None
    assessment_title = None

    for module in track.level.modules if track.level else []:
        for week in module.weeks:
            for session in week.sessions:
                if session.id != current_open.get("session_id"):
                    continue
                week_label = week.display_label
                session_title = session.title
                label_parts.append(week.display_label)
                label_parts.append(f"Session {session.session_number}")
                for assessment in session.assessments:
                    if assessment.id == current_open.get("assessment_id"):
                        assessment_title = assessment.title
                        label_parts.append(assessment.title)
                        break
                return {
                    "week_id": week.id,
                    "week_label": week_label,
                    "session_id": session.id,
                    "session_title": session_title,
                    "session_number": session.session_number,
                    "assessment_id": current_open.get("assessment_id"),
                    "assessment_title": assessment_title,
                    "module_kind": programme_module_kind_str(module.kind),
                    "level_slug": programme_level_slug_str(track.level.slug) if track.level else None,
                    "label": " · ".join(label_parts),
                }
    return None


def _serialize_last_progress(row: MemberSessionProgress | None) -> dict | None:
    if row is None:
        return None
    return {
        "id": row.id,
        "week_id": row.week_id,
        "week_name": row.week_name,
        "week_index": row.week_index,
        "session_id": row.session_id,
        "session_name": row.session_name,
        "session_number": row.session_number,
        "assigned_points": row.assigned_points,
        "feedback": row.feedback,
        "image_url": media_public_url(row.image_path),
        "is_completed": row.is_completed,
        "updated_at": row.updated_at,
    }


def _build_current_step_from_level(
    level: ProgrammeLevel | None,
    register_step: UserRegisterStep | None,
    *,
    completed_assessment_ids: set[int],
    completed_session_ids: set[int],
) -> dict | None:
    if level is None or register_step is None or not register_step.is_completed:
        return None

    modules = sorted(level.modules, key=lambda item: item.sort_order)
    if register_step.module_id is not None:
        modules = [module for module in modules if module.id == register_step.module_id]

    refs: list[tuple[ProgrammeModule, ProgrammeWeek, ProgrammeSession]] = []
    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((module, week, session))

    if not refs:
        return None

    anchor_index = 0
    if register_step.session_id is not None:
        for index, (_, _, session) in enumerate(refs):
            if session.id == register_step.session_id:
                anchor_index = index
                break

    for index, (module, week, session) in enumerate(refs):
        if index < anchor_index:
            continue

        assessments = sorted(session.assessments, key=lambda item: item.sort_order)
        assessment = next(
            (item for item in assessments if item.id not in completed_assessment_ids),
            None,
        )
        if assessment is None and assessments:
            continue
        if assessment is None and session.id in completed_session_ids:
            continue

        label_parts = [week.display_label, f"Session {session.session_number}"]
        if assessment is not None:
            label_parts.append(assessment.title)

        return {
            "week_id": week.id,
            "week_label": week.display_label,
            "session_id": session.id,
            "session_title": session.title,
            "session_number": session.session_number,
            "assessment_id": assessment.id if assessment else None,
            "assessment_title": assessment.title if assessment else None,
            "module_kind": programme_module_kind_str(module.kind),
            "level_slug": programme_level_slug_str(level.slug),
            "label": " · ".join(label_parts),
        }

    return None


def get_user_dashboard(db: Session, user: User) -> dict:
    register_step = get_user_register_step(db, user)
    level = _resolve_current_level(db, user, register_step=register_step)

    (
        completed_assessment_ids,
        completed_session_ids,
        completed_week_session_keys,
        completed_index_session_keys,
    ) = _load_completion_sets(db, user.id)

    latest_progress = (
        db.query(MemberSessionProgress)
        .filter(MemberSessionProgress.user_id == user.id)
        .order_by(MemberSessionProgress.updated_at.desc(), MemberSessionProgress.id.desc())
        .first()
    )

    points = get_member_points(db, user.id) or calculate_member_points(db, user.id)
    total_points = points["assigned_points_total"]
    available_points = points["available_points"]

    modules_out: list[dict] = []
    total_sessions = 0
    completed_sessions = 0
    session_rows: list[_DashboardSessionRow] = []

    if level is not None:
        module_filter = register_step.module_id if register_step and register_step.module_id else None
        session_rows = _load_level_sessions_flat(db, level.id, module_id=module_filter)
        modules_out, total_sessions, completed_sessions = _compute_module_progress(
            session_rows,
            completed_session_ids=completed_session_ids,
            completed_assessment_ids=completed_assessment_ids,
            completed_week_session_keys=completed_week_session_keys,
            completed_index_session_keys=completed_index_session_keys,
        )

    last_progress_percent = round((completed_sessions / total_sessions) * 100) if total_sessions else 0
    level_label = f"{level.display_name} Level" if level is not None else None

    track_current_step = _build_current_step_from_sessions(
        session_rows,
        register_step,
        level,
        completed_assessment_ids=completed_assessment_ids,
        completed_session_ids=completed_session_ids,
    )

    return {
        "user_id": user.id,
        "full_name": user.full_name,
        "level_id": level.id if level else None,
        "level_slug": programme_level_slug_str(level.slug) if level else None,
        "level_name": level_label,
        "total_points": total_points,
        "assigned_points_total": total_points,
        "available_points": available_points,
        "completed_sessions": completed_sessions,
        "total_sessions": total_sessions,
        "last_progress_percent": last_progress_percent,
        "stats": [
            {
                "label": "Points",
                "value": available_points,
                "helper_text": f"Total assigned: {total_points}",
            },
            {
                "label": "Session Complete",
                "value": f"{completed_sessions}/{total_sessions}",
                "helper_text": "Completed sessions in current level",
            },
            {
                "label": "Last Progress",
                "value": f"{last_progress_percent}%",
                "helper_text": "Current level completion",
            },
        ],
        "current_level_progress": modules_out,
        "last_progress": _serialize_last_progress(latest_progress),
        "track_current_step": track_current_step,
    }
