from __future__ import annotations

import re
import unicodedata
from collections.abc import Iterable
from dataclasses import dataclass
from decimal import Decimal, InvalidOperation
from typing import Any

from app.case_epicrisis.domain.models import (
    CostoFacturado,
    CostoFacturadoCoincidencia,
    CostoFacturadoOrigen,
)


COST_ORDER_VERSION = "soat_group_then_associated_surgical_cost_v4"

_CATEGORY_ALIASES = {
    "procedimientos_quirurgicos": "procedimiento",
    "procedimiento_quirurgico": "procedimiento",
    "procedimiento": "procedimiento",
    "procedimientos_no_quirurgicos": "procedimiento_no_quirurgico",
    "procedimiento_no_quirurgico": "procedimiento_no_quirurgico",
    "medicamentos": "medicamento",
    "medicamento": "medicamento",
    "imagenologia": "imagen",
    "imagen": "imagen",
    "radiologia": "imagen",
    "examenes_laboratorio": "laboratorio",
    "laboratorio": "laboratorio",
}
_SERVICE_SECTIONS = {
    "procedimientos_quirurgicos": "procedimiento",
    "procedimientos_no_quirurgicos": "procedimiento_no_quirurgico",
    "medicamentos": "medicamento",
    "imagenologia": "imagen",
    "examenes_laboratorio": "laboratorio",
}


def normalize_money(value: Any) -> Decimal | None:
    """Normaliza importes colombianos o decimales sin inferir valores inválidos."""
    if value is None or isinstance(value, bool):
        return None
    if isinstance(value, Decimal):
        result = value
    elif isinstance(value, (int, float)):
        try:
            result = Decimal(str(value))
        except InvalidOperation:
            return None
    else:
        text = str(value).strip()
        if not text or (text.startswith("(") and text.endswith(")")):
            return None
        text = re.sub(r"[^0-9,.-]", "", text)
        if not text or text.count("-") > 1 or ("-" in text and not text.startswith("-")):
            return None
        sign = "-" if text.startswith("-") else ""
        text = text.lstrip("-")
        if not text or not re.search(r"\d", text):
            return None
        separators = [position for position, char in enumerate(text) if char in ",."]
        if separators:
            last = separators[-1]
            decimal_digits = len(text) - last - 1
            decimal_separator = text[last] if decimal_digits in {1, 2} else ""
            if decimal_separator:
                if decimal_separator in text[:last]:
                    return None
                integer = re.sub(r"[,.]", "", text[:last]) or "0"
                fraction = re.sub(r"[,.]", "", text[last + 1 :])
                text = f"{integer}.{fraction}"
            else:
                text = re.sub(r"[,.]", "", text)
        try:
            result = Decimal(f"{sign}{text}")
        except InvalidOperation:
            return None
    if not result.is_finite() or result < 0:
        return None
    return result


def _normalize_text(value: Any) -> str:
    text = unicodedata.normalize("NFKD", str(value or ""))
    text = "".join(char for char in text if not unicodedata.combining(char)).casefold()
    return re.sub(r"[^a-z0-9]+", " ", text).strip()


def _normalize_code(value: Any) -> str:
    return re.sub(r"[^A-Z0-9]", "", str(value or "").upper())


def _first_value(item: dict[str, Any], *keys: str) -> Any:
    for key in keys:
        if item.get(key) not in (None, ""):
            return item[key]
    return None


def _line_cost(item: dict[str, Any]) -> tuple[Decimal | None, CostoFacturadoOrigen]:
    total = normalize_money(_first_value(item, "total", "tt", "valor_total"))
    if total is not None:
        return total, CostoFacturadoOrigen.TOTAL_FACTURA
    quantity = normalize_money(_first_value(item, "cantidad", "q", "dias"))
    unit_value = normalize_money(_first_value(item, "valor_unitario", "vu", "tarifa"))
    if quantity is None or unit_value is None:
        return None, CostoFacturadoOrigen.DESCONOCIDO
    return quantity * unit_value, CostoFacturadoOrigen.CANTIDAD_POR_VALOR_UNITARIO


@dataclass(frozen=True, slots=True)
class _InvoiceCostLine:
    identifier: str
    category: str
    codes: frozenset[str]
    texts: frozenset[str]
    value: Decimal | None
    origin: CostoFacturadoOrigen
    surgical_group: int | None

    @property
    def identity(self) -> tuple[str, ...]:
        return tuple(sorted(self.codes)) if self.codes else tuple(sorted(self.texts))


