from __future__ import annotations

import hashlib
import json
from dataclasses import dataclass
from datetime import datetime
from typing import Any

from bson import BSON


EPICRISIS_CACHE_SCHEMA_VERSION = "v4"
EPICRISIS_CACHE_MAX_BSON_BYTES = 14 * 1024 * 1024

_SOURCE_KEYS = {
    "historia": "historia_clinica",
    "historias_adicionales": "historia_clinica",
    "quirurgico": "quirurgico",
    "factura": "factura",
    "radiologia": "radiologia",
    "laboratorio": "laboratorio",
    "generico": "generico",
    "prescripcion": "prescripcion",
}
_EVIDENCE_KEYS = {
    "evidence",
    "evidencias",
    "evidencias_pop",
    "evidencias_cirugia",
    "evidencias_historia",
    "resolution_evidence",
    "surgical_resolution_evidence",
}
_EVIDENCE_FRAGMENT_SIZE = 4096


@dataclass(frozen=True)
class EpicrisisCacheProjection:
    context: dict[str, Any]
    evidence_records: tuple[dict[str, Any], ...]
    source_references: tuple[dict[str, Any], ...]
    bson_bytes: int


def _json_default(value: Any) -> str:
    if isinstance(value, datetime):
        return value.isoformat()
    return str(value)


def _stable_json(value: Any) -> str:
    return json.dumps(value, ensure_ascii=False, sort_keys=True, default=_json_default, separators=(",", ":"))


def source_hash(document: dict[str, Any]) -> str:
    explicit = str(document.get("source_hash") or document.get("content_hash") or "").strip()
    if explicit:
        return explicit
    digest_payload = {
        key: value
        for key, value in document.items()
        if key not in {"_id", "fecha_analisis", "fecha_guardado"}
    }
    return hashlib.sha256(_stable_json(digest_payload).encode("utf-8")).hexdigest()


def _string_id(value: Any) -> str:
    return str(value or "").strip()


def _document_reference(document_type: str, document: dict[str, Any], *, case_key: str = "") -> dict[str, Any]:
    document_id = _string_id(document.get("_id") or document.get("document_id") or document.get("analysis_document_id"))
    reference = {
        "document_uid": _string_id(document.get("document_uid") or document_id),
        "document_id": document_id,
        "document_type": document_type,
        "filename": _string_id(document.get("nombre_archivo")),
        "source_hash": source_hash(document),
    }
    reference_case_key = _string_id(document.get("case_key") or case_key)
    if reference_case_key:
        reference["case_key"] = reference_case_key
    return reference


def _document_stub(document_type: str, document: dict[str, Any], *, case_key: str = "") -> dict[str, Any]:
    reference = _document_reference(document_type, document, case_key=case_key)
    stub: dict[str, Any] = {
        key: value
        for key, value in reference.items()
        if value
    }
    for key in ("_id", "document_id", "analysis_document_id", "nombre_archivo", "nombre_paciente", "case_number", "case_key", "tipo_documento", "fecha_analisis"):
        value = document.get(key)
        if value is not None and value != "":
            stub[key] = value.isoformat() if isinstance(value, datetime) else _string_id(value) if key in {"_id", "document_id", "analysis_document_id"} else value
    stub.setdefault("tipo_documento", document_type)
    return stub


def _evidence_identity(record: dict[str, Any]) -> str:
    occurrence_id = _string_id(record.get("occurrence_id"))
    if occurrence_id:
        return f"occurrence:{occurrence_id}"
    return _stable_json(
        {
            key: record.get(key)
            for key in ("canonical_id", "document_uid", "document_id", "document_type", "page", "section", "source_hash", "excerpt", "canonical_term", "original_term")
            if record.get(key) not in (None, "")
        }
    )


def _evidence_id(record: dict[str, Any]) -> str:
    return hashlib.sha256(_evidence_identity(record).encode("utf-8")).hexdigest()[:32]


def _collect_evidence(
    value: Any,
    *,
    records: dict[str, dict[str, Any]],
    artifact_type: str = "clinical_evidence",
) -> str | None:
    if not isinstance(value, dict):
        return None
    record = dict(value)
    evidence_id = _evidence_id(record)
    if evidence_id in records:
        return evidence_id
    record["evidence_id"] = evidence_id
    record.setdefault("artifact_type", artifact_type)
    fragment_refs: list[str] = []
    for field, field_value in list(record.items()):
        if field in {"evidence_id", "artifact_type"} or not isinstance(field_value, str):
            continue
        if len(field_value) <= _EVIDENCE_FRAGMENT_SIZE:
            continue
        fragments = [
            field_value[index : index + _EVIDENCE_FRAGMENT_SIZE]
            for index in range(0, len(field_value), _EVIDENCE_FRAGMENT_SIZE)
        ]
        record[field] = ""
        for index, fragment in enumerate(fragments):
            fragment_id = hashlib.sha256(
                f"{evidence_id}:{field}:{index}".encode()
            ).hexdigest()[:32]
            records[fragment_id] = {
                "evidence_id": fragment_id,
                "artifact_type": "evidence_fragment",
                "parent_evidence_id": evidence_id,
                "field": field,
                "fragment_index": index,
                "fragment_total": len(fragments),
                "text": fragment,
            }
            fragment_refs.append(fragment_id)
    if fragment_refs:
        record["_fragment_refs"] = fragment_refs
    records.setdefault(evidence_id, record)
    return evidence_id


