#!/usr/bin/env python3
"""01_match_balance.py
Build matched analysis samples and balance diagnostics for WP speed plugin event studies.

Inputs:
  --panel:    unique origin-device-month panel exported from BigQuery, e.g. delay_js_analysis_panel
  --baseline: risk-set baseline table, e.g. delay_js_riskset_baseline. If omitted, the
              script falls back to rows with event_time == -1 in --panel.

Outputs:
  <prefix>_pairs.csv
  <prefix>_panel.csv
  <prefix>_balance_summary.csv

The script is intentionally conservative. It matches within exact adoption-month/device strata,
then uses propensity-score nearest neighbours for baseline covariates. Failed balance should be
reported as a design limitation, not hidden.
"""
import argparse
import numpy as np
import pandas as pd
from sklearn.compose import ColumnTransformer
from sklearn.impute import SimpleImputer
from sklearn.linear_model import LogisticRegression
from sklearn.neighbors import NearestNeighbors
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import OneHotEncoder, StandardScaler


def smd(x_t, x_c):
    vt = np.nanvar(x_t, ddof=1)
    vc = np.nanvar(x_c, ddof=1)
    denom = np.sqrt((vt + vc) / 2.0)
    if denom == 0 or np.isnan(denom):
        return 0.0
    return (np.nanmean(x_t) - np.nanmean(x_c)) / denom


def parse_csv_list(value):
    return [v.strip() for v in value.split(',') if v.strip()]