def _line_from_item(
    item: dict[str, Any],
    *,
    identifier: str,
    category: str,
    surgical_group: int | None = None,
) -> _InvoiceCostLine:
    codes = {
        _normalize_code(value)
        for value in (
            _first_value(item, "codigo_cups", "cc"),
            _first_value(item, "codigo_soat"),
            _first_value(item, "codigo_referencia", "cr", "codigo_medicamento"),
            _first_value(item, "codigo_facturacion", "cf", "codigo_interno", "codigo"),
        )
        if _normalize_code(value)
    }
    description = _first_value(
        item,
        "descripcion",
        "d",
        "medicamento",
        "nombre",
        "estudio",
        "prueba",
        "concepto",
    )
    presentation = _first_value(item, "presentacion")
    texts = {_normalize_text(description)} if _normalize_text(description) else set()
    combined = _normalize_text(" ".join(str(value) for value in (description, presentation) if value))
    if combined:
        texts.add(combined)
    associated_cost = item.get("costo_quirurgico_asociado")
    if category == "procedimiento" and isinstance(associated_cost, dict):
        value = normalize_money(associated_cost.get("valor_final"))
        origin = CostoFacturadoOrigen.TOTAL_FACTURA
    else:
        value, origin = _line_cost(item)
    return _InvoiceCostLine(
        identifier,
        category,
        frozenset(codes),
        frozenset(texts),
        value,
        origin,
        surgical_group,
    )


def _valid_surgical_group(value: Any) -> int | None:
    if isinstance(value, bool):
        return None
    try:
        group = int(value)
    except (TypeError, ValueError):
        return None
    if str(value).strip() != str(group) or not 2 <= group <= 23:
        return None
    return group


def _official_surgical_group(item: dict[str, Any]) -> int | None:
    valuation = item.get("valoracion_soat")
    if not isinstance(valuation, dict):
        return None
    rule = str(valuation.get("regla_aplicada") or "").strip()
    if rule in {
        "codigo_soat_no_encontrado",
        "crosswalk_ambiguo",
        "crosswalk_no_mapeado",
        "crosswalk_vigencia_no_disponible",
        "vigencia_no_determinada",
        "vigencia_no_disponible",
    }:
        return None
    crosswalk = valuation.get("cruce_cups_soat")
    explicit_valid = str(valuation.get("fuente_asignacion") or "") == "codigo_soat_explicito"
    if (
        not explicit_valid
        and isinstance(crosswalk, dict)
        and str(crosswalk.get("status") or "")
        in {
            "ambiguous",
            "unmapped",
            "unsupported_year",
        }
    ):
        return None
    return _valid_surgical_group(item.get("grupo_quirurgico_soat"))


def _canonical_procedure_groups(services: dict[str, Any], canonical_lines: list[Any]) -> dict[int, int]:
    candidates: dict[int, set[int]] = {}
    for item in services.get("procedimientos_quirurgicos") or []:
        if not isinstance(item, dict):
            continue
        valuation = item.get("valoracion_soat")
        group = _official_surgical_group(item)
        if group is None or not isinstance(valuation, dict):
            continue
        index = valuation.get("indice_linea_canonica")
        if isinstance(index, bool) or not isinstance(index, int) or not 0 <= index < len(canonical_lines):
            continue
        candidates.setdefault(index, set()).add(group)
    return {index: next(iter(groups)) for index, groups in candidates.items() if len(groups) == 1}


def _canonical_procedure_costs(services: dict[str, Any], canonical_lines: list[Any]) -> dict[int, Decimal]:
    result: dict[int, Decimal] = {}
    for item in services.get("procedimientos_quirurgicos") or []:
        if not isinstance(item, dict):
            continue
        valuation = item.get("valoracion_soat")
        cost = item.get("costo_quirurgico_asociado")
        if not isinstance(valuation, dict) or not isinstance(cost, dict):
            continue
        index = valuation.get("indice_linea_canonica")
        value = normalize_money(cost.get("valor_final"))
        if (
            isinstance(index, int)
            and not isinstance(index, bool)
            and 0 <= index < len(canonical_lines)
            and value is not None
        ):
            result[index] = value
    return result


