from __future__ import annotations

from fastapi import HTTPException, status
from sqlalchemy import func, or_
from sqlalchemy.orm import Session, selectinload

from app.models import User, UserCompletedAssessment, UserCompletedSession, UserRegisterStep, UserRole
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 (
    _complete_entire_week,
    load_user_programme_summary,
    resolve_session,
    user_programme_load_options,
    validate_user_programme_selection,
)
from app.schemas import (
    AdminMemberPersonalInfoUpdateRequest,
    MarkMemberWeekCompleteRequest,
    MemberWeekDetailOut,
    MemberWeekSessionDetailOut,
    MemberWeeksDetailResponse,
    RoleEnum,
    UpdateUserProgrammeRequest,
    UserProgrammeProgressPublic,
    UserPublic,
)
from app.services.admin_service import calculate_member_points, list_member_session_progress, media_public_url, refresh_member_points
from app.services.user_dashboard_service import get_user_dashboard


def _format_week_label(week_index: int | None) -> str | None:
    if week_index is None:
        return None
    n = week_index + 1 
    if n == 1:
        return "1st week"
    if n == 2:
        return "2nd week"
    if n == 3:
        return "3rd week"
    return f"{n}th week"


def compute_member_progress_percent(programme: UserProgrammeProgressPublic | None) -> int:
    if programme is None:
        return 0

    completed = (programme.prior_completed_assessment_count or 0) + (
        programme.prior_completed_session_markers_count or 0
    )
    week = programme.week_index or 0
    session = programme.session_number or 0

    pct = min(100, completed * 8 + week * 4 + session * 2)
    if pct == 0 and (programme.level_display_name or programme.module_display_name):
        pct = 12
    return pct


def _load_member_user(db: Session, user_id: int) -> User:
    user = (
        db.query(User)
        .options(*user_programme_load_options())
        .filter(User.id == user_id)
        .first()
    )
    if user is None:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="User not found")
    if user.role != UserRole.user:
        raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="User is not a member")
    return user


def _build_programme_public(db: Session, user: User) -> UserProgrammeProgressPublic | None:
    from app.services.user_track_service import resolve_user_programme_summary

    raw = resolve_user_programme_summary(db, user)
    if not raw:
        return None

    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)
    return UserProgrammeProgressPublic.model_validate(raw)


def _has_google_drive_link(user: User) -> bool:
    return bool((user.google_drive_link or "").strip())


def _build_programme_public_light(
    user: User,
    db: Session | None = None,
    *,
    programme_levels: list[ProgrammeLevel] | None = None,
) -> UserProgrammeProgressPublic | None:
    if db is not None:
        from app.services.user_track_service import resolve_user_programme_summary

        raw = resolve_user_programme_summary(db, user, programme_levels=programme_levels)
    else:
        raw = load_user_programme_summary(user)
    if not raw:
        return None
    return UserProgrammeProgressPublic.model_validate(raw)


def serialize_member_google_drive_entry(user: User, db: Session) -> dict:
    programme = _build_programme_public_light(user)
    return {
        "id": user.id,
        "full_name": user.full_name,
        "email": user.email,
        "phone_number": user.phone_number,
        "google_drive_link": user.google_drive_link,
        "has_google_drive_link": _has_google_drive_link(user),
        "week_label": programme.week_label if programme else None,
        "level_display_name": programme.level_display_name if programme else None,
        "module_display_name": programme.module_display_name if programme else None,
    }


def list_members_google_drive_links(
    db: Session,
    *,
    skip: int = 0,
    limit: int = 50,
    search: str | None = None,
    link_status: str | None = None,
) -> tuple[list[User], int]:
    query = db.query(User).filter(User.role == UserRole.user)

    if search and search.strip():
        term = f"%{search.strip()}%"
        query = query.filter(or_(User.full_name.ilike(term), User.email.ilike(term)))

    normalized_status = (link_status or "all").strip().lower()
    if normalized_status == "with_link":
        query = query.filter(
            User.google_drive_link.isnot(None),
            func.length(func.trim(User.google_drive_link)) > 0,
        )
    elif normalized_status in ("without_link", "missing"):
        query = query.filter(
            or_(
                User.google_drive_link.is_(None),
                func.length(func.trim(User.google_drive_link)) == 0,
            )
        )

    total = query.count()
    members = (
        query.options(*user_programme_load_options())
        .order_by(User.full_name.asc(), User.id.asc())
        .offset(skip)
        .limit(limit)
        .all()
    )
    return members, total


