"""Fixed offline Labs 19/20. No acquisition, generation, training or prompt search.

Every numeric pass runs in a supervised subprocess. Exclusive phase transitions,
an fsynced attempt ledger and per-pass non-pickle arrays retain partial evidence.
The parent kills stalled native computation at the persisted cumulative deadline.
"""
import argparse, contextlib, csv, hashlib, io, json, os, re, subprocess, sys, time, uuid
from collections import Counter
from datetime import datetime, timezone
from pathlib import Path
from inspection_common import safe_path, runtime, environment, load_specimen, source_hashes

ATOL = RTOL = 1e-6
SECONDS = 900.0
CAPS = {19: 90, 20: 80}
def source_names(lab):
    return ('contrast_experiment.py', 'inspection_common.py', 'inspect_residual_stream.py',
            f'run_lab{lab}.py', f'lab{lab}-fixtures.json')


def digest(data):
    return hashlib.sha256(data).hexdigest()


def json_bytes(value):
    return (json.dumps(value, indent=2, allow_nan=False) + '\n').encode()


def stamp():
    return datetime.now(timezone.utc).isoformat()


def read_json(path):
    def unique(items):
        result = {}
        for key, value in items:
            if key in result: raise ValueError('Duplicate JSON key')
            result[key] = value
        return result
    return json.loads(safe_path(path).read_text(), object_pairs_hook=unique,
                      parse_constant=lambda value: (_ for _ in ()).throw(ValueError(value)))


class Run:
    def __init__(self, directory):
        self.path = safe_path(directory)
        if not self.path.is_dir(): raise ValueError('Run directory missing')

    def file(self, name):
        if Path(name).name != name: raise ValueError('Flat evidence filename required')
        return safe_path(self.path / name)

    def put(self, name, value):
        with self.file(name).open('xb') as handle:
            handle.write(value); handle.flush(); os.fsync(handle.fileno())

    def json(self, name, value):
        self.put(name, json_bytes(value))

    def replace(self, name, value):
        target = self.file(name)
        # A watchdog can leave an interrupted predecessor's exclusive temporary.
        # Preserve it; a fresh owned filename permits failure finalization without
        # deleting or overwriting any unidentified/stale evidence.
        temporary = self.file(name + '.writing-' + uuid.uuid4().hex)
        with temporary.open('xb') as handle:
            handle.write(json_bytes(value)); handle.flush(); os.fsync(handle.fileno())
        os.replace(temporary, target)

    def log(self, value):
        with self.file('attempts.jsonl').open('ab') as handle:
            handle.write((json.dumps(value, allow_nan=False) + '\n').encode())
            handle.flush(); os.fsync(handle.fileno())

    def hash(self, name):
        return digest(self.file(name).read_bytes())

    def arrays(self, name, values):
        import numpy as np
        buffer = io.BytesIO()
        np.savez_compressed(buffer, **{k: v.detach().cpu().numpy() for k, v in values.items()})
        self.put(name, buffer.getvalue())


def fixture_rows(lab, fixture):
    rows = []
    for family, values in fixture.items():
        if not isinstance(values, list) or not values or not isinstance(values[0], dict): continue
        for value in values:
            row = dict(value); row['family'] = family
            row['text'] = value.get('prompt', value.get('context', '') +
                                    fixture.get('review_suffix', fixture.get('suffix', '')))
            rows.append(row)
    if len(rows) != (12 if lab == 19 else 24): raise ValueError('Fixture coverage')
    return rows


def tokenize(rows, fixture, tokenizer, lab):
    for row in rows:
        encoded = tokenizer(row['text'], add_special_tokens=False, padding=False, truncation=False)
        ids = encoded['input_ids']
        if not 1 <= len(ids) <= 64: raise ValueError('Token budget')
        row.update(token_ids=ids, attention_mask=encoded['attention_mask'],
                   positions=list(range(len(ids))), final_position=len(ids)-1,
                   decoded_pieces=[tokenizer.decode([i], clean_up_tokenization_spaces=False) for i in ids])
    aligned = [r for r in rows if r['family'] != 'heldout_offtask']
    if len({r['token_ids'][-1] for r in aligned}) != 1: raise ValueError('Final token mismatch')
    candidates = fixture.get('score_tokens', fixture.get('score_candidates'))
    score_ids = [tokenizer.encode(s, add_special_tokens=False) for s in candidates]
    if any(len(x) != 1 for x in score_ids) or score_ids[0] == score_ids[1]:
        raise ValueError('Score candidates must be distinct single tokens')
    if lab == 20:
        pairs = {r['id']: r for r in rows}
        for n in range(1, 5):
            a, b = (pairs[f'B{n}{sign}']['token_ids'] for sign in ('+', '-'))
            if len(a) != len(b) or Counter(a) != Counter(b): raise ValueError('Binding multiset mismatch')
    return [x[0] for x in score_ids]