def _project_value(value: Any, *, records: dict[str, dict[str, Any]], key: str = "") -> Any:
    if key in _EVIDENCE_KEYS and isinstance(value, list) and all(isinstance(item, dict) for item in value):
        return [
            evidence_id
            for item in value
            if (evidence_id := _collect_evidence(item, records=records, artifact_type=key))
        ]
    if isinstance(value, dict):
        return {child_key: _project_value(child, records=records, key=child_key) for child_key, child in value.items()}
    if isinstance(value, list):
        return [_project_value(child, records=records) for child in value]
    return value


def _project_source(value: Any, document_type: str, *, case_key: str) -> Any:
    if isinstance(value, dict):
        return _document_stub(document_type, value, case_key=case_key)
    if isinstance(value, list):
        return [
            _document_stub(document_type, item, case_key=case_key)
            for item in value
            if isinstance(item, dict)
        ]
    return [] if document_type in {"radiologia", "laboratorio", "generico", "prescripcion"} else None


def project_epicrisis_cache(
    context: dict[str, Any],
    *,
    case_key: str,
    schema_version: str = EPICRISIS_CACHE_SCHEMA_VERSION,
) -> EpicrisisCacheProjection:
    records: dict[str, dict[str, Any]] = {}
    projected = {key: _project_value(value, records=records, key=key) for key, value in dict(context or {}).items()}
    references: list[dict[str, Any]] = []
    for context_key, document_type in _SOURCE_KEYS.items():
        if context_key not in projected:
            continue
        raw = context.get(context_key)
        if isinstance(raw, dict):
            references.append(_document_reference(document_type, raw, case_key=case_key))
        elif isinstance(raw, list):
            references.extend(
                _document_reference(document_type, item, case_key=case_key)
                for item in raw
                if isinstance(item, dict)
            )
        projected[context_key] = _project_source(raw, document_type, case_key=case_key)

    previous_evidence_ids = [str(item) for item in projected.get("evidence_ref_ids") or [] if str(item)]
    canonical_processing = projected.get("canonical_processing")
    if isinstance(canonical_processing, dict):
        canonical_processing = dict(canonical_processing)
        raw_processing = context.get("canonical_processing")
        occurrences = raw_processing.get("occurrences") if isinstance(raw_processing, dict) else []
        incidents = raw_processing.get("incidents") if isinstance(raw_processing, dict) else []
        canonical_processing.pop("occurrences", None)
        canonical_processing.pop("incidents", None)
        canonical_processing["occurrence_refs"] = (
            [
                evidence_id
                for item in occurrences
                if isinstance(item, dict)
                if (evidence_id := _collect_evidence(item, records=records, artifact_type="canonical_occurrence"))
            ]
            if isinstance(occurrences, list) and any(isinstance(item, dict) for item in occurrences)
            else list(canonical_processing.get("occurrence_refs") or [])
        )
        canonical_processing["incident_refs"] = (
            [
                evidence_id
                for item in incidents
                if isinstance(item, dict)
                if (evidence_id := _collect_evidence(item, records=records, artifact_type="canonical_incident"))
            ]
            if isinstance(incidents, list) and any(isinstance(item, dict) for item in incidents)
            else list(canonical_processing.get("incident_refs") or [])
        )
        projected["canonical_processing"] = canonical_processing

    top_level_incidents = context.get("canonical_incidents")
    if isinstance(top_level_incidents, list) and any(isinstance(item, dict) for item in top_level_incidents):
        projected["canonical_incidents"] = [
            evidence_id
            for item in top_level_incidents if isinstance(item, dict)
            if (evidence_id := _collect_evidence(item, records=records, artifact_type="canonical_incident"))
        ]

    projected["schema_version"] = schema_version
    projected["cache_schema_version"] = schema_version
    projected["document_references"] = references
    projected["evidence_ref_ids"] = list(dict.fromkeys([*previous_evidence_ids, *records]))
    projected["cache_bson_bytes"] = BSON.encode({"contexto": projected}).__len__()
    bson_bytes = len(BSON.encode({"contexto": projected}))
    projected["cache_bson_bytes"] = bson_bytes
    return EpicrisisCacheProjection(
        context=projected,
        evidence_records=tuple(records.values()),
        source_references=tuple(references),
        bson_bytes=bson_bytes,
    )


