#!/usr/bin/env python3
"""Load fetched RBI payment-statistics Excel files into payments.duckdb.

Tables:
  rbi_psi        - Payment System Indicators, one row per (month, indicator).
  rbi_bankwise   - Bank-wise NEFT/RTGS volumes+values, one row per (month, rail, bank).
  rbi_atmposcard - Bank-wise ATM/PoS/MicroATM/BharatQR/Cards infrastructure.
  rbi_file_log   - one row per ingested workbook.
"""
import json
import re
from datetime import date
from pathlib import Path
from dateutil.relativedelta import relativedelta

import duckdb
import openpyxl

ROOT = Path("/home/workspace/Projects/payments-stat-hub")
RAW = ROOT / "raw" / "rbi"
DB = ROOT / "data" / "payments.duckdb"

NUM = re.compile(r"[,\s]")


def num(v):
    if v is None:
        return None
    if isinstance(v, (int, float)):
        return float(v)
    s = NUM.sub("", str(v))
    if not s or not re.search(r"\d", s):
        return None
    try:
        return float(s)
    except ValueError:
        return None


def month_from_name(path):
    m = re.match(r"(\d{4})-([A-Za-z]+)$", path.stem)
    if not m:
        return None
    months = {n: i for i, n in enumerate(
        ["January", "February", "March", "April", "May", "June", "July",
         "August", "September", "October", "November", "December"], 1)}
    mo = months.get(m.group(2))
    return date(int(m.group(1)), mo, 1) if mo else None


def rows_of(ws):
    if hasattr(ws, "iter_rows"):
        return list(ws.iter_rows(values_only=True))
    return list(ws)

def sheets(path):
    try:
        wb = openpyxl.load_workbook(path, read_only=True, data_only=True)
    except Exception:  # old BIFF .xls saved with .xlsx extension
        import xlrd
        book = xlrd.open_workbook(str(path))
        for name in book.sheet_names():
            sh = book.sheet_by_name(name)
            rows = [tuple(sh.cell_value(r, c) for c in range(sh.ncols))
                    for r in range(sh.nrows)]
            yield name, iter(rows)
        return
    for name in wb.sheetnames:
        yield name, wb[name]
    wb.close()


def norm_bank(name):
    if not name:
        return None
    s = re.sub(r"\s+", " ", str(name)).strip()
    return s or None


MONTHS_FULL = ["January","February","March","April","May","June","July",
               "August","September","October","November","December"]


def parse_series_type(label, pub_month):
    """Map a PSI series column label ('FY 2025-26', '2025 July', '2026 July')
    to a series_type relative to the publication month."""
    t = re.sub(r"\s+", " ", str(label)).strip()
    if not t:
        return None, ""
    if t.upper().startswith("FY"):
        return "fy_to_date", t
    m = re.match(r"^(\d{4})\s+(" + "|".join(MONTHS_FULL) + r")$", t, re.I)
    if not m:
        return "other", t
    try:
        dt = date(int(m.group(1)), MONTHS_FULL.index(m.group(2).title()) + 1, 1)
    except ValueError:
        return "other", t
    if dt == pub_month:
        return "current_month", t
    if dt == pub_month - relativedelta(months=1):
        return "previous_month", t
    if dt == pub_month - relativedelta(years=1):
        return "same_month_prior_year", t
    return "other", t