class CleanupCheck(Exception):
    pass


def finite(value, torch):
    if not torch.isfinite(value).all(): raise ValueError('Nonfinite numeric measurement')
    return value


def replacement(original, arm, u, q, dose, torch):
    """Double analysis, native float32 edit, no mutation of original."""
    if arm == 'zero':
        perturbation = torch.zeros(128, dtype=original.dtype, device=original.device)
        return original + perturbation
    if arm == 'remove':
        return (original.double() - original.double().dot(u) * u).float()
    direction = q if arm.startswith('random') else u
    sign = -1 if 'negative' in arm else 1
    scale = 2 if 'larger' in arm else 1
    perturbation = (direction * dose * sign * scale).float()
    return original + perturbation


@contextlib.contextmanager
def hook(model, row, arm, holder, torch, u=None, q=None, dose=None):
    module = model.gpt_neox.layers[2]
    before = tuple(module._forward_hooks)
    n = len(row['token_ids']); holder['calls'] = 0
    def callback(module, args, output):
        holder['calls'] += 1
        if holder['calls'] != 1: raise ValueError('Repeated hook invocation')
        if not isinstance(output, tuple) or not output or not isinstance(output[0], torch.Tensor):
            raise ValueError('Expected post-block tuple')
        original = output[0]
        if tuple(original.shape) != (1, n, 128): raise ValueError('Boundary shape')
        finite(original, torch)
        holder['original'] = original[0, n-1].detach().clone()
        if arm == 'abort': raise CleanupCheck('deliberate_cleanup_check')
        if arm == 'observe': return None
        changed = original.clone()
        changed[0, n-1] = finite(replacement(holder['original'], arm, u, q, dose, torch), torch)
        if not torch.equal(original[:, :n-1], changed[:, :n-1]): raise ValueError('Earlier rows changed')
        holder['replacement'] = changed[0, n-1].detach().clone()
        return (changed,) + output[1:]
    handle = module.register_forward_hook(callback)
    try: yield
    finally:
        handle.remove()
        holder['cleanup'] = tuple(module._forward_hooks) == before
        if not holder['cleanup']: raise ValueError('Hook cleanup failed')


def agreement(a, b, torch):
    return {'max_abs_error': float((a-b).abs().max()),
            'passed': bool(torch.allclose(a, b, atol=ATOL, rtol=RTOL))}


class Experiment:
    def __init__(self, run, state, model, tokenizer, torch):
        self.run, self.state = run, state
        self.model, self.tokenizer, self.torch = model, tokenizer, torch
        self.started = time.monotonic()
        self.previous_elapsed = state['elapsed_seconds']
        self.checks = []
        self.run.replace('active-phase.json', {'started_monotonic': self.started,
                         'remaining_seconds': SECONDS-self.previous_elapsed})

    def elapsed(self):
        return self.previous_elapsed + time.monotonic()-self.started

    def check(self):
        if self.elapsed() >= SECONDS: raise TimeoutError('time_budget_exceeded')

    def save_state(self):
        self.state['elapsed_seconds'] = self.elapsed()
        self.run.replace('state.json', self.state)

    def require_agreement(self, name, a, b):
        result = {'name': name, **agreement(a, b, self.torch)}
        self.checks.append(result); self.run.replace('checks.json', self.checks)
        if not result['passed']: raise ValueError('Invariance failed: '+name)

    def forward(self, row, arm='baseline', u=None, q=None, dose=None, phase=None):
        self.check()
        if self.state['attempts'] >= CAPS[self.state['lab']]: raise ValueError('forward_budget_exceeded')
        self.state['attempts'] += 1; self.save_state()
        number = self.state['attempts']; holder = {'calls': 0, 'cleanup': True}
        event = {'attempt': number, 'phase': phase or self.state['phase'], 'prompt': row['id'],
                 'arm': arm, 'timestamp': stamp(), 'status': 'started'}
        self.run.log(event); status = 'failed'; error = None; result = None
        try:
            context = contextlib.nullcontext() if arm == 'baseline' else hook(
                self.model, row, arm, holder, self.torch, u, q, dose)
            with context:
                out = self.model(input_ids=self.torch.tensor([row['token_ids']], dtype=self.torch.long),
                    attention_mask=self.torch.tensor([row['attention_mask']], dtype=self.torch.long),
                    use_cache=False, output_attentions=False, output_hidden_states=False, logits_to_keep=1)
            if tuple(out.logits.shape) != (1, 1, self.model.config.vocab_size): raise ValueError('Logits shape')
            if arm != 'baseline' and holder['calls'] != 1: raise ValueError('Missing hook')
            result = {'logits': finite(out.logits[0, 0].detach().clone(), self.torch)}
            result.update({k: holder[k] for k in ('original', 'replacement') if k in holder})
            self.run.arrays(f'attempt-{number:03d}.npz', result)
            self.check(); status = 'completed'
        except CleanupCheck as exc:
            if arm != 'abort': raise
            error = str(exc); status = 'expected_exception'
        except BaseException as exc:
            error = type(exc).__name__ + ': ' + str(exc)
            raise
        finally:
            self.save_state()
            self.run.log({**event, 'timestamp': stamp(), 'status': status, 'exception': error,
                          'hook_calls': holder['calls'], 'hook_cleanup': holder['cleanup'],
                          'array_file': f'attempt-{number:03d}.npz' if result is not None else None})
        return result