def extract_invoice_cost_lines(factura: Any) -> list[_InvoiceCostLine]:
    if not isinstance(factura, dict):
        return []
    factura_json = factura.get("factura_json")
    if not isinstance(factura_json, dict):
        return []
    result: list[_InvoiceCostLine] = []
    canonical_lines = factura_json.get("lineas_canonicas") or []
    services = factura_json.get("servicios_procedimientos")
    services = services if isinstance(services, dict) else {}
    canonical_groups = _canonical_procedure_groups(services, canonical_lines)
    canonical_costs = _canonical_procedure_costs(services, canonical_lines)
    canonical_categories: set[str] = set()
    for index, raw in enumerate(canonical_lines):
        if not isinstance(raw, dict):
            continue
        raw_category = _normalize_text(_first_value(raw, "categoria", "ct")).replace(" ", "_")
        category = _CATEGORY_ALIASES.get(raw_category, "")
        if not category:
            continue
        canonical_categories.add(category)
        identifier = str(_first_value(raw, "linea_factura_id", "linea_id", "id", "_id") or "").strip()
        cost_line = _line_from_item(
            raw,
            identifier=identifier or f"lineas_canonicas:{index}",
            category=category,
            surgical_group=(canonical_groups.get(index) if category == "procedimiento" else None),
        )
        if category == "procedimiento" and index in canonical_costs:
            cost_line = _InvoiceCostLine(
                cost_line.identifier,
                cost_line.category,
                cost_line.codes,
                cost_line.texts,
                canonical_costs[index],
                CostoFacturadoOrigen.TOTAL_FACTURA,
                cost_line.surgical_group,
            )
        result.append(cost_line)
    if not services:
        return result
    for section, category in _SERVICE_SECTIONS.items():
        if category in canonical_categories:
            continue
        for index, raw in enumerate(services.get(section) or []):
            if not isinstance(raw, dict):
                continue
            identifier = str(_first_value(raw, "linea_factura_id", "linea_id", "id", "_id") or "").strip()
            result.append(
                _line_from_item(
                    raw,
                    identifier=identifier or f"servicios_procedimientos:{section}:{index}",
                    category=category,
                    surgical_group=(
                        _official_surgical_group(raw)
                        if category == "procedimiento" and not canonical_lines
                        else None
                    ),
                )
            )
    return result


def _artifact_codes(item: Any, category: str) -> set[str]:
    if not isinstance(item, dict):
        match = re.match(
            r"^\s*((?=[A-Za-z0-9.-]*\d)[A-Za-z0-9.-]{3,})\s*(?:[-:]|\s)\s*(.+)$",
            str(item or ""),
        )
        return {_normalize_code(match.group(1))} if match and _normalize_code(match.group(1)) else set()
    if category == "procedimiento":
        values = (
            item.get("codigo_cups"),
            item.get("codigo_soat"),
            item.get("codigo_facturacion"),
            item.get("codigo"),
        )
    elif category == "medicamento":
        values = (
            item.get("codigo_referencia"),
            item.get("codigo_facturacion"),
            item.get("codigo_medicamento"),
            item.get("codigo"),
        )
    else:
        values = (
            item.get("codigo"),
            item.get("codigo_cups"),
            item.get("codigo_referencia"),
            item.get("codigo_facturacion"),
        )
    return {_normalize_code(value) for value in values if _normalize_code(value)}


def _artifact_texts(item: Any, category: str) -> set[str]:
    if not isinstance(item, dict):
        text = str(item or "")
        text = re.sub(
            r"^\s*(?=[A-Za-z0-9.-]*\d)[A-Za-z0-9.-]{3,}\s*(?:[-:]|\s)\s*",
            "",
            text,
        )
        return {_normalize_text(text)} if _normalize_text(text) else set()
    if category == "procedimiento":
        values = (
            item.get("descripcion"),
            item.get("description"),
            item.get("procedimiento"),
            item.get("concepto"),
        )
    elif category == "medicamento":
        name = item.get("nombre") or item.get("medicamento") or item.get("descripcion")
        presentation = item.get("presentacion")
        values = (name, " ".join(str(value) for value in (name, presentation) if value))
    else:
        name = item.get("nombre") or item.get("descripcion")
        values = (
            name,
            re.sub(r"^\s*[A-Za-z0-9.-]{2,}\s*[-:]\s*", "", str(name or "")),
            item.get("descripcion"),
            item.get("estudio"),
            item.get("prueba"),
        )
    return {_normalize_text(value) for value in values if _normalize_text(value)}


def _unknown_cost(match: CostoFacturadoCoincidencia) -> CostoFacturado:
    return CostoFacturado(coincidencia=match)


