"""Deterministic reconciliation of observed invoice values.

This module never invents or changes extracted monetary values.  It reports
whether the observed rows and their evidence support the document arithmetic.
"""

from __future__ import annotations

import re
from decimal import Decimal, InvalidOperation
from typing import Any


TOLERANCE = Decimal("0.01")
TAX_CODE_RE = re.compile(r"^[A-Z]$")
PERCENT_RE = re.compile(r"\b\d+(?:\.\d+)?\s*%")


def reconcile_invoice(
    *, items: list[dict[str, Any]], adjustments: list[dict[str, Any]],
    tax_summary: dict[str, Any], subtotal: Any, total: Any,
    observed_item_line_count: int | None = None,
) -> dict[str, Any]:
    item_amounts = [_decimal(item.get("line_total")) for item in items]
    known_item_amounts = [value for value in item_amounts if value is not None]
    items_sum = sum(known_item_amounts, Decimal("0")) if known_item_amounts else None
    subtotal_value = _decimal(subtotal)
    total_value = _decimal(total)

    normalized_adjustments, adjustment_issues = normalize_adjustment_hierarchy(adjustments)
    adjustment_values = [
        _decimal(value.get("amount")) for value in normalized_adjustments
        if value.get("participates_in_total") is True
    ]
    adjustment_sum = sum((value for value in adjustment_values if value is not None), Decimal("0"))

    item_count_status = _status(
        Decimal(len(items)), Decimal(observed_item_line_count)
    ) if observed_item_line_count is not None else "not_available"
    subtotal_status = _status(items_sum, subtotal_value)
    total_status = _status(
        subtotal_value + adjustment_sum if subtotal_value is not None else None,
        total_value,
    )

    line_checks = []
    evidence_failures = 0
    suspicious_rows = []
    for index, item in enumerate(items):
        quantity = _decimal(item.get("quantity"))
        unit_price = _decimal(item.get("unit_price"))
        line_total = _decimal(item.get("line_total"))
        arithmetic = _status(
            quantity * unit_price if quantity is not None and unit_price is not None else None,
            line_total,
        )
        evidence = item.get("evidence") if isinstance(item.get("evidence"), dict) else {}
        evidence_valid = _money_evidence_valid(evidence.get("line_total"))
        evidence_failures += 0 if evidence_valid else 1
        basis = str(item.get("unit_price_basis") or "").strip()
        tax_code = str(item.get("tax_code") or "").strip()
        suspicious = []
        if TAX_CODE_RE.fullmatch(basis):
            suspicious.append("tax_code_used_as_unit_price_basis")
        package = item.get("package_size") if isinstance(item.get("package_size"), dict) else {}
        package_unit = str(package.get("unit") or "").strip().lower()
        if package_unit in {"l", "litre", "litres", "liter", "liters"} and str(item.get("transaction_unit") or "").lower() in {"g", "gram", "grams"}:
            suspicious.append("litres_converted_to_grams")
        if suspicious:
            suspicious_rows.append({"item_index": index, "issues": suspicious})
        line_checks.append({
            "item_index": index,
            "arithmetic_valid": arithmetic,
            "source_evidence_valid": "passed" if evidence_valid else "failed",
            "issues": suspicious,
        })

    item_rows_status = "passed" if (
        item_count_status == "passed" and subtotal_status == "passed"
        and evidence_failures == 0 and not suspicious_rows
    ) else "failed"
    tax_validation = _reconcile_tax(tax_summary, items, item_count_status, item_rows_status)
    tax_status = tax_validation["status"]
    source_evidence_status = "passed" if (
        items and not evidence_failures
        and item_count_status == "passed"
        and subtotal_status == "passed"
        and not suspicious_rows
    ) else "failed"
    issues = list(adjustment_issues)
    if item_count_status == "failed":
        issues.append("item_line_count_mismatch")
    if subtotal_status == "failed":
        issues.append("subtotal_not_reconciled")
    if total_status == "failed":
        issues.append("adjustments_not_reconciled")
    if tax_status == "failed":
        issues.append("tax_not_reconciled")
    if source_evidence_status == "failed":
        issues.append("source_evidence_incomplete")
    if suspicious_rows:
        issues.append("suspicious_extraction")

    return {
        "observed_item_line_count": observed_item_line_count,
        "extracted_item_count": len(items),
        "items_sum": _float(items_sum),
        "adjustments_sum": _float(adjustment_sum),
        "reconciled_adjustments": normalized_adjustments,
        "item_count_reconciled": item_count_status,
        "subtotal_reconciled": subtotal_status,
        "adjustments_reconciled": total_status,
        "tax_reconciled": tax_status,
        "tax_validation": tax_validation,
        "source_evidence_valid": source_evidence_status,
        "line_checks": line_checks,
        "suspicious_rows": suspicious_rows,
        "issues": list(dict.fromkeys(issues)),
    }


