# Copyright (c) 2026 Jaime Yan. See LICENSE and CITATION.cff.
"""Complete-case ANCOVA teaching adapter; not a randomized-population estimand."""
import numpy as np
import pandas as pd
import statsmodels.api as sm
import matplotlib.pyplot as plt
from clinical_extensions import validate,DOMAINS,STRATA,ARMS


def prepare_adjusted(frames,**units):
    validate(frames,**units)
    s,q=frames['subjects'],frames['domains']
    pairs=q.pivot(index=['subject','domain'],columns='week',values='score').reset_index().merge(s[['subject','arm','sex']],on='subject',validate='many_to_one')
    rows=[]
    for st in STRATA:
        roster=s if st=='All' else s.loc[s.sex==st]
        for domain in DOMAINS:
            h=pairs.loc[pairs.subject.isin(roster.subject)&pairs.domain.eq(domain)].dropna(subset=[0,12])
            n=len(h);nr=int(h.arm.eq(ARMS[0]).sum());ni=int(h.arm.eq(ARMS[1]).sum())
            row=dict(stratum=st,arm='Investigational - Reference',domain=domain,n=n,expected=len(roster),missing=len(roster)-n,
                     n_reference=nr,n_investigational=ni,estimate=np.nan,se=np.nan,lower=np.nan,upper=np.nan,df=np.nan,status='Unavailable: insufficient complete cases')
            if nr>=2 and ni>=2 and n>3:
                x=np.column_stack([np.ones(n),h.arm.eq(ARMS[1]).astype(float),h[0].to_numpy()-h[0].mean()])
                if np.linalg.matrix_rank(x)<3 or np.linalg.cond(x)>1e6:
                    row['status']='Unavailable: singular or ill-conditioned design'
                else:
                    fit=sm.OLS(h[12].to_numpy(),x,missing='raise',hasconst=True).fit(use_t=True)
                    low,high=fit.conf_int(alpha=.05)[1]
                    row.update(estimate=float(fit.params[1]),se=float(fit.bse[1]),lower=float(low),upper=float(high),df=int(fit.df_resid),status='Estimable')
            rows.append(row)
    return pd.DataFrame(rows)


def draw_adjusted(rows,stratum='All'):
    if stratum not in STRATA:raise ValueError('Unknown stratum')
    g=rows.loc[rows.stratum==stratum]
    if len(g)!=5:raise ValueError('Expected five domain rows')
    plt.rcParams.update({'font.family':'DejaVu Sans','svg.fonttype':'none','pdf.fonttype':42})
    fig,ax=plt.subplots(figsize=(11,7));fig.subplots_adjust(left=.16,right=.78,top=.78,bottom=.25)
    for i,domain in enumerate(DOMAINS):
        r=g.loc[g.domain==domain].iloc[0]
        if np.isfinite(r.estimate):ax.errorbar(r.estimate,i,xerr=[[r.estimate-r.lower],[r.upper-r.estimate]],fmt='s',color='#0072B2',capsize=3)
        else:ax.text(.02,i,'Unavailable',transform=ax.get_yaxis_transform(),va='center',fontsize=10)
        ax.text(1.02,i,f'{int(r.n_reference)} / {int(r.n_investigational)}',transform=ax.get_yaxis_transform(),va='center',fontsize=10)
    ax.set(yticks=range(5),yticklabels=DOMAINS,ylim=(4.7,-.7),xlabel='Adjusted Week-12 contrast (points)\nInvestigational minus Reference; negative favors lower scores')
    ax.axvline(0,color='#555555',ls=':',lw=1);ax.grid(axis='x',alpha=.15)
    ax.text(1.02,1.04,'Complete cases\nReference / Investigational',transform=ax.transAxes,fontsize=9)
    fig.suptitle('Baseline-adjusted domain contrasts',x=.06,y=.965,ha='left',fontsize=20,weight='bold')
    fig.text(.06,.89,f'SYNTHETIC TEACHING DATA | Stratum: {stratum} | Independent Python computation',fontsize=10)
    fig.text(.06,.12,'OLS: Week 12 ~ intercept + treatment + baseline. Pointwise 95% t intervals; no multiplicity adjustment.',fontsize=9)
    fig.text(.06,.075,'Complete cases only; no missing-data correction. These invented scores do not establish clinical benefit.',fontsize=9)
    fig.text(.06,.03,'Jaime Yan · Clinical Figure Library · Personal noncommercial use · Attribution required',fontsize=8)
    return fig
