"""
Mismatch Check — RO4 LAPORAN PEMBAYARAN KLAIM 2026
1. Proses Bayar Klaim: Klaim Disetujui vs Klaim Dibayar + Klaim Proses Bayar
2. Klaim Disetujui (PBK YTD) vs Klaim Disetujui (sheet per bulan, cumulative)
Output: mismatch_YYYY-MM-DD_vN.xlsx (multi-sheet)
"""

from pathlib import Path
import argparse
from datetime import date, datetime
from collections import defaultdict, Counter
import time
import openpyxl
from openpyxl.styles import Font, PatternFill, Alignment, Border, Side
from openpyxl.utils import get_column_letter

PER = (datetime, date)  # openpyxl bisa return datetime atau date

REPO_ROOT = Path(__file__).resolve().parents[1]
DEFAULT_SRC = REPO_ROOT / "Data Dummy - RO4 - LAPORAN PEMBAYARAN KLAIM 2026.xlsx"
DEFAULT_OUT_DIR = REPO_ROOT / "mismatch-output"


def to_float(v):
    if v is None:
        return 0.0
    if isinstance(v, str):
        v = v.strip()
        if not v:
            return 0.0
        try:
            # Format: koma = pemisah ribuan, titik = desimal (mis. "1,234,567.89")
            return round(float(v.replace(",", "")), 2)
        except:
            return 0.0
    return round(float(v), 2)


RED_FILL = PatternFill("solid", fgColor="FCE4EC")
GREEN_FILL = PatternFill("solid", fgColor="E8F5E9")
ORANGE_FILL = PatternFill("solid", fgColor="FFE0B2")
HEADER_FILL = PatternFill("solid", fgColor="37474F")
HF = Font(bold=True, color="FFFFFF", size=11)
THIN_BORDER = Border(
    left=Side(style="thin"), right=Side(style="thin"),
    top=Side(style="thin"), bottom=Side(style="thin"),
)


def read_sheet(wb, sheet_name):
    ws = wb[sheet_name]
    rows = []
    for row in ws.iter_rows(min_row=2, values_only=True):
        if row[0] is None:
            continue
        rows.append({
            "branch": row[0], "periode": row[1], "jenis": row[2],
            "cob": row[3], "toc": row[4], "toc_ket": row[5],
            "jumlah": to_float(row[6]), "net_klaim": to_float(row[7]),
        })
    return rows


def style_cell(cell, is_number=False):
    cell.border = THIN_BORDER
    if is_number:
        cell.number_format = "#,##0.##########"
        cell.alignment = Alignment(horizontal="right")


def write_header(ws, headers):
    for c, h in enumerate(headers, 1):
        x = ws.cell(row=1, column=c, value=h)
        x.font = HF; x.fill = HEADER_FILL; x.alignment = Alignment(horizontal="center")
        x.border = THIN_BORDER


def build_kd_cumulative(kd_rows):
    """Build KD cumulative: {(branch, cob, toc): [(month, cum_val), ...]}"""
    kd_monthly = defaultdict(float)
    for r in kd_rows:
        if "disetujui" not in str(r["jenis"] or "").lower():
            continue
        per = r["periode"]
        mo = per.month if isinstance(per, PER) else 0
        kd_monthly[(r["branch"], mo, r["cob"], r["toc"])] += r["net_klaim"]

    kd_by = defaultdict(list)
    for (br, mo, cob, toc), val in kd_monthly.items():
        kd_by[(br, cob, toc)].append((mo, val))

    kd_cum = {}
    for key, months in kd_by.items():
        months.sort()
        cum_list = []
        c = 0.0
        for mo, val in months:
            c += val
            cum_list.append((mo, c))
        kd_cum[key] = cum_list
    return kd_cum


def lookup_kd_cum(kd_cum, branch, periode, cob, toc):
    """Look up KD cumulative for a branch/COB/TOC up to given month."""
    mo = periode.month if isinstance(periode, PER) else 0
    cum_list = kd_cum.get((branch, cob, toc), [])
    kd_v = 0.0
    for m, cv in cum_list:
        if m <= mo:
            kd_v = cv
        else:
            break
    return kd_v


def check1_internal(pbk_rows, kd_cum=None):
    JENIS_MAP = {
        "klaim dibayar": "dibayar",
        "klaim proses bayar": "proses_bayar",
        "klaim disetujui": "disetujui",
    }
    groups = defaultdict(lambda: defaultdict(float))
    for r in pbk_rows:
        key = (r["branch"], r["periode"], r["cob"], r["toc"])
        j = str(r["jenis"] or "").strip().lower()
        net = r["net_klaim"]
        bucket = JENIS_MAP.get(j)
        if bucket:
            groups[key][bucket] += net
    out = []
    for key, v in groups.items():
        dis = v.get("disetujui", 0.0)
        hitung = v.get("dibayar", 0.0) + v.get("proses_bayar", 0.0)
        if abs(dis - hitung) > 1:
            kd_cum_val = lookup_kd_cum(kd_cum, *key) if kd_cum else 0.0
            out.append({"key": key, "disetujui": dis, "dibayar": v.get("dibayar", 0.0),
                        "proses": v.get("proses_bayar", 0.0), "hitung": hitung, "diff": dis - hitung,
                        "kd_cum": kd_cum_val})
    out.sort(key=lambda x: (str(x["key"][0]), str(x["key"][1] or ""), str(x["key"][2])))
    return out