def ingest_psi(con, path, month):
    """Header-aware PSI parse (layout verified on 2026-July.xlsx):

        row 2: PART header (merged, col B)
        row 3: metric groups, forward-filled ('Volume (lakh)', 'Value (₹ crore)')
        rows 4-5: series labels ('FY 2025-26', '2025 July', '2026 June', '2026 July')
        row 6: column index ('1','2','3','4')
        row 7+: sections ('A. ...') and indicator rows (label col B)

    One output row per (indicator, metric, series, value):
        rbi_psi(month, part, section, indicator, metric, unit,
                series_type, series_label, value, src_file)
    """
    n = 0
    part = None
    for sheet_name, ws in sheets(path):
        rows = rows_of(ws)
        if not rows:
            continue

        # label column varies by vintage (col B in 2026, col C in 2022)
        label_col = 1
        for r in rows[:8]:
            for j, c in enumerate(r):
                if c is not None and "PART" in str(c):
                    label_col = j
                    break
            if label_col != 1:
                break

        # locate metric-group row and index row
        metric_row = idx_row = None
        for i, r in enumerate(rows[:10]):
            joined = " ".join(str(c) for c in r if c is not None)
            if metric_row is None and re.search(r"Volume|Value|Count", joined):
                metric_row = i
            cells = [str(c).strip() for c in r if c is not None]
            if cells and all(re.fullmatch(r"\d{1,2}", c) for c in cells) and len(cells) >= 3:
                idx_row = i
                break
        if metric_row is None or idx_row is None:
            continue

        # forward-fill metric labels across columns
        ncols = max((len(r) for r in rows), default=0)
        metric, unit = {}, {}
        cur_m = cur_u = None
        for c in range(ncols):
            v = rows[metric_row][c] if c < len(rows[metric_row]) else None
            if v is not None and str(v).strip():
                text = re.sub(r"\s+", " ", str(v)).strip()
                cur_m = text
                um = re.search(r"\(([^)]+)\)", text)
                cur_u = um.group(1) if um else ""
            metric[c], unit[c] = cur_m, cur_u

        # series labels = join of non-empty cells in each column between metric row and index row
        series = {}
        for c in range(ncols):
            bits = []
            for i in range(metric_row + 1, idx_row):
                v = rows[i][c] if c < len(rows[i]) else None
                if v is not None and str(v).strip():
                    bits.append(re.sub(r"\s+", " ", str(v)).strip())
            series[c] = " ".join(bits)

        label = section = None
        for r in rows[idx_row + 1:]:
            b = r[label_col] if len(r) > label_col else None
            if b is not None and str(b).strip():
                text = re.sub(r"\s+", " ", str(b)).strip()
                if text.upper().startswith("PART"):
                    part, section, label = text, None, None
                elif re.match(r"^[A-Z]\.\s", text):
                    section, label = text, None
                else:
                    label = text
            if not label or not re.search(r"Volume|Value|Count", " ".join(metric.get(c) or "" for c in range(ncols))):
                pass
            vals = [(c, num(r[c])) for c in range(ncols)
                    if c < len(r) and metric.get(c) and series.get(c)
                    and num(r[c]) is not None]
            if label and vals:
                for c, v in vals:
                    st, sl = parse_series_type(series[c], month)
                    con.execute(
                        "INSERT INTO rbi_psi VALUES (?,?,?,?,?,?,?,?,?,?)",
                        [month, part, section, label, metric[c], unit[c] or "",
                         st, sl, v, path.name])
                    n += 1
    return n


def ingest_bankwise(con, path, month):
    """Bankwise volumes: one sheet per rail (NEFT, RTGS, ...).
    Cols: Sl, BANK NAME, inward txns, inward amt, outward txns, outward amt."""
    n = 0
    errs = []
    for sheet_name, ws in sheets(path):
        rows = rows_of(ws)
        for r in rows:
            if len(r) < 6:
                continue
            bank = norm_bank(r[2])
            if not bank or str(r[1]).strip() in ("", "Sl. No."):
                continue
            if not re.search(r"[A-Za-z]", str(bank)):
                continue
            vals = [num(r[3]), num(r[4]), num(r[5]), num(r[6]) if len(r) > 6 else None]
            con.execute(
                "INSERT INTO rbi_bankwise VALUES (?,?,?,?,?,?,?,?,?)",
                [month, sheet_name, bank, vals[0], vals[1], vals[2], vals[3],
                 json.dumps({"raw": [str(x) for x in r[1:7]]}, ensure_ascii=False),
                 path.name])
            n += 1
    return n


