# Copyright (c) 2026 Jaime Yan. See LICENSE and CITATION.cff.
"""Scheduled observation matrix; no imputation or inferred dropout events."""
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from matplotlib.patches import Rectangle
from clinical_extensions import ARMS, STRATA, validate

WEEKS = [0, 4, 8, 12]


def prepare_matrix(frames, **units):
    """Return one row per scheduled subject/visit within each overlapping stratum.

    Input uses the existing teaching contract: full subjects/domains/tumor grids.
    Counts n/expected/missing describe the arm and visit, not individual risks.
    """
    validate(frames, **units)
    s, d = frames['subjects'], frames['tumor']
    result = []
    for st in STRATA:
        roster = s if st == 'All' else s.loc[s.sex == st]
        for arm in ARMS:
            ids = sorted(roster.loc[roster.arm == arm, 'subject'])
            values = d.loc[d.subject.isin(ids)].pivot(index='subject', columns='week', values='diameter').reindex(index=ids, columns=WEEKS)
            observed = values.notna()
            for sid in ids:
                flags = observed.loc[sid].to_numpy()
                if flags.all(): pattern = 'Complete'
                elif not flags.any(): pattern = 'All missing'
                elif (~flags[np.flatnonzero(flags)[-1]+1:]).any() and flags[:np.flatnonzero(flags)[-1]+1].all(): pattern = 'Trailing missing'
                else: pattern = 'Intermittent or early missing'
                for week in WEEKS:
                    n = int(observed[week].sum())
                    result.append(dict(stratum=st, arm=arm, subject=sid, week=week,
                                       status='Observed' if observed.loc[sid, week] else 'Missing', pattern=pattern,
                                       n=n, expected=len(ids), missing=len(ids)-n))
    return pd.DataFrame(result, columns=['stratum','arm','subject','week','status','pattern','n','expected','missing'])


def draw_matrix(rows, stratum='All'):
    """Return an editable Figure; O and X encode status independently of color."""
    if stratum not in STRATA: raise ValueError('Unknown stratum')
    h = rows.loc[rows.stratum == stratum]
    plt.rcParams.update({'font.family':'DejaVu Sans','svg.fonttype':'none','pdf.fonttype':42})
    fig, axes = plt.subplots(1, 2, figsize=(11, 9), gridspec_kw={'wspace':.45})
    fig.subplots_adjust(left=.09, right=.96, top=.80, bottom=.20)
    for ax, arm in zip(axes, ARMS):
        g = h.loc[h.arm == arm]; ids = sorted(g.subject.unique())
        ax.set_title(f'{arm} (N={len(ids)})', fontsize=13)
        if not ids:
            ax.text(.5,.5,'No participants',ha='center',va='center',transform=ax.transAxes)
            ax.set_axis_off(); continue
        for i, sid in enumerate(ids):
            for j, week in enumerate(WEEKS):
                r = g.loc[g.subject.eq(sid) & g.week.eq(week)].iloc[0]
                observed = r.status == 'Observed'
                ax.add_patch(Rectangle((j-.48,i-.46),.96,.92,facecolor='#0072B2' if observed else '#FFFFFF',edgecolor='#555555',lw=.6))
                ax.text(j,i,'O' if observed else 'X',ha='center',va='center',fontsize=9,color='white' if observed else '#333333')
        labels=[]
        for week in WEEKS:
            r=g.loc[g.week==week].iloc[0];labels.append(f'{week}\n{int(r.n)}/{int(r.expected)}')
        ax.set(xticks=range(4),xticklabels=labels,yticks=range(len(ids)),yticklabels=ids,xlim=(-.5,3.5),ylim=(len(ids)-.5,-.5),xlabel='Scheduled week\nObserved / roster',ylabel='Fictional subject')
        ax.tick_params(labelsize=9)
    fig.suptitle('Which scheduled observations are missing?',x=.06,y=.97,ha='left',fontsize=19,weight='bold')
    fig.text(.06,.90,f'SYNTHETIC TEACHING DATA | Stratum: {stratum} | Independent Python computation',fontsize=10)
    fig.text(.06,.11,'O = observed; X = missing. Subject order is fixed within treatment arm; no clustering.',fontsize=10)
    fig.text(.06,.075,'Trailing missingness is an observed pattern, not a withdrawal or censoring event.',fontsize=9)
    fig.text(.06,.035,'Jaime Yan · Clinical Figure Library · Personal noncommercial use · Attribution required',fontsize=8)
    return fig
