import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
from matplotlib.backends.backend_pdf import PdfPages
from matplotlib.patches import Patch
from matplotlib.lines import Line2D
COLORS={"sst":"#0072B2","wt":"#D55E00"}
plt.rcParams.update({"font.size":8,"axes.titlesize":8,"axes.labelsize":8,"xtick.labelsize":6,"ytick.labelsize":6,"legend.fontsize":7,"pdf.fonttype":42})
def shades(ax):
    ax.axvspan(0,10,color="#56B4E9",alpha=.14,zorder=0)
    for lo,hi in [(-2,2),(8,12)]:ax.axvspan(lo,hi,color="0.65",alpha=.23,zorder=0)
    ax.axvline(0,color="0.5",lw=.5)
def raster(ax,group,uid,level,wide=False):
    D=DATA[group];sp=D["trains"][int(uid)];tr=D["trials"].loc[D["trials"].level.eq(level)]
    xx=[];yy=[]
    for j,t in enumerate(tr.start_time):
        rel=(sp[np.searchsorted(sp,t-.1):np.searchsorted(sp,t+.1)]-t)*1000
        xx.extend(rel);yy.extend([j+1]*len(rel))
    xx=np.asarray(xx); yy=np.asarray(yy)
    excluded=((xx>=-2)&(xx<2))|((xx>=8)&(xx<12))
    shades(ax)
    ax.scatter(xx[~excluded],yy[~excluded],marker="|",s=7,linewidths=.5,color=COLORS[group],rasterized=True)
    ax.scatter(xx[excluded],yy[excluded],marker="|",s=7,linewidths=.5,color="0.55",rasterized=True)
    ax.set(xlim=(-100,100) if wide else (-10,25),ylim=(.2,len(tr)+.8),yticks=[1,len(tr)])
    ax.spines[["top","right"]].set_visible(False)
    return len(tr)
def savecheck(fig,name):
    for ax in fig.axes:
        lo,hi=sorted(ax.get_xlim());ax.set_xticks([v for v in ax.get_xticks() if lo<=v<=hi])
        lo,hi=sorted(ax.get_ylim());ax.set_yticks([v for v in ax.get_yticks() if lo<=v<=hi])
    fig.canvas.draw()
    renderer=fig.canvas.get_renderer()
    outside=[]
    for tx in fig.findobj(matplotlib.text.Text):
        if tx.get_visible() and tx.get_text().strip():
            box=tx.get_window_extent(renderer)
            if not fig.bbox.contains(box.x0,box.y0) or not fig.bbox.contains(box.x1,box.y1):outside.append(tx.get_text())
    assert not outside,(name,outside)
    fig.savefig(OUT/(name+".png"),dpi=180)
    fig.savefig(OUT/(name+".pdf"))
# Overview: matched visual axes, unequal trial counts labeled.
fig,axes=plt.subplots(2,2,figsize=(10,8),layout="constrained")
for j,group in enumerate(("sst","wt")):
    rr=R[R.group.eq(group)&R.primary]
    ax=axes[0,j]
    ax.scatter(rr.baseline_hz,rr.response_hz,s=9,alpha=.5,color=COLORS[group],edgecolors="none")
    ax.plot([0,560],[0,560],color=".5",ls="--",lw=.7)
    ax.set(xlim=(-8,560),ylim=(-8,560),xlabel="Before pulse (Hz)",ylabel="After pulse (Hz)",title=f"{group.upper()}: {len(rr)} units; {int(rr.n_trials.iloc[0])} trials/unit; level {rr.level.iloc[0]:g}")
    if group=="sst":
        r=rr[rr.responder.fillna(False)]
        ax.scatter(r.baseline_hz,r.response_hz,s=20,facecolors="none",edgecolors="black",lw=.6,label="Meets response rule")
        ax.legend(loc="lower right",frameon=False)
    else:ax.text(.04,.93,"Descriptive control",transform=ax.transAxes,fontsize=7)
ax=axes[1,0]
for group in ("sst","wt"):
    D=DATA[group];rr=R[R.group.eq(group)&R.primary];level=rr.level.iloc[0];tr=D["trials"][D["trials"].level.eq(level)]
    edges=np.arange(-10,26,dtype=float); hist=np.zeros(35)
    for uid in rr.unit_id:
        sp=D["trains"][int(uid)]
        for t in tr.start_time:
            rel=(sp[np.searchsorted(sp,t-.010):np.searchsorted(sp,t+.025)]-t)*1000
            hist+=np.histogram(rel,bins=edges)[0]
    rates=hist/(len(rr)*len(tr)*.001)
    ax.step(edges[:-1]+.5,rates,where="mid",color=COLORS[group],label=group.upper())
