from typing import List, Optional

from pymongo import ASCENDING

from app.core.database import get_db
from app.models.chunk import DocChunk


_INDEXES_READY = False


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


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


async def _ensure_indexes() -> None:
    global _INDEXES_READY
    if _INDEXES_READY:
        return

    db = get_db()
    await db.doc_chunks.create_index(
        [("repo_id", ASCENDING), ("snapshot_id", ASCENDING)]
    )
    await db.doc_chunks.create_index(
        [("repo_id", ASCENDING), ("snapshot_id", ASCENDING), ("doc_path", ASCENDING)]
    )
    await db.doc_chunks.create_index(
        [("repo_id", ASCENDING), ("snapshot_id", ASCENDING), ("chunk_id", ASCENDING)],
        unique=True,
        sparse=True,
    )
    _INDEXES_READY = True


async def delete_chunks_for_doc(repo_id: str, doc_path: str) -> None:
    """Delete all chunks for a document (HTTP ingest path without snapshot scope)."""
    await _ensure_indexes()
    db = get_db()
    await db.doc_chunks.delete_many({"repo_id": repo_id, "doc_path": doc_path})


async def delete_chunks_for_doc_snapshot(
    repo_id: str,
    snapshot_id: str,
    doc_path: str,
) -> None:
    await _ensure_indexes()
    db = get_db()
    await db.doc_chunks.delete_many(_snapshot_filter(repo_id, snapshot_id, doc_path))


async def insert_chunks(chunks: List[DocChunk]) -> int:
    await _ensure_indexes()
    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,
    snapshot_id: Optional[str] = None,
    skip: int = 0,
    limit: int = 50,
) -> List[DocChunk]:
    await _ensure_indexes()
    db = get_db()
    query: dict = {"repo_id": repo_id}
    if doc_path:
        query["doc_path"] = doc_path
    if snapshot_id:
        query["$or"] = [
            {"snapshot_id": snapshot_id},
            {"metadata.snapshot.snapshot_id": snapshot_id},
        ]
    cursor = db.doc_chunks.find(query).skip(skip).limit(limit)
    return [DocChunk(**_serialize(doc)) async for doc in cursor]


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