"""Deterministic invoice field extraction from OCR text.

This is NOT AI extraction. It uses regex/heuristics only.
Unknown fields remain null — never invent values.
Confidence scores are Local OCR/extraction confidence (0.0–1.0 per field).
"""

from __future__ import annotations

import re
from datetime import datetime

from app.services.ocr.base import ExtractedField, ExtractionResult

# Labels commonly found near invoice numbers
_INV_NUMBER_PATTERNS = [
    re.compile(
        r"(?:invoice\s*(?:no\.?|number|#)|factuur\s*(?:nr\.?|nummer)|inv\.?\s*(?:no\.?|#)|nummer)\s*[:.\-]?\s*([A-Z0-9][A-Z0-9\-/\.]{2,30})",
        re.IGNORECASE,
    ),
    re.compile(r"\b((?:INV|INVOICE|FAC)[-/\s]?\d{3,}\b)", re.IGNORECASE),
]

_DATE_PATTERNS = [
    # ISO / yyyy-mm-dd
    re.compile(r"\b(\d{4}[-/.]\d{1,2}[-/.]\d{1,2})\b"),
    # dd-mm-yyyy / dd/mm/yyyy / dd.mm.yyyy
    re.compile(r"\b(\d{1,2}[-/.]\d{1,2}[-/.]\d{2,4})\b"),
]

_INVOICE_DATE_LABEL = re.compile(
    r"(?:invoice\s*date|factuurdatum|date\s*of\s*invoice|datum)\s*[:.\-]?\s*",
    re.IGNORECASE,
)
_DUE_DATE_LABEL = re.compile(
    r"(?:due\s*date|vervaldatum|payment\s*due|betaal\s*vóór|betaal\s*voor)\s*[:.\-]?\s*",
    re.IGNORECASE,
)

_AMOUNT_TOKEN = re.compile(
    r"(?:€|EUR|\$|USD|GBP|£)?\s*"
    r"(\d{1,3}(?:[.\s]\d{3})*(?:[.,]\d{2})|\d+[.,]\d{2}|\d{3,})",
    re.IGNORECASE,
)

_TOTAL_LABEL = re.compile(
    r"(?:(?<!sub)\btotal\b\s*(?:amount|due|incl\.?\s*vat)?|\btotaal\b|amount\s*due|te\s*betalen)\s*[:.\-]?",
    re.IGNORECASE,
)
_VAT_LABEL = re.compile(
    r"(?:\bvat\b|\bbtw\b|\btax\b|\bmwst\b)\s*(?:\d+\s*%?)?\s*[:.\-]?",
    re.IGNORECASE,
)
_SUBTOTAL_LABEL = re.compile(
    r"(?:\bsubtotal\b|\bsub-total\b|\bsub\s*total\b|\bnetto\b|excl\.?\s*(?:vat|btw))\s*[:.\-]?",
    re.IGNORECASE,
)

_CURRENCY_PATTERNS = [
    (re.compile(r"\bEUR\b|€", re.IGNORECASE), "EUR"),
    (re.compile(r"\bUSD\b|\$", re.IGNORECASE), "USD"),
    (re.compile(r"\bGBP\b|£", re.IGNORECASE), "GBP"),
]

_PAYMENT_REF = re.compile(
    r"(?:payment\s*(?:ref(?:erence)?|reference)|kenmerk|betalingskenmerk|ocr)\s*[:.\-]?\s*([A-Z0-9][A-Z0-9\s\-]{4,40})",
    re.IGNORECASE,
)

_COMPANY_SUFFIX = re.compile(
    r"\b(?:B\.?V\.?|N\.?V\.?|Ltd\.?|Limited|Inc\.?|GmbH|LLC|PLC|VOF|CV)\b",
    re.IGNORECASE,
)

_NOISE_LINE = re.compile(
    r"^(?:page\s+\d+|invoice|factuur|tax\s*invoice|commercial\s*invoice)\s*$",
    re.IGNORECASE,
)


def _null_field() -> ExtractedField:
    return ExtractedField(value=None, confidence=0.0)


