from __future__ import annotations

from copy import deepcopy
from typing import Any


EPICRISIS_CONTEXT_PARTS_COLLECTION = "epicrisis_case_context_parts"


class InMemoryEpicrisisContextPartsRepository:
    def __init__(self) -> None:
        self._parts: dict[tuple[str, str, str, str], dict[str, Any]] = {}

    def ensure_indexes(self) -> None:
        return None

    def replace_parts(
        self,
        *,
        username: str,
        case_key: str,
        generation_id: str,
        parts: list[dict[str, Any]],
    ) -> None:
        self.delete_case(username=username, case_key=case_key, keep_generation_id=generation_id)
        for part in parts:
            part_id = str(part.get("part_id") or "").strip()
            if part_id:
                self._parts[(username, case_key, generation_id, part_id)] = deepcopy(part)

    def get_parts(
        self,
        *,
        username: str,
        case_key: str,
        generation_id: str,
    ) -> list[dict[str, Any]]:
        return [
            deepcopy(part)
            for (stored_username, stored_case_key, stored_generation_id, _), part in self._parts.items()
            if (stored_username, stored_case_key, stored_generation_id)
            == (username, case_key, generation_id)
        ]

    def delete_case(
        self,
        *,
        username: str,
        case_key: str,
        keep_generation_id: str = "",
    ) -> int:
        keys = [
            key
            for key in self._parts
            if key[0] == username
            and key[1] == case_key
            and (not keep_generation_id or key[2] != keep_generation_id)
        ]
        for key in keys:
            del self._parts[key]
        return len(keys)


class MongoEpicrisisContextPartsRepository:
    def __init__(self, mongo_analyses: Any):
        self.mongo_analyses = mongo_analyses
        self.collection = self._resolve_collection(mongo_analyses)
        self._fallback = InMemoryEpicrisisContextPartsRepository() if self.collection is None else None

    @staticmethod
    def _resolve_collection(mongo_analyses: Any) -> Any:
        try:
            database = mongo_analyses.collection.database
            return database[EPICRISIS_CONTEXT_PARTS_COLLECTION]
        except (AttributeError, KeyError, TypeError):
            return None

    def ensure_indexes(self) -> None:
        if self._fallback is not None:
            self._fallback.ensure_indexes()
            return
        self.collection.create_index(
            [("usuario", 1), ("case_key", 1), ("generation_id", 1), ("part_id", 1)],
            unique=True,
        )
        self.collection.create_index([("usuario", 1), ("case_key", 1), ("generation_id", 1)])

    def replace_parts(
        self,
        *,
        username: str,
        case_key: str,
        generation_id: str,
        parts: list[dict[str, Any]],
    ) -> None:
        if self._fallback is not None:
            self._fallback.replace_parts(
                username=username,
                case_key=case_key,
                generation_id=generation_id,
                parts=parts,
            )
            return
        for part in parts:
            part_id = str(part.get("part_id") or "").strip()
            if not part_id:
                continue
            self.collection.update_one(
                {
                    "usuario": username,
                    "case_key": case_key,
                    "generation_id": generation_id,
                    "part_id": part_id,
                },
                {"$set": {**part, "usuario": username, "case_key": case_key, "generation_id": generation_id}},
                upsert=True,
            )
        self.collection.delete_many(
            {
                "usuario": username,
                "case_key": case_key,
                "generation_id": {"$ne": generation_id},
            }
        )

    def get_parts(
        self,
        *,
        username: str,
        case_key: str,
        generation_id: str,
    ) -> list[dict[str, Any]]:
        if self._fallback is not None:
            return self._fallback.get_parts(
                username=username,
                case_key=case_key,
                generation_id=generation_id,
            )
        return [
            deepcopy(item)
            for item in self.collection.find(
                {
                    "usuario": username,
                    "case_key": case_key,
                    "generation_id": generation_id,
                },
                sort=[("field", 1), ("part_number", 1)],
            )
        ]

    def delete_case(
        self,
        *,
        username: str,
        case_key: str,
        keep_generation_id: str = "",
    ) -> int:
        if self._fallback is not None:
            return self._fallback.delete_case(
                username=username,
                case_key=case_key,
                keep_generation_id=keep_generation_id,
            )
        query: dict[str, Any] = {"usuario": username, "case_key": case_key}
        if keep_generation_id:
            query["generation_id"] = {"$ne": keep_generation_id}
        result = self.collection.delete_many(query)
        return int(getattr(result, "deleted_count", 0) or 0)
