import uuid

from fastapi import APIRouter, Depends, HTTPException
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.knowledge import KnowledgeSource
from app.models.org import User, UserRole
from app.models.stubs import VectorEmbedding
from app.schemas.rag import CollectionSummary, SearchRequest, SearchResult
from app.services.rag import collection_summary, hybrid_search, semantic_search
from app.workers.rag_tasks import index_knowledge_source_task

router = APIRouter(prefix="/rag", tags=["rag"])


def _get_source_or_404(source_id: uuid.UUID, db: Session, current_user: User) -> KnowledgeSource:
    source = db.get(KnowledgeSource, source_id)
    if source is None:
        raise HTTPException(status_code=404, detail="Knowledge source not found")
    get_project_or_404(source.project_id, db, current_user)
    return source


@router.post("/index/{source_id}", status_code=202)
def index_source(
    source_id: uuid.UUID,
    db: Session = Depends(get_db),
    current_user: User = Depends(require_role(UserRole.editor)),
):
    from app.models.knowledge import ExtractedText

    source = _get_source_or_404(source_id, db, current_user)
    extracted = db.query(ExtractedText).filter_by(knowledge_source_id=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 indexing")

    index_knowledge_source_task.delay(str(source_id))
    return {"status": "queued"}


@router.delete("/index/{source_id}", status_code=204)
def delete_index(
    source_id: uuid.UUID,
    db: Session = Depends(get_db),
    current_user: User = Depends(require_role(UserRole.editor)),
):
    _get_source_or_404(source_id, db, current_user)
    db.query(VectorEmbedding).filter(VectorEmbedding.knowledge_source_id == source_id).delete()
    db.commit()


@router.get("/collections", response_model=list[CollectionSummary])
def list_collections(
    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 collection_summary(db, project_id)


@router.post("/search", response_model=list[SearchResult])
def search(
    payload: SearchRequest,
    db: Session = Depends(get_db),
    current_user: User = Depends(get_current_principal),
):
    get_project_or_404(payload.project_id, db, current_user)
    search_fn = hybrid_search if payload.hybrid else semantic_search
    return search_fn(
        db, payload.project_id, payload.query, top_k=payload.top_k, similarity_threshold=payload.similarity_threshold
    )
