"""
normaliser.py — Column normalisation helpers for recon_engine_py.

The bank/MomentsPay/HIS exports have inconsistent column names across banks and
over time. This module resolves an expected canonical column name from a list
of known aliases using whatever columns are actually present in a DataFrame.
"""

import re
import pandas as pd


# ── Canonical → aliases ──────────────────────────────────────────────────────

COLUMN_ALIASES = {
    # Bank Card columns
    "CARD_NUMBER":   ["CARD NUMBER", "Card Number", "Card No", "CARD NO", "CARDNBR", "PAN"],
    "TERMINAL_ID":   ["TERMINAL NUMBER", "Terminal Number", "TERMINAL_NO", "TerminalID",
                      "EXTERNAL TID", "Terminal ID"],
    "DOMESTIC_AMT":  ["DOMESTIC AMT", "Domestic Amt", "DOMESTIC AMOUNT"],
    "INTL_AMT":      ["INTNL AMT", "INTL AMT", "International Amt", "INTERNATIONAL AMT"],
    "APPROVAL_CODE": ["APPROV CODE", "Approval Code", "APPROVAL CODE", "AUTH CODE",
                      "Authorisation Code"],
    "TRANS_DATE":    ["TRANS DATE", "Transaction Date", "Date", "TRANSACTION DATE", "TXN DATE"],

    # Bank UPI columns
    "PAYER_VPA":     ["PAYER VPA", "Payer VPA", "VPA", "UPI ID", "CARDNBR"],
    "RRN_NO":        ["Txn ref no.(RRN)", "RRN NO", "RRN", "Reference No",
                      "Transaction Reference", "RRN_NO"],
    "UPI_AMOUNT":    ["Transaction Amount", "AUTH_AMOUNT", "TransactionAmount",
                      "Amount", "DOMESTIC AMT"],

    # MomentsPay columns
    "MP_CARD_NUM":   ["card_num"],
    "MP_TERMINAL_ID":["terminal_id"],
    "MP_AMOUNT":     ["total_amount"],
    "MP_APPROVAL":   ["approval_code"],
    "MP_RRN":        ["rrn_no"],
    "MP_PROCESSING": ["processing_id"],
    "MP_TRANSACTION":["transaction_id"],
    "MP_EMAIL":      ["email"],

    # HIS columns
    "HIS_CARD":      ["Card No/Cheque", "Card No", "Card Number", "CARD NO"],
    "HIS_AMOUNT":    ["Amount", "Transaction Amount"],
    "HIS_APPROVAL":  ["Approval No", "Approval Code", "APPROVAL NO"],
    "HIS_DATE":      ["Doc Date", "Date", "Transaction Date"],
    "HIS_UNIT":      ["Unit/Location", "Unit", "Location", "Branch"],
    "HIS_PAYMENT":   ["payment mode", "Payment Mode", "Mode", "Payment Type"],
    "HIS_MR_NO":     ["IP/MRNO", "MRNO", "MR No", "Patient MR"],
    "HIS_DOC_NO":    ["Doc No", "Document No"],
}

EXCLUDE_BANK_CARD_PATTERNS = [
    "Bank PAN No", "Merchant PAN No.", "Curr Ac GSTN", "Acc No2 GSTN"
]


def resolve_column(df: pd.DataFrame, canonical: str, extra_aliases: list = None) -> str | None:
    """Return the first column in df that matches the alias list for `canonical`, or None."""
    aliases = COLUMN_ALIASES.get(canonical, [])
    if extra_aliases:
        aliases = list(extra_aliases) + aliases
    for alias in aliases:
        if alias in df.columns:
            return alias
    return None


def require_column(df: pd.DataFrame, canonical: str, extra_aliases: list = None) -> str:
    """Like resolve_column but raises ValueError if not found."""
    col = resolve_column(df, canonical, extra_aliases)
    if col is None:
        aliases = COLUMN_ALIASES.get(canonical, [])
        if extra_aliases:
            aliases = list(extra_aliases) + aliases
        raise ValueError(
            f"Cannot find column '{canonical}'. "
            f"Expected one of: {aliases}. "
            f"Actual columns: {list(df.columns)}"
        )
    return col


def normalise_str(series: pd.Series) -> pd.Series:
    """Strip whitespace, remove trailing .0, cast to str."""
    return (
        series.astype(str)
        .str.strip()
        .str.replace(r"\.0$", "", regex=True)
    )


def pad_approval(series: pd.Series, width: int = 6) -> pd.Series:
    """Zero-pad approval codes to `width` digits, strip leading quotes."""
    return (
        normalise_str(series)
        .str.replace(r"^'", "", regex=True)
        .apply(lambda x: x.zfill(width) if x != "nan" and len(x) < width else x)
    )


def pad_rrn(series: pd.Series, width: int = 12) -> pd.Series:
    """Zero-pad RRN numbers to `width` digits."""
    return (
        normalise_str(series)
        .str.replace(r"^'", "", regex=True)
        .apply(lambda x: x.zfill(width) if x != "nan" else x)
    )


def last4(series: pd.Series) -> pd.Series:
    """Return last 4 chars; if value contains '@' (VPA) return as-is."""
    return series.apply(lambda x: x if "@" in str(x) else str(x)[-4:])


def auth_amount(df: pd.DataFrame, domestic_col: str, intl_col: str) -> pd.Series:
    """Compute effective amount: DOMESTIC_AMT if INTL_AMT == '0', else INTL_AMT."""
    domestic = normalise_str(df[domestic_col])
    intl = normalise_str(df[intl_col])
    return intl.where(intl != "0", domestic)


def filter_bank_header_rows(df: pd.DataFrame) -> pd.DataFrame:
    """Remove summary/header rows injected by the bank portal."""
    first_col = df.columns[0]
    mask = df[first_col].astype(str).str.startswith(tuple(EXCLUDE_BANK_CARD_PATTERNS))
    return df[~mask].copy()


def summarise_by_location(matched_df: pd.DataFrame,
                           unmatched_df: pd.DataFrame,
                           location_col: str,
                           amount_col: str) -> pd.DataFrame:
    """
    Build a per-location summary DataFrame with columns:
      Unit/Location, total_count, total_amount, matched_count, matched_amount,
      unmatched_count, unmatched_amount
    """
    combined = pd.concat([matched_df, unmatched_df], ignore_index=True)

    total_grouped = combined.groupby(location_col).agg(
        total_count=pd.NamedAgg(column=location_col, aggfunc="count"),
        total_amount=pd.NamedAgg(column=amount_col, aggfunc="sum"),
    ).reset_index()

    matched_grouped = matched_df.groupby(location_col).agg(
        matched_count=pd.NamedAgg(column=location_col, aggfunc="count"),
        matched_amount=pd.NamedAgg(column=amount_col, aggfunc="sum"),
    ).reset_index() if len(matched_df) > 0 else pd.DataFrame(
        columns=[location_col, "matched_count", "matched_amount"]
    )

    unmatched_grouped = unmatched_df.groupby(location_col).agg(
        unmatched_count=pd.NamedAgg(column=location_col, aggfunc="count"),
        unmatched_amount=pd.NamedAgg(column=amount_col, aggfunc="sum"),
    ).reset_index() if len(unmatched_df) > 0 else pd.DataFrame(
        columns=[location_col, "unmatched_count", "unmatched_amount"]
    )

    summary = (
        total_grouped
        .merge(matched_grouped, on=location_col, how="left")
        .merge(unmatched_grouped, on=location_col, how="left")
        .fillna(0)
    )
    return summary