def match_billed_cost(
    item: Any,
    *,
    category: str,
    lines: Iterable[_InvoiceCostLine],
) -> CostoFacturado:
    relevant = [line for line in lines if line.category == category]
    if isinstance(item, dict):
        explicit_id = str(item.get("linea_factura_id") or "").strip()
        direct = [line for line in relevant if explicit_id and line.identifier == explicit_id]
        if direct:
            selected = max(direct, key=lambda line: line.value if line.value is not None else Decimal("-1"))
            return _cost_from_line(selected, CostoFacturadoCoincidencia.DIRECTA)
        own_value, own_origin = _line_cost(item)
        if explicit_id and own_value is not None:
            return CostoFacturado(
                valor_total=own_value,
                origen=own_origin,
                coincidencia=CostoFacturadoCoincidencia.DIRECTA,
                linea_factura_id=explicit_id,
            )

    codes = _artifact_codes(item, category)
    code_matches = [line for line in relevant if codes & set(line.codes)]
    if code_matches:
        selected = _max_known_line(code_matches)
        return (
            _cost_from_line(selected, CostoFacturadoCoincidencia.CODIGO_EXACTO)
            if selected
            else _unknown_cost(CostoFacturadoCoincidencia.CODIGO_EXACTO)
        )

    if category in {"procedimiento", "procedimiento_no_quirurgico"}:
        return _unknown_cost(CostoFacturadoCoincidencia.INEXISTENTE)

    texts = _artifact_texts(item, category)
    text_matches = [line for line in relevant if texts & set(line.texts)]
    coded_identities = {tuple(sorted(line.codes)) for line in text_matches if line.codes}
    if len(coded_identities) > 1:
        return _unknown_cost(CostoFacturadoCoincidencia.AMBIGUA)
    if text_matches:
        selected = _max_known_line(text_matches)
        return (
            _cost_from_line(selected, CostoFacturadoCoincidencia.TEXTO_NORMALIZADO_UNICO)
            if selected
            else _unknown_cost(CostoFacturadoCoincidencia.TEXTO_NORMALIZADO_UNICO)
        )
    return _unknown_cost(CostoFacturadoCoincidencia.INEXISTENTE)


def _max_known_line(lines: Iterable[_InvoiceCostLine]) -> _InvoiceCostLine | None:
    known = [line for line in lines if line.value is not None]
    return max(known, key=lambda line: line.value) if known else None


def _cost_from_line(line: _InvoiceCostLine, match: CostoFacturadoCoincidencia) -> CostoFacturado:
    return CostoFacturado(
        valor_total=line.value,
        origen=line.origin,
        coincidencia=match,
        linea_factura_id=line.identifier,
    )


def _unique_surgical_group(lines: Iterable[_InvoiceCostLine]) -> int | None:
    groups = {line.surgical_group for line in lines if line.surgical_group is not None}
    return next(iter(groups)) if len(groups) == 1 else None


def match_billed_surgical_group(item: Any, *, lines: Iterable[_InvoiceCostLine]) -> int | None:
    relevant = [line for line in lines if line.category == "procedimiento"]
    if isinstance(item, dict):
        explicit_id = str(item.get("linea_factura_id") or "").strip()
        direct = [line for line in relevant if explicit_id and line.identifier == explicit_id]
        if direct:
            return _unique_surgical_group(direct)
        if explicit_id:
            return None

    codes = _artifact_codes(item, "procedimiento")
    code_matches = [line for line in relevant if codes & set(line.codes)]
    if code_matches:
        return _unique_surgical_group(code_matches)

    # El grupo quirúrgico nunca se hereda por descripción: la llave es el CUPS exacto.
    return None


def attach_and_sort_costs(items: Any, *, category: str, lines: Iterable[_InvoiceCostLine]) -> list[Any]:
    lines = list(lines)
    result: list[Any] = []
    for raw in items if isinstance(items, list) else []:
        if isinstance(raw, dict):
            item = dict(raw)
            item["costo_facturado"] = match_billed_cost(item, category=category, lines=lines).model_dump(
                mode="json"
            )
            if category == "procedimiento":
                item["grupo_quirurgico_soat"] = match_billed_surgical_group(item, lines=lines)
            result.append(item)
        else:
            result.append(raw)

    if category == "procedimiento":
        return sort_procedures_by_soat_group_and_cost(result)
    return sort_by_billed_cost(result)


def sort_by_billed_cost(items: Iterable[Any]) -> list[Any]:
    def key(item: Any) -> tuple[bool, Decimal]:
        cost = item.get("costo_facturado") if isinstance(item, dict) else None
        value = normalize_money(cost.get("valor_total")) if isinstance(cost, dict) else None
        return value is None, -(value or Decimal(0))

    return sorted(items, key=key)


