# Copyright (c) 2026 Jaime Yan. See LICENSE and CITATION.cff.
"""Descriptive, denominator-aware adapters over SciPy and Matplotlib."""
import numpy as np
import pandas as pd
from scipy import stats
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt

ARMS=['Reference','Investigational']
DOMAINS=['Fatigue','Pain','Sleep','Appetite','Mobility']
KINDS=['ecdf','completion','domain-intervals']
STRATA=['All','F','M']
COLORS=['#0072B2','#D55E00']

def validate(frames, score_unit='points_0_100', diameter_unit='mm'):
    """Validate explicit units, unique keys, finite ranges and retained schedules."""
    if score_unit!='points_0_100' or diameter_unit!='mm':
        raise ValueError('Expected points_0_100 and mm; convert explicitly before analysis')
    required={'subjects':['subject','arm','sex'],'domains':['subject','domain','week','score'],
              'tumor':['subject','week','diameter']}
    for k,cols in required.items():
        if k not in frames or not set(cols)<=set(frames[k].columns):raise ValueError(f'Missing table/columns: {k} {cols}')
    s,q,d=(frames[k] for k in ('subjects','domains','tumor'))
    if len(s)==0:raise ValueError('Empty subject roster')
    if s[['subject','arm','sex']].isna().any().any() or s.subject.astype(str).str.strip().eq('').any():raise ValueError('Missing roster field')
    if not s.arm.isin(ARMS).all() or not s.sex.isin(['F','M']).all():raise ValueError('Unknown arm or sex')
    for f,key in [(s,['subject']),(q,['subject','domain','week']),(d,['subject','week'])]:
        if f[key].isna().any().any() or f.duplicated(key).any():raise ValueError('Missing or duplicate key')
    for f,col in [(q,'score'),(d,'diameter')]:
        if pd.api.types.is_bool_dtype(f[col]) or not pd.api.types.is_numeric_dtype(f[col]):raise ValueError(f'{col} must be numeric, not boolean')
        x=f[col].dropna().to_numpy()
        if not np.isfinite(x).all():raise ValueError('Nonfinite measurement')
    if not q.score.dropna().between(0,100).all() or (d.diameter.dropna()<=0).any():raise ValueError('Measurement outside permitted range')
    expected_q={(sid,dom,w) for sid in s.subject for dom in DOMAINS for w in [0,12]}
    expected_d={(sid,w) for sid in s.subject for w in [0,4,8,12]}
    if set(q[['subject','domain','week']].itertuples(index=False,name=None))!=expected_q:raise ValueError('Domain grid mismatch / unknown subject')
    if set(d[['subject','week']].itertuples(index=False,name=None))!=expected_d:raise ValueError('Visit grid mismatch / unknown subject')

def prepare(kind,frames,**units):
    """Return unrounded rows for All/F/M; each stratum has its own denominators."""
    validate(frames,**units)
    if kind not in KINDS:raise ValueError('Unknown template')
    s,q,d=(frames[k].copy() for k in ('subjects','domains','tumor'))
    pairs=q.pivot(index=['subject','domain'],columns='week',values='score').reset_index()
    pairs['change']=pairs[12]-pairs[0]
    pairs=pairs.merge(s[['subject','arm','sex']],on='subject',validate='many_to_one')
    rows=[]
    for stratum in STRATA:
        roster=s if stratum=='All' else s.loc[s.sex==stratum]
        for arm in ARMS:
            ids=roster.loc[roster.arm==arm,'subject'];expected=len(ids)
            group=pairs.loc[pairs.subject.isin(ids)]
            if kind=='completion':
                for week in [0,4,8,12]:
                    n=int(d.loc[d.subject.isin(ids)&d.week.eq(week),'diameter'].count())
                    rows.append(dict(stratum=stratum,arm=arm,week=week,n=n,expected=expected,missing=expected-n,percent=100*n/expected if expected else np.nan))
            else:
                for dom in (['Fatigue'] if kind=='ecdf' else DOMAINS):
                    x=group.loc[group.domain==dom,'change'].dropna().to_numpy();n=len(x)
                    meta=dict(stratum=stratum,arm=arm,domain=dom,n=n,expected=expected,missing=expected-n)
                    if kind=='ecdf':
                        if n:
                            result=stats.ecdf(x).cdf
                            rows.extend(dict(**meta,change=float(a),probability=float(b)) for a,b in zip(result.quantiles,result.probabilities))
                        else:rows.append(dict(**meta,change=np.nan,probability=np.nan))
                    else:
                        mean=float(x.mean()) if n else np.nan
                        sd=float(x.std(ddof=1)) if n>=2 else np.nan
                        half=float(stats.t.ppf(.975,n-1)*sd/np.sqrt(n)) if n>=2 else np.nan
                        rows.append(dict(**meta,estimate=mean,sd=sd,lower=mean-half,upper=mean+half))
    return pd.DataFrame(rows)