def direction(vectors, torch, lab):
    values = torch.stack([vectors[f'D{i}+'].double()-vectors[f'D{i}-'].double() for i in range(1, 4)])
    d = finite(values.mean(0), torch); norm = float(d.norm())
    if not torch.isfinite(d.norm()): raise ValueError('Nonfinite direction norm')
    valid = norm > 1e-8
    if not valid and lab == 19: raise ValueError('degenerate_direction')
    u = finite(d/norm, torch) if valid else None
    R = float(torch.stack([v.double().norm() for v in vectors.values()]).mean())
    if not torch.isfinite(torch.tensor(R)): raise ValueError('Nonfinite reference norm')
    if lab == 19 and R <= 1e-8: raise ValueError('Degenerate reference norm')
    raw = torch.randn(128, generator=torch.Generator(device='cpu').manual_seed(lab), dtype=torch.float64)
    q = finite(raw/raw.norm(), torch)
    return {'d': d.tolist(), 'norm': norm, 'u': u.tolist() if valid else None,
            'status': 'defined' if valid else 'degenerate_direction', 'R': R,
            'dose': R*(.05 if lab == 19 else .02), 'random_raw': raw.tolist(), 'q': q.tolist(),
            'u_dot_q': float(u.dot(q)) if valid else None,
            'cast_u_perturbation_norm': float((u*R*(.05 if lab == 19 else .02)).float().norm()) if valid else None,
            'cast_q_perturbation_norm': float((q*R*(.05 if lab == 19 else .02)).float().norm())}


def measurement(result, ids, tokenizer, torch, baseline=None, u=None):
    logits = result['logits'].double(); logp = finite(torch.log_softmax(logits, 0), torch)
    greedy = int(logits.argmax()); probabilities = logp[ids].exp()
    record = {'y': float(logits[ids[0]]-logits[ids[1]]), 'score_logits': logits[ids].tolist(),
              'score_probabilities': probabilities.tolist(), 'greedy_id': greedy,
              'greedy_piece': tokenizer.decode([greedy], clean_up_tokenization_spaces=False)
                              if greedy < len(tokenizer) else None}
    if baseline is not None:
        reference = baseline['logits'].double(); reference_logp = torch.log_softmax(reference, 0)
        record.update(delta_y=record['y']-float(reference[ids[0]]-reference[ids[1]]),
                      kl_baseline_to_arm=float(finite((reference_logp.exp()*(reference_logp-logp)).sum(), torch)))
    if 'original' in result:
        original = result['original'].double(); changed = result.get('replacement', result['original']).double()
        norm = float(original.norm()); perturbation = float((changed-original).norm())
        record.update(activation_norm=norm, perturbation_norm=perturbation,
                      relative_perturbation_norm=perturbation/norm if norm else None,
                      relative_norm_status='defined' if norm else 'undefined_zero_norm',
                      projection=float(original.dot(u)) if u is not None else None,
                      replacement_projection=float(changed.dot(u)) if u is not None else None,
                      projection_status='defined' if u is not None else 'degenerate_direction')
    return record


def shortcut_controls(rows, model, torch):
    embeddings = model.gpt_neox.embed_in.weight
    means = {}; lexical = {}; final = {}; dictionary = {r['id']: r for r in rows}
    for row in rows:
        counts = Counter(row['token_ids'])
        mean = torch.zeros(128, dtype=torch.float64)
        for token in sorted(counts): mean += embeddings[token].double()*counts[token]
        means[row['id']] = mean/len(row['token_ids'])
        final[row['id']] = embeddings[row['token_ids'][-1]].clone()
        words = Counter(re.findall(r'\b[a-z]+\b', row['text'].lower()))
        lexical[row['id']] = sum(words[x] for x in ('cheerful','delighted','pleased')) - sum(
            words[x] for x in ('miserable','gloomy','upset'))
    reference = next(iter(final.values()))
    if not all(torch.equal(x, reference) for x in final.values()): raise ValueError('Final embedding control')
    binding = []
    for n in range(1, 5):
        equal = torch.equal(means[f'B{n}+'], means[f'B{n}-'])
        if not equal: raise ValueError('Bag embedding control')
        binding.append({'pair': f'B{n}', 'equal': equal, 'token_counts': len(dictionary[f'B{n}+']['token_ids'])})
    return {'final_embeddings_identical': True, 'binding_bag_means': binding, 'word_count_scores': lexical}, {
        **{'bag_'+k: v for k,v in means.items()}, **{'final_'+k: v for k,v in final.items()}}


