#!/usr/bin/env python3
"""Adversarial regression tests for T^GPT v0.8.5 engineering repair."""
from __future__ import annotations
import importlib.util, random
from pathlib import Path

HERE=Path(__file__).resolve().parent

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 close(a,b,tol=1e-9): return abs(a-b)<=tol

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 main():
    rec=load('rec','compute_recovery_metrics_v0.8.5.py')
    agg=load('agg','evaluate_aggregation_v0.8.5.py')
    cov=load('cov','coverage_baseline_v0.8.5.py')
    plan=load('plan','plan_task_precision_v0.8.5.py')

    # A01: same semantic distribution + differential UNMAPPED must never become a primary semantic effect.
    rng=random.Random(123); cats=[f'C{i}' for i in range(1,11)]; weights=[.25,.18,.14,.11,.09,.07,.06,.04,.035,.025]; rel={c:1 for c in cats}
    insufficient=0
    for j in range(80):
        A=rng.choices(cats,weights,k=128)
        B=[]
        for _ in range(128):
            if rng.random()<.20: B.append('UNMAPPED')
            else: B.append(rng.choices(cats,weights,k=1)[0])
        o=rec.one_task({'task_id':f'MAP{j}','relevance':rel,'reference_A':A,'contracted_B':B},120,50,9000+j)
        insufficient += (o['measurement_status']=='MEASUREMENT_INSUFFICIENT')
        assert not o['shift']['PRIMARY_ENDPOINT_AVAILABLE'], o
        assert o['shift']['SEMANTIC_DEFICIT_EXCESS'] is None, o
    assert insufficient==80, insufficient

    # Exact-null with no mapping gap remains estimable and centered after conditional permutation calibration.
    ex=[]
    for j in range(80):
        A=rng.choices(cats,weights,k=128); B=rng.choices(cats,weights,k=128)
        o=rec.one_task({'task_id':f'NULL{j}','relevance':rel,'reference_A':A,'contracted_B':B},150,50,12000+j)
        assert o['shift']['PRIMARY_ENDPOINT_AVAILABLE']
        ex.append(o['shift']['SEMANTIC_DEFICIT_EXCESS'])
    assert abs(sum(ex)/len(ex)) < .025, sum(ex)/len(ex)

    # A02: A-prime mode rejects unequal total N.
    expect_raises(lambda: rec.one_task({'task_id':'UNEQ','relevance':{'C1':1,'C2':1},'reference_A':['C1']*32+['C2']*32,'reference_A_prime':['C1']*32+['C2']*32,'contracted_B':['C1']*8+['C2']*8},100,50,7), 'N_A == N_A_prime == N_B')

    # Equal-N A-prime null remains available.
    eq=rec.one_task({'task_id':'EQ','relevance':{'C1':1,'C2':1},'reference_A':['C1']*32+['C2']*32,'reference_A_prime':['C1']*32+['C2']*32,'contracted_B':['C1']*32+['C2']*32},100,50,8)
    assert close(eq['shift']['SEMANTIC_DEFICIT_EXCESS'],0.0), eq

    # R01: differential mapping between recovery/control hard-gates RECOVERY_LIFT.
    pairs=[]
    for i in range(100):
        base='C1' if i<60 else 'C2'
        recovery='C1' if i<50 else ('C2' if i<85 else 'C3')
        control='UNMAPPED' if i<20 else base
        pairs.append({'pair_id':str(i),'base':base,'recovery':recovery,'control':control})
    ro=rec.one_task({'task_id':'RMAP','relevance':{'C1':1,'C2':1,'C3':1},'reference_A':['C1']*50+['C2']*30+['C3']*20,'paired_probes':pairs},100,200,9)
    assert not ro['recovery']['PRIMARY_ENDPOINT_AVAILABLE'], ro
    assert ro['recovery']['RECOVERY_LIFT'] is None, ro
    assert ro['mapping']['recovery_control_gate']['status']=='MEASUREMENT_INSUFFICIENT', ro

    # G01: high FPR has an actual consequence; adjusted primary endpoint disappears.
    anns=[]
    for cid,rep in [('C1',1),('C2',1)]:
        anns.append({'aggregate_id':'A1','candidate_cluster_id':cid,'candidate_type':'TRUE_POOL','represented':rep,'aggregate_output_tokens':80,'condition_blind_id':'X'})
    for i in range(8):
        anns.append({'aggregate_id':'A1','candidate_cluster_id':f'D{i}','candidate_type':'DISTRACTOR','distractor_stratum':'NEAR' if i<4 else 'FAR','represented':1 if i<4 else 0,'aggregate_output_tokens':80,'condition_blind_id':'X'})
    ao=agg.analyze({'max_output_tokens':100,'true_pool_clusters':[{'id':'C1','weight':.7},{'id':'C2','weight':.3}],'annotations':anns})['aggregates'][0]
    assert close(ao['REPRESENTATION_FPR'],.5), ao
    assert ao['measurement_status']=='MEASUREMENT_INSUFFICIENT' and ao['AGG_COVERAGE_ADJUSTED_PRIMARY'] is None, ao

    # Low-FPR adjusted coverage available.
    anns2=[]
    for cid,rep in [('C1',1),('C2',0)]:
        anns2.append({'aggregate_id':'A2','candidate_cluster_id':cid,'candidate_type':'TRUE_POOL','represented':rep,'aggregate_output_tokens':80,'condition_blind_id':'Y'})
    for i in range(8):
        anns2.append({'aggregate_id':'A2','candidate_cluster_id':f'D{i}','candidate_type':'DISTRACTOR','distractor_stratum':'NEAR' if i<4 else 'FAR','represented':1 if i==0 else 0,'aggregate_output_tokens':80,'condition_blind_id':'Y'})
    a2=agg.analyze({'max_output_tokens':100,'true_pool_clusters':[{'id':'C1','weight':.7},{'id':'C2','weight':.3}],'annotations':anns2})['aggregates'][0]
    assert a2['measurement_status']=='OK' and a2['AGG_COVERAGE_ADJUSTED_PRIMARY'] is not None, a2

    # Coverage baseline remains exact for small pool and reports a bound.
    b=cov.solve({'token_budget':240,'clusters':[{'id':'C1','weight':.4},{'id':'C2','weight':.35},{'id':'C3','weight':.25}], 'candidates':[{'id':'S1','token_cost':100,'covers':{'C1':1}},{'id':'S2','token_cost':120,'covers':{'C2':1}},{'id':'S3','token_cost':180,'covers':{'C1':.5,'C3':1}}]})
    assert b['selected_ids']==['S1','S2'] and close(b['objective_normalized'],.75) and close(b['upper_bound_normalized'],.75), b

    # P01/P03: critical planner constants cannot be overridden; planning/confirmation N must match.
    families=sorted(plan.EXPECTED_FAMILIES); rows=[]
    for fi,f in enumerate(families):
        for j in range(4): rows.append({'task_id':f'{fi}-{j}','family_id':f,'effect':.05*fi + (-.03,-.01,.02,.04)[j],'n_per_task':64})
    valid={'planning_tasks':rows,'confirmation_n_per_task':64,'minimum_relevant_effect':.08,'bootstrap_reps':2500,'seed':77}
    po=plan.plan(valid)
    assert po['required_confirmation_tasks_balanced']%5==0, po
    expect_raises(lambda: plan.plan({**valid,'redesign_cap':500}), 'not caller-overridable')
    expect_raises(lambda: plan.plan({**valid,'confirmation_n_per_task':128}), 'planning n_per_task == confirmation_n_per_task')
    zero=[{**r,'effect':.1} for r in rows]
    expect_raises(lambda: plan.plan({**valid,'planning_tasks':zero}), 'residual SD is zero')

    print('T_GPT_V0_8_5_ENGINEERING_SELFTEST_OK')

if __name__=='__main__': main()
