from __future__ import annotations

from typing import List

from pydantic import ValidationError

from app.core.database import get_code_chunks_db, get_upstream_collection
from app.core.logger import logger
from app.models.upstream.code_chunk import CodeChunkDocument
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


async def list_by_snapshot(repo_id: str, snapshot_id: str) -> List[CodeChunkDocument]:
    col = get_upstream_collection(get_code_chunks_db(), "code_chunks")
    cursor = col.find({"repo_id": repo_id, "snapshot_id": snapshot_id})

    chunks: list[CodeChunkDocument] = []
    async for doc in cursor:
        try:
            chunk = CodeChunkDocument.model_validate(_serialize(doc))
        except ValidationError as exc:
            logger.warning(
                "Skipping invalid code 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.text, file_path=chunk.file_path):
            logger.debug(
                "Skipping non-embeddable code chunk chunk_id={chunk_id} file_path={file_path}",
                chunk_id=chunk.chunk_id,
                file_path=chunk.file_path,
            )
            continue

        chunks.append(chunk)
    return chunks


async def count_by_snapshot(repo_id: str, snapshot_id: str) -> int:
    db = get_code_chunks_db()
    return await db.code_chunks.count_documents(
        {"repo_id": repo_id, "snapshot_id": snapshot_id}
    )


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

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