from __future__ import annotations

import hashlib
import re
from copy import deepcopy
from dataclasses import dataclass
from datetime import UTC, datetime
from typing import Any, Protocol

from pymongo import ReturnDocument
from pymongo.errors import OperationFailure

from app.config import config


def _now_utc() -> datetime:
    return datetime.now(UTC)


def _normalize_spaces(value: str) -> str:
    return re.sub(r"\s+", " ", str(value or "").strip())


def _normalize_lower(value: str) -> str:
    return _normalize_spaces(value).lower()


def _is_unknown_name(value: str) -> bool:
    normalized = _normalize_lower(value)
    return normalized in {"", "desconocido", "paciente sin nombre"}


def _sha256_text(value: str) -> str:
    return hashlib.sha256(value.encode("utf-8")).hexdigest()


def _safe_int(value: Any) -> int:
    try:
        return int(value)
    except (TypeError, ValueError):
        return 0


@dataclass(frozen=True)
class DemoIdentityView:
    demo_case_key: str
    demo_patient_id: str
    demo_patient_name: str
    real_case_key: str
    real_patient_id: str
    real_patient_name: str


class DemoIdentityServiceProtocol(Protocol):
    enabled: bool

    def ensure_indexes(self) -> None: ...

    def resolve_case_key(self, *, username: str, visible_case_key: str) -> str: ...

    def resolve_patient_name(self, *, username: str, visible_patient_name: str) -> str: ...

    def resolve_patient_id(self, *, username: str, visible_patient_id: str) -> str: ...

    def project_case_identity(
        self,
        *,
        username: str,
        case_key: str,
        case_number: str = "",
        patient_id: str = "",
        patient_name: str = "",
    ) -> dict[str, str]: ...

    def project_document(self, *, username: str, document: dict[str, Any] | None) -> dict[str, Any] | None: ...

    def project_snapshot(self, *, username: str, snapshot: dict[str, Any]) -> dict[str, Any]: ...

    def project_epicrisis_context(
        self,
        *,
        username: str,
        context: dict[str, Any],
        case_key: str,
        case_number: str = "",
        patient_id: str = "",
        patient_name: str = "",
    ) -> dict[str, Any]: ...


class NoopDemoIdentityService:
    enabled = False

    def ensure_indexes(self) -> None:
        return

    def resolve_case_key(self, *, username: str, visible_case_key: str) -> str:
        return str(visible_case_key or "").strip()

    def resolve_patient_name(self, *, username: str, visible_patient_name: str) -> str:
        return str(visible_patient_name or "").strip()

    def resolve_patient_id(self, *, username: str, visible_patient_id: str) -> str:
        return str(visible_patient_id or "").strip()

    def project_case_identity(
        self,
        *,
        username: str,
        case_key: str,
        case_number: str = "",
        patient_id: str = "",
        patient_name: str = "",
    ) -> dict[str, str]:
        return {
            "case_key": str(case_key or "").strip(),
            "case_number": str(case_number or "").strip(),
            "patient_id": str(patient_id or "").strip(),
            "patient_name": str(patient_name or "").strip(),
            "nombre_paciente": str(patient_name or "").strip(),
        }

    def project_document(self, *, username: str, document: dict[str, Any] | None) -> dict[str, Any] | None:
        if document is None:
            return None
        return dict(document)

    def project_snapshot(self, *, username: str, snapshot: dict[str, Any]) -> dict[str, Any]:
        return deepcopy(snapshot)

    def project_epicrisis_context(
        self,
        *,
        username: str,
        context: dict[str, Any],
        case_key: str,
        case_number: str = "",
        patient_id: str = "",
        patient_name: str = "",
    ) -> dict[str, Any]:
        return deepcopy(context)


_NOOP_DEMO_IDENTITY_SERVICE = NoopDemoIdentityService()


