#!/usr/bin/env python3
"""
magnitude_variance_analysis.py
===============================
Spatially Local Pre-Seismic Magnitude Variance Analysis
Electronic Supplement to BSSA submission

USAGE:
    python magnitude_variance_analysis.py --catalog jma_M3plus_2000_2023.csv \
        --mc 3.0 --mthresh 6.0 --iso-days 60 --iso-km 200 --radius 100

REQUIREMENTS: Python 3.11+, numpy, pandas, scipy, matplotlib

METHOD:
    For each isolated mainshock (M >= mthresh, spatiotemporally isolated):
    1. Select all events within R km of the mainshock epicenter
    2. Compute rolling std(magnitude) in W-event windows
    3. Extract indicator at pre-seismic lags (-21, -14, -7, -3, -1 days)
    4. Test suppression via Wilcoxon signed-rank with BH-FDR
    5. Classify windows into S1-S4 via training-period quartiles
    6. Measure S4 fraction, persistence (lock-in), and post-event reset

No model fitting. No free parameters beyond W and R. Deterministic (seed=42).
"""

import numpy as np
import pandas as pd
from scipy import stats
import argparse, os, sys
np.random.seed(42)

LAGS = [-21, -14, -7, -3, -1]
FDR_ALPHA = 0.05
S4_LOOKBACK = 20

def haversine_km(lat1, lon1, lat2, lon2):
    """Great-circle distance in km."""
    R = 6371.0
    dlat, dlon = np.radians(lat2-lat1), np.radians(lon2-lon1)
    a = np.sin(dlat/2)**2 + np.cos(np.radians(lat1))*np.cos(np.radians(lat2))*np.sin(dlon/2)**2
    return R * 2 * np.arctan2(np.sqrt(a), np.sqrt(1-a))

def find_clean_mainshocks(df, mthresh, iso_days, iso_km):
    """Spatiotemporally isolated mainshocks."""
    big = df[df['mag'] >= mthresh]
    if len(big) == 0: return []
    bt, blat, blon = big['time'].values, big['lat'].values, big['lon'].values
    clean = []
    for i, idx in enumerate(big.index):
        t, la, lo = bt[i], blat[i], blon[i]
        isolated = True
        for j in range(len(bt)):
            if i == j: continue
            if abs((bt[j]-t) / np.timedelta64(1,'D')) <= iso_days:
                if haversine_km(la, lo, blat[j], blon[j]) <= iso_km:
                    isolated = False; break
        if isolated: clean.append(idx)
    return clean

def benjamini_hochberg(pvals, alpha=FDR_ALPHA):
    """BH-FDR correction."""
    n = len(pvals)
    valid = [(i, p) for i, p in enumerate(pvals) if not np.isnan(p)]
    if not valid: return [False] * n
    sv = sorted(valid, key=lambda x: x[1])
    sig = [False] * n; mx = -1
    for rank, (oi, p) in enumerate(sv, 1):
        if p <= (rank / len(sv)) * alpha: mx = rank
    if mx > 0:
        for rank, (oi, _) in enumerate(sv, 1):
            if rank <= mx: sig[oi] = True
    return sig