def _hydrated_evidence_record(
    value: dict[str, Any],
    *,
    evidence_by_id: dict[str, dict[str, Any]],
) -> dict[str, Any]:
    hydrated = dict(value)
    fragment_ids = hydrated.pop("_fragment_refs", [])
    fragments = [
        evidence_by_id[item]
        for item in fragment_ids
        if item in evidence_by_id and isinstance(evidence_by_id[item], dict)
    ]
    for fragment in fragments:
        field = str(fragment.get("field") or "")
        if field:
            hydrated[field] = f"{hydrated.get(field) or ''}{fragment.get('text') or ''}"
    return {
        field: field_value
        for field, field_value in hydrated.items()
        if field not in {"evidence_id", "artifact_type"}
    }


def _hydrate_value(value: Any, *, evidence_by_id: dict[str, dict[str, Any]], key: str = "") -> Any:
    if key in _EVIDENCE_KEYS and isinstance(value, list) and all(isinstance(item, str) for item in value):
        return [_hydrated_evidence_record(evidence_by_id[item], evidence_by_id=evidence_by_id) for item in value if item in evidence_by_id]
    if isinstance(value, dict):
        return {child_key: _hydrate_value(child, evidence_by_id=evidence_by_id, key=child_key) for child_key, child in value.items()}
    if isinstance(value, list):
        return [_hydrate_value(child, evidence_by_id=evidence_by_id) for child in value]
    return value


def hydrate_epicrisis_context(
    context: dict[str, Any],
    *,
    documents: dict[str, Any] | None = None,
    evidence_by_id: dict[str, dict[str, Any]] | None = None,
) -> dict[str, Any]:
    hydrated = _hydrate_value(dict(context or {}), evidence_by_id=dict(evidence_by_id or {}))
    source_documents = documents or {}
    for key in _SOURCE_KEYS:
        if key not in hydrated or key not in source_documents:
            continue
        source = source_documents.get(key)
        if isinstance(source, (dict, list)):
            hydrated[key] = source

    canonical_processing = hydrated.get("canonical_processing")
    evidence_map = dict(evidence_by_id or {})
    if isinstance(canonical_processing, dict):
        canonical_processing = dict(canonical_processing)
        canonical_processing["occurrences"] = [
            _hydrated_evidence_record(evidence_map[item], evidence_by_id=evidence_map)
            for item in canonical_processing.pop("occurrence_refs", [])
            if item in evidence_map
        ]
        canonical_processing["incidents"] = [
            _hydrated_evidence_record(evidence_map[item], evidence_by_id=evidence_map)
            for item in canonical_processing.pop("incident_refs", [])
            if item in evidence_map
        ]
        hydrated["canonical_processing"] = canonical_processing
    if isinstance(hydrated.get("canonical_incidents"), list):
        hydrated["canonical_incidents"] = [
            _hydrated_evidence_record(evidence_map[item], evidence_by_id=evidence_map)
            if isinstance(item, str) and item in evidence_map
            else item
            for item in hydrated["canonical_incidents"]
        ]
    return hydrated


def bson_size(value: dict[str, Any]) -> int:
    return len(BSON.encode(value))


def compact_context_for_task(
    context: dict[str, Any],
    *,
    max_characters: int,
    text_fields: tuple[str, ...] = ("resumen_clinico_integrado", "hallazgos_quirurgicos", "descripcion_procedimiento"),
) -> tuple[str, dict[str, int]]:
    parts: list[str] = []
    for field in text_fields:
        value = context.get(field)
        if isinstance(value, dict):
            value = value.get("texto") or value.get("resumen")
        text = " ".join(str(value or "").split())
        if text:
            parts.append(f"{field}: {text}")
    prompt = "\n".join(parts)
    bounded = prompt[: max(0, max_characters)]
    return bounded, {
        "input_characters": len(prompt),
        "output_characters": len(bounded),
        "truncated_characters": max(0, len(prompt) - len(bounded)),
    }


def prepare_epicrisis_task_context(
    documents: dict[str, Any],
    *,
    task: str,
    max_characters: int,
    fragment_characters: int = 12000,
) -> tuple[str, dict[str, int]]:
    """Construye contexto acotado por tarea sin enviar el expediente completo."""
    labels = (
        ("historia", "historia_clinica"),
        ("quirurgico", "quirurgico"),
        ("factura", "factura"),
    )
    source_parts: list[str] = []
    input_characters = 0
    fragments = 0
    for key, label in labels:
        value = documents.get(key)
        if isinstance(value, dict):
            value = value.get("analisis_html") or value.get("descripcion") or value.get("texto") or ""
        text = " ".join(str(value or "").split())
        input_characters += len(text)
        if not text:
            continue
        for index in range(0, len(text), max(1, fragment_characters)):
            fragments += 1
            source_parts.append(
                f"[{label} fragmento {fragments}] {text[index:index + max(1, fragment_characters)]}"
            )
    prompt = f"tarea={task}\n" + "\n".join(source_parts)
    bounded = prompt[: max(0, max_characters)]
    return bounded, {
        "input_characters": input_characters,
        "output_characters": len(bounded),
        "fragments": fragments,
        "truncated_characters": max(0, len(prompt) - len(bounded)),
    }
