from __future__ import annotations

from typing import List

from pydantic import ValidationError

from app.core.database import get_docs_db, get_upstream_collection
from app.core.logger import logger
from app.models.upstream.doc_chunk import DocChunkDocument
from app.utils.text_safety import is_embeddable_text


def _serialize(doc: dict) -> dict:
    if doc and "_id" in doc:
        doc = dict(doc)
        doc.pop("_id", None)
    return doc


def _snapshot_query(repo_id: str, snapshot_id: str) -> dict:
    return {
        "repo_id": repo_id,
        "$or": [
            {"metadata.snapshot_id": snapshot_id},
            {"snapshot_id": snapshot_id},
        ],
    }


async def list_by_snapshot(repo_id: str, snapshot_id: str) -> List[DocChunkDocument]:
    col = get_upstream_collection(get_docs_db(), "doc_chunks")
    cursor = col.find(_snapshot_query(repo_id, snapshot_id))

    chunks: list[DocChunkDocument] = []
    async for doc in cursor:
        try:
            chunk = DocChunkDocument.model_validate(_serialize(doc))
        except ValidationError as exc:
            logger.warning(
                "Skipping invalid doc chunk repo_id={repo_id} snapshot_id={snapshot_id} chunk_id={chunk_id} error={error}",
                repo_id=repo_id,
                snapshot_id=snapshot_id,
                chunk_id=doc.get("chunk_id"),
                error=str(exc),
            )
            continue

        if not is_embeddable_text(chunk.chunk_text, file_path=chunk.doc_path):
            logger.debug(
                "Skipping non-embeddable doc chunk chunk_id={chunk_id} doc_path={doc_path}",
                chunk_id=chunk.chunk_id,
                doc_path=chunk.doc_path,
            )
            continue

        chunks.append(chunk)
    return chunks


async def count_by_snapshot(repo_id: str, snapshot_id: str) -> int:
    db = get_docs_db()
    return await db.doc_chunks.count_documents(_snapshot_query(repo_id, snapshot_id))


async def get_by_chunk_ids(chunk_ids: list[str]) -> dict[str, DocChunkDocument]:
    """Batch-load doc chunks keyed by chunk_id."""
    if not chunk_ids:
        return {}

    col = get_upstream_collection(get_docs_db(), "doc_chunks")
    unique_ids = list(dict.fromkeys(chunk_ids))
    cursor = col.find({"chunk_id": {"$in": unique_ids}})
    chunks: dict[str, DocChunkDocument] = {}
    async for doc in cursor:
        try:
            chunk = DocChunkDocument.model_validate(_serialize(doc))
        except ValidationError:
            continue
        if not is_embeddable_text(chunk.chunk_text, file_path=chunk.doc_path):
            continue
        chunks[chunk.chunk_id] = chunk
    return chunks
