from __future__ import annotations

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

from app.models import User, UserRegisterStep
from app.programme.models_programme import ProgrammeSession
from app.programme.services.programme_service import seed_prior_programme_completions
from sqlalchemy.orm import selectinload
from app.schemas import (
    RegisterStepItemOut,
    RegisterStepStatus,
    UserRegisterStepAcceptRequest,
    UserRegisterStepSaveRequest,
    UserRegisterStepStateResponse,
)

REGISTER_STEP_ORDER: tuple[str, ...] = ("level", "module", "week", "session")
REGISTER_STEP_LABELS: dict[str, str] = {
    "level": "Level",
    "module": "Module",
    "week": "Week",
    "session": "Session",
    "completed": "Completed",
}

_STEP_FIELDS: dict[str, tuple[str, str]] = {
    "level": ("level_id", "level_name"),
    "module": ("module_id", "module_name"),
    "week": ("week_id", "week_name"),
    "session": ("session_id", "session_name"),
}


def get_user_register_step(db: Session, user: User) -> UserRegisterStep | None:
    return db.query(UserRegisterStep).filter(UserRegisterStep.user_id == user.id).first()


def delete_user_register_step(db: Session, user: User) -> None:
    row = get_user_register_step(db, user)
    if row is not None:
        db.delete(row)
        user.programme_session_id = None
        user.programme_assessment_id = None
        db.commit()


def _step_index(step: str) -> int:
    normalized = step.strip().lower()
    if normalized == "completed":
        return len(REGISTER_STEP_ORDER)
    try:
        return REGISTER_STEP_ORDER.index(normalized)
    except ValueError as exc:
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="current_step must be one of: level, module, week, session, completed",
        ) from exc


def _selected_for_step(row: UserRegisterStep | None, step: str) -> tuple[int | None, str | None]:
    if row is None:
        return None, None
    mapping = {
        "level": (row.level_id, row.level_name),
        "module": (row.module_id, row.module_name),
        "week": (row.week_id, row.week_name),
        "session": (row.session_id, row.session_name),
    }
    return mapping.get(step, (None, None))


def _step_has_selection(row: UserRegisterStep | None, step: str) -> bool:
    selected_id, selected_name = _selected_for_step(row, step)
    return selected_id is not None or bool(selected_name and selected_name.strip())


def _validate_step_selection(row: UserRegisterStep, step: str) -> None:
    if not _step_has_selection(row, step):
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail=f"Select a {step} before continuing.",
        )


def build_register_step_state(row: UserRegisterStep | None) -> UserRegisterStepStateResponse:
    if row is None:
        current_step = "level"
        current_idx = 0
        is_completed = False
    elif row.is_completed or row.current_step == "completed":
        current_step = "completed"
        current_idx = len(REGISTER_STEP_ORDER)
        is_completed = True
    else:
        current_step = row.current_step.strip().lower()
        current_idx = _step_index(current_step)
        is_completed = False

    steps: list[RegisterStepItemOut] = []
    for idx, step in enumerate(REGISTER_STEP_ORDER):
        if is_completed or idx < current_idx:
            step_status = RegisterStepStatus.completed
        elif idx == current_idx:
            step_status = RegisterStepStatus.open
        else:
            step_status = RegisterStepStatus.locked

        selected_id, selected_name = _selected_for_step(row, step)
        steps.append(
            RegisterStepItemOut(
                step=step,
                label=REGISTER_STEP_LABELS[step],
                status=step_status,
                selected_id=selected_id,
                selected_name=selected_name,
            )
        )

    can_proceed = (
        not is_completed
        and row is not None
        and _step_has_selection(row, current_step if current_step in REGISTER_STEP_ORDER else "session")
    )

    return UserRegisterStepStateResponse(
        current_step=current_step,
        is_completed=is_completed,
        can_proceed=can_proceed,
        steps=steps,
        selection=row,
    )


def _step_data_in_body(
    body: UserRegisterStepSaveRequest | UserRegisterStepAcceptRequest,
    step: str,
) -> bool:
    id_field, name_field = _STEP_FIELDS[step]
    sent_fields = body.model_fields_set
    if id_field in sent_fields and getattr(body, id_field, None) is not None:
        return True
    if name_field in sent_fields:
        value = getattr(body, name_field, None)
        return bool(value and str(value).strip())
    return False


