from __future__ import annotations

import re
from dataclasses import dataclass, field
from decimal import ROUND_DOWN, Decimal, InvalidOperation
from typing import Any

from app.soat_tariffs.domain.surgical_costs import (
    ComponenteCostoQuirurgico,
    CostoQuirurgicoAsociado,
    EstadoAsociacionCosto,
)


_COMPONENT_NAMES = ("cirujano", "anestesia", "ayudantia", "sala", "materiales")


def _money(value: Any) -> Decimal | None:
    if value is None or isinstance(value, bool):
        return None
    text = re.sub(r"[^0-9,.-]", "", str(value).strip())
    if not text or text.startswith("-"):
        return None
    separators = [index for index, char in enumerate(text) if char in ",."]
    if separators:
        last = separators[-1]
        decimal_separator = text[last] if len(text) - last - 1 in {1, 2} else ""
        if decimal_separator:
            text = f"{re.sub(r'[,.]', '', text[:last]) or '0'}.{re.sub(r'[,.]', '', text[last + 1 :])}"
        else:
            text = re.sub(r"[,.]", "", text)
    try:
        result = Decimal(text)
    except InvalidOperation:
        return None
    return result if result.is_finite() and result >= 0 else None


def _line_value(line: dict[str, Any]) -> Decimal | None:
    total = _money(line.get("total"))
    if total is not None:
        return total
    quantity = _money(line.get("cantidad"))
    unit = _money(line.get("valor_unitario"))
    return quantity * unit if quantity is not None and unit is not None else None


def _cents(value: Decimal) -> int:
    return int((value * 100).quantize(Decimal("1")))


def _from_cents(value: int) -> Decimal:
    return (Decimal(value) / 100).quantize(Decimal("0.01"))


@dataclass
class _ProcedureAllocation:
    index: int
    item: dict[str, Any]
    direct: Decimal
    direct_line_index: int | None
    expected: dict[str, Decimal]
    codes: dict[str, str]
    allocated: dict[str, int] = field(default_factory=dict)
    line_indices: dict[str, list[int]] = field(default_factory=dict)
    proportional: bool = False
    duplicate_direct_line: bool = False


def _allocate_largest_remainder(total_cents: int, weights: list[tuple[int, Decimal]]) -> dict[int, int]:
    weight_total = sum((weight for _, weight in weights), Decimal("0"))
    if total_cents <= 0 or weight_total <= 0:
        return {}
    quotas = [(index, Decimal(total_cents) * weight / weight_total) for index, weight in weights]
    assigned = {index: int(quota.to_integral_value(rounding=ROUND_DOWN)) for index, quota in quotas}
    remainder = total_cents - sum(assigned.values())
    order = sorted(quotas, key=lambda pair: (-(pair[1] - assigned[pair[0]]), pair[0]))
    for index, _quota in order[:remainder]:
        assigned[index] += 1
    return assigned


