import uuid

from app.core.celery_app import celery_app
from app.db.session import SessionLocal
from app.models.stubs import FineTuningJob, FineTuningJobStatus
from app.services import finetune
from app.services.model_registry import register_model_version
from app.ws.events import publish_job_event


@celery_app.task(name="training.run_finetune_job")
def run_finetune_job_task(job_id: str) -> None:
    db = SessionLocal()
    try:
        job = db.get(FineTuningJob, uuid.UUID(job_id))
        if job is None:
            return

        def set_progress(stage: str, **extra) -> None:
            # Merge rather than replace so fields like train_loss/val_loss survive
            # later stage transitions (fusing, exporting_gguf, ...) instead of being
            # wiped the moment the stage moves past "training".
            job.progress = {**job.progress, "stage": stage, **extra}
            db.commit()
            publish_job_event(job_id, {"status": stage, **extra})

        job.status = FineTuningJobStatus.running
        db.commit()

        backend = finetune.get_backend()
        config = job.config
        work_dir = finetune.job_dir(job.id)
        data_dir = work_dir / "data"
        data_dir.mkdir(exist_ok=True)
        adapter_path = work_dir / "adapters"
        fused_dir = work_dir / "fused"
        gguf_path = work_dir / "model.gguf"

        set_progress("preparing")
        train_count, valid_count = finetune.export_training_data(job.dataset_id, data_dir, db)
        # Every backend's trainer requires each split to have at least `batch_size`
        # examples — clamp down for small datasets rather than crashing mid-run.
        config["batch_size"] = max(1, min(config["batch_size"], train_count, valid_count))

        set_progress("downloading_base_model", train_examples=train_count, valid_examples=valid_count)
        finetune.ensure_base_model_downloaded(job.base_model)

        set_progress("training", iter=0, iters_total=config["iters"], train_examples=train_count, valid_examples=valid_count)

        def on_line(line: str) -> None:
            parsed = backend.parse_progress_line(line)
            if parsed is None:
                return
            if parsed["type"] == "train":
                set_progress("training", iter=parsed["iter"], iters_total=config["iters"], train_loss=parsed["loss"])
            else:
                set_progress("training", val_loss=parsed["loss"])

        backend.run_lora_training(job.base_model, data_dir, adapter_path, config, on_line)

        set_progress("fusing")
        backend.run_fuse(job.base_model, adapter_path, fused_dir)

        set_progress("exporting_gguf")
        finetune.run_gguf_convert(fused_dir, gguf_path)

        tag = config.get("ollama_tag") or f"jozie-ft-{str(job.id)[:8]}"
        set_progress("importing_ollama", ollama_tag=tag)
        finetune.import_into_ollama(tag, gguf_path, config.get("system_prompt"))

        set_progress("registering", ollama_tag=tag)
        version = register_model_version(
            db,
            job.project_id,
            config.get("name") or tag,
            job.base_model,
            tag,
            job.created_by,
            notes=f"Fine-tuned via LoRA ({config['iters']} iters) from dataset {job.dataset_id}",
            metrics={
                "iters": config["iters"],
                "train_examples": train_count,
                "valid_examples": valid_count,
                "train_loss": job.progress.get("train_loss"),
                "val_loss": job.progress.get("val_loss"),
            },
            adapter_path=str(adapter_path),
            task_type="fine-tuning",
        )

        job.model_version_id = version.id
        job.status = FineTuningJobStatus.completed
        job.progress = {**job.progress, "stage": "completed", "ollama_tag": tag}
        db.commit()
        publish_job_event(job_id, {"status": "completed", "ollama_tag": tag})
    except Exception as exc:  # noqa: BLE001
        db.rollback()
        job = db.get(FineTuningJob, uuid.UUID(job_id))
        if job is not None:
            job.status = FineTuningJobStatus.failed
            job.progress = {"stage": "failed", "error": str(exc)}
            db.commit()
        publish_job_event(job_id, {"status": "failed", "error": str(exc)})
        raise
    finally:
        db.close()