def _parse_amount(raw: str) -> float | None:
    """Parse European/US amount strings. Returns None if ambiguous/invalid."""
    s = raw.strip().replace(" ", "").replace("€", "").replace("$", "").replace("£", "")
    s = re.sub(r"(?i)^(EUR|USD|GBP)", "", s).strip()
    if not s:
        return None
    # 1.234,56 → European
    if re.fullmatch(r"\d{1,3}(\.\d{3})+,\d{2}", s):
        s = s.replace(".", "").replace(",", ".")
    # 1,234.56 → US
    elif re.fullmatch(r"\d{1,3}(,\d{3})+\.\d{2}", s):
        s = s.replace(",", "")
    # 1234,56
    elif re.fullmatch(r"\d+,\d{2}", s):
        s = s.replace(",", ".")
    # 1234.56
    elif re.fullmatch(r"\d+\.\d{2}", s):
        pass
    # OCR often drops the decimal: "30250" near a total label → try 302.50
    elif re.fullmatch(r"\d{4,8}", s):
        whole, cents = s[:-2], s[-2:]
        try:
            value = float(f"{int(whole)}.{cents}")
        except ValueError:
            return None
        if 0 < value <= 1_000_000_000:
            return round(value, 2)
        return None
    else:
        if "," in s and "." in s:
            return None
        if s.count(",") == 1 and len(s.split(",")[-1]) != 2:
            return None
        if s.count(".") == 1 and len(s.split(".")[-1]) != 2 and len(s.split(".")[-1]) != 3:
            try:
                return float(s)
            except ValueError:
                return None
        try:
            return float(s.replace(",", ""))
        except ValueError:
            return None
    try:
        value = float(s)
    except ValueError:
        return None
    if value < 0 or value > 1_000_000_000:
        return None
    return round(value, 2)


def _normalize_date(raw: str) -> str | None:
    raw = raw.strip()
    formats = (
        "%Y-%m-%d",
        "%Y/%m/%d",
        "%Y.%m.%d",
        "%d-%m-%Y",
        "%d/%m/%Y",
        "%d.%m.%Y",
        "%d-%m-%y",
        "%d/%m/%y",
        "%d.%m.%y",
    )
    for fmt in formats:
        try:
            dt = datetime.strptime(raw, fmt)
            # Reject clearly wrong 2-digit year expansions for invoices
            if dt.year < 1990 or dt.year > 2100:
                continue
            return dt.date().isoformat()
        except ValueError:
            continue
    return None


def _amount_after_label(text: str, label: re.Pattern[str]) -> tuple[float | None, float]:
    # Normalize glued OCR tokens a bit for matching ("Totalamount")
    soft = re.sub(r"(?i)(total)\s*(amount)", r"\1 \2", text)
    soft = re.sub(r"(?i)(due)\s*(date)", r"\1 \2", soft)
    for match in label.finditer(soft):
        window = soft[match.end() : match.end() + 48]
        amt = _AMOUNT_TOKEN.search(window)
        if not amt:
            continue
        raw = amt.group(1)
        value = _parse_amount(raw)
        if value is None:
            continue
        # Lower confidence when we had to invent a decimal point from bare digits
        conf = 0.7 if re.fullmatch(r"\d{4,8}", raw.replace(" ", "")) else 0.85
        return value, conf
    return None, 0.0


def _date_after_label(text: str, label: re.Pattern[str]) -> tuple[str | None, float]:
    for match in label.finditer(text):
        window = text[match.end() : match.end() + 24]
        for pat in _DATE_PATTERNS:
            dm = pat.search(window)
            if not dm:
                continue
            normalized = _normalize_date(dm.group(1))
            if normalized:
                return normalized, 0.88
    return None, 0.0


def _guess_supplier(lines: list[str]) -> ExtractedField:
    """Heuristic: first substantial letter line, preferably with a company suffix."""
    candidates: list[tuple[str, float]] = []
    for line in lines[:12]:
        cleaned = re.sub(r"\s+", " ", line).strip(" -•·|")
        if len(cleaned) < 3 or len(cleaned) > 80:
            continue
        if _NOISE_LINE.match(cleaned):
            continue
        if re.fullmatch(r"[\d\s\-/.€$£,:]+", cleaned):
            continue
        if cleaned.lower().startswith(("tel", "phone", "email", "www", "http", "iban", "kvk", "btw", "vat")):
            continue
        conf = 0.55
        if _COMPANY_SUFFIX.search(cleaned):
            conf = 0.82
        elif cleaned[:1].isupper() and sum(c.isalpha() for c in cleaned) >= 4:
            conf = 0.62
        else:
            continue
        candidates.append((cleaned, conf))
    if not candidates:
        return _null_field()
    # Prefer highest confidence, then earliest
    candidates.sort(key=lambda c: (-c[1],))
    best, conf = candidates[0]
    return ExtractedField(value=best, confidence=conf)