def check2_ytd_vs_monthly(pbk_rows, kd_cum):
    pbk_ytd = defaultdict(float)
    for r in pbk_rows:
        if "disetujui" not in str(r["jenis"] or "").lower():
            continue
        per = r["periode"]
        mo = per.month if isinstance(per, PER) else 0
        pbk_ytd[(r["branch"], mo, r["cob"], r["toc"])] += r["net_klaim"]

    out = []
    for key, pbk_val in pbk_ytd.items():
        br, mo, cob, toc = key
        cum_list = kd_cum.get((br, cob, toc), [])
        kd_v = 0.0
        for m, cv in cum_list:
            if m <= mo:
                kd_v = cv
            else:
                break
        d = round(pbk_val - kd_v, 2)
        if abs(d) > 1:
            # Detect if PBK is 0 (not yet input)
            mistype = "Belum Input" if pbk_val == 0.0 else "Tidak Sesuai"
            out.append({"key": key, "pbk_ytd": pbk_val, "kd_sum": kd_v, "diff": d, "type": mistype})

    # Check KD cumulative values missing in PBK
    for (br, cob, toc), cum_list in kd_cum.items():
        for mo, kd_v in cum_list:
            if (br, mo, cob, toc) not in pbk_ytd and kd_v > 1:
                out.append({"key": (br, mo, cob, toc), "pbk_ytd": 0.0, "kd_sum": kd_v, "diff": -kd_v, "type": "Belum Input"})

    out.sort(key=lambda x: (str(x["key"][0]), x["key"][1], str(x["key"][2])))
    return out


def save_output(m1, m2, out_dir):
    wb = openpyxl.Workbook()
    ws1 = wb.active
    ws1.title = "Internal Mismatch"
    write_header(ws1, ["Branch Office", "Periode", "COB", "TOC",
                        "Klaim Dibayar", "Klaim Proses Bayar",
                        "Disetujui (Sheet)", "Disetujui (Hitung)",
                        "Disetujui (KD Cumulative)", "Selisih", "Type"])
    for i, m in enumerate(m1, 2):
        br, per, cob, toc = m["key"]
        per_s = per.strftime("%Y-%m") if isinstance(per, datetime) else str(per)[:7]
        vals = [br, per_s, cob, toc, m["dibayar"], m["proses"], m["disetujui"], m["hitung"], m.get("kd_cum", 0.0), m["diff"], "Belum Input" if m["dibayar"] == 0 and m["proses"] == 0 else "Tidak Sesuai"]
        for c, v in enumerate(vals, 1):
            style_cell(ws1.cell(row=i, column=c, value=v), is_number=isinstance(v, float))
        ws1.cell(row=i, column=7).fill = RED_FILL
        ws1.cell(row=i, column=7).font = Font(bold=True, color="C62828")
        ws1.cell(row=i, column=8).fill = GREEN_FILL
        ws1.cell(row=i, column=8).font = Font(bold=True, color="2E7D32")
        ws1.cell(row=i, column=9).fill = ORANGE_FILL
        ws1.cell(row=i, column=10).fill = ORANGE_FILL
    for i, w in enumerate([18, 12, 10, 14, 18, 18, 20, 20, 22, 16], 1):
        ws1.column_dimensions[get_column_letter(i)].width = w
    ws1.freeze_panes = "A2"
    if m1:
        ws1.auto_filter.ref = f"A1:K{len(m1) + 1}"

    ws2 = wb.create_sheet("YTD vs KD Cumulative")
    write_header(ws2, ["Branch Office", "Bulan", "COB", "TOC",
                        "PBK Disetujui (YTD)", "KD Cumulative (s/d Bulan)", "Selisih", "Type"])
    for i, m in enumerate(m2, 2):
        br, mo, cob, toc = m["key"]
        vals = [br, mo, cob, toc, m["pbk_ytd"], m["kd_sum"], m["diff"], m.get("type", "")]
        for c, v in enumerate(vals, 1):
            style_cell(ws2.cell(row=i, column=c, value=v), is_number=isinstance(v, float))
        ws2.cell(row=i, column=5).fill = RED_FILL
        ws2.cell(row=i, column=6).fill = GREEN_FILL
        ws2.cell(row=i, column=7).fill = ORANGE_FILL
    for i, w in enumerate([18, 12, 10, 14, 22, 26, 16], 1):
        ws2.column_dimensions[get_column_letter(i)].width = w
    ws2.freeze_panes = "A2"
    if m2:
        ws2.auto_filter.ref = f"A1:G{len(m2) + 1}"

    today_str = date.today().strftime("%Y-%m-%d")
    existing = list(out_dir.glob(f"mismatch_{today_str}_v*.xlsx"))
    version = 1
    for f in existing:
        try:
            version = max(version, int(f.stem.rsplit("_v", 1)[1]) + 1)
        except (ValueError, IndexError):
            pass
    out = out_dir / f"mismatch_{today_str}_v{version}.xlsx"
    try:
        wb.save(str(out))
    except PermissionError:
        out = out_dir / f"mismatch_{today_str}_v{version}_{int(time.time())}.xlsx"
        wb.save(str(out))
    return out