def build_preprocessor(num_cols, cat_cols):
    return ColumnTransformer([
        ('num', Pipeline([('impute', SimpleImputer(strategy='median')), ('scale', StandardScaler())]), num_cols),
        ('cat', Pipeline([('impute', SimpleImputer(strategy='most_frequent')), ('oh', OneHotEncoder(handle_unknown='ignore'))]), cat_cols),
    ])


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--panel', required=True, help='unique origin-device-month analysis panel CSV')
    ap.add_argument('--baseline', default=None, help='risk-set baseline CSV; defaults to panel rows with event_time == -1')
    ap.add_argument('--treatment', default='treated_cohort', help='baseline treated indicator column')
    ap.add_argument('--exact_cols', default='candidate_event_month,device', help='comma-separated exact matching strata')
    ap.add_argument('--controls_per_treated', type=int, default=5)
    ap.add_argument('--output_prefix', default='matched')
    args = ap.parse_args()

    panel = pd.read_csv(args.panel).replace([np.inf, -np.inf], np.nan)
    if args.baseline:
        base = pd.read_csv(args.baseline).replace([np.inf, -np.inf], np.nan)
    else:
        if 'event_time' not in panel.columns:
            raise SystemExit('No --baseline supplied and panel has no event_time column.')
        base = panel[panel['event_time'].fillna(-999).eq(-1)].copy()

    base = base.copy().reset_index(drop=True)
    base['_row_id'] = np.arange(len(base))
    exact_cols = [c for c in parse_csv_list(args.exact_cols) if c in base.columns]
    if not exact_cols:
        exact_cols = ['_all']
        base['_all'] = 'all'

    if args.treatment not in base.columns:
        raise SystemExit(f'Treatment column {args.treatment!r} not found in baseline table.')
    base['_treated'] = base[args.treatment].fillna(0).astype(int)

    cov_num = [c for c in [
        'fast_lcp', 'fast_inp', 'small_cls', 'p75_lcp', 'p75_inp', 'p75_cls',
        'bytes_total', 'request_total', 'bytes_js', 'bytes_img', 'lighthouse_tbt',
        'lighthouse_performance_score', 'rank'
    ] if c in base.columns]
    cov_cat = [c for c in [
        'rank_bucket', 'host_cdn_server', 'theme_builder_ecom', 'device', 'client'
    ] if c in base.columns and c not in exact_cols]
    use_cols = cov_num + cov_cat
    if not use_cols:
        raise SystemExit('No usable covariates found; check input schema.')

    pairs = []
    skipped = []
    for strata_key, g in base.groupby(exact_cols, dropna=False):
        treated = g[g['_treated'] == 1].copy()
        controls = g[g['_treated'] == 0].copy()
        if len(treated) == 0 or len(controls) == 0:
            skipped.append({'stratum': str(strata_key), 'treated': len(treated), 'controls': len(controls), 'reason': 'empty arm'})
            continue

        X = g[use_cols]
        y = g['_treated'].astype(int)
        pre = build_preprocessor(cov_num, cov_cat)
        try:
            model = Pipeline([('pre', pre), ('logit', LogisticRegression(max_iter=2000, class_weight='balanced'))])
            model.fit(X, y)
            g = g.copy()
            g['_score'] = model.predict_proba(X)[:, 1]
            treated = g[g['_treated'] == 1].copy()
            controls = g[g['_treated'] == 0].copy()
            nn_features_t = treated[['_score']]
            nn_features_c = controls[['_score']]
        except Exception:
            # Fallback: nearest-neighbour distance in transformed covariate space.
            Xt = pre.fit_transform(X)
            g = g.copy()
            g['_vec_row'] = np.arange(len(g))
            treated = g[g['_treated'] == 1].copy()
            controls = g[g['_treated'] == 0].copy()
            nn_features_t = Xt[treated['_vec_row'].to_numpy()]
            nn_features_c = Xt[controls['_vec_row'].to_numpy()]

        n_neighbors = min(args.controls_per_treated, len(controls))
        nn = NearestNeighbors(n_neighbors=n_neighbors, metric='euclidean')
        nn.fit(nn_features_c)
        dist, idx = nn.kneighbors(nn_features_t)
        controls_reset = controls.reset_index(drop=True)
        for i, (_, tr) in enumerate(treated.reset_index(drop=True).iterrows()):
            for j, d in zip(idx[i], dist[i]):
                co = controls_reset.iloc[j]
                rec = {c: tr[c] for c in exact_cols}
                rec.update({
                    'treated_origin': tr['origin'],
                    'control_origin': co['origin'],
                    'treated_row_id': int(tr['_row_id']),
                    'control_row_id': int(co['_row_id']),
                    'distance': float(d),
                })
                if 'device' in tr.index:
                    rec['treated_device'] = tr['device']
                    rec['control_device'] = co['device']
                pairs.append(rec)

    pairs = pd.DataFrame(pairs)
    if pairs.empty:
        pd.DataFrame(skipped).to_csv(f'{args.output_prefix}_skipped_strata.csv', index=False)
        raise SystemExit('No matched pairs produced; inspect skipped strata and loosen exact matching only if justified.')

    matched_row_ids = set(pairs['treated_row_id']).union(set(pairs['control_row_id']))
    matched_base = base[base['_row_id'].isin(matched_row_ids)].copy()

    # Keep only matched origin-device units in the unique analysis panel.
    key_cols = ['origin'] + (['device'] if 'device' in panel.columns and 'device' in matched_base.columns else [])
    matched_keys = matched_base[key_cols].drop_duplicates()
    matched_panel = panel.merge(matched_keys, on=key_cols, how='inner')
    matched_panel['matched_weight'] = 1.0

    bal = []
    for c in cov_num:
        bal.append({
            'covariate': c,
            'smd': smd(matched_base.loc[matched_base['_treated'] == 1, c], matched_base.loc[matched_base['_treated'] == 0, c]),
            'treated_mean': np.nanmean(matched_base.loc[matched_base['_treated'] == 1, c]),
            'control_mean': np.nanmean(matched_base.loc[matched_base['_treated'] == 0, c]),
        })
    bal = pd.DataFrame(bal).sort_values('smd', key=lambda s: s.abs(), ascending=False)

    pairs.to_csv(f'{args.output_prefix}_pairs.csv', index=False)
    matched_panel.to_csv(f'{args.output_prefix}_panel.csv', index=False)
    bal.to_csv(f'{args.output_prefix}_balance_summary.csv', index=False)
    if skipped:
        pd.DataFrame(skipped).to_csv(f'{args.output_prefix}_skipped_strata.csv', index=False)
    print('Wrote matched files. Matched units:', len(matched_keys), 'Max |SMD|:', bal['smd'].abs().max() if len(bal) else 'NA')


if __name__ == '__main__':
    main()