class DemoIdentityService:
    def __init__(
        self,
        collection: Any,
        *,
        enabled: bool | None = None,
    ) -> None:
        self.collection = getattr(collection, "collection", collection)
        self.enabled = config.DEMO_ANONYMIZATION_ENABLED if enabled is None else bool(enabled)

    def ensure_indexes(self) -> None:
        collection = self.collection
        if not self.enabled or collection is None or not hasattr(collection, "create_index"):
            return
        try:
            collection.create_index(
                [("kind", 1), ("scope", 1)],
                unique=True,
                partialFilterExpression={"kind": "counter", "scope": {"$exists": True, "$type": "string"}},
            )
            collection.create_index(
                [("kind", 1), ("real_patient_key", 1)],
                unique=True,
                partialFilterExpression={
                    "kind": "patient_alias",
                    "real_patient_key": {"$exists": True, "$type": "string"},
                },
            )
            collection.create_index(
                [("kind", 1), ("real_case_key", 1)],
                unique=True,
                partialFilterExpression={
                    "kind": "case_alias",
                    "real_case_key": {"$exists": True, "$type": "string"},
                },
            )
            collection.create_index(
                [("kind", 1), ("demo_case_key", 1)],
                unique=True,
                partialFilterExpression={
                    "kind": "case_alias",
                    "demo_case_key": {"$exists": True, "$type": "string"},
                },
            )
            collection.create_index(
                [("kind", 1), ("demo_patient_name", 1)],
                partialFilterExpression={
                    "kind": "patient_alias",
                    "demo_patient_name": {"$exists": True, "$type": "string"},
                },
            )
        except OperationFailure as exc:
            if getattr(exc, "code", None) not in {85, 86}:
                raise

    def resolve_case_key(self, *, username: str, visible_case_key: str) -> str:
        normalized = str(visible_case_key or "").strip()
        if not self.enabled or not normalized or self.collection is None or not hasattr(self.collection, "find_one"):
            return normalized
        alias = self.collection.find_one({"kind": "case_alias", "demo_case_key": normalized})
        return str((alias or {}).get("real_case_key") or normalized).strip()

    def resolve_patient_name(self, *, username: str, visible_patient_name: str) -> str:
        normalized = str(visible_patient_name or "").strip()
        if not self.enabled or not normalized or self.collection is None or not hasattr(self.collection, "find_one"):
            return normalized
        alias = self.collection.find_one({"kind": "patient_alias", "demo_patient_name": normalized})
        return str((alias or {}).get("real_patient_name") or normalized).strip()

    def resolve_patient_id(self, *, username: str, visible_patient_id: str) -> str:
        normalized = str(visible_patient_id or "").strip()
        if not self.enabled or not normalized or self.collection is None or not hasattr(self.collection, "find_one"):
            return normalized
        alias = self.collection.find_one({"kind": "patient_alias", "demo_patient_id": normalized})
        return str((alias or {}).get("real_patient_id") or normalized).strip()

    def project_case_identity(
        self,
        *,
        username: str,
        case_key: str,
        case_number: str = "",
        patient_id: str = "",
        patient_name: str = "",
    ) -> dict[str, str]:
        normalized_case_key = str(case_key or "").strip()
        if not self.enabled or not normalized_case_key:
            return _NOOP_DEMO_IDENTITY_SERVICE.project_case_identity(
                username=username,
                case_key=case_key,
                case_number=case_number,
                patient_id=patient_id,
                patient_name=patient_name,
            )
        aliases = self._resolve_aliases(
            case_key=normalized_case_key,
            case_number=case_number,
            patient_id=patient_id,
            patient_name=patient_name,
        )
        return {
            "case_key": aliases.demo_case_key,
            "case_number": str(case_number or "").strip(),
            "patient_id": aliases.demo_patient_id,
            "patient_name": aliases.demo_patient_name,
            "nombre_paciente": aliases.demo_patient_name,
        }

    def project_document(self, *, username: str, document: dict[str, Any] | None) -> dict[str, Any] | None:
        if document is None:
            return None
        projected = deepcopy(document)
        aliases = self._resolve_aliases_from_payload(projected)
        if aliases is None:
            return projected
        projected["case_key"] = aliases.demo_case_key
        if "patient_id" in projected:
            projected["patient_id"] = aliases.demo_patient_id
        if "patient_name" in projected:
            projected["patient_name"] = aliases.demo_patient_name
        if "nombre_paciente" in projected:
            projected["nombre_paciente"] = aliases.demo_patient_name
        return self._sanitize_payload(projected, aliases)

    def project_snapshot(self, *, username: str, snapshot: dict[str, Any]) -> dict[str, Any]:
        if not self.enabled:
            return deepcopy(snapshot)
        projected = deepcopy(snapshot)
        projected_cases: dict[str, Any] = {}
        projected_patients: dict[str, list[dict[str, Any]]] = {}

        for case_group in (snapshot.get("historias_por_caso") or {}).values():
            if not isinstance(case_group, dict):
                continue
            aliases = self._resolve_aliases_from_payload(case_group)
            visible_case_key = aliases.demo_case_key if aliases else str(case_group.get("case_key") or "").strip()
            projected_group = self._project_case_group(case_group, aliases=aliases)
            projected_cases[visible_case_key or str(case_group.get("group_key") or "")] = projected_group

        for documents in (snapshot.get("historias_por_paciente") or {}).values():
            if not isinstance(documents, list):
                continue
            projected_documents = []
            visible_patient_name = "desconocido"
            for document in documents:
                if not isinstance(document, dict):
                    continue
                projected_document = self.project_document(username=username, document=document)
                if not isinstance(projected_document, dict):
                    continue
                visible_patient_name = str(projected_document.get("nombre_paciente") or visible_patient_name)
                projected_documents.append(projected_document)
            projected_patients.setdefault(visible_patient_name, []).extend(projected_documents)

        projected["historias_por_caso"] = projected_cases
        projected["historias_por_paciente"] = projected_patients
        return projected

    def project_epicrisis_context(
        self,
        *,
        username: str,
        context: dict[str, Any],
        case_key: str,
        case_number: str = "",
        patient_id: str = "",
        patient_name: str = "",
    ) -> dict[str, Any]:
        projected = deepcopy(context)
        if not self.enabled:
            return projected
        aliases = self._resolve_aliases(
            case_key=case_key,
            case_number=case_number,
            patient_id=patient_id,
            patient_name=patient_name,
        )
        projected["case_key"] = aliases.demo_case_key
        projected["regen_case_key"] = aliases.demo_case_key
        projected["nombre_paciente"] = aliases.demo_patient_name
        metadatos_hc = projected.get("metadatos_hc")
        if isinstance(metadatos_hc, dict):
            metadatos_hc["nombre_paciente"] = aliases.demo_patient_name
            projected["metadatos_hc"] = metadatos_hc
        return self._sanitize_payload(projected, aliases)

    def _project_case_group(self, case_group: dict[str, Any], *, aliases: DemoIdentityView | None) -> dict[str, Any]:
        projected = deepcopy(case_group)
        if aliases is None:
            return projected
        projected["case_key"] = aliases.demo_case_key
        projected["patient_name"] = aliases.demo_patient_name
        projected["patient_id"] = aliases.demo_patient_id
        projected["group_key"] = aliases.demo_case_key
        return self._sanitize_payload(projected, aliases)

    def _resolve_aliases_from_payload(self, payload: dict[str, Any]) -> DemoIdentityView | None:
        normalized_case_key = str(payload.get("case_key") or "").strip()
        if not self.enabled or not normalized_case_key:
            return None
        patient_name = str(payload.get("patient_name") or payload.get("nombre_paciente") or "").strip()
        patient_id = str(payload.get("patient_id") or "").strip()
        case_number = str(payload.get("case_number") or "").strip()
        return self._resolve_aliases(
            case_key=normalized_case_key,
            case_number=case_number,
            patient_id=patient_id,
            patient_name=patient_name,
        )

    def _resolve_aliases(
        self,
        *,
        case_key: str,
        case_number: str,
        patient_id: str,
        patient_name: str,
    ) -> DemoIdentityView:
        normalized_case_key = str(case_key or "").strip()
        normalized_patient_id = str(patient_id or "").strip()
        normalized_patient_name = _normalize_spaces(patient_name)
        patient_alias = self._get_or_create_patient_alias(
            patient_id=normalized_patient_id,
            patient_name=normalized_patient_name,
        )
        case_alias = self._get_or_create_case_alias(
            case_key=normalized_case_key,
            case_number=str(case_number or "").strip(),
            patient_alias=patient_alias,
        )
        return DemoIdentityView(
            demo_case_key=str(case_alias.get("demo_case_key") or normalized_case_key).strip(),
            demo_patient_id=str(patient_alias.get("demo_patient_id") or normalized_patient_id).strip(),
            demo_patient_name=str(patient_alias.get("demo_patient_name") or normalized_patient_name or "paciente").strip(),
            real_case_key=normalized_case_key,
            real_patient_id=normalized_patient_id,
            real_patient_name=normalized_patient_name,
        )

    def _get_or_create_patient_alias(self, *, patient_id: str, patient_name: str) -> dict[str, Any]:
        collection = self.collection
        normalized_patient_id = str(patient_id or "").strip()
        normalized_patient_name = _normalize_spaces(patient_name)
        patient_key = self._build_real_patient_key(normalized_patient_id, normalized_patient_name)
        if (
            not self.enabled
            or not patient_key
            or collection is None
            or not hasattr(collection, "find_one_and_update")
        ):
            return {
                "demo_patient_id": normalized_patient_id,
                "demo_patient_name": normalized_patient_name or "paciente",
            }

        existing = collection.find_one({"kind": "patient_alias", "real_patient_key": patient_key})
        if existing:
            return dict(existing)

        sequence = self._next_counter("patient")
        demo_patient_id = f"{sequence:05d}"
        demo_patient_name = f"paciente-{sequence:02d}"
        created = collection.find_one_and_update(
            {"kind": "patient_alias", "real_patient_key": patient_key},
            {
                "$setOnInsert": {
                    "kind": "patient_alias",
                    "real_patient_key": patient_key,
                    "real_patient_id": normalized_patient_id,
                    "real_patient_name": normalized_patient_name,
                    "demo_patient_id": demo_patient_id,
                    "demo_patient_name": demo_patient_name,
                    "sequence": sequence,
                    "created_at": _now_utc(),
                }
            },
            upsert=True,
            return_document=ReturnDocument.AFTER,
        )
        return dict(created or {})

    def _get_or_create_case_alias(
        self,
        *,
        case_key: str,
        case_number: str,
        patient_alias: dict[str, Any],
    ) -> dict[str, Any]:
        collection = self.collection
        normalized_case_key = str(case_key or "").strip()
        if (
            not self.enabled
            or not normalized_case_key
            or collection is None
            or not hasattr(collection, "find_one_and_update")
        ):
            return {"demo_case_key": normalized_case_key}

        existing = collection.find_one({"kind": "case_alias", "real_case_key": normalized_case_key})
        if existing:
            return dict(existing)

        sequence = self._next_counter("case")
        demo_case_key = self._build_demo_case_key(
            case_sequence=sequence,
            demo_patient_id=str(patient_alias.get("demo_patient_id") or "").strip(),
            demo_patient_name=str(patient_alias.get("demo_patient_name") or "").strip(),
            case_number=case_number,
        )
        created = collection.find_one_and_update(
            {"kind": "case_alias", "real_case_key": normalized_case_key},
            {
                "$setOnInsert": {
                    "kind": "case_alias",
                    "real_case_key": normalized_case_key,
                    "case_number": str(case_number or "").strip(),
                    "demo_case_key": demo_case_key,
                    "demo_patient_id": str(patient_alias.get("demo_patient_id") or "").strip(),
                    "demo_patient_name": str(patient_alias.get("demo_patient_name") or "").strip(),
                    "case_sequence": sequence,
                    "created_at": _now_utc(),
                }
            },
            upsert=True,
            return_document=ReturnDocument.AFTER,
        )
        return dict(created or {})

    def _next_counter(self, scope: str) -> int:
        collection = self.collection
        if collection is None or not hasattr(collection, "find_one_and_update"):
            return 1
        result = collection.find_one_and_update(
            {"kind": "counter", "scope": scope},
            {
                "$inc": {"value": 1},
                "$setOnInsert": {"kind": "counter", "scope": scope, "created_at": _now_utc()},
            },
            upsert=True,
            return_document=ReturnDocument.AFTER,
        )
        return max(1, _safe_int((result or {}).get("value")))

    def _build_real_patient_key(self, patient_id: str, patient_name: str) -> str:
        normalized_patient_id = str(patient_id or "").strip()
        normalized_patient_name = _normalize_lower(patient_name)
        if normalized_patient_id:
            return _sha256_text(f"id::{normalized_patient_id}")
        if normalized_patient_name and not _is_unknown_name(normalized_patient_name):
            return _sha256_text(f"name::{normalized_patient_name}")
        return ""

    def _build_demo_case_key(
        self,
        *,
        case_sequence: int,
        demo_patient_id: str,
        demo_patient_name: str,
        case_number: str,
    ) -> str:
        safe_patient_name = re.sub(r"[^a-z0-9-]+", "-", _normalize_lower(demo_patient_name)).strip("-")
        safe_patient_id = re.sub(r"\D+", "", str(demo_patient_id or "").strip())
        safe_case_number = re.sub(r"[^A-Za-z0-9]+", "-", str(case_number or "").strip()).strip("-")
        parts = [safe_patient_id or "00000"]
        if safe_case_number:
            parts.append(safe_case_number)
        parts.extend([safe_patient_name or "paciente", f"caso-{case_sequence:05d}"])
        return "-".join(part for part in parts if part)

    def _sanitize_payload(self, value: Any, aliases: DemoIdentityView) -> Any:
        if isinstance(value, dict):
            return {key: self._sanitize_payload(item, aliases) for key, item in value.items()}
        if isinstance(value, list):
            return [self._sanitize_payload(item, aliases) for item in value]
        if isinstance(value, str):
            return self._sanitize_text(value, aliases)
        return value

    def _sanitize_text(self, text: str, aliases: DemoIdentityView) -> str:
        current = str(text or "")
        replacements = [
            (aliases.real_case_key, aliases.demo_case_key, False),
            (aliases.real_patient_id, aliases.demo_patient_id, False),
            (aliases.real_patient_name, aliases.demo_patient_name, True),
        ]
        for source, target, ignore_case in replacements:
            source_text = str(source or "").strip()
            target_text = str(target or "").strip()
            if not source_text or not target_text:
                continue
            if ignore_case:
                current = re.sub(re.escape(source_text), target_text, current, flags=re.IGNORECASE)
                continue
            if source_text.isdigit():
                current = re.sub(
                    rf"(?<!\d){re.escape(source_text)}(?!\d)",
                    target_text,
                    current,
                )
                continue
            current = current.replace(source_text, target_text)
        return current


def get_demo_identity_service(services: Any) -> DemoIdentityServiceProtocol:
    service = getattr(services, "demo_identity_service", None)
    return service if service is not None else _NOOP_DEMO_IDENTITY_SERVICE