def print_summary(m1, m2):
    c1 = Counter()
    for m in m1:
        c1[m["key"][0]] += 1
    c2 = Counter()
    for m in m2:
        c2[m["key"][0]] += 1
    all_branches = sorted(set(list(c1.keys()) + list(c2.keys())))
    total = len(m1) + len(m2)
    print(f"\n{'='*60}")
    print(f"  RINGKASAN MISMATCH — Total: {total}")
    print(f"{'='*60}")
    print(f"  {'Cabang':12} {'Check1':>7} {'Check2':>7} {'Total':>7}")
    print(f"  {'-'*35}")
    for br in all_branches:
        ch1 = c1.get(br, 0)
        ch2 = c2.get(br, 0)
        print(f"  {br:12} {ch1:>7} {ch2:>7} {ch1+ch2:>7}")
    print(f"  {'-'*35}")
    print(f"  {'TOTAL':12} {len(m1):>7} {len(m2):>7} {total:>7}")

    if m2:
        print(f"\n  Penyebab utama CHECK 2 (YTD vs KD):")
        m2s = sorted(m2, key=lambda x: abs(x["diff"]), reverse=True)[:5]
        for m in m2s:
            br, mo, cob, toc = m["key"]
            print(f"    {br:12} Bulan={mo} | {cob:6} {toc:12} | PBK={m['pbk_ytd']:>16,.2f} KD={m['kd_sum']:>16,.2f} Selisih={m['diff']:>14,.2f} [{m['type']}]")

    if m1:
        print(f"\n  Penyebab utama CHECK 1 (Internal):")
        m1s = sorted(m1, key=lambda x: abs(x["diff"]), reverse=True)[:5]
        for m in m1s:
            br, per, cob, toc = m["key"]
            per_s = per.strftime("%Y-%m") if isinstance(per, datetime) else str(per)[:7]
            print(f"    {br:12} {per_s} | {cob:6} {toc:12} | Disetujui={m['disetujui']:>16,.2f} Hitung={m['hitung']:>16,.2f} Selisih={m['diff']:>14,.2f}")


def main(argv=None):
    parser = argparse.ArgumentParser(description="Check mismatch laporan pembayaran klaim.")
    parser.add_argument(
        "--input",
        type=Path,
        default=DEFAULT_SRC,
        help="Path workbook input; default: Data Dummy - RO4 - LAPORAN PEMBAYARAN KLAIM 2026.xlsx",
    )
    parser.add_argument(
        "--output-dir",
        type=Path,
        default=DEFAULT_OUT_DIR,
        help="Folder output; default: mismatch-output",
    )
    args = parser.parse_args(argv)
    src = args.input.expanduser().resolve()
    out_dir = args.output_dir.expanduser().resolve()
    if not src.exists():
        parser.error(f"file input tidak ditemukan: {src}")
    out_dir.mkdir(parents=True, exist_ok=True)

    print(f"[INFO] Reading: {src}")
    wb = openpyxl.load_workbook(str(src), data_only=True)
    pbk = read_sheet(wb, "Proses Bayar Klaim")
    print(f"[INFO] PBK rows: {len(pbk)}")

    # Check duplicates in PBK
    dup_counter = Counter()
    for r in pbk:
        per = r["periode"]
        mo = per.month if isinstance(per, PER) else 0
        dup_counter[(r["branch"], mo, r["cob"], r["toc"], r["jenis"], r["jumlah"], r["net_klaim"])] += 1
    dups = {k: v for k, v in dup_counter.items() if v > 1}
    if dups:
        print(f"[WARN] {len(dups)} duplicate row(s) found in PBK:")
        for k, v in dups.items():
            print(f"       {k[0]} Bulan={k[1]} | {k[2]} | {k[3]} | {k[4]} | Net={k[6]:,.2f} ({v}x)")
    else:
        print("[OK] No duplicates in PBK")

    m2 = []
    kd_cum = {}
    if "Klaim Disetujui" in wb.sheetnames:
        kd = read_sheet(wb, "Klaim Disetujui")
        print(f"[INFO] KD rows: {len(kd)}")
        kd_cum = build_kd_cumulative(kd)
        m2 = check2_ytd_vs_monthly(pbk, kd_cum)
        print(f"[CHECK 2] PBK YTD vs KD Cumulative: {len(m2)}")

    m1 = check1_internal(pbk, kd_cum)
    print(f"[CHECK 1] Disetujui vs Dibayar+Proses: {len(m1)}")
    wb.close()

    if m1 or m2:
        p = save_output(m1, m2, out_dir)
        print(f"[OK] Saved: {p}")
        print_summary(m1, m2)
    else:
        print("[OK] No mismatches")
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