def run_local_analysis(df, mthresh, iso_days, iso_km, radius, window=30):
    """Complete local magnitude variance analysis."""
    times = df['time'].values
    lats, lons, mags = df['lat'].values, df['lon'].values, df['mag'].values

    clean = find_clean_mainshocks(df, mthresh, iso_days, iso_km)
    print(f"  Found {len(clean)} clean M>={mthresh} mainshocks ({iso_days}d, {iso_km}km)")

    lag_ratios = {l: [] for l in LAGS}
    pre_s4_fracs, post_s4_fracs, bg_s4_fracs = [], [], []
    pre_persist, bg_persist = [], []
    pre_states = {1:0, 2:0, 3:0, 4:0}
    post_states = {1:0, 2:0, 3:0, 4:0}
    n_usable = 0

    for mi in clean:
        mlat, mlon, mtime = lats[mi], lons[mi], times[mi]
        lo = max(0, mi - 8000); hi = min(len(df), mi + 2000)
        dists = np.array([haversine_km(mlat, mlon, lats[j], lons[j])
                          for j in range(lo, hi)])
        local_idx = np.arange(lo, hi)[dists <= radius]

        if len(local_idx) < window + 30: continue
        local_mags = mags[local_idx]
        local_times = times[local_idx]
        nl = len(local_mags)

        # Rolling std(magnitude)
        rstd = np.full(nl, np.nan)
        for k in range(window-1, nl):
            rstd[k] = np.std(local_mags[k-window+1:k+1])
        bg = np.nanmean(rstd)
        if np.isnan(bg) or bg <= 0: continue

        # Lag extraction
        for lag in LAGS:
            target = mtime + np.timedelta64(lag, 'D')
            dt = np.abs((local_times - target) / np.timedelta64(1, 'D'))
            c = np.argmin(dt)
            if dt[c] < 3.0 and not np.isnan(rstd[c]):
                lag_ratios[lag].append(rstd[c] / bg)

        # S1-S4 classification
        valid_csd = rstd[~np.isnan(rstd)]
        if len(valid_csd) < 40: continue
        train_end = int(0.40 * len(valid_csd))
        Q1 = np.percentile(valid_csd[:train_end], 25)
        Q2 = np.percentile(valid_csd[:train_end], 50)
        Q3 = np.percentile(valid_csd[:train_end], 75)

        states = np.zeros(nl, dtype=int)
        for k in range(nl):
            if np.isnan(rstd[k]): states[k] = 0
            elif rstd[k] > Q3: states[k] = 1
            elif rstd[k] > Q2: states[k] = 2
            elif rstd[k] > Q1: states[k] = 3
            else: states[k] = 4

        mpos = np.argmin(np.abs((local_times - mtime) / np.timedelta64(1, 'D')))
        pre_w = states[max(0, mpos-S4_LOOKBACK):mpos]
        pre_v = pre_w[pre_w > 0]
        if len(pre_v) < 5: continue

        pre_s4_fracs.append(np.sum(pre_v == 4) / len(pre_v))
        for s in [1,2,3,4]: pre_states[s] += np.sum(pre_v == s)

        post_w = states[mpos+1:min(nl, mpos+S4_LOOKBACK+1)]
        post_v = post_w[post_w > 0]
        if len(post_v) >= 5:
            post_s4_fracs.append(np.sum(post_v == 4) / len(post_v))
            for s in [1,2,3,4]: post_states[s] += np.sum(post_v == s)

        s4t = s4s = 0
        for k in range(max(0, mpos-S4_LOOKBACK)+1, mpos):
            if states[k-1] == 4: s4t += 1; s4s += (states[k] == 4)
        if s4t >= 2: pre_persist.append(s4s / s4t)

        train_st = states[window-1:window-1+train_end]
        tv = train_st[train_st > 0]
        if len(tv) > 10: bg_s4_fracs.append(np.sum(tv == 4) / len(tv))
        bt = bs = 0
        for k in range(1, len(train_st)):
            if train_st[k-1] == 4: bt += 1; bs += (train_st[k] == 4)
        if bt >= 2: bg_persist.append(bs / bt)
        n_usable += 1

    # Statistical tests
    pvals = []; lag_stats = {}
    for lag in LAGS:
        vals = np.array(lag_ratios[lag])
        if len(vals) < 5:
            lag_stats[lag] = (np.nan, np.nan, len(vals)); pvals.append(np.nan)
        else:
            shifted = vals - 1.0
            try: _, p = stats.wilcoxon(shifted, alternative='less')
            except: p = np.nan
            lag_stats[lag] = ((np.mean(vals)-1.0)*100, p, len(vals)); pvals.append(p)
    fdr = benjamini_hochberg(pvals)

    # Print results
    print(f"\n  RESULTS (R={radius}km, W={window}, n={n_usable} usable):")
    print(f"  {'Lag':>8} {'Change':>10} {'p-value':>12} {'n':>5} {'FDR':>5}")
    print(f"  {'-'*45}")
    for i, lag in enumerate(LAGS):
        pct, p, n = lag_stats[lag]
        f_str = " *" if fdr[i] else ""
        if not np.isnan(pct):
            print(f"  {lag:>+5d} d {pct:>+8.1f}% {p:>12.6f} {n:>5}{f_str}")
    print(f"  FDR significant: {sum(fdr)}/5")

    if pre_s4_fracs:
        print(f"\n  S4 STATE ANALYSIS:")
        print(f"    Pre-event S4:  {np.mean(pre_s4_fracs)*100:.1f}% (bg={np.mean(bg_s4_fracs)*100:.1f}%)")
        if post_s4_fracs:
            print(f"    Post-event S4: {np.mean(post_s4_fracs)*100:.1f}%")
        if pre_persist and bg_persist and len(pre_persist) >= 5:
            _, p_lock = stats.mannwhitneyu(pre_persist, bg_persist, alternative='greater')
            print(f"    Lock-in p:     {p_lock:.8f}")

if __name__ == '__main__':
    parser = argparse.ArgumentParser(description='Magnitude Variance Analysis')
    parser.add_argument('--catalog', required=True, help='CSV catalog file')
    parser.add_argument('--mc', type=float, default=3.0)
    parser.add_argument('--mthresh', type=float, default=6.0)
    parser.add_argument('--iso-days', type=int, default=60)
    parser.add_argument('--iso-km', type=int, default=200)
    parser.add_argument('--radius', type=int, default=100)
    parser.add_argument('--window', type=int, default=30)
    args = parser.parse_args()

    df = pd.read_csv(args.catalog)
    df['time'] = pd.to_datetime(df['time'], utc=True, errors='coerce')
    df = df.dropna(subset=['time', 'mag'])
    df = df[df['mag'] >= args.mc].sort_values('time').reset_index(drop=True)
    print(f"Loaded {len(df):,} events (M>={args.mc})")

    run_local_analysis(df, args.mthresh, args.iso_days, args.iso_km, args.radius, args.window)