def _extract_invoice_number(text: str) -> ExtractedField:
    for pat in _INV_NUMBER_PATTERNS:
        m = pat.search(text)
        if not m:
            continue
        value = re.sub(r"\s+", "", m.group(1)).strip(" .-")
        if len(value) < 3:
            continue
        # Avoid matching bare years
        if re.fullmatch(r"20\d{2}", value):
            continue
        conf = 0.9 if "invoice" in pat.pattern.lower() or "factuur" in pat.pattern.lower() else 0.75
        return ExtractedField(value=value, confidence=conf)
    return _null_field()


def _extract_currency(text: str) -> ExtractedField:
    for pat, code in _CURRENCY_PATTERNS:
        if pat.search(text):
            return ExtractedField(value=code, confidence=0.8)
    return _null_field()


def _extract_payment_ref(text: str) -> ExtractedField:
    m = _PAYMENT_REF.search(text)
    if not m:
        return _null_field()
    value = re.sub(r"\s+", " ", m.group(1)).strip(" .-")
    if len(value) < 5:
        return _null_field()
    return ExtractedField(value=value, confidence=0.7)


def extract_invoice_fields(text: str) -> ExtractionResult:
    """Extract invoice fields from OCR text. Missing values stay null."""
    cleaned = text or ""
    lines = [ln.strip() for ln in cleaned.splitlines() if ln.strip()]

    supplier = _guess_supplier(lines)
    invoice_number = _extract_invoice_number(cleaned)

    inv_date, inv_conf = _date_after_label(cleaned, _INVOICE_DATE_LABEL)
    due_date, due_conf = _date_after_label(cleaned, _DUE_DATE_LABEL)

    # Fallback: first/second date in document if labels missing — low confidence only
    if not inv_date:
        dates: list[str] = []
        for pat in _DATE_PATTERNS:
            for m in pat.finditer(cleaned):
                normalized = _normalize_date(m.group(1))
                if normalized and normalized not in dates:
                    dates.append(normalized)
        if dates:
            inv_date, inv_conf = dates[0], 0.45
            if not due_date and len(dates) > 1:
                due_date, due_conf = dates[1], 0.4

    total_val, total_conf = _amount_after_label(cleaned, _TOTAL_LABEL)
    vat_val, vat_conf = _amount_after_label(cleaned, _VAT_LABEL)
    sub_val, sub_conf = _amount_after_label(cleaned, _SUBTOTAL_LABEL)

    # Soft cross-check: if we have subtotal+VAT and no total, sum them only when both confident
    if total_val is None and sub_val is not None and vat_val is not None and sub_conf >= 0.7 and vat_conf >= 0.7:
        total_val = round(sub_val + vat_val, 2)
        total_conf = min(sub_conf, vat_conf) * 0.85

    currency = _extract_currency(cleaned)
    payment_ref = _extract_payment_ref(cleaned)

    fields = [
        supplier,
        invoice_number,
        ExtractedField(inv_date, inv_conf),
        ExtractedField(due_date, due_conf),
        ExtractedField(sub_val, sub_conf),
        ExtractedField(vat_val, vat_conf),
        ExtractedField(total_val, total_conf),
        currency,
        payment_ref,
    ]

    # Overall Local OCR/extraction confidence (0–100).
    # Weight key bookkeeping fields more heavily. No text → 0.
    if not cleaned.strip():
        overall = 0.0
    else:
        weights = {
            "supplier": 0.15,
            "invoice_number": 0.2,
            "invoice_date": 0.15,
            "total": 0.25,
            "vat": 0.1,
            "currency": 0.05,
            "due_date": 0.05,
            "payment_ref": 0.05,
        }
        keyed = {
            "supplier": supplier,
            "invoice_number": invoice_number,
            "invoice_date": ExtractedField(inv_date, inv_conf),
            "due_date": ExtractedField(due_date, due_conf),
            "total": ExtractedField(total_val, total_conf),
            "vat": ExtractedField(vat_val, vat_conf),
            "currency": currency,
            "payment_ref": payment_ref,
        }
        score = 0.0
        # Base credit for having substantial OCR text
        text_credit = min(0.25, len(cleaned.strip()) / 4000)
        score += text_credit
        for key, weight in weights.items():
            f = keyed[key]
            if f.value is not None and f.confidence > 0:
                score += weight * f.confidence
        overall = round(min(100.0, score * 100), 1)

    return ExtractionResult(
        supplier=supplier,
        invoice_number=invoice_number,
        invoice_date=ExtractedField(inv_date, inv_conf),
        due_date=ExtractedField(due_date, due_conf),
        subtotal=ExtractedField(sub_val, sub_conf),
        vat=ExtractedField(vat_val, vat_conf),
        total=ExtractedField(total_val, total_conf),
        currency=currency,
        payment_ref=payment_ref,
        overall_confidence=overall,
    )
