from __future__ import annotations

import os
from collections.abc import Callable
from datetime import datetime
from threading import Thread
from typing import Any

from app.batch_processing.domain.models import (
    EPICRISIS_STATUS_COMPLETADO,
    EPICRISIS_STATUS_EN_COLA,
    EPICRISIS_STATUS_FALLIDO,
    EPICRISIS_STATUS_PENDIENTE,
    EPICRISIS_STATUS_PROCESANDO,
)
from app.batch_processing.infrastructure.mongo_repositories import MongoBatchCaseRepository
from app.case_epicrisis.domain import EpicrisisGenerationBlockedError
from app.core.logging import (
    bind_log_context,
    clear_log_context,
    get_audit_logger,
    get_log_context,
    set_log_context,
)


def _canonical_case_epicrisis_url(case_key: str) -> str:
    return f"/epicrisis?case_key={case_key}"


def _normalize_rule_findings(value: Any) -> list[dict[str, Any]]:
    return [item for item in (value or []) if isinstance(item, dict)]


def _context_rule_payload(context: dict[str, Any] | None) -> dict[str, Any]:
    context = context if isinstance(context, dict) else {}
    raw_evaluation = context.get("rule_evaluation")
    evaluation: dict[str, Any] = raw_evaluation if isinstance(raw_evaluation, dict) else {}
    blocking_reason = str(context.get("blocking_reason") or "").strip()
    if not blocking_reason:
        for item in _normalize_rule_findings(evaluation.get("findings")):
            if bool(item.get("blocking")):
                blocking_reason = str(item.get("message") or "").strip()
                break
    return {
        "epicrisis_rule_status": str(evaluation.get("status") or "").strip(),
        "epicrisis_rule_findings": _normalize_rule_findings(evaluation.get("findings")),
        "epicrisis_blocking_reason": blocking_reason,
        "epicrisis_missing_documents": [
            str(item).strip()
            for item in (context.get("missing_documents") or evaluation.get("missing_documents") or [])
            if str(item).strip()
        ],
        "epicrisis_last_rule_evaluation_at": str(evaluation.get("evaluated_at") or "").strip(),
        "ready_for_epicrisis": bool(context.get("epicrisis_generation_allowed", True)),
    }


