"""Document type classification — rules first, extensible to ML later."""

from __future__ import annotations

import re
from dataclasses import dataclass

from app.fdi.types import DOC_TYPES, DocType


@dataclass(frozen=True)
class Classification:
    doc_type: DocType
    confidence: float
    alternatives: tuple[tuple[str, float], ...] = ()
    reasons: tuple[str, ...] = ()


_RULES: list[tuple[DocType, float, re.Pattern[str]]] = [
    ("upi_statement", 0.9, re.compile(r"(?i)\b(upi\s+statement|paytm|phonepe|google\s*pay|gpay|bhim)\b")),
    ("credit_card_statement", 0.88, re.compile(r"(?i)\b(credit\s+card\s+statement|card\s+statement|minimum\s+amount\s+due)\b")),
    ("bank_statement", 0.85, re.compile(r"(?i)\b(bank\s+statement|passbook|account\s+statement|ifsc|neft|imps|savings\s+account)\b")),
    ("gst_return", 0.92, re.compile(r"(?i)\b(gstr[\s\-]?[1239][a-z]?|gst\s+return|goods\s+and\s+services\s+tax)\b")),
    ("invoice", 0.86, re.compile(r"(?i)\b(tax\s+invoice|invoice\s+no|bill\s+to|place\s+of\s+supply|hsn)\b")),
    ("bill", 0.75, re.compile(r"(?i)\b(utility\s+bill|electricity\s+bill|water\s+bill)\b")),
    ("payslip", 0.9, re.compile(r"(?i)\b(payslip|pay\s+slip|salary\s+slip|net\s+pay|earneds?\s+leave|pf\s+contribution)\b")),
    ("tax_document", 0.88, re.compile(r"(?i)\b(form\s*16|form\s*26as|income\s+tax\s+return|itr[\s\-]?[123vv]|tds\s+certificate)\b")),
    ("mutual_fund_cas", 0.92, re.compile(r"(?i)\b(consolidated\s+account\s+statement|cas\b|folio\s*(?:no|number)|amfi|mutual\s+fund)\b")),
    ("insurance_policy", 0.88, re.compile(r"(?i)\b(policy\s+(?:no|number|schedule)|sum\s+assured|premium\s+due|insurance\s+policy)\b")),
    ("loan_statement", 0.88, re.compile(r"(?i)\b(loan\s+account|emi\s+schedule|outstanding\s+principal|amorti[sz]ation)\b")),
    ("balance_sheet", 0.9, re.compile(r"(?i)\b(balance\s+sheet|assets\s+and\s+liabilit|statement\s+of\s+financial\s+position)\b")),
    ("profit_and_loss", 0.9, re.compile(r"(?i)\b(profit\s+(?:and|&)\s+loss|statement\s+of\s+profit|income\s+statement|p\s*&\s*l)\b")),
    ("cash_flow", 0.88, re.compile(r"(?i)\b(cash\s+flow\s+statement|statement\s+of\s+cash\s+flows)\b")),
    ("general_ledger", 0.85, re.compile(r"(?i)\b(general\s+ledger|g\/?l\s+report|account\s+ledger)\b")),
    ("trial_balance", 0.9, re.compile(r"(?i)\b(trial\s+balance)\b")),
    ("purchase_order", 0.86, re.compile(r"(?i)\b(purchase\s+order|p\.?\s*o\.?\s*number)\b")),
    ("expense_report", 0.84, re.compile(r"(?i)\b(expense\s+report|reimbursement\s+claim)\b")),
    ("receipt", 0.8, re.compile(r"(?i)\b(payment\s+receipt|official\s+receipt|received\s+with\s+thanks)\b")),
    ("annual_report", 0.82, re.compile(r"(?i)\b(annual\s+report|form\s*10[\-\s]?k|sec\s+filing)\b")),
]


def classify_document(text: str, filename: str = "") -> Classification:
    blob = f"{filename}\n{(text or '')[:12000]}"
    scores: dict[str, float] = {t: 0.0 for t in DOC_TYPES}
    reasons: list[str] = []

    for doc_type, weight, pattern in _RULES:
        if pattern.search(blob):
            scores[doc_type] = max(scores[doc_type], weight)
            reasons.append(f"{doc_type}:{pattern.pattern[:40]}")

    # Filename hints
    fn = (filename or "").lower()
    if "upi" in fn or "paytm" in fn or "phonepe" in fn:
        scores["upi_statement"] = max(scores["upi_statement"], 0.8)
    if "invoice" in fn:
        scores["invoice"] = max(scores["invoice"], 0.75)
    if "payslip" in fn or "salary" in fn:
        scores["payslip"] = max(scores["payslip"], 0.8)
    if "gstr" in fn or "gst" in fn:
        scores["gst_return"] = max(scores["gst_return"], 0.75)

    ranked = sorted(
        ((t, s) for t, s in scores.items() if s > 0 and t != "unknown"),
        key=lambda x: -x[1],
    )
    if not ranked:
        return Classification("unknown", 0.2, (), ("no_rule_match",))

    best_type, best_score = ranked[0]
    alts = tuple(ranked[1:4])
    # Soften if close competitors
    if alts and alts[0][1] >= best_score - 0.05:
        best_score = max(0.55, best_score - 0.1)
    return Classification(best_type, float(best_score), alts, tuple(reasons[:8]))  # type: ignore[arg-type]