def update_member_google_drive_link(db: Session, user_id: int, google_drive_link: str | None) -> User:
    from app.google_drive_link import validate_google_drive_folder_link

    user = _load_member_user(db, user_id)
    link = validate_google_drive_folder_link(google_drive_link)
    user.google_drive_link = link
    db.commit()
    db.refresh(user)
    return (
        db.query(User)
        .options(*user_programme_load_options())
        .filter(User.id == user.id)
        .first()
        or user
    )


def serialize_admin_member(user: User, db: Session) -> UserPublic:
    programme = _build_programme_public(db, user)
    points = calculate_member_points(db, user.id)
    # Align admin "Progress" with the user dashboard calculation.
    # Dashboard uses completed_sessions / total_sessions for the current level.
    try:
        dash = get_user_dashboard(db, user)
        progress_percent = int(dash.get("last_progress_percent") or 0)
    except Exception:
        progress_percent = compute_member_progress_percent(programme)
    role = user.role.value if isinstance(user.role, UserRole) else str(user.role)

    profile_image = getattr(user, "profile_image", None)
    profile_image_url = media_public_url(profile_image)

    return UserPublic(
        id=user.id,
        email=user.email,
        full_name=user.full_name,
        phone_number=user.phone_number,
        address=user.address,
        postal_code=user.postal_code,
        profile_image=profile_image,
        profile_image_url=profile_image_url,
        has_profile_image=bool(profile_image),
        role=RoleEnum(role),
        is_active=user.is_active,
        email_verified=bool(user.email_verified),
        created_at=user.created_at,
        programme=programme,
        total_points=points["available_points"],
        assigned_points_total=points["assigned_points_total"],
        available_points=points["available_points"],
        progress_percent=progress_percent,
        google_drive_link=user.google_drive_link,
    )


def serialize_admin_member_light(
    user: User,
    *,
    db: Session | None = None,
    programme_levels: list[ProgrammeLevel] | None = None,
    points: dict | None = None,
) -> UserPublic:
    """Fast member serializer for list endpoints (no dashboard + no per-user count queries)."""
    programme = _build_programme_public_light(user, db=db, programme_levels=programme_levels)
    progress_percent = compute_member_progress_percent(programme)
    role = user.role.value if isinstance(user.role, UserRole) else str(user.role)

    profile_image = getattr(user, "profile_image", None)
    profile_image_url = media_public_url(profile_image)

    pts = points or {"assigned_points_total": 0, "available_points": 0}
    available_points = int(pts.get("available_points") or 0)
    assigned_points_total = int(pts.get("assigned_points_total") or 0)

    return UserPublic(
        id=user.id,
        email=user.email,
        full_name=user.full_name,
        phone_number=user.phone_number,
        address=user.address,
        postal_code=user.postal_code,
        profile_image=profile_image,
        profile_image_url=profile_image_url,
        has_profile_image=bool(profile_image),
        role=RoleEnum(role),
        is_active=user.is_active,
        email_verified=bool(user.email_verified),
        created_at=user.created_at,
        programme=programme,
        total_points=available_points,
        assigned_points_total=assigned_points_total,
        available_points=available_points,
        progress_percent=progress_percent,
        google_drive_link=user.google_drive_link,
    )


def serialize_admin_member_minimal(user: User) -> UserPublic:
    """Super fast serializer for members list (no programme/progress/points)."""
    role = user.role.value if isinstance(user.role, UserRole) else str(user.role)
    profile_image = getattr(user, "profile_image", None)
    profile_image_url = media_public_url(profile_image)
    return UserPublic(
        id=user.id,
        email=user.email,
        full_name=user.full_name,
        phone_number=user.phone_number,
        address=user.address,
        postal_code=user.postal_code,
        profile_image=profile_image,
        profile_image_url=profile_image_url,
        has_profile_image=bool(profile_image),
        role=RoleEnum(role),
        is_active=user.is_active,
        email_verified=bool(user.email_verified),
        created_at=user.created_at,
        programme=None,
        total_points=None,
        assigned_points_total=None,
        available_points=None,
        progress_percent=None,
        google_drive_link=user.google_drive_link,
    )


def _parse_level_slug(value: str) -> str:
    try:
        return normalize_programme_level_slug(value)
    except ValueError as exc:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="level slug is required",
        ) from exc


def _parse_module_kind(value: str) -> str:
    try:
        return normalize_programme_module_kind(value)
    except ValueError as exc:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="module kind is required",
        ) from exc