def discover(experiment, rows, ids, options):
    e = experiment; torch = e.torch; run = e.run; lab = e.state['lab']
    discovery = [r for r in rows if r['family'] == 'discovery']
    baselines = {}; vectors = {}; measured = []
    for row in discovery: baselines[row['id']] = e.forward(row)
    for row in discovery:
        result = e.forward(row, 'observe'); vectors[row['id']] = result['original']
        e.require_agreement(row['id']+' observer', result['logits'], baselines[row['id']]['logits'])
        measured.append({'prompt': row['id'], **measurement(result, ids, e.tokenizer, torch)})
    first = discovery[0]
    result = e.forward(first); e.require_agreement('discovery repeat', result['logits'], baselines[first['id']]['logits'])
    result = e.forward(first, 'zero'); e.require_agreement('discovery zero', result['logits'], baselines[first['id']]['logits'])
    if not torch.equal(result['original'], result['replacement']): raise ValueError('Nonzero zero control')
    e.forward(first, 'abort')
    result = e.forward(first); e.require_agreement('discovery recovery', result['logits'], baselines[first['id']]['logits'])
    d = direction(vectors, torch, lab)
    run.json('direction.json', d)
    run.arrays('discovery-vectors.npz' if lab == 19 else 'discovery-arrays.npz', vectors)
    pilot = []
    u = torch.tensor(d['u'], dtype=torch.float64) if d['u'] is not None else None
    q = torch.tensor(d['q'], dtype=torch.float64)
    for record in measured:
        record['projection'] = float(vectors[record['prompt']].double().dot(u)) if u is not None else None
        record['projection_status'] = d['status']
    if lab == 20 and options['optional_intervention'] and u is not None and d['R'] <= 1e-8:
        raise ValueError('Degenerate optional reference norm')
    if lab == 19:
        for row in discovery:
            result = e.forward(row, 'positive_small', u, q, d['dose'])
            pilot.append({'prompt': row['id'], **measurement(result, ids, e.tokenizer, torch, baselines[row['id']], u)})
        run.json('discovery-pilot.json', pilot)
    else:
        controls, arrays = shortcut_controls(rows, e.model, torch)
        run.json('shortcut-controls.json', controls); run.arrays('shortcut-embeddings.npz', arrays)
        run.put('evidence-audit.md', b'# Evidence audit worksheet\n\nResearch reading is not executed by this numerical runner.\nComplete the two source/version/method/limitation records required by Lab 20 before claiming the full coursework audit is complete.\n')
    run.json('discovery-measurements.json', measured)
    frozen_files = ['environment.json','artifact-manifest.json','fixtures.json','tokens.json',
                    'direction.json','sources.json','discovery-measurements.json',
                    'discovery-vectors.npz' if lab == 19 else 'discovery-arrays.npz']
    if lab == 19: frozen_files.append('discovery-pilot.json')
    else: frozen_files += ['shortcut-controls.json','shortcut-embeddings.npz']
    plan = {'lab': lab, 'timestamp': stamp(), 'files': {name: run.hash(name) for name in frozen_files},
            'sources': source_hashes(source_names(lab)), 'block': 2, 'position': 'final input token only',
            'boundary': 'raw post-block residual after additions before block 3',
            'metric': 'first candidate logit minus second candidate logit; full-vocabulary probabilities and KL',
            'score_ids': ids, 'direction': d, 'expectations': None,
            'discovery_pilot_summary': pilot,
            'atol': ATOL, 'rtol': RTOL, 'near_zero': 1e-5, 'degenerate_norm': 1e-8,
            'optional_intervention': bool(options['optional_intervention'] and u is not None),
            'optional_requested': options['optional_intervention'],
            'grouping': 'H three pairs; B prize B1/B2 and gift/fine B3/B4 averaged within scenario; C separate',
            'fixed_hypotheses': {'lab19_positive_mean': '>0 over four H prompts', 'lab19_negative_mean': '<0 over four H prompts',
                'lab20_H': 'representation and output means >1e-5 separately and conjunction',
                'lab20_B': 'both scenario means >1e-5 separately for representation/output and conjunction',
                'lab20_extension': 'four selected prefixes mean delta_y positive >1e-5 and negative <-1e-5'},
            'selected_extension_inputs': ['H1+','H1-','B1+','B1-'],
            'controls': ['observer','zero','restoration','opposite signs','one fixed random direction; not a null distribution'],
            'confounders': ['lexical cues','positions and prefix lengths','syntax','topics','small constructed sample','distribution shift']}
    run.json('plan-draft.json', plan)
    e.state['plan_draft_sha256'] = run.hash('plan-draft.json')
    e.state['phase'] = 'awaiting_predictions'; e.save_state()


