"""Type-specific enrichers — extend base extract with domain assertions/lines."""

from __future__ import annotations

import re
from typing import Callable

from app.fdi.extract import ExtractedAssertion, ExtractedLine, StructuredExtract, _norm_desc
from app.fdi.identifiers import FoundIdentifier, extract_identifiers, normalize_id

_MONEY = re.compile(
    r"(?i)(?:₹|rs\.?|inr|usd|\$)?\s*([\d,]+\.\d{2}|[\d,]+)"
)


def _money(text: str) -> float | None:
    m = _MONEY.search(text or "")
    if not m:
        return None
    try:
        return float(m.group(1).replace(",", ""))
    except ValueError:
        return None


def _add_assertion(
    out: list[ExtractedAssertion],
    *,
    kind: str,
    label: str,
    value: float | None = None,
    text: str | None = None,
    unit: str | None = None,
    conf: float = 0.8,
) -> None:
    if value is None and not text:
        return
    out.append(
        ExtractedAssertion(
            kind=kind,
            label=label,
            value_numeric=value,
            value_text=text,
            unit=unit,
            currency="INR" if value is not None and unit != "count" else None,
            source="typed_extract",
            confidence=conf,
        )
    )


def enrich_invoice(extract: StructuredExtract) -> StructuredExtract:
    text = extract.document_text or ""
    assertions = list(extract.assertions)
    lines = list(extract.lines)

    inv = re.search(r"(?i)invoice\s*(?:no|num|number|#)?\s*[:\-]?\s*([A-Z0-9][A-Z0-9\-/]{3,30})", text)
    if inv:
        _add_assertion(assertions, kind="invoice_number", label="Invoice number", text=inv.group(1).strip(), conf=0.9)

    for label, kind in (
        (r"(?i)(?:grand\s+)?total\s*(?:amount)?\s*[:\-]?\s*", "invoice_total"),
        (r"(?i)(?:taxable\s+value|taxable\s+amount)\s*[:\-]?\s*", "taxable_value"),
        (r"(?i)(?:cgst|sgst|igst|gst)\s*(?:amount)?\s*[:\-]?\s*", "tax_total"),
        (r"(?i)(?:amount\s+due|balance\s+due)\s*[:\-]?\s*", "amount_due"),
    ):
        m = re.search(label + r"((?:₹|rs\.?|inr|\$)?\s*[\d,]+\.?\d*)", text)
        if m:
            _add_assertion(assertions, kind=kind, label=kind.replace("_", " ").title(), value=_money(m.group(1)))

    # Line-ish rows: description .... amount
    if len(lines) < 2:
        for i, m in enumerate(
            re.finditer(
                r"(?m)^(.{8,80}?)\s{2,}((?:₹|rs\.?)?\s*[\d,]+\.\d{2})\s*$",
                text[:20000],
            )
        ):
            amt = _money(m.group(2))
            if amt is None or amt <= 0:
                continue
            desc = _norm_desc(m.group(1))
            if re.search(r"(?i)total|subtotal|gst|tax|invoice", desc):
                continue
            lines.append(
                ExtractedLine(
                    row_idx=1000 + i,
                    description_raw=desc,
                    description_norm=desc,
                    amount=amt,
                    debit=amt,
                    direction="out",
                    currency="INR",
                    event_class="invoice_line",
                    confidence=0.65,
                )
            )

    seller = re.search(r"(?i)(?:sold\s+by|from|supplier|vendor)\s*[:\-]\s*([A-Za-z0-9][A-Za-z0-9 .,&']{2,60})", text)
    buyer = re.search(r"(?i)(?:bill\s+to|buyer|customer)\s*[:\-]\s*([A-Za-z0-9][A-Za-z0-9 .,&']{2,60})", text)
    parties = list(extract.parties)
    if seller:
        parties.append(("vendor", seller.group(1).strip()))
        extract.issuer_name = extract.issuer_name or seller.group(1).strip()
    if buyer:
        parties.append(("customer", buyer.group(1).strip()))
        extract.holder_name = extract.holder_name or buyer.group(1).strip()

    extract.assertions = assertions
    extract.lines = lines
    extract.parties = parties
    return extract


