from __future__ import annotations

import uuid

from fastapi import APIRouter, HTTPException, Query

from app.api.schemas.pipeline import (
    CommitEmbedRequest,
    EmbeddingRecordCountResponse,
    EmbeddingRecordSummary,
    EmbeddingRunResponse,
    PipelineJobResponse,
    SnapshotEmbedRequest,
)
from app.models.enums import EmbedStatus, SourceType
from app.repositories import embedding_record_repo, embedding_run_repo
from app.workers.pipeline_tasks import (
    process_commit_analysis_embedding_task,
    process_snapshot_embedding_task,
)

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


def _to_run_response(run) -> EmbeddingRunResponse:
    return EmbeddingRunResponse.model_validate(run.model_dump())


def _to_record_summary(record) -> EmbeddingRecordSummary:
    text = record.text or ""
    preview = text if len(text) <= 240 else f"{text[:237]}..."
    return EmbeddingRecordSummary(
        record_id=record.record_id,
        repo_id=record.repo_id,
        snapshot_id=record.snapshot_id,
        source_type=record.source_type,
        chunk_type=record.chunk_type,
        embed_status=record.embed_status,
        commit_sha=record.commit_sha,
        file_path=record.file_path,
        doc_path=record.doc_path,
        symbol_name=record.symbol_name,
        section_title=record.section_title,
        graph_node_id=record.graph_node_id,
        qdrant_point_id=record.qdrant_point_id,
        upstream_chunk_id=record.upstream_chunk_id,
        upstream_analysis_id=record.upstream_analysis_id,
        text_preview=preview,
        updated_at=record.updated_at,
    )


@router.get("/runs/{run_id}", response_model=EmbeddingRunResponse)
async def get_embedding_run(run_id: str) -> EmbeddingRunResponse:
    run = await embedding_run_repo.get_run(run_id)
    if run is None:
        raise HTTPException(status_code=404, detail="Embedding run not found")
    return _to_run_response(run)


@router.get("/runs", response_model=list[EmbeddingRunResponse])
async def list_embedding_runs(
    repo_id: str = Query(...),
    snapshot_id: str = Query(...),
    limit: int = Query(50, ge=1, le=200),
) -> list[EmbeddingRunResponse]:
    runs = await embedding_run_repo.list_runs_by_snapshot(
        repo_id,
        snapshot_id,
        limit=limit,
    )
    return [_to_run_response(run) for run in runs]


@router.get("/records/{record_id}")
async def get_embedding_record(record_id: str):
    record = await embedding_record_repo.get_record(record_id)
    if record is None:
        raise HTTPException(status_code=404, detail="Embedding record not found")
    return record


@router.get("/records", response_model=list[EmbeddingRecordSummary])
async def list_embedding_records(
    repo_id: str = Query(...),
    snapshot_id: str = Query(...),
    source_type: SourceType | None = Query(None),
    embed_status: EmbedStatus | None = Query(None),
    skip: int = Query(0, ge=0),
    limit: int = Query(100, ge=1, le=500),
) -> list[EmbeddingRecordSummary]:
    records = await embedding_record_repo.list_records_by_snapshot(
        repo_id,
        snapshot_id,
        source_type=source_type,
        embed_status=embed_status,
        skip=skip,
        limit=limit,
    )
    return [_to_record_summary(record) for record in records]


@router.get("/records/count/summary", response_model=EmbeddingRecordCountResponse)
async def count_embedding_records(
    repo_id: str = Query(...),
    snapshot_id: str = Query(...),
) -> EmbeddingRecordCountResponse:
    total = await embedding_record_repo.count_records_by_snapshot(repo_id, snapshot_id)
    embedded = await embedding_record_repo.count_records_by_snapshot(
        repo_id,
        snapshot_id,
        embed_status=EmbedStatus.EMBEDDED,
    )
    pending = await embedding_record_repo.count_records_by_snapshot(
        repo_id,
        snapshot_id,
        embed_status=EmbedStatus.PENDING,
    )
    failed = await embedding_record_repo.count_records_by_snapshot(
        repo_id,
        snapshot_id,
        embed_status=EmbedStatus.FAILED,
    )
    return EmbeddingRecordCountResponse(
        repo_id=repo_id,
        snapshot_id=snapshot_id,
        total=total,
        embedded=embedded,
        pending=pending,
        failed=failed,
    )


@router.post("/pipeline/snapshots/embed", response_model=PipelineJobResponse, status_code=202)
async def trigger_snapshot_embedding(request: SnapshotEmbedRequest) -> PipelineJobResponse:
    event_id = request.event_id or f"evt_manual_{request.snapshot_id}_{uuid.uuid4().hex[:8]}"
    run = await embedding_run_repo.upsert_run(
        repo_id=request.repo_id,
        snapshot_id=request.snapshot_id,
        event_id=event_id,
        run_type="snapshot_batch",
    )
    task = process_snapshot_embedding_task.delay(
        run.run_id,
        request.repo_id,
        request.snapshot_id,
        request.commit_sha,
        request.artifact_uri or "",
        event_id,
    )
    await embedding_run_repo.update_run_state(
        run.run_id,
        EmbedStatus.PENDING,
        celery_task_id=task.id,
    )
    return PipelineJobResponse(run_id=run.run_id, celery_task_id=task.id)


@router.post("/pipeline/commits/embed", response_model=PipelineJobResponse, status_code=202)
async def trigger_commit_embedding(request: CommitEmbedRequest) -> PipelineJobResponse:
    event_id = request.event_id or f"evt_manual_{request.analysis_id}_{uuid.uuid4().hex[:8]}"
    snapshot_id = request.snapshot_id or request.commit_sha
    run = await embedding_run_repo.upsert_run(
        repo_id=request.repo_id,
        snapshot_id=snapshot_id,
        event_id=event_id,
        run_type="commit_analysis",
    )
    task = process_commit_analysis_embedding_task.delay(
        run.run_id,
        request.repo_id,
        request.commit_sha,
        request.analysis_id,
        event_id,
        snapshot_id,
    )
    await embedding_run_repo.update_run_state(
        run.run_id,
        EmbedStatus.PENDING,
        celery_task_id=task.id,
    )
    return PipelineJobResponse(run_id=run.run_id, celery_task_id=task.id)