def _apply_step_fields(
    row: UserRegisterStep,
    body: UserRegisterStepSaveRequest | UserRegisterStepAcceptRequest,
    step: str,
    *,
    only_sent_fields: bool,
) -> None:
    """Update id/name fields for one step only; never touches current_step or is_completed."""
    id_field, name_field = _STEP_FIELDS[step]
    sent_fields = body.model_fields_set

    if not only_sent_fields or id_field in sent_fields:
        value = getattr(body, id_field, None)
        if value is not None or (only_sent_fields and id_field in sent_fields):
            setattr(row, id_field, value)

    if not only_sent_fields or name_field in sent_fields:
        value = getattr(body, name_field, None)
        if isinstance(value, str):
            value = value.strip() or None
        if value is not None or (only_sent_fields and name_field in sent_fields):
            setattr(row, name_field, value)


def reseed_track_from_register_step(db: Session, user: User) -> None:
    """Mark everything before the user's registration anchor as completed."""
    register_step = get_user_register_step(db, user)
    if register_step is None or not register_step.is_completed or register_step.session_id is None:
        return

    session = (
        db.query(ProgrammeSession)
        .options(selectinload(ProgrammeSession.assessments))
        .filter(ProgrammeSession.id == register_step.session_id)
        .first()
    )
    if session is None:
        return

    assessments = sorted(session.assessments, key=lambda item: item.sort_order)
    first_assessment = assessments[0] if assessments else None
    seed_prior_programme_completions(db, user.id, session, first_assessment)
    user.programme_session_id = session.id
    user.programme_assessment_id = first_assessment.id if first_assessment else None
    db.flush()


def _sync_user_programme_pointer(db: Session, user: User, row: UserRegisterStep) -> None:
    if row.session_id is None:
        return
    reseed_track_from_register_step(db, user)


def _advance_register_step(db: Session, user: User, row: UserRegisterStep) -> None:
    _validate_step_selection(row, row.current_step)
    current_idx = _step_index(row.current_step)

    if current_idx >= len(REGISTER_STEP_ORDER) - 1:
        row.current_step = "completed"
        row.is_completed = True
        _sync_user_programme_pointer(db, user, row)
        return

    row.current_step = REGISTER_STEP_ORDER[current_idx + 1]
    row.is_completed = False


def _get_or_create_row(db: Session, user: User) -> UserRegisterStep:
    row = get_user_register_step(db, user)
    if row is None:
        row = UserRegisterStep(user_id=user.id, current_step="level", is_completed=False)
        db.add(row)
        db.flush()
    return row


def _finish_register_step_from_body(
    db: Session, user: User, row: UserRegisterStep, body: UserRegisterStepAcceptRequest
) -> UserRegisterStepStateResponse:
    """Session submit with level/module/week/session ids — complete onboarding in one call."""
    for step in REGISTER_STEP_ORDER:
        if _step_data_in_body(body, step):
            _apply_step_fields(row, body, step, only_sent_fields=True)

    for step in REGISTER_STEP_ORDER:
        _validate_step_selection(row, step)

    row.current_step = "completed"
    row.is_completed = True
    _sync_user_programme_pointer(db, user, row)
    db.commit()
    db.refresh(row)
    return build_register_step_state(row)


def save_user_register_step(
    db: Session, user: User, body: UserRegisterStepSaveRequest
) -> UserRegisterStepStateResponse:
    row = _get_or_create_row(db, user)

    if row.is_completed or row.current_step == "completed":
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="Registration is already completed.",
        )

    for step in REGISTER_STEP_ORDER:
        if _step_data_in_body(body, step):
            _apply_step_fields(row, body, step, only_sent_fields=True)

    if _step_data_in_body(body, "session"):
        return _finish_register_step_from_body(
            db, user, row, UserRegisterStepAcceptRequest.model_validate(body.model_dump())
        )

    db.commit()
    db.refresh(row)
    return build_register_step_state(row)


def accept_user_register_step(
    db: Session, user: User, body: UserRegisterStepAcceptRequest
) -> UserRegisterStepStateResponse:
    row = _get_or_create_row(db, user)

    if row.is_completed or row.current_step == "completed":
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="Registration steps are already completed.",
        )

    if _step_data_in_body(body, "session"):
        return _finish_register_step_from_body(db, user, row, body)

    step = row.current_step.strip().lower()
    if step not in _STEP_FIELDS:
        raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="Invalid current step.")

    _apply_step_fields(row, body, step, only_sent_fields=False)
    _validate_step_selection(row, step)
    _advance_register_step(db, user, row)

    db.commit()
    db.refresh(row)
    return build_register_step_state(row)


def complete_user_register_step(db: Session, user: User) -> UserRegisterStepStateResponse:
    row = _get_or_create_row(db, user)

    if row.is_completed or row.current_step == "completed":
        raise HTTPException(
            status_code=status.HTTP_400_BAD_REQUEST,
            detail="Registration steps are already completed.",
        )

    _advance_register_step(db, user, row)
    db.commit()
    db.refresh(row)
    return build_register_step_state(row)