def _sync_register_step_from_programme(
    db: Session,
    user: User,
    *,
    level_slug: str,
    module_kind: str,
    week_index: int,
    session: ProgrammeSession,
) -> None:
    row = db.query(UserRegisterStep).filter(UserRegisterStep.user_id == user.id).first()
    if row is None:
        return

    level = (
        db.query(ProgrammeLevel)
        .filter(ProgrammeLevel.slug == level_slug)
        .first()
    )
    module = (
        db.query(ProgrammeModule)
        .filter(ProgrammeModule.level_id == level.id, ProgrammeModule.kind == module_kind)
        .first()
        if level
        else None
    )
    week = (
        db.query(ProgrammeWeek)
        .filter(ProgrammeWeek.module_id == module.id, ProgrammeWeek.week_index == week_index)
        .first()
        if module
        else None
    )

    if level:
        row.level_id = level.id
        row.level_name = level.display_name
    if module:
        row.module_id = module.id
        row.module_name = module.display_name
    if week:
        row.week_id = week.id
        row.week_name = week.display_label
    row.session_id = session.id
    row.session_name = session.title
    row.current_step = "completed"
    row.is_completed = True


def assign_user_programme(db: Session, body: UpdateUserProgrammeRequest) -> User:
    user = _load_member_user(db, body.user_id)
    level_slug = _parse_level_slug(body.programme_level_slug)
    module_kind = _parse_module_kind(body.programme_module_kind)

    session = resolve_session(
        db,
        level_slug=level_slug,
        module_kind=module_kind,
        week_index=body.programme_week_index,
        session_number=body.programme_session_number,
    )
    if session is None:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="No matching session for the given programme path",
        )

    first_assessment = (
        db.query(ProgrammeAssessment)
        .filter(ProgrammeAssessment.session_id == session.id)
        .order_by(ProgrammeAssessment.sort_order.asc())
        .first()
    )

    try:
        session, assessment = validate_user_programme_selection(
            db,
            level_slug=level_slug,
            module_kind=module_kind,
            week_index=body.programme_week_index,
            session_number=body.programme_session_number,
            assessment_id=first_assessment.id if first_assessment else None,
        )
    except ValueError as exc:
        raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail=str(exc)) from exc

    user.programme_session_id = session.id
    user.programme_assessment_id = assessment.id if assessment else None
    _sync_register_step_from_programme(
        db,
        user,
        level_slug=level_slug,
        module_kind=module_kind,
        week_index=body.programme_week_index,
        session=session,
    )
    db.commit()
    db.refresh(user)
    return (
        db.query(User)
        .options(*user_programme_load_options())
        .filter(User.id == user.id)
        .first()
        or user
    )


def update_member_personal_info(db: Session, body: AdminMemberPersonalInfoUpdateRequest) -> User:
    user = _load_member_user(db, body.user_id)

    if body.full_name is not None:
        cleaned = body.full_name.strip()
        if not cleaned:
            raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="full_name cannot be empty")
        user.full_name = cleaned

    if body.google_drive_link is not None:
        from app.google_drive_link import validate_google_drive_folder_link

        user.google_drive_link = validate_google_drive_folder_link(body.google_drive_link)

    if body.phone_number is not None:
        user.phone_number = body.phone_number.strip() or None

    programme_fields = (
        body.programme_level_slug,
        body.programme_module_kind,
        body.programme_week_index,
        body.programme_session_number,
    )
    if any(value is not None for value in programme_fields):
        if body.programme_level_slug is None or body.programme_module_kind is None:
            raise HTTPException(
                status_code=status.HTTP_400_BAD_REQUEST,
                detail="programme_level_slug and programme_module_kind are required when updating programme",
            )

        programme = _build_programme_public(db, user)
        week_index = (
            body.programme_week_index
            if body.programme_week_index is not None
            else (programme.week_index if programme and programme.week_index is not None else 0)
        )
        session_number = (
            body.programme_session_number
            if body.programme_session_number is not None
            else (programme.session_number if programme and programme.session_number is not None else 1)
        )

        assign_user_programme(
            db,
            UpdateUserProgrammeRequest(
                user_id=body.user_id,
                programme_level_slug=body.programme_level_slug,
                programme_module_kind=body.programme_module_kind,
                programme_week_index=week_index,
                programme_session_number=session_number,
            ),
        )
        db.refresh(user)
        user = _load_member_user(db, body.user_id)
    else:
        db.commit()
        db.refresh(user)

    return user