def sort_procedures_by_soat_group_and_cost(items: Iterable[Any]) -> list[Any]:
    def key(item: Any) -> tuple[bool, int, bool, Decimal]:
        group = _valid_surgical_group(item.get("grupo_quirurgico_soat")) if isinstance(item, dict) else None
        cost = item.get("costo_facturado") if isinstance(item, dict) else None
        value = normalize_money(cost.get("valor_total")) if isinstance(cost, dict) else None
        return group is None, -(group or 0), value is None, -(value or Decimal(0))

    return sorted(items, key=key)


def sort_publishable_procedures(items: Iterable[Any], *, lines: Iterable[_InvoiceCostLine]) -> list[Any]:
    """Orden contractual: quirúrgicos por grupo/costo y no quirúrgicos por total."""
    source = list(items)
    invoice_lines = list(lines)
    surgical = [
        item for item in source if isinstance(item, dict) and item.get("clasificacion") == "quirurgico"
    ]
    non_surgical = [
        item
        for item in source
        if isinstance(item, dict) and item.get("clasificacion") == "no_quirurgico"
    ]
    ordered_surgical = attach_and_sort_costs(surgical, category="procedimiento", lines=invoice_lines)
    ordered_non_surgical = attach_and_sort_costs(
        non_surgical,
        category="procedimiento_no_quirurgico",
        lines=invoice_lines,
    )
    return [*ordered_surgical, *ordered_non_surgical]


def sort_artifacts_by_cost(
    items: Iterable[Any], *, category: str, lines: Iterable[_InvoiceCostLine]
) -> list[Any]:
    lines = list(lines)
    decorated = [
        (
            item,
            match_billed_cost(item, category=category, lines=lines),
            match_billed_surgical_group(item, lines=lines) if category == "procedimiento" else None,
        )
        for item in items
    ]
    return [
        item
        for item, _cost, _group in sorted(
            decorated,
            key=lambda pair: (
                pair[2] is None if category == "procedimiento" else False,
                -(pair[2] or 0) if category == "procedimiento" else 0,
                pair[1].valor_total is None,
                -(pair[1].valor_total or Decimal(0)),
            ),
        )
    ]


def apply_factura_cost_ordering(context: dict[str, Any]) -> dict[str, Any]:
    """Adjunta costo interno y ordena colecciones sin cambiar sus proyecciones visibles."""
    lines = extract_invoice_cost_lines(context.get("factura"))
    if isinstance(context.get("medicamentos_caso"), list):
        context["medicamentos_caso"] = attach_and_sort_costs(
            context["medicamentos_caso"], category="medicamento", lines=lines
        )
    if isinstance(context.get("ayudas_diagnosticas"), list):
        # Imagen y laboratorio comparten la colección, pero se asocian contra su categoría financiera.
        ordered_aids = []
        for item in context["ayudas_diagnosticas"]:
            if not isinstance(item, dict):
                ordered_aids.append(item)
                continue
            aid_category = "laboratorio" if str(item.get("tipo") or "") == "laboratorio" else "imagen"
            enriched = dict(item)
            enriched["costo_facturado"] = match_billed_cost(
                enriched, category=aid_category, lines=lines
            ).model_dump(mode="json")
            ordered_aids.append(enriched)
        context["ayudas_diagnosticas"] = sort_by_billed_cost(ordered_aids)
    if isinstance(context.get("procedimientos_clinicos"), list):
        ordered = attach_and_sort_costs(
            context["procedimientos_clinicos"], category="procedimiento", lines=lines
        )
        context["procedimientos_clinicos"] = [_project_clinical_procedure(item) for item in ordered]
    if isinstance(context.get("procedimientos_curados"), list):
        context["procedimientos_curados"] = attach_and_sort_costs(
            context["procedimientos_curados"], category="procedimiento", lines=lines
        )
    if isinstance(context.get("pdf_primary_procedimientos"), list):
        context["pdf_primary_procedimientos"] = sort_artifacts_by_cost(
            context["pdf_primary_procedimientos"], category="procedimiento", lines=lines
        )
    return context


def _project_clinical_procedure(item: Any) -> Any:
    if not isinstance(item, dict):
        return item
    return {
        key: value
        for key, value in item.items()
        if key
        not in {
            "codigo_soat",
            "codigo_facturacion",
            "codigo_referencia",
            "grupo_quirurgico_soat",
            "costo_facturado",
            "costo_quirurgico_asociado",
            "valoracion_soat",
            "formula",
            "evidencia",
        }
    }