def freeze_predictions(run, state, expectations):
    """Separate post-discovery step: no model import/load/forward, no clock charge."""
    lab = state['lab']
    if state['phase'] != 'awaiting_predictions': raise ValueError('Discovery must finish before predictions')
    if run.hash('plan-draft.json') != state['plan_draft_sha256']: raise ValueError('Discovery draft changed')
    plan = read_json(run.file('plan-draft.json'))
    if plan['lab'] != lab or plan['sources'] != source_hashes(source_names(lab)):
        raise ValueError('Source changed since discovery')
    for name, value in plan['files'].items():
        if run.hash(name) != value: raise ValueError('Discovery input changed: '+name)
    required = {'expected_outcome','confidence','larger_dose','projection_removal'} if lab==19 else {
        'H_representation','H_output','B_representation','B_output','C_specificity','confidence','optional_signs'}
    if not isinstance(expectations, dict) or set(expectations) != required or any(
        not isinstance(v,str) or not v.strip() for v in expectations.values()):
        raise ValueError('Post-discovery predictions require exactly these nonempty text fields: '+', '.join(sorted(required)))
    run.put('freeze-claimed', stamp().encode())
    plan.update(timestamp=stamp(), expectations=expectations, predictions_recorded_after_discovery=True,
                attempts_at_freeze=state['attempts'])
    name = 'predictions.json' if lab==19 else 'frozen-plan.json'
    run.json(name, plan); run.put(name+'.sha256', (run.hash(name)+'\n').encode())
    state['plan_sha256'] = run.hash(name); state['phase'] = 'discovered'; run.replace('state.json', state)


def verify_plan(run, state):
    name = 'predictions.json' if state['lab'] == 19 else 'frozen-plan.json'
    expected = state['plan_sha256']
    if run.hash(name) != expected or run.file(name+'.sha256').read_text().strip() != expected:
        raise ValueError('Frozen plan changed')
    plan = read_json(run.file(name))
    if plan['lab'] != state['lab'] or plan['sources'] != source_hashes(source_names(state['lab'])): raise ValueError('Runner source changed')
    for name, value in plan['files'].items():
        if run.hash(name) != value: raise ValueError('Frozen input changed: '+name)
    return plan


def evaluate(e, rows, plan):
    lab = e.state['lab']; run = e.run; torch = e.torch; ids = plan['score_ids']; d = plan['direction']
    u = torch.tensor(d['u'], dtype=torch.float64) if d['u'] is not None else None
    q = torch.tensor(d['q'], dtype=torch.float64); dose = d['dose']
    heldout = [r for r in rows if r['family'] != 'discovery']; base = {}; records = []
    run.json('heldout-start.json', {'timestamp': stamp(), 'plan_sha256': e.state['plan_sha256'],
             'attempt_count_before': e.state['attempts']})
    def record(row, arm, result, baseline=None):
        value = {'prompt': row['id'], 'family': row['family'], 'arm': arm,
                 **measurement(result, ids, e.tokenizer, torch, baseline, u)}
        records.append(value); run.replace('measurements.json', records)
    if lab == 19:
        arms = ['baseline','zero','positive_small','negative_small','positive_larger','negative_larger',
                'random_positive','random_negative','remove']
        for row in heldout:
            for arm in arms:
                result = e.forward(row, arm, u, q, dose)
                if arm == 'baseline': base[row['id']] = result
                if arm == 'zero': e.require_agreement(row['id']+' zero', result['logits'], base[row['id']]['logits'])
                if arm in ('positive_small','negative_small','random_positive','random_negative'):
                    expected = (u*dose).float().norm()
                    actual = (q*dose).float().norm() if arm.startswith('random') else expected
                    if not torch.allclose(actual, expected, atol=ATOL, rtol=RTOL): raise ValueError('Matched norm failed')
                record(row, arm, result, base[row['id']])
                if arm == 'zero':
                    # Baseline has no hook; the verified zero control supplies its
                    # original internal row, explicitly annotated rather than invented.
                    baseline_record = next(r for r in records if r['prompt']==row['id'] and r['arm']=='baseline')
                    baseline_record.update({k: v for k,v in records[-1].items() if k in
                        ('activation_norm','projection','projection_status','relative_norm_status')})
                    baseline_record['internal_capture_source'] = 'verified zero-control pass'
                    run.replace('measurements.json', records)
    else:
        for row in heldout:
            result = e.forward(row, 'observe'); base[row['id']] = result
            record(row, 'observe', result)
    for row in heldout:
        result = e.forward(row)
        e.require_agreement(row['id']+' restoration', result['logits'], base[row['id']]['logits'])
    if lab == 20 and plan['optional_intervention']:
        selected = [{r['id']: r for r in heldout}[key] for key in plan['selected_extension_inputs']]
        for row in selected:
            for arm in ['zero','positive_small','negative_small','random_positive','random_negative']:
                result = e.forward(row, arm, u, q, dose)
                if arm == 'zero': e.require_agreement(row['id']+' extension zero', result['logits'], base[row['id']]['logits'])
                if not torch.allclose((u*dose).float().norm(), (q*dose).float().norm(), atol=ATOL, rtol=RTOL):
                    raise ValueError('Matched norm failed')
                record(row, arm, result, base[row['id']])
        for row in selected:
            result = e.forward(row)
            e.require_agreement(row['id']+' extension restoration', result['logits'], base[row['id']]['logits'])
    expected = 82 if lab == 19 else (76 if plan['optional_intervention'] else 52)
    if e.state['attempts'] != expected: raise ValueError('Planned attempt coverage differs')
    e.state['phase'] = 'completed'; e.save_state()