def evidence(raw_text: Any, source_line: Any, confidence: Any, origin: str = "observed", bounding_box=None) -> dict[str, Any]:
    box = bounding_box if isinstance(bounding_box, list) and len(bounding_box) == 4 else None
    return {
        "raw_text": None if raw_text is None else str(raw_text),
        "source_line": None if source_line is None else str(source_line),
        "bounding_box": box,
        "bounding_box_origin": "model_estimated" if box is not None else "not_available",
        "confidence": _confidence(confidence),
        "origin": origin if origin in {"observed", "derived", "inferred"} else "inferred",
    }


def normalize_adjustment_hierarchy(values: list[dict[str, Any]]) -> tuple[list[dict[str, Any]], list[str]]:
    """Select payable parent adjustments and exclude tax-category components."""
    accepted: list[dict[str, Any]] = []
    issues = []
    surcharge_components: list[dict[str, Any]] = []
    surcharge_parents: list[dict[str, Any]] = []
    for value in values:
        source = str(value.get("source_text") or "")
        amount = _decimal(value.get("amount"))
        if PERCENT_RE.search(source) and amount is not None and "$" not in source and "amount" not in source.lower():
            issues.append("percentage_interpreted_as_money")
            continue
        decorated = dict(value)
        if decorated.get("reconciliation_role") == "component":
            decorated["participates_in_total"] = False
            surcharge_components.append(decorated)
            continue
        if decorated.get("reconciliation_role") == "parent" and value.get("type") == "payment_surcharge":
            decorated["participates_in_total"] = True
            surcharge_parents.append(decorated)
            continue
        if value.get("type") == "payment_surcharge":
            if re.search(r"(?:^|\s)[A-Z]\s*$", source.strip()) and not re.search(r"\btotal\b", source, re.I):
                decorated["reconciliation_role"] = "component"
                decorated["participates_in_total"] = False
                surcharge_components.append(decorated)
            else:
                decorated["reconciliation_role"] = "parent"
                decorated["participates_in_total"] = True
                surcharge_parents.append(decorated)
            continue
        decorated["reconciliation_role"] = "parent"
        decorated["participates_in_total"] = True
        accepted.append(decorated)
    if surcharge_parents:
        accepted.append(max(surcharge_parents, key=lambda row: abs(_decimal(row.get("amount")) or Decimal("0"))))
    elif surcharge_components:
        derived = sum((_decimal(row.get("amount")) or Decimal("0") for row in surcharge_components), Decimal("0"))
        accepted.append({
            "type": "payment_surcharge", "description": "Surcharge total derived from tax-category components",
            "amount": _float(derived), "source_text": " + ".join(str(row.get("source_text") or "") for row in surcharge_components),
            "confidence": min(float(row.get("confidence") or 0) for row in surcharge_components),
            "origin": "derived", "reconciliation_role": "parent", "participates_in_total": True,
        })
    accepted.extend({**row, "reconciliation_role": "component", "participates_in_total": False} for row in surcharge_components)
    return accepted, issues