def _resolve_member_module(
    db: Session, user: User
) -> tuple[ProgrammeLevel, ProgrammeModule, UserProgrammeProgressPublic | None]:
    from app.services.user_track_service import resolve_user_programme_summary

    raw = resolve_user_programme_summary(db, user)
    programme = UserProgrammeProgressPublic.model_validate(raw) if raw else None
    level_slug = (
        _parse_level_slug(programme.level_slug)
        if programme and programme.level_slug
        else "beginner"
    )
    module_kind = (
        _parse_module_kind(programme.module_kind)
        if programme and programme.module_kind
        else "foundations"
    )

    level = (
        db.query(ProgrammeLevel)
        .options(
            selectinload(ProgrammeLevel.modules)
            .selectinload(ProgrammeModule.weeks)
            .selectinload(ProgrammeWeek.sessions)
            .selectinload(ProgrammeSession.assessments)
        )
        .filter(ProgrammeLevel.slug == level_slug)
        .first()
    )
    if level is None:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Programme level not found")

    module = next((m for m in level.modules if m.kind == module_kind), None)
    if module is None and level.modules:
        module = sorted(level.modules, key=lambda item: item.sort_order)[0]
    if module is None:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Programme module not found")

    return level, module, programme


def _session_is_completed(
    session: ProgrammeSession,
    week: ProgrammeWeek,
    *,
    completed_session_ids: set[int],
    completed_assessment_ids: set[int],
    progress_by_session_id: dict[int, MemberSessionProgress],
    progress_by_week_session: dict[tuple[int, int], MemberSessionProgress],
) -> bool:
    progress = progress_by_session_id.get(session.id) or progress_by_week_session.get(
        (week.week_index, session.session_number)
    )
    if progress is not None and progress.is_completed:
        return True

    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


def _progress_lookup(
    rows: list[MemberSessionProgress],
) -> tuple[dict[int, MemberSessionProgress], dict[tuple[int, int], MemberSessionProgress]]:
    by_session_id: dict[int, MemberSessionProgress] = {}
    by_week_session: dict[tuple[int, int], MemberSessionProgress] = {}
    for row in rows:
        if row.session_id is not None:
            by_session_id[row.session_id] = row
        if row.week_index is not None and row.session_number is not None:
            by_week_session[(row.week_index, row.session_number)] = row
    return by_session_id, by_week_session


def get_member_weeks_detail(db: Session, user_id: int) -> MemberWeeksDetailResponse:
    user = _load_member_user(db, user_id)
    level, module, programme = _resolve_member_module(db, user)

    progress_rows = list_member_session_progress(db, user_id)
    progress_by_session_id, progress_by_week_session = _progress_lookup(progress_rows)

    completed_session_ids = {
        row.session_id
        for row in db.query(UserCompletedSession).filter(UserCompletedSession.user_id == user_id).all()
    }
    completed_assessment_ids = {
        row.assessment_id
        for row in db.query(UserCompletedAssessment).filter(UserCompletedAssessment.user_id == user_id).all()
    }

    current_week_index = programme.week_index if programme else None
    weeks_out: list[MemberWeekDetailOut] = []

    for idx, week in enumerate(sorted(module.weeks, key=lambda item: item.week_index)):
        sessions_out: list[MemberWeekSessionDetailOut] = []
        completed_count = 0

        for session in sorted(week.sessions, key=lambda item: item.session_number):
            progress = progress_by_session_id.get(session.id) or progress_by_week_session.get(
                (week.week_index, session.session_number)
            )
            is_completed = _session_is_completed(
                session,
                week,
                completed_session_ids=completed_session_ids,
                completed_assessment_ids=completed_assessment_ids,
                progress_by_session_id=progress_by_session_id,
                progress_by_week_session=progress_by_week_session,
            )
            if is_completed:
                completed_count += 1

            assessments = sorted(session.assessments, key=lambda item: item.sort_order)
            first_assessment = assessments[0] if assessments else None
            has_feedback = bool(progress and (progress.feedback or "").strip())

            sessions_out.append(
                MemberWeekSessionDetailOut(
                    session_id=session.id,
                    session_number=session.session_number,
                    title=session.title,
                    status="video_submitted" if is_completed else "pending_submission",
                    video_submitted=is_completed,
                    is_completed=is_completed,
                    has_feedback=has_feedback,
                    feedback=progress.feedback if progress else None,
                    assigned_points=int(progress.assigned_points) if progress else 0,
                    progress_id=progress.id if progress else None,
                    image_url=media_public_url(progress.image_path) if progress and progress.image_path else None,
                    image_urls=(
                        [media_public_url(img.image_path) for img in (progress.images or []) if img.image_path]
                        if progress
                        else []
                    ),
                    first_assessment_id=first_assessment.id if first_assessment else None,
                    first_assessment_title=first_assessment.title if first_assessment else None,
                )
            )

        total_sessions = len(sessions_out)
        week_completed = total_sessions > 0 and completed_count == total_sessions
        bonus_row = _week_completion_bonus_row(db, user.id, week)
        week_point_awarded = bonus_row is not None and int(bonus_row.assigned_points or 0) > 0
        weeks_out.append(
            MemberWeekDetailOut(
                week_id=week.id,
                week_index=week.week_index,
                week_number=week.week_index + 1,
                display_label=week.display_label,
                total_sessions=total_sessions,
                completed_sessions=completed_count,
                is_completed=week_completed or week_point_awarded,
                week_point_awarded=week_point_awarded,
                expanded=week.week_index == current_week_index or idx == 0,
                sessions=sessions_out,
            )
        )

    points = calculate_member_points(db, user.id)
    return MemberWeeksDetailResponse(
        user_id=user.id,
        level_slug=programme_level_slug_str(level.slug),
        level_display_name=level.display_name,
        module_kind=programme_module_kind_str(module.kind),
        module_display_name=module.display_name,
        current_week_index=current_week_index,
        assigned_points_total=points["assigned_points_total"],
        available_points=points["available_points"],
        weeks=weeks_out,
    )