def enrich_gst(extract: StructuredExtract) -> StructuredExtract:
    text = extract.document_text or ""
    assertions = list(extract.assertions)
    for label, kind in (
        (r"(?i)total\s+taxable\s+value\s*[:\-]?\s*", "taxable_value"),
        (r"(?i)total\s+(?:tax|gst)\s*(?:liability|amount)?\s*[:\-]?\s*", "tax_total"),
        (r"(?i)(?:igst|cgst|sgst)\s*[:\-]?\s*", "tax_component"),
    ):
        m = re.search(label + r"((?:₹|rs\.?)?\s*[\d,]+\.?\d*)", text)
        if m:
            _add_assertion(assertions, kind=kind, label=kind.replace("_", " ").title(), value=_money(m.group(1)))
    period = re.search(r"(?i)(?:tax\s+period|return\s+period)\s*[:\-]?\s*([A-Za-z0-9 /\-]{4,20})", text)
    if period:
        _add_assertion(assertions, kind="return_period", label="Return period", text=period.group(1).strip())
    extract.assertions = assertions
    return extract


def enrich_payslip(extract: StructuredExtract) -> StructuredExtract:
    text = extract.document_text or ""
    assertions = list(extract.assertions)
    for label, kind in (
        (r"(?i)(?:gross\s+(?:pay|salary|earnings))\s*[:\-]?\s*", "gross_pay"),
        (r"(?i)(?:net\s+(?:pay|salary))\s*[:\-]?\s*", "net_pay"),
        (r"(?i)(?:total\s+deductions?)\s*[:\-]?\s*", "deductions"),
        (r"(?i)(?:basic\s+(?:pay|salary))\s*[:\-]?\s*", "basic_pay"),
        (r"(?i)(?:pf|provident\s+fund)\s*[:\-]?\s*", "pf_contribution"),
    ):
        m = re.search(label + r"((?:₹|rs\.?)?\s*[\d,]+\.?\d*)", text)
        if m:
            _add_assertion(assertions, kind=kind, label=kind.replace("_", " ").title(), value=_money(m.group(1)))
    emp = re.search(r"(?i)employee\s*(?:id|code|no)?\s*[:\-]?\s*([A-Z0-9\-/]{3,20})", text)
    if emp:
        _add_assertion(assertions, kind="employee_id", label="Employee ID", text=emp.group(1).strip())
    extract.assertions = assertions
    if any(a.kind == "net_pay" for a in assertions):
        net = next(a for a in assertions if a.kind == "net_pay")
        extract.lines = list(extract.lines) + [
            ExtractedLine(
                description_raw="Net salary",
                description_norm="net salary",
                amount=net.value_numeric,
                credit=net.value_numeric,
                direction="in",
                event_class="salary",
                confidence=0.85,
            )
        ]
    return extract


