"""Matplotlib exports from independently calculated Python results.
Run: python figures.py outputs/python figures/python
"""
import sys,json
from pathlib import Path
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import numpy as np
source,dest=map(Path,sys.argv[1:3]);dest.mkdir(parents=True,exist_ok=True)
colors={'Placebo':'#0072B2','Xanomeline Low Dose':'#009E73','Xanomeline High Dose':'#CC79A7'}
for case in ['baseline','disposition','longitudinal','adverse-events','laboratory','subgroups','time-to-event']:
    data=[r for r in json.loads((source/(case+'.json')).read_text(encoding='utf8')) if r['stratum']=='All']
    if case=='baseline':data=[r for r in data if r['measure']=='Age (years)']
    if case=='longitudinal':data=[r for r in data if r['measure']=='Change from baseline']
    if case=='adverse-events':data=[r for r in data if r['level']=='SOC']
    if case=='disposition':data=[r for r in data if r['measure']!='Analysis cohort']
    fig,ax=plt.subplots(figsize=(10,9 if case=='adverse-events' else 6),layout='constrained')
    labels=list(dict.fromkeys(r['measure'] for r in data))
    for i,(arm,color) in enumerate(colors.items()):
        rows=[r for r in data if r['arm']==arm]
        if not rows:continue
        if case in ['longitudinal','time-to-event']:
            rows=[r for r in rows if r['estimate'] is not None]
            rows.sort(key=lambda r:r['time']);x=[r['time'] for r in rows];y=[r['estimate'] for r in rows]
            if case=='time-to-event':
                ax.step(x,y,where='post',label=arm,color=color)
                c=[r for r in rows if r['censored']];ax.scatter([r['time'] for r in c],[r['estimate'] for r in c],marker='+',color=color)
                ax.set(xlabel='Days since treatment start',ylabel='Event-free probability',ylim=(0,1))
            else:
                err=np.array([[r['estimate']-r['lower'] for r in rows],[r['upper']-r['estimate'] for r in rows]])
                ax.errorbar(x,y,yerr=err,marker='o',capsize=3,label=arm,color=color)
                ax.set(xlabel='Analysis visit (weeks)',ylabel='ADAS-Cog change (points)')
        else:
            ys=np.array([labels.index(r['measure']) for r in rows])+(i-1)*.21
            valid=[(r,y) for r,y in zip(rows,ys) if r['estimate'] is not None]
            xx=[r['estimate'] for r,y in valid]; yy=[y for r,y in valid]
            error=np.array([[r['estimate']-r['lower'] for r,y in valid],[r['upper']-r['estimate'] for r,y in valid]]) if all(r.get('lower') is not None for r,y in valid) else None
            ax.errorbar(xx,yy,xerr=error,fmt='o',capsize=3,label=arm,color=color)
            ax.set_yticks(range(len(labels)),labels)
            ax.set_xlabel('Mean age (years), 95% CI' if case=='baseline' else 'Active - placebo change (points), 95% CI' if case=='subgroups' else 'Participants (%)')
    ax.set_title(f'CDISC Pilot | {case}\nAll participants | independently computed in Python',loc='left',pad=18)
    if case not in ['longitudinal','time-to-event']:ax.invert_yaxis()
    if case=='subgroups':ax.axvline(0,color='#78716c',ls='--',lw=1)
    ax.spines[['top','right']].set_visible(False);ax.grid(axis='x',alpha=.18);ax.legend(loc='upper center',bbox_to_anchor=(.5,-.14),ncol=3,frameon=False,fontsize=8)
    fig.text(.01,.005,'Educational reanalysis. See case methods. Source: CDISC Pilot, 667511d.',fontsize=8,color='#57534e')
    fig.savefig(dest/(case+'.svg'));plt.close(fig)
print('Python figures: 7 SVG exports; matplotlib',matplotlib.__version__)