def _reconcile_tax(summary: dict[str, Any], items: list[dict[str, Any]], item_count_status: str,
                   item_rows_status: str) -> dict[str, Any]:
    derived = _decimal(summary.get("tax_total"))
    declared = _decimal(summary.get("declared_tax_total"))
    lines = summary.get("lines") if isinstance(summary.get("lines"), list) else []
    rows = [line for line in lines if isinstance(line, dict)]
    amounts = [_decimal(line.get("tax_amount")) for line in rows]
    known = [value for value in amounts if value is not None]
    row_sum = sum(known, Decimal("0")) if known and len(known) == len(rows) else None
    codes = [str(row.get("tax_code") or "").strip() for row in rows]
    item_codes = {str(item.get("tax_code") or "").strip() for item in items if str(item.get("tax_code") or "").strip()}
    summary_codes = {code for code in codes if code}
    summary_checks = {
        "required_tax_codes_valid": "passed" if summary_codes and len(codes) == len(summary_codes) else "failed",
        "rates_valid": "passed" if rows and all(_valid_rate(row.get("rate")) for row in rows) else "failed",
        "source_evidence_valid": "passed" if rows and all((row.get("evidence") or {}).get("status") == "confirmed" for row in rows) else "failed",
        "row_arithmetic_valid": "passed" if rows and all(_tax_row_valid(row) for row in rows) else "failed",
        "declared_total_valid": _status(row_sum, declared),
        "derived_total_valid": _status(row_sum, derived),
    }
    summary_status = "passed" if all(value == "passed" for value in summary_checks.values()) else "failed"
    item_checks = {
        "all_items_extracted": item_count_status,
        "all_item_rows_reconciled": item_rows_status,
        "all_items_have_tax_code": "passed" if items and all(str(item.get("tax_code") or "") in {"A", "B"} for item in items) else "failed",
        "all_item_tax_codes_evidenced": "passed" if items and all(
            (item.get("tax_code_evidence") or {}).get("evidence_state") == "confirmed" for item in items
        ) else "failed",
    }
    item_status = "passed" if all(value == "passed" for value in item_checks.values()) else "failed"
    comparison_checks = {
        "all_items_extracted": item_count_status,
        "all_item_rows_reconciled": item_rows_status,
        "tax_summary_extraction": summary_status,
        "item_tax_classification": item_status,
        "item_codes_covered_by_summary": "passed" if item_codes and item_codes.issubset(summary_codes) else "failed",
    }
    comparison_status = "passed" if all(value == "passed" for value in comparison_checks.values()) else "failed"
    return {
        "status": comparison_status,
        "tax_summary_extraction": {"status": summary_status, "checks": summary_checks},
        "item_tax_classification": {"status": item_status, "checks": item_checks},
        "tax_summary_vs_items": {"status": comparison_status, "checks": comparison_checks},
        "checks": {**summary_checks, **item_checks, "item_class_consistency": comparison_checks["item_codes_covered_by_summary"]},
        "required_tax_codes": sorted(item_codes),
        "observed_tax_codes": codes, "row_tax_sum": _float(row_sum),
        "declared_tax_total": _float(declared), "derived_tax_total": _float(derived),
    }


def _valid_rate(value: Any) -> bool:
    rate = _decimal(value)
    return rate is not None and Decimal("0") <= rate <= Decimal("1")


def _tax_row_valid(row: dict[str, Any]) -> bool:
    rate = _decimal(row.get("rate")); net = _decimal(row.get("net_amount")); tax = _decimal(row.get("tax_amount"))
    return rate is not None and net is not None and tax is not None and abs((net * rate) - tax) <= TOLERANCE


def _money_evidence_valid(value: Any) -> bool:
    if not isinstance(value, dict):
        return False
    return bool(
        value.get("raw_text") and value.get("source_line")
        and value.get("origin") in {"observed", "derived"}
        and value.get("verification") == "ocr_verified"
        and _confidence(value.get("confidence")) > 0
    )


def _status(actual: Decimal | None, expected: Decimal | None) -> str:
    if actual is None or expected is None:
        return "not_available"
    return "passed" if abs(actual - expected) <= TOLERANCE else "failed"


def _decimal(value: Any) -> Decimal | None:
    if isinstance(value, dict):
        value = value.get("computed") if value.get("computed") is not None else value.get("verbatim")
    if value is None or isinstance(value, bool):
        return None
    try:
        return Decimal(str(value).replace("AUD", "").replace("$", "").replace(",", "").strip())
    except (InvalidOperation, ValueError):
        return None


def _float(value: Decimal | None) -> float | None:
    return None if value is None else float(value.quantize(Decimal("0.01")))


def _confidence(value: Any) -> float:
    try:
        return min(1.0, max(0.0, float(value or 0)))
    except (TypeError, ValueError):
        return 0.0