class CaseEpicrisisRuntimeService:
    """Orquesta la cola y el estado runtime de epicrisis para casos batch y manuales."""

    def __init__(
        self,
        *,
        case_repository: MongoBatchCaseRepository,
        case_epicrisis_service: Any,
        clinical_document_service: Any,
        colombia_tz: Any,
        recompute_batch_bulk_epicrisis: Callable[[str], Any] | None = None,
    ) -> None:
        self.audit_logger = get_audit_logger()
        self.case_repository = case_repository
        self.case_epicrisis_service = case_epicrisis_service
        self.clinical_document_service = clinical_document_service
        self.colombia_tz = colombia_tz
        self.recompute_batch_bulk_epicrisis = recompute_batch_bulk_epicrisis

    def queue_case_epicrisis(self, username: str, case_key: str, *, regen: bool = False) -> dict[str, object]:
        normalized_case_key = str(case_key or "").strip()
        if not normalized_case_key:
            raise ValueError("El parámetro 'case_key' es obligatorio.")

        bind_log_context(username=username, case_key=normalized_case_key, document_type="epicrisis")
        case = self._ensure_runtime_case(username, normalized_case_key)
        cached_context = (
            None
            if regen
            else self.case_epicrisis_service.get_cached_case_context(username, normalized_case_key)
        )
        if self._should_reset_stale_manual_runtime(case, cached_context):
            case = self._persist_case_state(
                username,
                normalized_case_key,
                self._manual_runtime_pending_payload(normalized_case_key),
            )
        if not bool(case.get("ready_for_epicrisis")):
            blocking_reason = str(case.get("epicrisis_blocking_reason") or "").strip()
            if cached_context and isinstance(cached_context.get("contexto"), dict) and blocking_reason:
                return self._queue_response(
                    case_key=normalized_case_key,
                    job_id=str(case.get("epicrisis_job_id") or ""),
                    status_value=EPICRISIS_STATUS_FALLIDO,
                    reused=True,
                    epicrisis_url=str(
                        case.get("epicrisis_url") or _canonical_case_epicrisis_url(normalized_case_key)
                    ),
                )
            raise ValueError("El caso aún no esta listo para generar epicrisis.")

        if cached_context and isinstance(cached_context.get("contexto"), dict):
            cached_payload = _context_rule_payload(cached_context["contexto"])
            if cached_payload["epicrisis_rule_status"] == "blocked":
                self._persist_case_state(
                    username,
                    normalized_case_key,
                    {
                        "ready_for_epicrisis": False,
                        "epicrisis_status": EPICRISIS_STATUS_FALLIDO,
                        "epicrisis_error": cached_payload["epicrisis_blocking_reason"],
                        "epicrisis_url": _canonical_case_epicrisis_url(normalized_case_key),
                        "updated_at": self._now_iso(),
                        **cached_payload,
                    },
                )
                return self._queue_response(
                    case_key=normalized_case_key,
                    job_id=str(case.get("epicrisis_job_id") or ""),
                    status_value=EPICRISIS_STATUS_FALLIDO,
                    reused=True,
                )
            updated = self._persist_case_state(
                username,
                normalized_case_key,
                {
                    "ready_for_epicrisis": True,
                    "epicrisis_status": EPICRISIS_STATUS_COMPLETADO,
                    "epicrisis_url": _canonical_case_epicrisis_url(normalized_case_key),
                    "epicrisis_error": "",
                    "updated_at": self._now_iso(),
                    **cached_payload,
                },
            )
            self._recompute_batch(updated)
            return self._queue_response(
                case_key=normalized_case_key,
                job_id=str(updated.get("epicrisis_job_id") or ""),
                status_value=EPICRISIS_STATUS_COMPLETADO,
                reused=True,
            )

        if not regen:
            current_status = str(case.get("epicrisis_status") or "").strip().lower()
            current_job_id = str(case.get("epicrisis_job_id") or "").strip()
            if current_status in {EPICRISIS_STATUS_EN_COLA, EPICRISIS_STATUS_PROCESANDO} and current_job_id:
                return self._queue_response(
                    case_key=normalized_case_key,
                    job_id=current_job_id,
                    status_value=current_status,
                    reused=True,
                    epicrisis_url=str(
                        case.get("epicrisis_url") or _canonical_case_epicrisis_url(normalized_case_key)
                    ),
                )

        dispatcher_mode = os.getenv("BATCH_DISPATCHER", "inprocess").strip().lower()
        if dispatcher_mode == "celery":
            from app.batch_processing.celery_app import generate_epicrisis_job

            result = generate_epicrisis_job.delay(
                username,
                normalized_case_key,
                regen,
                audit_context=get_log_context(),
            )
            job_id = result.id
        else:
            job_id = f"inprocess-{normalized_case_key}-{int(datetime.now().timestamp())}"

        self._persist_case_state(
            username,
            normalized_case_key,
            {
                "epicrisis_status": EPICRISIS_STATUS_EN_COLA,
                "epicrisis_job_id": job_id,
                "epicrisis_error": "",
                "epicrisis_url": _canonical_case_epicrisis_url(normalized_case_key),
                "updated_at": self._now_iso(),
            },
        )

        if dispatcher_mode != "celery":
            self._spawn_inprocess_job(
                username=username,
                case_key=normalized_case_key,
                regen=regen,
                job_id=job_id,
            )

        return self._queue_response(
            case_key=normalized_case_key,
            job_id=job_id,
            status_value=EPICRISIS_STATUS_EN_COLA,
            reused=False,
        )

    def get_case_epicrisis_status(self, username: str, case_key: str) -> dict[str, Any]:
        normalized_case_key = str(case_key or "").strip()
        if not normalized_case_key:
            raise ValueError("El parámetro 'case_key' es obligatorio.")

        bind_log_context(username=username, case_key=normalized_case_key, document_type="epicrisis")
        case = self.case_repository.get_user_case(username, normalized_case_key) or {}
        if not case:
            metadata = (
                self.clinical_document_service.get_user_case_context(username, normalized_case_key) or {}
            )
            if not metadata:
                raise ValueError("Caso no encontrado")
            case = {
                "case_key": normalized_case_key,
                "case_number": str(metadata.get("case_number") or "").strip(),
                "patient_id": str(metadata.get("patient_id") or "").strip(),
                "patient_name": str(metadata.get("patient_name") or "").strip(),
                "ready_for_epicrisis": True,
                "epicrisis_status": EPICRISIS_STATUS_PENDIENTE,
                "epicrisis_job_id": "",
                "epicrisis_error": "",
                "epicrisis_url": _canonical_case_epicrisis_url(normalized_case_key),
            }

        cache_doc = self.case_epicrisis_service.get_cached_case_context(username, normalized_case_key)
        if self._should_reset_stale_manual_runtime(case, cache_doc):
            case = {
                **case,
                **self._manual_runtime_pending_payload(normalized_case_key),
            }
        epicrisis_status = str(case.get("epicrisis_status") or EPICRISIS_STATUS_PENDIENTE)
        epicrisis_url = str(case.get("epicrisis_url") or _canonical_case_epicrisis_url(normalized_case_key))
        raw_missing_documents = case.get("epicrisis_missing_documents")
        missing_documents = raw_missing_documents if isinstance(raw_missing_documents, list) else []
        rule_payload = {
            "epicrisis_rule_status": str(case.get("epicrisis_rule_status") or "").strip(),
            "epicrisis_rule_findings": _normalize_rule_findings(case.get("epicrisis_rule_findings")),
            "epicrisis_blocking_reason": str(case.get("epicrisis_blocking_reason") or "").strip(),
            "epicrisis_missing_documents": [
                str(item).strip() for item in missing_documents if str(item).strip()
            ],
            "epicrisis_last_rule_evaluation_at": str(
                case.get("epicrisis_last_rule_evaluation_at") or ""
            ).strip(),
        }
        if cache_doc and isinstance(cache_doc.get("contexto"), dict):
            context_payload = _context_rule_payload(cache_doc["contexto"])
            rule_payload = {**rule_payload, **context_payload}
            epicrisis_status = (
                EPICRISIS_STATUS_FALLIDO
                if context_payload["epicrisis_rule_status"] == "blocked"
                else EPICRISIS_STATUS_COMPLETADO
            )
            epicrisis_url = _canonical_case_epicrisis_url(normalized_case_key)

        return {
            "case_key": normalized_case_key,
            "case_number": str(
                case.get("case_number") or ((cache_doc or {}).get("contexto") or {}).get("case_number") or ""
            ).strip(),
            "ready_for_epicrisis": bool(case.get("ready_for_epicrisis", False))
            if not cache_doc
            else bool(rule_payload.get("ready_for_epicrisis", True)),
            "epicrisis_status": epicrisis_status,
            "epicrisis_job_id": str(case.get("epicrisis_job_id") or ""),
            "epicrisis_error": str(
                rule_payload.get("epicrisis_blocking_reason") or case.get("epicrisis_error") or ""
            ),
            "epicrisis_url": epicrisis_url,
            **rule_payload,
        }

    def run_generation_job(
        self,
        username: str,
        case_key: str,
        *,
        regen: bool = False,
        job_id: str = "",
        audit_context: dict[str, object] | None = None,
    ) -> str:
        normalized_case_key = str(case_key or "").strip()
        if not normalized_case_key:
            raise ValueError("El parámetro 'case_key' es obligatorio.")

        set_log_context(dict(audit_context or {}))
        bind_log_context(
            username=username,
            case_key=normalized_case_key,
            job_id=job_id,
            document_type="epicrisis",
        )
        case = self._ensure_runtime_case(username, normalized_case_key)
        batch_id = str(case.get("batch_id") or "").strip()

        self._persist_case_state(
            username,
            normalized_case_key,
            {
                "epicrisis_status": EPICRISIS_STATUS_PROCESANDO,
                "epicrisis_job_id": job_id or str(case.get("epicrisis_job_id") or ""),
                "epicrisis_error": "",
                "updated_at": self._now_iso(),
            },
        )
        self._recompute_batch(case)

        try:
            if not regen:
                cached = self.case_epicrisis_service.get_cached_case_context(username, normalized_case_key)
                if cached and isinstance(cached.get("contexto"), dict):
                    self._persist_case_state(
                        username,
                        normalized_case_key,
                        {
                            "epicrisis_status": EPICRISIS_STATUS_COMPLETADO,
                            "epicrisis_url": _canonical_case_epicrisis_url(normalized_case_key),
                            "epicrisis_error": "",
                            "updated_at": self._now_iso(),
                        },
                    )
                    self._recompute_batch(case)
                    self.audit_logger.business_event(
                        event_type="epicrisis.cache_hit",
                        action="run_generation_job",
                        outcome="success",
                        service="case_epicrisis_runtime_service",
                        resource={"case_key": normalized_case_key, "batch_id": batch_id},
                    )
                    return normalized_case_key

            context = self.case_epicrisis_service.cache_case_context(
                username,
                normalized_case_key,
                regen=regen,
            )
            rule_payload = _context_rule_payload(context)
            self._persist_case_state(
                username,
                normalized_case_key,
                {
                    "ready_for_epicrisis": True,
                    "epicrisis_status": EPICRISIS_STATUS_COMPLETADO,
                    "epicrisis_url": _canonical_case_epicrisis_url(normalized_case_key),
                    "epicrisis_error": "",
                    "updated_at": self._now_iso(),
                    **rule_payload,
                },
            )
            self._recompute_batch(case)
            self.audit_logger.business_event(
                event_type="epicrisis.generated",
                action="run_generation_job",
                outcome="success",
                service="case_epicrisis_runtime_service",
                resource={"case_key": normalized_case_key, "batch_id": batch_id, "regen": bool(regen)},
            )
            return normalized_case_key
        except EpicrisisGenerationBlockedError as exc:
            rule_payload = _context_rule_payload(exc.context)
            self._persist_case_state(
                username,
                normalized_case_key,
                {
                    "ready_for_epicrisis": False,
                    "epicrisis_status": EPICRISIS_STATUS_FALLIDO,
                    "epicrisis_error": str(exc),
                    "epicrisis_url": _canonical_case_epicrisis_url(normalized_case_key),
                    "updated_at": self._now_iso(),
                    **rule_payload,
                },
            )
            self._recompute_batch(case)
            self.audit_logger.business_event(
                event_type="epicrisis.generated",
                action="run_generation_job",
                outcome="blocked",
                service="case_epicrisis_runtime_service",
                resource={"case_key": normalized_case_key, "batch_id": batch_id, "regen": bool(regen)},
                error={"class": exc.__class__.__name__, "message": str(exc)},
            )
            raise
        except Exception as exc:
            self._persist_case_state(
                username,
                normalized_case_key,
                {
                    "epicrisis_status": EPICRISIS_STATUS_FALLIDO,
                    "epicrisis_error": str(exc),
                    "updated_at": self._now_iso(),
                },
            )
            self._recompute_batch(case)
            self.audit_logger.business_event(
                event_type="epicrisis.generated",
                action="run_generation_job",
                outcome="error",
                service="case_epicrisis_runtime_service",
                resource={"case_key": normalized_case_key, "batch_id": batch_id, "regen": bool(regen)},
                error={"class": exc.__class__.__name__, "message": str(exc)},
            )
            raise
        finally:
            clear_log_context()

    def _spawn_inprocess_job(self, *, username: str, case_key: str, regen: bool, job_id: str) -> None:
        audit_context = get_log_context()

        def _runner() -> None:
            self.run_generation_job(
                username,
                case_key,
                regen=regen,
                job_id=job_id,
                audit_context=audit_context,
            )

        Thread(target=_runner, daemon=True).start()

    def _ensure_runtime_case(self, username: str, case_key: str) -> dict[str, Any]:
        existing = self.case_repository.get_user_case(username, case_key) or {}
        if existing:
            return existing

        metadata = self.clinical_document_service.get_user_case_context(username, case_key) or {}
        if not metadata:
            raise ValueError("Caso no encontrado")

        return self.case_repository.upsert_user_case(
            username,
            case_key,
            {
                "batch_id": "",
                "case_number": str(metadata.get("case_number") or "").strip(),
                "patient_name": str(metadata.get("patient_name") or "").strip(),
                "patient_id": str(metadata.get("patient_id") or "").strip(),
                "ready_for_epicrisis": True,
                "epicrisis_status": EPICRISIS_STATUS_PENDIENTE,
                "epicrisis_job_id": "",
                "epicrisis_error": "",
                "epicrisis_url": _canonical_case_epicrisis_url(case_key),
                "runtime_origin": "manual",
                "updated_at": self._now_iso(),
            },
        )

    def _persist_case_state(self, username: str, case_key: str, payload: dict[str, Any]) -> dict[str, Any]:
        return self.case_repository.upsert_user_case(username, case_key, payload)

    def _recompute_batch(self, case: dict[str, Any]) -> None:
        batch_id = str(case.get("batch_id") or "").strip()
        if batch_id and self.recompute_batch_bulk_epicrisis is not None:
            self.recompute_batch_bulk_epicrisis(batch_id)

    def _is_manual_runtime_case(self, case: dict[str, Any]) -> bool:
        runtime_origin = str(case.get("runtime_origin") or "").strip().lower()
        batch_id = str(case.get("batch_id") or "").strip()
        return not batch_id and runtime_origin in {"", "manual", "individual"}

    def _has_runtime_rule_state(self, case: dict[str, Any]) -> bool:
        missing_documents = case.get("epicrisis_missing_documents")
        return bool(
            str(case.get("epicrisis_rule_status") or "").strip()
            or str(case.get("epicrisis_blocking_reason") or "").strip()
            or (isinstance(missing_documents, list) and any(str(item).strip() for item in missing_documents))
        )

    def _should_reset_stale_manual_runtime(
        self,
        case: dict[str, Any],
        cached_context: dict[str, Any] | None,
    ) -> bool:
        if not case or not self._is_manual_runtime_case(case):
            return False
        if cached_context and isinstance(cached_context.get("contexto"), dict):
            return False
        if self._has_runtime_rule_state(case):
            return True
        return (
            not bool(case.get("ready_for_epicrisis", True))
            and str(case.get("epicrisis_status") or "").strip().lower() != EPICRISIS_STATUS_FALLIDO
        )

    def _manual_runtime_pending_payload(self, case_key: str) -> dict[str, Any]:
        return {
            "ready_for_epicrisis": True,
            "epicrisis_status": EPICRISIS_STATUS_PENDIENTE,
            "epicrisis_job_id": "",
            "epicrisis_error": "",
            "epicrisis_url": _canonical_case_epicrisis_url(case_key),
            "epicrisis_rule_status": "",
            "epicrisis_rule_findings": [],
            "epicrisis_blocking_reason": "",
            "epicrisis_missing_documents": [],
            "epicrisis_last_rule_evaluation_at": "",
            "runtime_origin": "manual",
            "updated_at": self._now_iso(),
        }

    def _now_iso(self) -> str:
        return datetime.now(self.colombia_tz).isoformat()

    def _queue_response(
        self,
        *,
        case_key: str,
        job_id: str,
        status_value: str,
        reused: bool,
        epicrisis_url: str | None = None,
    ) -> dict[str, object]:
        return {
            "job_id": job_id,
            "status": status_value,
            "reused": reused,
            "epicrisis_url": epicrisis_url or _canonical_case_epicrisis_url(case_key),
        }
