"""Transaction matching script for Data Clean Room ML Jobs.

Reads pre-filtered candidate pairs from cleanroom.match_candidates (written by
the SQL pre-filter step), distributes a custom matching function across Ray
workers, and writes results to cleanroom.match_results.

Replace the custom_match function body with your matching algorithm.
"""

import argparse
import json
import traceback

import pandas as pd
import ray
from snowflake.snowpark.context import get_active_session


def custom_match(segment, candidates_df):
    """
    Replace this function body with your matching algorithm.

    Input:  pandas DataFrame of candidates for one segment.
            Columns: HASHED_EMAIL_SHA256, PROVIDER_AMOUNT, PARTNER_AMOUNT,
                     SEGMENT, PROVIDER_DATE, PARTNER_DATE
    Output: dict with segment, candidates, matched, matches keys

    Examples:
      Fuzzy amount + date:
        abs(row["PROVIDER_AMOUNT"] - row["PARTNER_AMOUNT"]) / row["PROVIDER_AMOUNT"] < 0.05
        and abs((row["PROVIDER_DATE"] - row["PARTNER_DATE"]).days) < 7

      Statistical scoring:
        score = model.predict_proba(features)[0][1]
        if score > threshold

      Clustering:
        labels = DBSCAN(eps=0.3).fit_predict(features)
    """
    matches = []
    for _, row in candidates_df.iterrows():
        amount_ratio = (
            abs(row["PROVIDER_AMOUNT"] - row["PARTNER_AMOUNT"])
            / max(row["PROVIDER_AMOUNT"], 0.01)
        )
        date_diff = (
            abs((row["PROVIDER_DATE"] - row["PARTNER_DATE"]).days)
            if pd.notna(row["PROVIDER_DATE"]) and pd.notna(row["PARTNER_DATE"])
            else 999
        )
        if amount_ratio < 0.05 and date_diff <= 7:
            matches.append({
                "hashed_email": str(row["HASHED_EMAIL_SHA256"]),
                "segment": str(row["SEGMENT"]),
                "provider_amount": float(row["PROVIDER_AMOUNT"]),
                "partner_amount": float(row["PARTNER_AMOUNT"]),
                "amount_diff": float(abs(row["PROVIDER_AMOUNT"] - row["PARTNER_AMOUNT"])),
                "date_diff_days": int(date_diff),
            })

    return {
        "segment": segment,
        "candidates": len(candidates_df),
        "matched": len(matches),
        "matches": matches,
    }


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--args", type=str, default="{}")
    args = json.loads(parser.parse_args().args)

    session = get_active_session()

    print("Reading match candidates from cleanroom.match_candidates...")
    candidates = session.table("cleanroom.match_candidates").to_pandas()
    candidates.columns = [c.upper() for c in candidates.columns]
    print(f"Candidates loaded: {len(candidates):,} rows")

    seg_col = "SEGMENT"
    segments = candidates[seg_col].unique().tolist()
    print(f"Segments: {segments}")

    ray.init(address="auto", ignore_reinit_error=True)
    print(f"Ray nodes: {len(ray.nodes())}")

    match_remote = ray.remote(custom_match)
    futures = []
    for seg in segments:
        seg_df = candidates[candidates[seg_col] == seg].copy()
        futures.append(match_remote.remote(seg, seg_df))

    results = ray.get(futures)
    ray.shutdown()
    print("Ray shutdown complete")

    total_matched = sum(r["matched"] for r in results)
    all_matches = []
    for r in results:
        all_matches.extend(r["matches"])

    print("\n=== RESULTS ===")
    print(f"Total matched: {total_matched:,}")
    for r in results:
        print(f"  {r['segment']:15s} {r['matched']:,} / {r['candidates']:,} candidates")

    if all_matches:
        result_df = session.create_dataframe(pd.DataFrame(all_matches))
        result_df.write.save_as_table(
            "cleanroom.match_results", mode="overwrite"
        )
        print("\nResults written to cleanroom.match_results")

    output = json.dumps({
        "total_matched": total_matched,
        "segments": len(results),
        "details": [
            {k: v for k, v in r.items() if k != "matches"} for r in results
        ],
    })
    print(output)
    return output


if __name__ == "__main__":
    try:
        main()
    except Exception as e:
        print(f"FATAL: {e}")
        traceback.print_exc()
        raise
