#!/usr/bin/env python3
"""Adversarial regression tests for T^GPT v0.8.7 P0 Safety Repair."""
from __future__ import annotations
import importlib.util, json, random
from pathlib import Path

HERE=Path(__file__).resolve().parent
DATA=HERE.parent/'data'

def load(name, fn):
    spec=importlib.util.spec_from_file_location(name,HERE/fn); mod=importlib.util.module_from_spec(spec); spec.loader.exec_module(mod); return mod

def expect_raises(fn, contains=None):
    try: fn()
    except Exception as e:
        if contains and contains not in str(e): raise AssertionError((contains,str(e)))
        return
    raise AssertionError('expected exception')

def cb_task(rec, tid, fam, A, B, cats, relevance=None, pairs=None):
    relevance=relevance or {c:1 for c in cats}
    cb={'task_id':tid,'task_version':'1','clusters':[{'cluster_id':c,'definition':f'frozen definition {c}','relevance':int(relevance[c])} for c in cats]}
    sha=rec.codebook_sha256(cb)
    t={'task_id':tid,'family_id':fam,'frozen_codebook':cb,'codebook_sha256':sha,'reference_A':A}
    if pairs is None: t['contracted_B']=B
    else: t['paired_probes']=pairs
    return t

def hash_map(rec,tasks):
    m={t['task_id']:t['codebook_sha256'] for t in tasks}
    return m,rec.frozen_codebook_hash_map_sha256(m)

def make_24_distractors(aid,tokens,fp_near=0,fp_mid=0,fp_far=0):
    rows=[]
    for s,nfp in [('NEAR',fp_near),('MID',fp_mid),('FAR',fp_far)]:
        for i in range(8): rows.append({'aggregate_id':aid,'candidate_cluster_id':f'D_{s}_{i}','candidate_type':'DISTRACTOR','distractor_stratum':s,'represented':1 if i<nfp else 0,'aggregate_output_tokens':tokens,'condition_blind_id':aid})
    return rows

