#!/usr/bin/env python3
"""Generate proposed context controls; does not run a model or simulate biology."""
import argparse
from collections import Counter
import hashlib
import json
from pathlib import Path
import random

LENGTH, START, END, CHUNK = 131072, 61440, 69632, 1024
SEEDS = (17, 29, 43)

def digest(sequence):
    return hashlib.sha256(sequence.encode()).hexdigest()

def controls(sequence, seed):
    if len(sequence) != LENGTH or set(sequence) - set('ACGTN'):
        raise ValueError('Expected exactly 131072 uppercase A/C/G/T/N bases')
    rng = random.Random(seed)
    permuted, destroyed, orders = [], [], []
    for flank in (sequence[:START], sequence[END:]):
        chunks = [flank[i:i+CHUNK] for i in range(0, len(flank), CHUNK)]
        order = list(range(len(chunks)))
        rng.shuffle(order)
        permuted.append(''.join(chunks[i] for i in order))
        shuffled_chunks = []
        for chunk in chunks:
            bases = list(chunk)
            rng.shuffle(bases)
            shuffled_chunks.append(''.join(bases))
        destroyed.append(''.join(shuffled_chunks))
        orders.append(order)
    focal = sequence[START:END]
    chunk_order = permuted[0] + focal + permuted[1]
    within_chunk = destroyed[0] + focal + destroyed[1]
    # These are the actual invariants the proposed experiment relies on.
    for output in (chunk_order, within_chunk):
        assert len(output) == LENGTH
        assert output[START:END] == focal
        assert Counter(output) == Counter(sequence)
    for old, new in ((sequence[:START], permuted[0]), (sequence[END:], permuted[1])):
        assert Counter(old[i:i+CHUNK] for i in range(0,len(old),CHUNK)) == Counter(new[i:i+CHUNK] for i in range(0,len(new),CHUNK))
    return chunk_order, within_chunk, orders

def main():
    parser = argparse.ArgumentParser(description=__doc__)
    inputs = parser.add_mutually_exclusive_group(required=True)
    inputs.add_argument('--fasta', type=Path, help='Exactly one sequence record')
    inputs.add_argument('--synthetic', action='store_true', help='Seeded random DNA: bookkeeping demonstration only')
    parser.add_argument('--out', type=Path, required=True)
    args = parser.parse_args()
    if args.synthetic:
        rng = random.Random(20260923)
        sequence = ''.join(rng.choice('ACGT') for _ in range(LENGTH))
        provenance = 'Synthetic independent bases; no biological endpoint or model output'
    else:
        lines = args.fasta.read_text().splitlines()
        if sum(line.startswith('>') for line in lines) != 1 or not lines[0].startswith('>'):
            raise ValueError('Input must be a single-record FASTA')
        sequence = ''.join(line.strip() for line in lines[1:]).upper()
        provenance = str(args.fasta)
    args.out.mkdir(parents=True, exist_ok=True)
    outputs = [('original', sequence), ('local_8192', sequence[START:END])]
    report = dict(provenance=provenance, original_sha256=digest(sequence),
                  focal_interval_zero_based_half_open=[START, END], chunk_size=CHUNK,
                  coordinate_note='Centre the crop on the declared TSS before invoking this script.',
                  model_inference=False, controls=[])
    for seed in SEEDS:
        chunk_order, within_chunk, orders = controls(sequence, seed)
        outputs.extend([(f'chunk_order_seed_{seed}',chunk_order), (f'within_chunk_seed_{seed}',within_chunk)])
        report['controls'].append(dict(seed=seed, flank_permutations=orders,
            chunk_order_sha256=digest(chunk_order), within_chunk_sha256=digest(within_chunk),
            invariants_passed=True))
    with (args.out/'sequences.fasta').open('w') as f:
        for label, seq in outputs:
            f.write('>'+label+'\n')
            f.write('\n'.join(seq[i:i+80] for i in range(0,len(seq),80))+'\n')
    (args.out/'manifest.json').write_text(json.dumps(report,indent=2)+'\n')
    print(json.dumps(dict(records=len(outputs), seeds=list(SEEDS), invariants_passed=True, model_inference=False)))

if __name__ == '__main__':
    main()