def statistics(values):
    if not values or any(v is None for v in values): return {'status': 'undefined', 'mean': None}
    mean = sum(values)/len(values)
    return {'status': 'defined', 'values': values, 'mean': mean,
            'sample_variance': sum((v-mean)**2 for v in values)/(len(values)-1) if len(values)>1 else None,
            'minimum': min(values), 'maximum': max(values),
            'positive': sum(v>1e-5 for v in values), 'negative': sum(v<-1e-5 for v in values),
            'near_zero': sum(abs(v)<=1e-5 for v in values)}


def persona_decisions(pairs):
    groups = {}
    for score in ('a','b'):
        h = [pairs[f'H{i}'][score] for i in range(1,4)]
        b = [None if any(pairs[f'B{i}'][score] is None for i in indices) else
             sum(pairs[f'B{i}'][score] for i in indices)/2 for indices in ((1,2),(3,4))]
        groups[score] = {'H': statistics(h), 'B_scenarios': statistics(b),
                        'H_holds': None if any(v is None for v in h) else sum(h)/3>1e-5,
                        'B_holds': None if any(v is None for v in b) else all(v>1e-5 for v in b)}
    for family in ('H','B'):
        decisions = [groups[k][family+'_holds'] for k in ('a','b')]
        groups[family+'_conjunction'] = None if None in decisions else all(decisions)
    return groups


def summarize(run, state):
    """Saved-only analysis; torch import for arrays never loads or forwards a model."""
    import numpy as np
    import torch
    torch.set_num_threads(1)
    plan = verify_plan(run, state); lab = state['lab']; records = read_json(run.file('measurements.json'))
    events = [json.loads(line) for line in run.file('attempts.jsonl').read_text().splitlines()]
    complete = [r for r in events if r['status'] == 'completed']
    logits = {}; vectors = {}; index = []
    for event in complete:
        with np.load(run.file(event['array_file']), allow_pickle=False) as arrays:
            key = f"attempt_{event['attempt']:03d}"
            logits[key] = arrays['logits']; index.append({**event, 'key': key})
            for name in ('original','replacement'):
                if name in arrays: vectors[key+'_'+name] = arrays[name]
    for name, arrays in [('logits.npz',logits),('vectors.npz',vectors)]:
        if not run.file(name).exists():
            buffer = io.BytesIO(); np.savez_compressed(buffer, **arrays); run.put(name, buffer.getvalue())
    if not run.file('array-index.json').exists(): run.json('array-index.json', index)
    report = {'lab': lab, 'run_phase': state['phase'], 'attempts': state['attempts'],
              'elapsed_seconds': state['elapsed_seconds'], 'plan_sha256': state['plan_sha256'],
              'scope': 'Fixed fixtures only; no population inference or subjective-experience claim.'}
    if lab == 19:
        arms = sorted({r['arm'] for r in records})
        report['reviews'] = {arm: statistics([r['delta_y'] for r in records if r['family']=='heldout_reviews' and r['arm']==arm]) for arm in arms}
        report['primary_positive_mean_holds'] = report['reviews']['positive_small']['mean']>0
        report['opposite_negative_mean_holds'] = report['reviews']['negative_small']['mean']<0
        report['pairs'] = {prefix: {arm: statistics([r['delta_y'] for r in records if r['prompt'].startswith(prefix) and r['arm']==arm]) for arm in arms} for prefix in ('H1','H2')}
        report['offtask'] = [r for r in records if r['family']=='heldout_offtask']
    else:
        by_id = {r['prompt']: r for r in records if r['arm']=='observe'}
        originals = {}
        for event in complete:
            if event['phase']=='evaluating' and event['arm']=='observe':
                with np.load(run.file(event['array_file']), allow_pickle=False) as a: originals[event['prompt']] = torch.from_numpy(a['original']).double()
        d = torch.tensor(plan['direction']['d'], dtype=torch.float64); norm = float(d.norm()); pairs = {}
        for prefix, count in [('H',3),('B',4),('C',2)]:
            for n in range(1,count+1):
                key = prefix+str(n); plus, minus = by_id[key+'+'], by_id[key+'-']; contrast = originals[key+'+']-originals[key+'-']; cn = float(contrast.norm())
                pairs[key] = {'a': plus['projection']-minus['projection'] if plus['projection'] is not None else None,
                    'b': plus['y']-minus['y'], 'contrast_norm': cn,
                    'cosine': float(contrast.dot(d)/(cn*norm)) if cn>1e-8 and norm>1e-8 else None,
                    'cosine_status': 'defined' if cn>1e-8 and norm>1e-8 else 'undefined_degenerate_norm'}
        report.update(pairs=pairs, decisions=persona_decisions(pairs), lexical_controls=read_json(run.file('shortcut-controls.json')),
            extension_status='completed' if plan['optional_intervention'] else 'skipped',
            extension={arm: statistics([r['delta_y'] for r in records if r['arm']==arm]) for arm in ('zero','positive_small','negative_small','random_positive','random_negative')})
        report['extension_positive_mean_holds'] = report['extension']['positive_small']['mean']>1e-5 if plan['optional_intervention'] else None
        report['extension_negative_mean_holds'] = report['extension']['negative_small']['mean']<-1e-5 if plan['optional_intervention'] else None
    if not run.file('summary.json').exists(): run.json('summary.json', report)
    if not run.file('metrics.csv').exists():
        columns = sorted(set().union(*(r.keys() for r in records)))
        buffer = io.StringIO(); writer = csv.DictWriter(buffer, fieldnames=columns); writer.writeheader(); writer.writerows(records)
        run.put('metrics.csv', buffer.getvalue().encode())
    if not run.file('report.md').exists():
        run.put('report.md', ('# Saved numerical report\n\nThis mechanically generated report is not reviewed teaching prose. '
            'Interpret all rows, including failed hypotheses and null scores; complete the Lab\'s written evidence audit separately.\n\n'
            '```json\n'+json.dumps(report,indent=2,allow_nan=False)+'\n```\n').encode())
    print(json.dumps(report, allow_nan=False))