def main():
    rec=load('rec','compute_recovery_metrics_v0.8.7.py')
    agg=load('agg','evaluate_aggregation_v0.8.7.py')
    plan=load('plan','plan_task_precision_v0.8.7.py')
    rng=random.Random(807)
    cats=[f'C{i}' for i in range(10)]; w=[.25,.18,.14,.11,.09,.07,.06,.04,.035,.025]

    for N in (64,128,512):
        for j in range(20):
            A=rng.choices(cats,w,k=N); B=rng.choices(cats,w,k=N)
            t=cb_task(rec,f'NULL-{N}-{j}','F1_POLYSEMY_UNDERSPECIFICATION',A,B,cats)
            o=rec.one_task(t,80,40,100+j)
            assert o['partial_identification']['EMPIRICAL_TOTAL_DEFICIT_IDENTIFICATION_INTERVAL'][0] > 0
            assert o['sampling_aware_outer_identification']['TOTAL_DEFICIT_OUTER_REGION'][0] == 0
            assert o['shift']['TOTAL_SEMANTIC_ZERO_EXCLUSION_STATUS']=='NOT_AUTHORIZED_IN_P0'
            assert not o['shift']['CONFIRMATORY_TOTAL_SEMANTIC_CLAIM_AVAILABLE']

    A=['C0']*38+['C1']*24+['C2']*18+['C3']*16+['C4']*12+['C5']*8+['C6']*4+['C7']*3+['C8']*3+['C9']*2
    B=A.copy(); moved=0
    for i,x in enumerate(B):
        if x=='C1' and moved<5: B[i]='UNMAPPED'; moved+=1
    t=cb_task(rec,'MNAR','F1_POLYSEMY_UNDERSPECIFICATION',A,B,cats); t['unmapped_affinity']={'B':['C1']*5}
    o=rec.one_task(t,200,50,9)
    assert o['measurement_status']=='OK'
    assert o['sampling_aware_outer_identification']['TOTAL_DEFICIT_OUTER_REGION'][0]==0
    assert o['unmapped_affinity_diagnostic']['role']=='MNAR_DIAGNOSTIC_ONLY_DOES_NOT_REMAP_UNMAPPED'

    t=cb_task(rec,'BIND','F1_POLYSEMY_UNDERSPECIFICATION',['C0']*32+['C1']*32,['C0']*32+['C1']*32,['C0','C1'])
    bad={**t,'relevance':{'C0':1,'C1':0}}
    expect_raises(lambda: rec.one_task(bad,50,20,1),'free relevance field diverges')
    expect_raises(lambda: rec.one_task(t,50,20,1,'0'*64),'preregistered expected hash')

    small=cb_task(rec,'SMALL','F1_POLYSEMY_UNDERSPECIFICATION',['C0']*31,['C0']*31,['C0'])
    expect_raises(lambda: rec.one_task(small,50,20,1),'MIN_N_PER_ARM')

    pairs=[]
    for i in range(64):
        base='C0' if i<40 else 'C1'; pairs.append({'pair_id':str(i),'base':base,'recovery':base,'control':'UNMAPPED' if i<12 else base})
    t=cb_task(rec,'REC','F2_CAUSAL_ALTERNATIVES',['C0']*40+['C1']*24,[],['C0','C1'],pairs=pairs)
    ro=rec.one_task(t,50,80,2)
    assert not ro['recovery']['PRIMARY_ENDPOINT_AVAILABLE'] and ro['recovery']['RECOVERY_LIFT_CONDITIONAL'] is None
    assert 'COMPLETE_CASE_RECOVERY_LIFT_SENSITIVITY' not in ro['recovery']

    true=[{'id':'C0','weight':.7},{'id':'C1','weight':.3}]
    anns=[{'aggregate_id':'A','candidate_cluster_id':cid,'candidate_type':'TRUE_POOL','represented':1,'aggregate_output_tokens':90,'condition_blind_id':'A'} for cid in ('C0','C1')]
    anns+=make_24_distractors('A',90,fp_near=3)
    a=agg.analyze({'max_output_tokens':100,'true_pool_clusters':true,'annotations':anns})['aggregates'][0]
    assert a['measurement_status']=='MEASUREMENT_INSUFFICIENT'
    assert a['AGG_COVERAGE_ADJUSTED_PRIMARY_NEAR_FPR'] is None
    assert a['AGG_COVERAGE_ADJUSTED_GLOBAL_FPR_SENSITIVITY'] is None

    families=sorted(rec.EXPECTED_FAMILIES); tasks=[]
    for fi,f in enumerate(families):
        n=20 if fi==4 else 10
        for j in range(n):
            valid = not (fi==4 and j>=2)
            A=['C0']*32+['C1']*32
            B=A[:] if valid else ['C0']*20+['C1']*20+['UNMAPPED']*24
            tasks.append(cb_task(rec,f'{f}-{j}',f,A,B,['C0','C1']))
    m,msha=hash_map(rec,tasks)
    out=rec.analyze({'tasks':tasks,'frozen_codebook_hashes':m,'frozen_codebook_hashes_sha256':msha,'permutation_reps':50,'task_bootstrap_reps':2000,'seed':4})
    assert out['across_tasks']['experiment_status']=='MEASUREMENT_INSUFFICIENT_FAMILY_ATTRITION'
    assert not out['across_tasks']['PRIMARY_GENERALIZED_CLAIM_AVAILABLE']
    assert 'F5_MULTI_SOLUTION_REASONING' in out['measurement_summary']['family_attrition_failures']
    assert out['across_tasks']['TOTAL_SEMANTIC_IDENTIFICATION_DIAGNOSTIC']['claim_available'] is False

    expect_raises(lambda: rec.analyze({'tasks':tasks,'frozen_codebook_hashes':m,'frozen_codebook_hashes_sha256':'0'*64,'permutation_reps':50,'task_bootstrap_reps':2000}), 'hashes_sha256 mismatch')

    rows=[]; p0=[]
    for fi,f in enumerate(families):
        for j in range(4):
            rows.append({'task_id':f'{fi}-{j}','family_id':f,'effect':.03*fi+(-.02,-.01,.01,.03)[j],'n_per_task':64,'measurement_status':'OK'})
            ok = not (fi==4 and j>=2)
            p0.append({'task_id':f'P0-{fi}-{j}','family_id':f,'measurement_status':'OK' if ok else 'MEASUREMENT_INSUFFICIENT'})
    inp={'experiment_id':'E1_STAGE','planning_started_utc':'2026-09-12T06:00:00Z','planning_tasks':rows,'p0_registry':p0,'confirmation_n_per_task':64,'bootstrap_reps':2500,'seed':7}
    po=plan.plan(inp,DATA/'minimum-effects-v0.8.6.json')
    assert po['decision']=='REDESIGN_NO_GO' and 'F5_MULTI_SOLUTION_REASONING' in po['p0_family_attrition_failures']

    print('T_GPT_V0_8_7_P0_SAFETY_SELFTEST_OK')

if __name__=='__main__': main()