def draw(kind,rows,stratum='All'):
    """Return an editable Figure for one exact precomputed stratum."""
    if kind not in KINDS or stratum not in STRATA:raise ValueError('Unknown template or stratum')
    g=rows.loc[rows.stratum==stratum]
    if g.empty:raise ValueError('Selected stratum absent from result')
    plt.rcParams.update({'font.family':'DejaVu Sans','font.size':11,'svg.fonttype':'none','pdf.fonttype':42,
                         'axes.spines.top':False,'axes.spines.right':False,'figure.facecolor':'white'})
    fig,ax=plt.subplots(figsize=(11,7));fig.subplots_adjust(left=.14,right=.96,top=.80,bottom=.23)
    titles={'ecdf':'How broadly is improvement distributed?','completion':'Who contributed at each scheduled visit?',
            'domain-intervals':'Symptom changes, with their precision.'}
    notes={'ecdf':'Complete-pair distribution; no confidence band. Missing pairs are excluded, not censored.',
           'completion':'Observed / full scheduled roster. A missing measurement is not a withdrawal.',
           'domain-intervals':'Pointwise 95% t intervals; no multiplicity adjustment. These are not treatment contrasts.'}
    for i,arm in enumerate(ARMS):
        a=g.loc[g.arm==arm];color=COLORS[i];marker=['o','s'][i]
        if kind=='ecdf':
            n=int(a.n.iloc[0]);total=int(a.expected.iloc[0]);valid=a.dropna(subset=['change'])
            if len(valid):
                x=valid.change.to_numpy();y=valid.probability.to_numpy()
                if x[0]>-100:x=np.r_[-100,x];y=np.r_[0,y]
                x=np.r_[x,100];y=np.r_[y,1]
                ax.step(x,y,where='post',color=color,linestyle=['-','--'][i],lw=2,label=f'{arm}: {n}/{total} complete')
            else:ax.plot([],[],color=color,label=f'{arm}: 0/{total}, unavailable')
            ax.set(xlim=(-40,15),ylim=(0,1.04),xlabel='Week-12 minus baseline Fatigue (points); negative = improvement',ylabel='Proportion with change ≤ x')
            # Expand to include every observed value; never crop extreme changes.
            finite=g.change.dropna()
            if len(finite):ax.set_xlim(min(-40,float(finite.min())-2),max(15,float(finite.max())+2))
            ax.axvline(0,color='#666666',ls=':',lw=1)
        elif kind=='completion':
            ax.plot(a.week,a.percent,color=color,marker=marker,linestyle=['-','--'][i],lw=2,label=arm)
            for r in a.itertuples():
                if np.isfinite(r.percent):
                    other=g.loc[g.arm.ne(arm)&g.week.eq(r.week),'percent'].iloc[0]
                    above=(not np.isfinite(other)) or r.percent>other or (r.percent==other and i==0)
                    ax.annotate(f'{r.n}/{r.expected}',(r.week,r.percent),xytext=(0,10 if above else -18),textcoords='offset points',ha='center',fontsize=10,color=color)
            ax.set(xticks=[0,4,8,12],ylim=(-5,110),xlabel='Scheduled week',ylabel='Observed measurements (%)')
        else:
            labelled=False
            for j,dom in enumerate(DOMAINS):
                r=a.loc[a.domain==dom].iloc[0];y=j+(i-.5)*.25
                if np.isfinite(r.estimate):
                    ax.plot(r.estimate,y,marker=marker,color=color,ls='',label=arm if not labelled else None)
                    labelled=True
                    if np.isfinite(r.lower):ax.hlines(y,r.lower,r.upper,color=color,lw=2)
                ax.text(1.01,y,f'{int(r.n)}/{int(r.expected)}',transform=ax.get_yaxis_transform(),va='center',color=color,fontsize=9)
            if not labelled:
                ax.plot([],[],marker=marker,color=color,ls='',label=f'{arm}: no complete pairs')
            ax.set(yticks=range(5),yticklabels=DOMAINS,xlabel='Mean within-person change (points); negative = improvement')
            ax.set_ylim(4.6,-.6);ax.axvline(0,color='#666666',ls=':',lw=1);fig.subplots_adjust(right=.87)
            ax.text(1.01,1.03,'Pairs / roster',transform=ax.transAxes,fontsize=9)
    ax.grid(axis='x' if kind=='domain-intervals' else 'y',alpha=.15)
    ax.legend(loc='upper center',bbox_to_anchor=(.5,-.20),ncol=2,frameon=False,fontsize=10)
    fig.suptitle(titles[kind],x=.06,y=.965,ha='left',fontsize=20,weight='bold')
    fig.text(.06,.895,f'SYNTHETIC TEACHING DATA  |  Stratum: {stratum}  |  Independent Python computation',fontsize=10)
    fig.text(.06,.06,notes[kind],fontsize=9)
    fig.text(.06,.025,'Jaime Yan · Clinical Figure Library · Personal noncommercial use · Attribution and citation required',fontsize=8)
    return fig