def worker(lab, phase, directory, model_directory):
    run = Run(directory); state = read_json(run.file('state.json')); e = None
    try:
        if os.environ.get('LLMCOURSE_SUPERVISOR_PID') != str(os.getppid()):
            raise ValueError('Internal worker requires the owning supervisor')
        if state['lab'] != lab or state['phase'] != ('discovering' if phase=='discover' else 'evaluating'):
            raise ValueError('Internal worker phase mismatch')
        plan = verify_plan(run, state) if phase == 'evaluate' else None
        torch = runtime(); torch.manual_seed(lab); torch.use_deterministic_algorithms(True)
        model, tokenizer, artifacts = load_specimen(model_directory, torch)
        actual_environment = environment(torch)
        if phase == 'discover':
            run.json('environment.json', actual_environment); run.json('artifact-manifest.json', artifacts)
            fixture = read_json(Path(__file__).with_name(f'lab{lab}-fixtures.json'))
            run.json('fixtures.json', fixture); rows = fixture_rows(lab, fixture)
            ids = tokenize(rows, fixture, tokenizer, lab); run.json('tokens.json', {'rows':rows,'score_ids':ids})
            run.json('sources.json', source_hashes(source_names(lab))); run.json('checks.json', [])
        else:
            if artifacts != read_json(run.file('artifact-manifest.json')): raise ValueError('Model artifact mismatch')
            saved_environment = read_json(run.file('environment.json'))
            if any(actual_environment[k] != saved_environment[k] for k in actual_environment if k != 'timestamp_utc'):
                raise ValueError('Environment differs across phases')
            token_record = read_json(run.file('tokens.json')); rows = token_record['rows']; ids = token_record['score_ids']
        e = Experiment(run, state, model, tokenizer, torch)
        if phase == 'evaluate': e.checks = read_json(run.file('checks.json'))
        with torch.inference_mode():
            if phase == 'discover': discover(e, rows, ids, read_json(run.file('options.json')))
            else: evaluate(e, rows, plan)
    except BaseException as exc:
        state['phase'] = 'failed'; state['failure'] = type(exc).__name__+': '+str(exc)
        if e is not None: e.save_state()
        else: run.replace('state.json', state)
        run.json('failure-'+phase+'.json', {'timestamp':stamp(),'reason':state['failure'],'attempts':state['attempts']})
        raise


def watch_worker(child, run, prior_elapsed, setup_limit=120):
    """Parent-owned monotonic deadline; worker cannot reset the remaining allowance."""
    setup_started = time.monotonic(); active_record = None; reason = None
    try:
        while child.poll() is None:
            if active_record is None and run.file('active-phase.json').exists():
                active_record = read_json(run.file('active-phase.json'))
            now = time.monotonic()
            if active_record and now >= active_record['started_monotonic']+SECONDS-prior_elapsed:
                reason = 'time_budget_exceeded'; break
            if active_record is None and now-setup_started > setup_limit:
                reason = 'setup_timeout'; break
            time.sleep(.02)
    finally:
        if child.poll() is None:
            child.kill(); child.wait()
    return reason, active_record