shades(ax);ax.set(xlim=(-10,25),xlabel="Time from pulse onset (ms)",ylabel="Mean rate across units (Hz)",title="Unsmoothened rates include masked artifact intervals")
ax.legend(frameon=False)
ax=axes[1,1]
resp=R[R.group.eq("sst")&R.primary&R.responder.fillna(False)].sort_values("first_spike_median_ms")
for k,r in enumerate(resp.itertuples()):
    ax.plot([r.first_spike_ci_low_ms,r.first_spike_ci_high_ms],[k,k],color=COLORS["sst"],lw=.7)
    ax.plot(r.first_spike_median_ms,k,"o",color=COLORS["sst"],ms=3)
    if np.isfinite(r.onset_ms):ax.plot(r.onset_ms,k,"s",color="black",ms=3)
ax.set(xlim=(1.7,8.2),ylim=(-1,len(resp)),xlabel="Latency from pulse onset (ms)",ylabel="SST responders, ordered by median",title="First-spike timing and resolved onset differ")
ax.legend(handles=[Line2D([],[],marker="o",color=COLORS["sst"],lw=.7,label="Conditional median; 95% bootstrap CI"),Line2D([],[],marker="s",color="black",lw=0,label="Resolved onset (7 of 31)")],frameon=False,loc="upper left")
savecheck(fig,"response-overview");plt.close(fig)
# Show every unit with resolved onset, not a selected best example.
resolved=resp[resp.onset_ms.notna()].sort_values(["onset_ms","unit_id"])
fig,axes=plt.subplots(len(resolved),2,figsize=(10,12),layout="constrained",gridspec_kw={"width_ratios":[2,1]})
for j,row in enumerate(resolved.itertuples()):
    ax=axes[j,0];raster(ax,"sst",row.unit_id,row.level)
    ax.axvline(row.onset_ms,color="black",ls="--",lw=.8)
    ax.set_title(f"Unit {row.unit_id} ({row.region}); onset {row.onset_ms:g} ms",loc="left")
    ax.set_ylabel("Trial")
    tx=T[T.group.eq("sst")&T.unit_id.eq(row.unit_id)&T.level.eq(row.level)]
    first=tx.first_spike_ms.dropna()
    ax=axes[j,1];ax.hist(first,bins=np.arange(2,8.001,.5),color=COLORS["sst"],alpha=.75)
    ax.axvline(row.first_spike_median_ms,color=COLORS["sst"],lw=1.3)
    ax.axvline(row.onset_ms,color="black",ls="--",lw=.8)
    ax.set(xlim=(1.8,8.2),ylabel="Trials",title=f"First spike: {len(first)}/25 trials; median {row.first_spike_median_ms:.2f} ms")
    if j==len(resolved)-1:
        axes[j,0].set_xlabel("Time from pulse onset (ms)");axes[j,1].set_xlabel("First spike in [2,8) ms")
fig.suptitle("All seven resolved SST onsets; level 2.0 | blue: light; gray: excluded intervals",fontsize=8)
savecheck(fig,"resolved-onset-rasters");plt.close(fig)
# Fixed quartile examples from ALL eligible primary units.
fig,axes=plt.subplots(3,2,figsize=(10,7),layout="constrained")
example_ids={}
for j,group in enumerate(("sst","wt")):
    rr=R[R.group.eq(group)&R.primary].sort_values(["difference_hz","unit_id"])
    ids=[]
    for k,q in enumerate((.25,.5,.75)):
        target=rr.difference_hz.quantile(q)
        row=rr.assign(distance=abs(rr.difference_hz-target)).sort_values(["distance","unit_id"]).iloc[0]
        ids.append(int(row.unit_id));ax=axes[k,j];raster(ax,group,int(row.unit_id),row.level)
        ax.set_title(f"{group.upper()} {int(q*100)}th-percentile target: unit {int(row.unit_id)}",loc="left")
        ax.set_ylabel("Trial")
        if k==2:ax.set_xlabel("Time from pulse onset (ms)")
    example_ids[group]=ids
fig.suptitle("Prespecified quartile examples; ties can select the same unit",fontsize=8)
savecheck(fig,"quartile-example-rasters");plt.close(fig)
(OUT/"figure-selections.json").write_text(json.dumps({"quartile_examples":example_ids,"resolved_onset_units":resolved.unit_id.tolist()},indent=2))
print("Overview and individual-trial figures saved.")