def ingest_atmpos(con, path, month):
    """ATM/POS/Card: cols = Sr, Bank Name, then infrastructure counts in
    first data column of each group (ATMs, PoS, MicroATMs, BharatQR, Cards)."""
    n = 0
    errs = []
    for sheet_name, ws in sheets(path):
        rows = rows_of(ws)
        # map column index -> group header by scanning rows 0-8
        colgroup = {}
        last = None
        for i in range(0, min(9, len(rows))):
            r = rows[i]
            for c, v in enumerate(r):
                if v and str(v).strip() and c >= 3:
                    t = str(v).replace("\n", " ").strip()
                    if re.search(r"[A-Za-z]", t) and len(t) > 2:
                        last = (c, t)
                if last and last[0] <= c:
                    colgroup.setdefault(c, last[1])
        for r in rows:
            if len(r) < 5:
                continue
            bank = norm_bank(r[2])
            if not bank or not re.search(r"[A-Za-z]", str(bank)):
                continue
            vals = {}
            for c in range(3, len(r)):
                v = num(r[c])
                if v is not None and c in colgroup:
                    vals[colgroup[c]] = v
            if vals:
                con.execute(
                    "INSERT INTO rbi_atmposcard VALUES (?,?,?,?,?,?)",
                    [month, sheet_name, bank, json.dumps(vals, ensure_ascii=False),
                     json.dumps({"n_cols": len(r)}, ensure_ascii=False), path.name])
                n += 1
    return n


def main():
    DB.parent.mkdir(parents=True, exist_ok=True)
    con = duckdb.connect(str(DB))
    con.execute("DROP TABLE IF EXISTS rbi_psi")
    con.execute("""
        CREATE TABLE IF NOT EXISTS rbi_psi (
            month DATE, part TEXT, section TEXT, indicator TEXT,
            metric TEXT, unit TEXT, series_type TEXT, series_label TEXT,
            value DOUBLE, src_file TEXT)""")
    con.execute("""
        CREATE TABLE IF NOT EXISTS rbi_bankwise (
            month DATE, rail TEXT, bank TEXT, inward_txn DOUBLE, inward_cr DOUBLE,
            outward_txn DOUBLE, outward_cr DOUBLE, payload JSON, src_file TEXT)""")
    con.execute("""
        CREATE TABLE IF NOT EXISTS rbi_atmposcard (
            month DATE, sheet TEXT, bank TEXT, infra JSON, meta JSON, src_file TEXT)""")
    con.execute("""
        CREATE TABLE IF NOT EXISTS rbi_file_log (
            kind TEXT, src_file TEXT, n_rows INT)""")

    jobs = {
        "psi": (RAW / "psi", ingest_psi),
        "bankwise-volumes": (RAW / "bankwise-volumes", ingest_bankwise),
        "atm-pos-card": (RAW / "atm-pos-card", ingest_atmpos),
    }
    try:
        con.execute("DELETE FROM rbi_file_log WHERE kind='psi'")
    except Exception:
        pass
    done_all = set(con.execute(
        "SELECT kind, src_file FROM rbi_file_log").fetchall())
    for kind, (d, fn) in jobs.items():
        files = sorted(d.glob("*.xlsx"))
        n_new = 0
        for f in files:
            if (kind, f.name) in done_all:
                continue
            month = month_from_name(f)
            if not month:
                continue
            try:
                n = fn(con, f, month)
            except Exception as e:  # noqa: BLE001
                print(f"  ERROR {kind}/{f.name}: {e}")
                continue
            con.execute("INSERT INTO rbi_file_log VALUES (?,?,?)", [kind, f.name, n])
            n_new += 1
        print(f"{kind}: +{n_new} files (total {len(files)})")
    con.close()
    print("rbi normalize done")


if __name__ == "__main__":
    main()
