from __future__ import annotations

from dataclasses import dataclass

from app.rda.application.builders import (
    RDA_MAPPER_VERSION,
    RDA_SCHEMA_VERSION,
    build_rda_emergency_summary,
    build_rda_patient_summary,
    build_source_context_fingerprint,
)
from app.rda.application.models import GenerateRdaCommand, RdaGenerationResult, RdaStatusResult
from app.rda.domain.models import (
    RdaArtifact,
    RdaArtifactType,
    RdaEmergencySummary,
    RdaGenerationStatus,
    RdaJobState,
    RdaPatientSummary,
    RdaSectionTrace,
)
from app.rda.domain.ports import CaseEpicrisisContextProvider, RdaArtifactRepository, RdaJobStatusRepository


@dataclass
class GenerateCaseRdaUseCase:
    context_provider: CaseEpicrisisContextProvider
    artifact_repository: RdaArtifactRepository
    job_status_repository: RdaJobStatusRepository

    def _normalize_artifact_type(self, artifact_type: str | RdaArtifactType) -> RdaArtifactType:
        if isinstance(artifact_type, RdaArtifactType):
            return artifact_type
        return RdaArtifactType(str(artifact_type).strip().lower())

    def _build_artifact_result(self, document: dict, *, source: str) -> RdaGenerationResult:
        return RdaGenerationResult(
            case_key=document.get("case_key", ""),
            artifact_type=document.get("artifact_type", RdaArtifactType.PATIENT.value),
            version=document.get("version"),
            schema_version=int(document.get("schema_version") or RDA_SCHEMA_VERSION),
            mapper_version=str(document.get("mapper_version") or RDA_MAPPER_VERSION),
            completeness_status=document.get("completeness_status") or "partial",
            missing_fields=list(document.get("missing_fields") or []),
            section_trace=[
                RdaSectionTrace(**item) for item in document.get("section_trace", []) if isinstance(item, dict)
            ],
            payload=dict(document.get("payload") or {}),
            source_context_fingerprint=document.get("source_context_fingerprint"),
            stored=True,
            source=source,  # type: ignore[arg-type]
        )

    def _build_status_result(self, document: dict, *, source: str = "status") -> RdaStatusResult:
        return RdaStatusResult(
            case_key=document.get("case_key", ""),
            artifact_type=document.get("artifact_type", RdaArtifactType.PATIENT.value),
            status=document.get("status", RdaJobState.QUEUED.value),
            job_id=document.get("job_id"),
            error=document.get("error"),
            source_context_fingerprint=document.get("source_context_fingerprint"),
            resolved_version=document.get("resolved_version"),
            reused=bool(document.get("reused", False)),
            updated_at=document.get("updated_at"),
            source=source,  # type: ignore[arg-type]
        )

    def _build_payload(
        self,
        command: GenerateRdaCommand,
        context: dict,
    ) -> tuple[dict, list[RdaSectionTrace], list[str], str]:
        normalized_artifact_type = self._normalize_artifact_type(command.artifact_type)
        if normalized_artifact_type == RdaArtifactType.PATIENT:
            payload, section_trace, missing_fields = build_rda_patient_summary(context)
            payload = RdaPatientSummary(**payload).model_dump(mode="python", exclude_none=True)
        else:
            payload, section_trace, missing_fields = build_rda_emergency_summary(context)
            payload = RdaEmergencySummary(**payload).model_dump(mode="python", exclude_none=True)
        completeness_status = "partial" if missing_fields else "complete"
        return payload, section_trace, missing_fields, completeness_status

    def get_latest(
        self,
        username: str,
        case_key: str,
        artifact_type: RdaArtifactType,
    ) -> RdaGenerationResult | None:
        document = self.artifact_repository.get_latest(username, case_key, artifact_type)
        if not document:
            return None
        return self._build_artifact_result(document, source="cache")

    def get_status(
        self,
        username: str,
        case_key: str,
        artifact_type: RdaArtifactType,
    ) -> RdaStatusResult | None:
        document = self.job_status_repository.get_status(username, case_key, artifact_type)
        if document:
            return self._build_status_result(document)

        latest = self.artifact_repository.get_latest(username, case_key, artifact_type)
        if not latest:
            return None

        derived = RdaGenerationStatus(
            case_key=case_key,
            artifact_type=artifact_type,
            status=RdaJobState.COMPLETED,
            job_id="",
            error=None,
            source_context_fingerprint=latest.get("source_context_fingerprint"),
            resolved_version=latest.get("version"),
            reused=True,
            updated_at=latest.get("fecha_analisis"),
        )
        return self._build_status_result(derived.model_dump(mode="python", exclude_none=True), source="artifact")

    def execute(self, command: GenerateRdaCommand) -> RdaGenerationResult:
        normalized_artifact_type = self._normalize_artifact_type(command.artifact_type)
        context = self.context_provider.get_case_context(command.username, command.case_key, regen=True)
        fingerprint = build_source_context_fingerprint(
            context,
            artifact_type=normalized_artifact_type,
            schema_version=RDA_SCHEMA_VERSION,
            mapper_version=RDA_MAPPER_VERSION,
        )
        latest = self.artifact_repository.get_latest(command.username, command.case_key, normalized_artifact_type)
        if (
            latest
            and not command.force
            and str(latest.get("source_context_fingerprint") or "").strip() == fingerprint
        ):
            return self._build_artifact_result(latest, source="cache")

        payload, section_trace, missing_fields, completeness_status = self._build_payload(command, context)
        version = None
        if command.persist:
            saved = self.artifact_repository.save_version(
                username=command.username,
                case_key=command.case_key,
                artifact_type=normalized_artifact_type,
                payload=payload,
                completeness_status=completeness_status,
                missing_fields=missing_fields,
                section_trace=[item.model_dump(mode="python", exclude_none=True) for item in section_trace],
                source_context=context,
                source_context_fingerprint=fingerprint,
                schema_version=RDA_SCHEMA_VERSION,
                mapper_version=RDA_MAPPER_VERSION,
                regen_requested=command.force,
            )
            version = saved.get("version")

        artifact = RdaArtifact(
            case_key=command.case_key,
            artifact_type=normalized_artifact_type,
            version=version,
            schema_version=RDA_SCHEMA_VERSION,
            mapper_version=RDA_MAPPER_VERSION,
            completeness_status=completeness_status,  # type: ignore[arg-type]
            missing_fields=missing_fields,
            section_trace=section_trace,
            payload=payload,
            source_context_fingerprint=fingerprint,
        )
        return RdaGenerationResult(**artifact.model_dump(mode="python"), stored=command.persist, source="generated")