def enrich_mutual_fund_cas(extract: StructuredExtract) -> StructuredExtract:
    text = extract.document_text or ""
    assertions = list(extract.assertions)
    lines = list(extract.lines)
    folios = re.findall(r"(?i)folio\s*(?:no|number|#)?\s*[:\-]?\s*([A-Z0-9/\-]{4,20})", text)
    for f in folios[:20]:
        extract.identifiers = list(extract.identifiers) + [
            FoundIdentifier("FOLIO", f, normalize_id("FOLIO", f), 0.85)
        ]
    for i, m in enumerate(
        re.finditer(
            r"(?i)([A-Za-z][A-Za-z0-9 &.\-]{5,50})\s+(?:growth|dividend|direct|regular).{0,40}?"
            r"([\d,]+\.\d{2,4})\s+(?:units?)?.{0,20}?([\d,]+\.\d{2})",
            text[:30000],
        )
    ):
        units = _money(m.group(2))
        value = _money(m.group(3))
        lines.append(
            ExtractedLine(
                row_idx=2000 + i,
                description_raw=m.group(1).strip(),
                description_norm=_norm_desc(m.group(1)),
                amount=value,
                direction="neutral",
                event_class="nav",
                confidence=0.6,
                extras={"units": units},
            )
        )
    total = re.search(r"(?i)(?:total\s+(?:cost|value)|portfolio\s+value)\s*[:\-]?\s*((?:₹|rs\.?)?\s*[\d,]+\.?\d*)", text)
    if total:
        _add_assertion(assertions, kind="portfolio_value", label="Portfolio value", value=_money(total.group(1)))
    extract.assertions = assertions
    extract.lines = lines
    return extract


def enrich_insurance(extract: StructuredExtract) -> StructuredExtract:
    text = extract.document_text or ""
    assertions = list(extract.assertions)
    for label, kind in (
        (r"(?i)sum\s+assured\s*[:\-]?\s*", "sum_assured"),
        (r"(?i)(?:premium\s+(?:amount|due)|installment\s+premium)\s*[:\-]?\s*", "premium"),
        (r"(?i)policy\s*(?:no|number|#)?\s*[:\-]?\s*", "policy_number_text"),
    ):
        m = re.search(label + r"([A-Z0-9/,\-₹Rs.\s]{3,40})", text)
        if not m:
            continue
        raw = m.group(1).strip()
        if kind == "policy_number_text":
            _add_assertion(assertions, kind="policy_number", label="Policy number", text=raw.split()[0], conf=0.85)
        else:
            _add_assertion(assertions, kind=kind, label=kind.replace("_", " ").title(), value=_money(raw))
    extract.assertions = assertions
    return extract


def enrich_loan(extract: StructuredExtract) -> StructuredExtract:
    text = extract.document_text or ""
    assertions = list(extract.assertions)
    for label, kind in (
        (r"(?i)(?:outstanding\s+(?:principal|balance)|principal\s+outstanding)\s*[:\-]?\s*", "outstanding_principal"),
        (r"(?i)(?:emi\s+amount|installment\s+amount)\s*[:\-]?\s*", "emi_amount"),
        (r"(?i)(?:interest\s+(?:rate|amount))\s*[:\-]?\s*", "interest"),
        (r"(?i)(?:loan\s+amount|sanctioned\s+amount)\s*[:\-]?\s*", "loan_amount"),
    ):
        m = re.search(label + r"((?:₹|rs\.?)?\s*[\d,]+\.?\d*%?)", text)
        if m:
            val = _money(m.group(1))
            _add_assertion(assertions, kind=kind, label=kind.replace("_", " ").title(), value=val, text=m.group(1).strip() if val is None else None)
    extract.assertions = assertions
    return extract


