import io
import json
import uuid

import pandas as pd
from fastapi import APIRouter, Depends, HTTPException, Query
from fastapi.responses import StreamingResponse
from sqlalchemy.orm import Session

from app.api.deps import get_current_principal, get_project_or_404, require_role
from app.db.session import get_db
from app.models.dataset import Dataset, DatasetExample, ExampleStatus
from app.models.knowledge import ExtractedText
from app.models.org import User, UserRole
from app.schemas.dataset import (
    DatasetCreate,
    DatasetExampleOut,
    DatasetExampleUpdate,
    DatasetOut,
    ExpandExampleRequest,
    GenerateExamplesRequest,
)
from app.services.audit import record_audit
from app.services.ollama_client import VARIANT_META
from app.workers.dataset_tasks import expand_example_task, generate_examples_task, regenerate_example_task

router = APIRouter(tags=["datasets"])


def _get_dataset_or_404(dataset_id: uuid.UUID, db: Session, current_user: User) -> Dataset:
    dataset = db.get(Dataset, 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)
    return dataset


@router.get("/projects/{project_id}/datasets", response_model=list[DatasetOut])
def list_datasets(
    project_id: uuid.UUID, db: Session = Depends(get_db), current_user: User = Depends(get_current_principal)
):
    get_project_or_404(project_id, db, current_user)
    return db.query(Dataset).filter(Dataset.project_id == project_id).order_by(Dataset.created_at.desc()).all()


@router.post("/projects/{project_id}/datasets", response_model=DatasetOut, status_code=201)
def create_dataset(
    project_id: uuid.UUID,
    payload: DatasetCreate,
    db: Session = Depends(get_db),
    current_user: User = Depends(require_role(UserRole.editor)),
):
    get_project_or_404(project_id, db, current_user)
    dataset = Dataset(project_id=project_id, name=payload.name, description=payload.description)
    db.add(dataset)
    db.commit()
    db.refresh(dataset)
    record_audit(db, current_user, "dataset.create", "dataset", dataset.id, {"name": dataset.name})
    return dataset


@router.get("/datasets/{dataset_id}", response_model=DatasetOut)
def get_dataset(
    dataset_id: uuid.UUID, db: Session = Depends(get_db), current_user: User = Depends(get_current_principal)
):
    return _get_dataset_or_404(dataset_id, db, current_user)


@router.get("/datasets/{dataset_id}/examples", response_model=list[DatasetExampleOut])
def list_examples(
    dataset_id: uuid.UUID,
    status: ExampleStatus | None = Query(default=None),
    db: Session = Depends(get_db),
    current_user: User = Depends(get_current_principal),
):
    _get_dataset_or_404(dataset_id, db, current_user)
    q = db.query(DatasetExample).filter(DatasetExample.dataset_id == dataset_id)
    if status is not None:
        q = q.filter(DatasetExample.status == status)
    return q.order_by(DatasetExample.created_at.desc()).all()


@router.post("/datasets/{dataset_id}/generate", status_code=202)
def generate_examples(
    dataset_id: uuid.UUID,
    payload: GenerateExamplesRequest,
    db: Session = Depends(get_db),
    current_user: User = Depends(require_role(UserRole.editor)),
):
    dataset = _get_dataset_or_404(dataset_id, db, current_user)
    extracted = db.query(ExtractedText).filter_by(knowledge_source_id=payload.knowledge_source_id).one_or_none()
    if extracted is None or not extracted.cleaned_text or not extracted.is_confirmed:
        raise HTTPException(status_code=400, detail="Document must have confirmed cleaned text before generating examples")

    generate_examples_task.delay(
        str(dataset.id),
        str(extracted.id),
        payload.count,
        [d.value for d in payload.difficulty_mix],
        payload.model,
    )
    record_audit(db, current_user, "dataset.generate_examples", "dataset", dataset.id, {"count": payload.count})
    return {"status": "queued"}


@router.get("/synthetic-expansion/variant-types")
def list_variant_types(current_user: User = Depends(get_current_principal)):
    return [{"key": key, "guidance": meta["guidance"]} for key, meta in VARIANT_META.items()]


@router.post("/examples/{example_id}/expand", status_code=202)
def expand_example(
    example_id: uuid.UUID,
    payload: ExpandExampleRequest,
    db: Session = Depends(get_db),
    current_user: User = Depends(require_role(UserRole.editor)),
):
    example = db.get(DatasetExample, example_id)
    if example is None:
        raise HTTPException(status_code=404, detail="Example not found")
    _get_dataset_or_404(example.dataset_id, db, current_user)
    if example.status != ExampleStatus.approved:
        raise HTTPException(status_code=400, detail="Only approved examples can be expanded")
    variant_types = [v for v in payload.variant_types if v in VARIANT_META]
    if not variant_types:
        raise HTTPException(status_code=400, detail="Select at least one valid variant type")

    expand_example_task.delay(str(example.id), variant_types, payload.model)
    record_audit(db, current_user, "example.expand", "dataset_example", example.id, {"variant_types": variant_types})
    return {"status": "queued"}