def supervise(lab, phase, directory, model_directory):
    run = Run(directory)
    # Invalid or premature requests must not consume the exclusive phase claim.
    state = read_json(run.file('state.json'))
    if state['lab'] != lab or state['phase'] != ('new' if phase=='discover' else 'discovered'):
        raise ValueError('Phase cannot start from current state')
    if phase == 'evaluate': verify_plan(run, state)
    # The exclusive marker arbitrates concurrently validated requests. Re-read
    # phase/plan after claiming, before any transition or child process starts.
    run.put(phase+'-claimed', stamp().encode())
    state = read_json(run.file('state.json'))
    if state['lab'] != lab or state['phase'] != ('new' if phase=='discover' else 'discovered'):
        raise ValueError('Phase changed before claimed transition')
    if phase == 'evaluate': verify_plan(run, state)
    state['phase'] = 'discovering' if phase=='discover' else 'evaluating'; run.replace('state.json', state)
    active = run.file('active-phase.json')
    if active.exists(): active.unlink()
    killed = None; active_record = None
    with run.file(phase+'-worker.log').open('xb') as log:
        child = subprocess.Popen([sys.executable, str(Path(__file__).with_name(f'run_lab{lab}.py')),
            '_worker', '--phase', phase, '--run', str(run.path), '--model-dir', str(safe_path(model_directory))],
            stdout=log, stderr=subprocess.STDOUT, start_new_session=True,
            env={**os.environ,'LLMCOURSE_SUPERVISOR_PID':str(os.getpid())})
        killed, active_record = watch_worker(child, run, state['elapsed_seconds'])
    if killed or child.returncode:
        finalize_failure(run, phase, killed or 'worker_failed', active_record, state['elapsed_seconds'])
        raise SystemExit('Bounded phase failed; evidence retained in '+str(run.path))
    print(phase+' completed; '+str(run.path))


def finalize_failure(run, phase, reason, active_record, prior_elapsed):
    state = read_json(run.file('state.json'))
    if state['phase'] != 'failed':
        state['phase'] = 'failed'; state['failure'] = reason
        if active_record: state['elapsed_seconds'] = min(SECONDS, prior_elapsed+time.monotonic()-active_record['started_monotonic'])
        run.replace('state.json', state)
        run.json('supervisor-failure-'+phase+'.json', {'timestamp':stamp(),'reason':state['failure'],
            'attempts':state['attempts'],'hook_cleanup':'unknown after forced process termination; no hooks survive process'})
    return state


def main(lab):
    parser = argparse.ArgumentParser(description=__doc__)
    sub = parser.add_subparsers(dest='command', required=True)
    discovery = sub.add_parser('discover'); discovery.add_argument('--model-dir', required=True); discovery.add_argument('--out', required=True)
    discovery.add_argument('--optional-intervention', action='store_true')
    freeze = sub.add_parser('freeze'); freeze.add_argument('--run', required=True)
    freeze.add_argument('--expectations-json', required=True)
    evaluation = sub.add_parser('evaluate'); evaluation.add_argument('--model-dir', required=True); evaluation.add_argument('--run', required=True)
    summary = sub.add_parser('summarize'); summary.add_argument('--run', required=True)
    internal = sub.add_parser('_worker'); internal.add_argument('--phase', choices=['discover','evaluate'], required=True)
    internal.add_argument('--run', required=True); internal.add_argument('--model-dir', required=True)
    args = parser.parse_args()
    if args.command == '_worker':
        worker(lab, args.phase, args.run, args.model_dir)
    elif args.command == 'freeze':
        run = Run(args.run); state = read_json(run.file('state.json'))
        if state['lab'] != lab: raise ValueError('Wrong Lab run')
        freeze_predictions(run, state, read_json(args.expectations_json))
        print('Post-discovery predictions frozen; evaluate may now run.')
    elif args.command == 'summarize':
        run = Run(args.run); state = read_json(run.file('state.json'))
        if state['phase'] != 'completed': raise SystemExit('Summary requires completed run; inspect partial attempt evidence directly')
        summarize(run, state)
    elif args.command == 'discover':
        if lab == 19 and args.optional_intervention: raise SystemExit('Optional extension belongs to Lab 20 only')
        path = safe_path(args.out); path.mkdir(mode=0o700, parents=True, exist_ok=False); run = Run(path)
        run.json('options.json', {'optional_intervention':args.optional_intervention})
        run.json('state.json', {'lab':lab,'phase':'new','attempts':0,'elapsed_seconds':0.0})
        supervise(lab, 'discover', path, args.model_dir)
    else:
        supervise(lab, 'evaluate', args.run, args.model_dir)


if __name__ == '__main__':
    raise SystemExit('Use run_lab19.py or run_lab20.py')
