import uuid

from fastapi import APIRouter, Depends, HTTPException, Query
from sqlalchemy.orm import Session

from app.api.deps import get_current_principal, get_project_or_404, require_role
from app.core.config import get_settings
from app.db.session import get_db
from app.models.dataset import Dataset
from app.models.org import User, UserRole
from app.models.prompt import Prompt
from app.models.stubs import FineTuningJob, FineTuningJobStatus
from app.schemas.training import FineTuningJobCreate, FineTuningJobOut
from app.services.audit import record_audit
from app.services.finetune import MIN_APPROVED_EXAMPLES, detect_backend
from app.workers.training_tasks import run_finetune_job_task

router = APIRouter(prefix="/training", tags=["training"])
settings = get_settings()


def _active_base_models() -> tuple[str, list[str]]:
    backend = detect_backend()
    if backend == "mlx":
        return backend, settings.finetune_base_models_mlx
    if backend == "cuda":
        return backend, settings.finetune_base_models_cuda
    return backend, []


@router.get("/base-models")
def list_base_models(current_user: User = Depends(get_current_principal)):
    backend, base_models = _active_base_models()
    return {"backend": backend, "base_models": base_models, "min_approved_examples": MIN_APPROVED_EXAMPLES}


@router.get("/jobs", response_model=list[FineTuningJobOut])
def list_training_jobs(
    project_id: uuid.UUID = Query(...),
    db: Session = Depends(get_db),
    current_user: User = Depends(get_current_principal),
):
    get_project_or_404(project_id, db, current_user)
    return (
        db.query(FineTuningJob)
        .filter(FineTuningJob.project_id == project_id)
        .order_by(FineTuningJob.created_at.desc())
        .all()
    )


@router.get("/jobs/{job_id}", response_model=FineTuningJobOut)
def get_training_job(
    job_id: uuid.UUID, db: Session = Depends(get_db), current_user: User = Depends(get_current_principal)
):
    job = db.get(FineTuningJob, job_id)
    if job is None:
        raise HTTPException(status_code=404, detail="Training job not found")
    get_project_or_404(job.project_id, db, current_user)
    return job


@router.post("/jobs", response_model=FineTuningJobOut, status_code=202)
def create_training_job(
    payload: FineTuningJobCreate,
    db: Session = Depends(get_db),
    current_user: User = Depends(require_role(UserRole.editor)),
):
    dataset = db.get(Dataset, payload.dataset_id)
    if dataset is None:
        raise HTTPException(status_code=404, detail="Dataset not found")
    get_project_or_404(dataset.project_id, db, current_user)

    _backend, base_models = _active_base_models()
    if payload.base_model not in base_models:
        raise HTTPException(status_code=400, detail="Unsupported base model")

    system_prompt = None
    if payload.system_prompt_id is not None:
        prompt = db.get(Prompt, payload.system_prompt_id)
        if prompt is None or prompt.org_id != current_user.org_id:
            raise HTTPException(status_code=404, detail="System prompt not found")
        system_prompt = prompt.content

    job = FineTuningJob(
        project_id=dataset.project_id,
        dataset_id=dataset.id,
        base_model=payload.base_model,
        config={
            "name": payload.name,
            "ollama_tag": payload.ollama_tag or None,
            "system_prompt": system_prompt,
            "iters": payload.iters,
            "learning_rate": payload.learning_rate,
            "batch_size": payload.batch_size,
            "num_layers": payload.num_layers,
            "lora_r": payload.lora_r,
            "lora_alpha": payload.lora_alpha,
            "lora_dropout": payload.lora_dropout,
        },
        status=FineTuningJobStatus.queued,
        progress={"stage": "queued"},
        created_by=current_user.id,
    )
    db.add(job)
    db.commit()
    db.refresh(job)

    run_finetune_job_task.delay(str(job.id))
    record_audit(
        db, current_user, "training.start", "fine_tuning_job", job.id, {"base_model": payload.base_model}
    )
    return job