def enrich_financial_statement(extract: StructuredExtract) -> StructuredExtract:
    """Balance sheet / P&L / cash flow line labels."""
    text = extract.document_text or ""
    assertions = list(extract.assertions)
    lines = list(extract.lines)
    patterns = [
        (r"(?i)total\s+assets\s*[:\-]?\s*((?:₹|rs\.?)?\s*[\d,]+\.?\d*)", "total_assets"),
        (r"(?i)total\s+liabilit(?:y|ies)\s*[:\-]?\s*((?:₹|rs\.?)?\s*[\d,]+\.?\d*)", "total_liabilities"),
        (r"(?i)(?:shareholders?'?\s+equity|total\s+equity)\s*[:\-]?\s*((?:₹|rs\.?)?\s*[\d,]+\.?\d*)", "total_equity"),
        (r"(?i)(?:revenue|total\s+income|turnover)\s*[:\-]?\s*((?:₹|rs\.?)?\s*[\d,]+\.?\d*)", "revenue"),
        (r"(?i)(?:net\s+profit|profit\s+after\s+tax|pat)\s*[:\-]?\s*((?:₹|rs\.?)?\s*[\d,]+\.?\d*)", "net_profit"),
        (r"(?i)(?:net\s+loss)\s*[:\-]?\s*((?:₹|rs\.?)?\s*[\d,]+\.?\d*)", "net_loss"),
        (r"(?i)(?:operating\s+cash\s+flow|cash\s+from\s+operations)\s*[:\-]?\s*((?:₹|rs\.?)?\s*[\d,]+\.?\d*)", "operating_cash_flow"),
    ]
    for pat, kind in patterns:
        m = re.search(pat, text)
        if m:
            _add_assertion(assertions, kind=kind, label=kind.replace("_", " ").title(), value=_money(m.group(1)), conf=0.85)

    # Statement lines: Label .... amount
    if len(lines) < 5:
        for i, m in enumerate(
            re.finditer(
                r"(?m)^([A-Za-z][A-Za-z0-9 &()/.\-]{3,60})\s{2,}((?:₹|rs\.?)?\s*[\d,]+\.\d{2})\s*$",
                text[:25000],
            )
        ):
            amt = _money(m.group(2))
            if not amt:
                continue
            desc = _norm_desc(m.group(1))
            lines.append(
                ExtractedLine(
                    row_idx=3000 + i,
                    description_raw=desc,
                    description_norm=desc,
                    amount=amt,
                    direction="neutral",
                    event_class="journal",
                    confidence=0.6,
                    extras={"statement_line": True},
                )
            )
    extract.assertions = assertions
    extract.lines = lines
    return extract


def enrich_receipt_or_po(extract: StructuredExtract) -> StructuredExtract:
    text = extract.document_text or ""
    assertions = list(extract.assertions)
    total = re.search(r"(?i)(?:total|amount\s+paid|order\s+total)\s*[:\-]?\s*((?:₹|rs\.?)?\s*[\d,]+\.?\d*)", text)
    if total:
        kind = "po_total" if extract.doc_type == "purchase_order" else "receipt_total"
        _add_assertion(assertions, kind=kind, label=kind.replace("_", " ").title(), value=_money(total.group(1)))
    extract.assertions = assertions
    return extract


_ENRICHERS: dict[str, Callable[[StructuredExtract], StructuredExtract]] = {
    "invoice": enrich_invoice,
    "bill": enrich_invoice,
    "gst_return": enrich_gst,
    "tax_document": enrich_gst,
    "payslip": enrich_payslip,
    "mutual_fund_cas": enrich_mutual_fund_cas,
    "insurance_policy": enrich_insurance,
    "loan_statement": enrich_loan,
    "balance_sheet": enrich_financial_statement,
    "profit_and_loss": enrich_financial_statement,
    "cash_flow": enrich_financial_statement,
    "trial_balance": enrich_financial_statement,
    "general_ledger": enrich_financial_statement,
    "receipt": enrich_receipt_or_po,
    "purchase_order": enrich_receipt_or_po,
    "expense_report": enrich_receipt_or_po,
}


def apply_type_enrichment(extract: StructuredExtract) -> StructuredExtract:
    fn = _ENRICHERS.get(extract.doc_type)
    if not fn:
        # Still harvest any extra identifiers from full text
        extra = extract_identifiers(extract.document_text or "")
        have = {(i.id_type, i.value_norm) for i in extract.identifiers}
        for i in extra:
            if (i.id_type, i.value_norm) not in have:
                extract.identifiers.append(i)
        return extract
    enriched = fn(extract)
    extra = extract_identifiers(enriched.document_text or "")
    have = {(i.id_type, i.value_norm) for i in enriched.identifiers}
    for i in extra:
        if (i.id_type, i.value_norm) not in have:
            enriched.identifiers.append(i)
    return enriched
