from typing import List, Optional

from app.core.database import get_db
from app.models.chunk import DocChunk, IngestionChunkContract, to_ingestion_chunk_contract


def _serialize(doc: dict) -> dict:
    if doc and "_id" in doc:
        doc["id"] = str(doc.pop("_id"))
    return doc


async def delete_chunks_for_doc(repo_id: str, doc_path: str) -> None:
    db = get_db()
    await db.doc_chunks.delete_many({"repo_id": repo_id, "doc_path": doc_path})


async def insert_chunks(chunks: List[DocChunk]) -> int:
    db = get_db()
    if not chunks:
        return 0
    result = await db.doc_chunks.insert_many(
        [c.model_dump(exclude={"id"}) for c in chunks]
    )
    return len(result.inserted_ids)


async def get_chunks_by_repo(
    repo_id: str,
    doc_path: Optional[str] = None,
    skip: int = 0,
    limit: int = 50,
) -> List[IngestionChunkContract]:
    db = get_db()
    query: dict = {"repo_id": repo_id}
    if doc_path:
        query["doc_path"] = doc_path
    cursor = db.doc_chunks.find(query).skip(skip).limit(limit)
    return [to_ingestion_chunk_contract(_serialize(doc)) async for doc in cursor]


async def get_chunk_by_chunk_id(chunk_id: str) -> Optional[IngestionChunkContract]:
    db = get_db()
    doc = await db.doc_chunks.find_one({"chunk_id": chunk_id})
    if not doc:
        return None
    return to_ingestion_chunk_contract(_serialize(doc))


async def count_chunks(repo_id: str, doc_path: Optional[str] = None) -> int:
    db = get_db()
    query: dict = {"repo_id": repo_id}
    if doc_path:
        query["doc_path"] = doc_path
    return await db.doc_chunks.count_documents(query)
