from __future__ import annotations

from datetime import datetime

from pymongo import ASCENDING

from app.core.database import get_db
from app.models.workspace import StoredChunkDocument, WorkspaceRepoDocument


class WorkspaceRepoRepository:
    async def ensure_indexes(self) -> None:
        db = get_db()
        await db.workspace_repos.create_index(
            [("org_id", ASCENDING), ("repo_id", ASCENDING)],
            unique=True,
        )

    async def get(self, org_id: str, repo_id: str) -> WorkspaceRepoDocument | None:
        doc = await get_db().workspace_repos.find_one(
            {"org_id": org_id, "repo_id": repo_id}
        )
        if not doc:
            return None
        doc.pop("_id", None)
        return WorkspaceRepoDocument.model_validate(doc)

    async def count_for_org(self, org_id: str) -> int:
        return await get_db().workspace_repos.count_documents({"org_id": org_id})

    async def create(self, repo: WorkspaceRepoDocument) -> WorkspaceRepoDocument:
        payload = repo.model_dump()
        await get_db().workspace_repos.insert_one(payload)
        return repo

    async def update_stats(
        self,
        org_id: str,
        repo_id: str,
        *,
        chunks_total: int | None = None,
        files_indexed: int | None = None,
        stale_count_delta: int = 0,
        state: str | None = None,
        commit_sha: str | None = None,
        workspace_id: str | None = None,
    ) -> None:
        update: dict = {
            "$set": {
                "updated_at": datetime.utcnow(),
                "last_indexed_at": datetime.utcnow(),
            }
        }
        set_fields = update["$set"]
        if chunks_total is not None:
            set_fields["chunks_total"] = chunks_total
        if files_indexed is not None:
            set_fields["files_indexed"] = files_indexed
        if state is not None:
            set_fields["state"] = state
        if commit_sha is not None:
            set_fields["commit_sha"] = commit_sha
        if workspace_id is not None:
            set_fields["workspace_id"] = workspace_id
        if stale_count_delta:
            update["$inc"] = {"stale_count": stale_count_delta}

        await get_db().workspace_repos.update_one(
            {"org_id": org_id, "repo_id": repo_id},
            update,
        )


class ChunkRepository:
    async def ensure_indexes(self) -> None:
        db = get_db()
        await db.lucos_chunks.create_index(
            [("org_id", ASCENDING), ("repo_id", ASCENDING), ("chunk_id", ASCENDING)],
            unique=True,
        )
        await db.lucos_chunks.create_index(
            [("org_id", ASCENDING), ("repo_id", ASCENDING), ("chunk_hash", ASCENDING)]
        )

    async def get(
        self, org_id: str, repo_id: str, chunk_id: str
    ) -> StoredChunkDocument | None:
        doc = await get_db().lucos_chunks.find_one(
            {"org_id": org_id, "repo_id": repo_id, "chunk_id": chunk_id}
        )
        if not doc:
            return None
        doc.pop("_id", None)
        return StoredChunkDocument.model_validate(doc)

    async def upsert(self, chunk: StoredChunkDocument) -> None:
        payload = chunk.model_dump()
        # created_at must only appear in $setOnInsert — Mongo rejects the same
        # path in both $set and $setOnInsert ("would create a conflict").
        created_at = payload.pop("created_at", chunk.created_at)
        payload["updated_at"] = datetime.utcnow()
        await get_db().lucos_chunks.update_one(
            {
                "org_id": chunk.org_id,
                "repo_id": chunk.repo_id,
                "chunk_id": chunk.chunk_id,
            },
            {
                "$set": payload,
                "$setOnInsert": {"created_at": created_at},
            },
            upsert=True,
        )

    async def delete_many(
        self, org_id: str, repo_id: str, chunk_ids: list[str]
    ) -> int:
        if not chunk_ids:
            return 0
        result = await get_db().lucos_chunks.delete_many(
            {
                "org_id": org_id,
                "repo_id": repo_id,
                "chunk_id": {"$in": chunk_ids},
            }
        )
        return result.deleted_count

    async def count_for_repo(self, org_id: str, repo_id: str) -> int:
        return await get_db().lucos_chunks.count_documents(
            {"org_id": org_id, "repo_id": repo_id}
        )

    async def count_distinct_files(self, org_id: str, repo_id: str) -> int:
        files = await get_db().lucos_chunks.distinct(
            "file_path", {"org_id": org_id, "repo_id": repo_id}
        )
        return len(files)

    async def get_by_chunk_ids(
        self,
        org_id: str,
        repo_id: str,
        chunk_ids: list[str],
    ) -> dict[str, StoredChunkDocument]:
        if not chunk_ids:
            return {}

        cursor = get_db().lucos_chunks.find(
            {
                "org_id": org_id,
                "repo_id": repo_id,
                "chunk_id": {"$in": chunk_ids},
            }
        )
        chunks: dict[str, StoredChunkDocument] = {}
        async for doc in cursor:
            doc.pop("_id", None)
            chunk = StoredChunkDocument.model_validate(doc)
            chunks[chunk.chunk_id] = chunk
        return chunks

    async def get_embedding_record_ids(
        self, org_id: str, repo_id: str, chunk_ids: list[str]
    ) -> list[str]:
        if not chunk_ids:
            return []
        cursor = get_db().lucos_chunks.find(
            {
                "org_id": org_id,
                "repo_id": repo_id,
                "chunk_id": {"$in": chunk_ids},
            },
            {"embedding_record_id": 1},
        )
        record_ids: list[str] = []
        async for doc in cursor:
            record_id = doc.get("embedding_record_id")
            if record_id:
                record_ids.append(record_id)
        return record_ids