@router.patch("/examples/{example_id}", response_model=DatasetExampleOut)
def update_example(
    example_id: uuid.UUID,
    payload: DatasetExampleUpdate,
    db: Session = Depends(get_db),
    current_user: User = Depends(require_role(UserRole.editor)),
):
    example = db.get(DatasetExample, example_id)
    if example is None:
        raise HTTPException(status_code=404, detail="Example not found")
    _get_dataset_or_404(example.dataset_id, db, current_user)

    content_changed = False
    for field in ("instruction", "input", "output", "example_type", "difficulty"):
        value = getattr(payload, field)
        if value is not None:
            setattr(example, field, value)
            content_changed = True

    if payload.status is not None:
        example.status = payload.status
    elif content_changed and example.status == ExampleStatus.pending:
        example.status = ExampleStatus.edited

    if payload.review_notes is not None:
        example.review_notes = payload.review_notes

    example.reviewer_id = current_user.id
    db.commit()
    db.refresh(example)
    return example


@router.post("/examples/{example_id}/approve", response_model=DatasetExampleOut)
def approve_example(
    example_id: uuid.UUID, db: Session = Depends(get_db), current_user: User = Depends(require_role(UserRole.editor))
):
    example = db.get(DatasetExample, example_id)
    if example is None:
        raise HTTPException(status_code=404, detail="Example not found")
    _get_dataset_or_404(example.dataset_id, db, current_user)
    example.status = ExampleStatus.approved
    example.reviewer_id = current_user.id
    db.commit()
    db.refresh(example)
    record_audit(db, current_user, "example.approve", "dataset_example", example.id)
    return example


@router.post("/examples/{example_id}/reject", response_model=DatasetExampleOut)
def reject_example(
    example_id: uuid.UUID, db: Session = Depends(get_db), current_user: User = Depends(require_role(UserRole.editor))
):
    example = db.get(DatasetExample, example_id)
    if example is None:
        raise HTTPException(status_code=404, detail="Example not found")
    _get_dataset_or_404(example.dataset_id, db, current_user)
    example.status = ExampleStatus.rejected
    example.reviewer_id = current_user.id
    db.commit()
    db.refresh(example)
    record_audit(db, current_user, "example.reject", "dataset_example", example.id)
    return example


@router.post("/examples/{example_id}/regenerate", status_code=202)
def regenerate_example(
    example_id: uuid.UUID, db: Session = Depends(get_db), current_user: User = Depends(require_role(UserRole.editor))
):
    example = db.get(DatasetExample, example_id)
    if example is None:
        raise HTTPException(status_code=404, detail="Example not found")
    _get_dataset_or_404(example.dataset_id, db, current_user)
    regenerate_example_task.delay(str(example.id))
    return {"status": "queued"}


@router.delete("/examples/{example_id}", status_code=204)
def delete_example(
    example_id: uuid.UUID, db: Session = Depends(get_db), current_user: User = Depends(require_role(UserRole.editor))
):
    example = db.get(DatasetExample, example_id)
    if example is None:
        raise HTTPException(status_code=404, detail="Example not found")
    _get_dataset_or_404(example.dataset_id, db, current_user)
    db.delete(example)
    db.commit()
    record_audit(db, current_user, "example.delete", "dataset_example", example_id)


@router.get("/datasets/{dataset_id}/export")
def export_dataset(
    dataset_id: uuid.UUID,
    format: str = Query(default="jsonl", pattern="^(jsonl|json|csv|parquet)$"),
    status: ExampleStatus = Query(default=ExampleStatus.approved),
    db: Session = Depends(get_db),
    current_user: User = Depends(get_current_principal),
):
    dataset = _get_dataset_or_404(dataset_id, db, current_user)
    examples = (
        db.query(DatasetExample)
        .filter(DatasetExample.dataset_id == dataset_id, DatasetExample.status == status)
        .order_by(DatasetExample.created_at)
        .all()
    )

    records = [
        {
            "instruction": e.instruction,
            "input": e.input or "",
            "output": e.output,
            "example_type": e.example_type.value,
            "difficulty": e.difficulty.value,
        }
        for e in examples
    ]

    filename_base = f"{dataset.name.replace(' ', '_').lower()}_{status.value}"
    record_audit(db, current_user, "dataset.export", "dataset", dataset.id, {"format": format, "status": status.value})

    if format == "jsonl":
        body = "\n".join(json.dumps(r, ensure_ascii=False) for r in records)
        return StreamingResponse(
            io.BytesIO(body.encode("utf-8")),
            media_type="application/jsonl",
            headers={"Content-Disposition": f"attachment; filename={filename_base}.jsonl"},
        )

    if format == "json":
        body = json.dumps(records, ensure_ascii=False, indent=2)
        return StreamingResponse(
            io.BytesIO(body.encode("utf-8")),
            media_type="application/json",
            headers={"Content-Disposition": f"attachment; filename={filename_base}.json"},
        )

    df = pd.DataFrame(records)

    if format == "csv":
        buf = io.StringIO()
        df.to_csv(buf, index=False)
        return StreamingResponse(
            io.BytesIO(buf.getvalue().encode("utf-8")),
            media_type="text/csv",
            headers={"Content-Disposition": f"attachment; filename={filename_base}.csv"},
        )

    # parquet
    buf = io.BytesIO()
    df.to_parquet(buf, index=False)
    buf.seek(0)
    return StreamingResponse(
        buf,
        media_type="application/octet-stream",
        headers={"Content-Disposition": f"attachment; filename={filename_base}.parquet"},
    )