def associate_surgical_costs(factura: dict[str, Any]) -> dict[str, Any]:
    """Asocia cargos quirúrgicos sin modificar líneas ni importes originales."""
    result = dict(factura)
    services = dict(result.get("servicios_procedimientos") or {})
    procedures = [dict(item) for item in services.get("procedimientos_quirurgicos") or []]
    lines = [dict(item) for item in result.get("lineas_canonicas") or []]
    allocations: list[_ProcedureAllocation] = []
    procedure_line_indices: set[int] = set()
    claimed_direct_indices: set[int] = set()

    for index, item in enumerate(procedures):
        valuation = item.get("valoracion_soat") if isinstance(item.get("valoracion_soat"), dict) else {}
        canonical_index = valuation.get("indice_linea_canonica")
        valid_canonical_index = (
            canonical_index
            if isinstance(canonical_index, int)
            and not isinstance(canonical_index, bool)
            and 0 <= canonical_index < len(lines)
            else None
        )
        direct_line = lines[canonical_index] if valid_canonical_index is not None else item
        duplicate_direct_line = valid_canonical_index in claimed_direct_indices
        direct = Decimal("0") if duplicate_direct_line else (_line_value(direct_line) or Decimal("0"))
        if valid_canonical_index is not None:
            procedure_line_indices.add(valid_canonical_index)
            claimed_direct_indices.add(valid_canonical_index)
        expected = {
            name: _money(value) or Decimal("0")
            for name, value in (valuation.get("componentes_liquidados") or {}).items()
            if name in _COMPONENT_NAMES
        }
        codes = {
            str(name): str(code)
            for name, code in (valuation.get("codigos_componentes") or {}).items()
            if name in _COMPONENT_NAMES and str(code or "").strip()
        }
        allocations.append(
            _ProcedureAllocation(
                index,
                item,
                direct,
                valid_canonical_index if not duplicate_direct_line else None,
                expected,
                codes,
                duplicate_direct_line=duplicate_direct_line,
            )
        )

    unmatched_component_lines: list[tuple[int, str, Decimal]] = []
    for line_index, line in enumerate(lines):
        if line_index in procedure_line_indices:
            continue
        code = str(line.get("codigo_facturacion") or "").strip()
        value = _line_value(line)
        if not code or value is None:
            continue
        line_date = str(line.get("fecha_servicio") or "").strip()
        candidates: list[tuple[_ProcedureAllocation, str, Decimal]] = []
        for allocation in allocations:
            procedure_date = str(allocation.item.get("fecha_servicio") or "").strip()
            if line_date and procedure_date and line_date != procedure_date:
                continue
            for component, component_code in allocation.codes.items():
                weight = allocation.expected.get(component, Decimal("0"))
                if component_code == code and weight > 0:
                    candidates.append((allocation, component, weight))
        if not candidates:
            if re.fullmatch(r"39[0-3]\d{2}", code) and value > 0:
                unmatched_component_lines.append((line_index, line_date, value))
            continue
        shares = _allocate_largest_remainder(
            _cents(value), [(allocation.index, weight) for allocation, _component, weight in candidates]
        )
        shared = len(candidates) > 1
        for allocation, component, _weight in candidates:
            allocation.allocated[component] = (
                allocation.allocated.get(component, 0) + shares[allocation.index]
            )
            allocation.line_indices.setdefault(component, []).append(line_index)
            allocation.proportional = allocation.proportional or shared

    for allocation in allocations:
        components: list[ComponenteCostoQuirurgico] = []
        missing = []
        for component, expected in allocation.expected.items():
            assigned_cents = allocation.allocated.get(component, 0)
            indices = allocation.line_indices.get(component, [])
            if expected > 0 and not indices:
                missing.append(component)
            method = (
                "proporcional_peso_tarifario"
                if allocation.proportional and indices
                else "directo_codigo_oficial"
            )
            if not indices:
                method = "sin_asignacion"
            components.append(
                ComponenteCostoQuirurgico(
                    componente=component,
                    valor_cobrado=_from_cents(assigned_cents),
                    indices_lineas_canonicas=indices,
                    codigos_facturacion=[allocation.codes[component]]
                    if component in allocation.codes
                    else [],
                    conceptos_lineas=[
                        str(lines[line_index].get("descripcion") or "").strip()
                        for line_index in indices
                        if str(lines[line_index].get("descripcion") or "").strip()
                    ],
                    metodo=method,
                    proporcion=(
                        (
                            _from_cents(assigned_cents)
                            / sum(
                                (_line_value(lines[line_index]) or Decimal("0") for line_index in indices),
                                Decimal("0"),
                            )
                        ).quantize(Decimal("0.000001"))
                        if assigned_cents
                        and sum(
                            (_line_value(lines[line_index]) or Decimal("0") for line_index in indices),
                            Decimal("0"),
                        )
                        > 0
                        else None
                    ),
                    formula=(
                        "cargo × peso tarifario liquidado / suma de pesos elegibles"
                        if method == "proporcional_peso_tarifario"
                        else (
                            "cargo de línea canónica por código oficial"
                            if method == "directo_codigo_oficial"
                            else "sin cargo facturado elegible"
                        )
                    ),
                    redondeo="centavos por mayor residuo; desempate por orden canónico"
                    if method == "proporcional_peso_tarifario"
                    else None,
                    evidencia=[
                        f"Línea canónica {line_index}; código {allocation.codes.get(component, 'N/D')}."
                        for line_index in indices
                    ],
                )
            )
        component_total = sum((component.valor_cobrado for component in components), Decimal("0"))
        final = allocation.direct + component_total
        procedure_date = str(allocation.item.get("fecha_servicio") or "").strip()
        unmatched_for_procedure = []
        if not any(
            str(previous.item.get("fecha_servicio") or "").strip() == procedure_date
            for previous in allocations[: allocation.index]
        ):
            unmatched_for_procedure = [
                (line_index, value)
                for line_index, line_date, value in unmatched_component_lines
                if not line_date or not procedure_date or line_date == procedure_date
            ]
        unmatched_balance = sum((value for _line_index, value in unmatched_for_procedure), Decimal("0"))
        if missing or allocation.duplicate_direct_line or unmatched_balance:
            state = EstadoAsociacionCosto.INCOMPLETA
        elif allocation.proportional:
            state = EstadoAsociacionCosto.PROPORCIONAL
        else:
            state = EstadoAsociacionCosto.DIRECTA
        all_indices = sorted(
            ({allocation.direct_line_index} if allocation.direct_line_index is not None else set())
            | {line_index for indices in allocation.line_indices.values() for line_index in indices}
        )
        cost = CostoQuirurgicoAsociado(
            estado=state,
            valor_directo=allocation.direct,
            componentes=components,
            valor_final=final,
            indices_lineas_canonicas=all_indices,
            codigos_facturacion=list(dict.fromkeys(allocation.codes.values())),
            metodo="asignacion_proporcional" if allocation.proportional else "asignacion_directa",
            formula="valor directo + suma de componentes cobrados asignados",
            redondeo="centavos por mayor residuo; desempate por orden canónico"
            if allocation.proportional
            else None,
            saldo_no_asignado=unmatched_balance,
            indices_saldo_no_asignado=[line_index for line_index, _value in unmatched_for_procedure],
            evidencia=(
                ([f"Componentes sin línea facturada elegible: {', '.join(missing)}."] if missing else [])
                + (
                    ["La línea directa ya fue consumida por otro procedimiento canónico."]
                    if allocation.duplicate_direct_line
                    else []
                )
                + (
                    [f"Cargos quirúrgicos sin candidato elegible: {unmatched_balance:.2f}."]
                    if unmatched_balance
                    else []
                )
            ),
        )
        procedures[allocation.index]["costo_quirurgico_asociado"] = cost.model_dump(mode="json")

    services["procedimientos_quirurgicos"] = procedures
    result["servicios_procedimientos"] = services
    return result