def _week_completion_bonus_row(
    db: Session, user_id: int, week: ProgrammeWeek
) -> MemberSessionProgress | None:
    return (
        db.query(MemberSessionProgress)
        .filter(
            MemberSessionProgress.user_id == user_id,
            MemberSessionProgress.week_index == week.week_index,
            MemberSessionProgress.session_number == 0,
        )
        .first()
    )


def _award_week_completion_point(
    db: Session,
    admin: User,
    user: User,
    week: ProgrammeWeek,
) -> None:
    row = _week_completion_bonus_row(db, user.id, week)
    if row is None:
        row = MemberSessionProgress(user_id=user.id, session_number=0)
        db.add(row)
    row.week_id = week.id
    row.week_name = week.display_label
    row.week_index = week.week_index
    row.session_id = None
    row.session_name = f"{week.display_label} completion"
    row.assigned_points = 1
    row.feedback = "Week marked as completed"
    row.is_completed = True
    row.assigned_by_id = admin.id
    db.flush()
    refresh_member_points(db, user.id)


def _remove_week_completion_point(db: Session, user_id: int, week: ProgrammeWeek) -> None:
    row = _week_completion_bonus_row(db, user_id, week)
    if row is not None:
        db.delete(row)
        db.flush()
        refresh_member_points(db, user_id)


def mark_member_week_complete(
    db: Session, body: MarkMemberWeekCompleteRequest, admin: User
) -> MemberWeeksDetailResponse:
    user = _load_member_user(db, body.user_id)
    _level, module, _programme = _resolve_member_module(db, user)

    week = next((wk for wk in module.weeks if wk.week_index == body.week_index), None)
    if week is None:
        raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="Week not found")

    if body.is_completed:
        _complete_entire_week(db, user.id, week)
        for session in sorted(week.sessions, key=lambda item: item.session_number):
            progress = (
                db.query(MemberSessionProgress)
                .filter(
                    MemberSessionProgress.user_id == user.id,
                    MemberSessionProgress.session_id == session.id,
                )
                .first()
            )
            if progress is None:
                progress = MemberSessionProgress(user_id=user.id, session_id=session.id)
                db.add(progress)
            progress.week_id = week.id
            progress.week_name = week.display_label
            progress.week_index = week.week_index
            progress.session_id = session.id
            progress.session_name = session.title
            progress.session_number = session.session_number
            progress.is_completed = True
        _award_week_completion_point(db, admin, user, week)
    else:
        session_ids = [session.id for session in week.sessions]
        if session_ids:
            db.query(UserCompletedSession).filter(
                UserCompletedSession.user_id == user.id,
                UserCompletedSession.session_id.in_(session_ids),
            ).delete(synchronize_session=False)
            assessment_ids = [
                row[0]
                for row in db.query(ProgrammeAssessment.id)
                .filter(ProgrammeAssessment.session_id.in_(session_ids))
                .all()
            ]
            if assessment_ids:
                db.query(UserCompletedAssessment).filter(
                    UserCompletedAssessment.user_id == user.id,
                    UserCompletedAssessment.assessment_id.in_(assessment_ids),
                ).delete(synchronize_session=False)
        db.query(MemberSessionProgress).filter(
            MemberSessionProgress.user_id == user.id,
            MemberSessionProgress.week_index == week.week_index,
        ).update({MemberSessionProgress.is_completed: False}, synchronize_session=False)
        _remove_week_completion_point(db, user.id, week)

    from app.services.user_track_service import sync_user_programme_pointer

    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
    return get_member_weeks_detail(db, user.id)
