#!/usr/bin/env python3
"""Extract DNA C1'-C1' distances at all 20 target positions from AF3 ZIPs."""
import argparse
import json
import re
import tempfile
import zipfile
from pathlib import Path

import gemmi
import matplotlib.pyplot as plt
import numpy as np
import pandas as pd

parser = argparse.ArgumentParser()
parser.add_argument('--manifest', required=True)
parser.add_argument('--zip-dir', required=True)
parser.add_argument('--outdir', required=True)
parser.add_argument('--reference-table', help='Previous pos3-10 target distance table for cross-check')
args = parser.parse_args()

manifest = pd.read_csv(args.manifest, sep='\t', dtype=str)
required = {'af3_id', 'target_id', 'length_group_id', 'target_20mer', 'DNA1_af3', 'DNA2_af3'}
if missing := required - set(manifest):
    raise ValueError(f'Manifest missing columns: {sorted(missing)}')
zip_dir = Path(args.zip_dir)
outdir = Path(args.outdir)
outdir.mkdir(parents=True, exist_ok=True)
dna_letters = {'DA': 'A', 'DC': 'C', 'DG': 'G', 'DT': 'T'}
rows = []

for condition in manifest.itertuples(index=False):
    af3_id = condition.af3_id
    zip_path = zip_dir / f'{af3_id}.zip'
    if not zip_path.is_file():
        raise FileNotFoundError(zip_path)
    target = condition.target_20mer.upper()
    dna1 = condition.DNA1_af3.upper()
    dna2 = condition.DNA2_af3.upper()
    if len(target) != 20 or dna1.count(target) != 1:
        raise ValueError(f'{af3_id}: target_20mer must occur once in DNA1')
    start = dna1.index(target)

    with zipfile.ZipFile(zip_path) as archive:
        jobs = [n for n in archive.namelist() if n.endswith('_job_request.json')]
        models = sorted((n for n in archive.namelist() if re.search(r'_model_\d+\.cif$', n)),
                        key=lambda n: int(re.search(r'_model_(\d+)\.cif$', n).group(1)))
        if len(jobs) != 1 or len(models) != 5:
            raise ValueError(f'{af3_id}: expected one job request and five models, found {len(jobs)} and {len(models)}')
        job = json.loads(archive.read(jobs[0]))
        if isinstance(job, list):
            job = job[0]
        job_dnas = [s['dnaSequence']['sequence'].upper() for s in job['sequences'] if 'dnaSequence' in s]
        if sorted(job_dnas) != sorted([dna1, dna2]):
            raise ValueError(f'{af3_id}: job DNA sequences do not match manifest')

        with tempfile.TemporaryDirectory() as tempdir:
            for name in models:
                model_number = int(re.search(r'_model_(\d+)\.cif$', name).group(1))
                cif_path = Path(tempdir) / Path(name).name
                cif_path.write_bytes(archive.read(name))
                structure = gemmi.read_structure(str(cif_path))
                dna_chains = {}
                for chain in structure[0]:
                    if all(res.name in dna_letters for res in chain):
                        sequence = ''.join(dna_letters[res.name] for res in chain)
                        dna_chains[sequence] = chain
                if dna1 not in dna_chains or dna2 not in dna_chains:
                    raise ValueError(f'{af3_id} model {model_number}: DNA chains do not match manifest')
                chain1, chain2 = dna_chains[dna1], dna_chains[dna2]
                for position in range(1, 21):
                    i1 = start + position
                    i2 = len(dna2) - i1 + 1
                    atom1 = chain1[i1 - 1].find_atom("C1'", '*')
                    atom2 = chain2[i2 - 1].find_atom("C1'", '*')
                    if atom1 is None or atom2 is None:
                        raise ValueError(f'{af3_id} model {model_number} position {position}: missing C1\' atom')
                    distance = float(np.linalg.norm(
                        [atom1.pos.x - atom2.pos.x, atom1.pos.y - atom2.pos.y, atom1.pos.z - atom2.pos.z]
                    ))
                    rows.append((af3_id, condition.target_id, condition.length_group_id,
                                 model_number, position, i1, i2, distance))
    print(f'{af3_id}: five models, positions 1-20')

all_models = pd.DataFrame(rows, columns=['af3_id', 'target_id', 'length_group_id', 'model',
                                          'position', 'dna1_residue', 'dna2_residue', 'distance_A'])
all_models.to_csv(outdir / 'distance_all_models_pos1_20.tsv', sep='\t', index=False)
target_means = (all_models.groupby(['af3_id', 'target_id', 'length_group_id', 'position'], as_index=False)
                .agg(distance_A=('distance_A', 'mean'), model_sd_A=('distance_A', 'std'), n_models=('model', 'nunique')))
target_means.to_csv(outdir / 'distance_by_target_pos1_20.tsv', sep='\t', index=False)

if args.reference_table:
    reference = pd.read_csv(args.reference_table, sep='\t')
    diffs = []
    for row in target_means[target_means.position.between(3, 10)].itertuples(index=False):
        match = reference.loc[reference.af3_id.eq(row.af3_id)]
        if len(match) != 1:
            raise ValueError(f'{row.af3_id}: missing or duplicate reference row')
        diffs.append(abs(row.distance_A - match.iloc[0][f'distance_pos{row.position}']))
    print(f'Max difference from previous positions 3-10: {max(diffs):.8g} Å')
    if max(diffs) > 1e-5:
        raise ValueError('Previously published distances did not reproduce; inspect chain alignment')

baseline = target_means[target_means.length_group_id.eq('L20')][['target_id', 'position', 'distance_A']]
baseline = baseline.rename(columns={'distance_A': 'distance_L20_A'})
paired = target_means.merge(baseline, on=['target_id', 'position'], validate='many_to_one')
paired['delta_from_L20_A'] = paired.distance_A - paired.distance_L20_A
paired.to_csv(outdir / 'distance_and_paired_delta_by_target_pos1_20.tsv', sep='\t', index=False)
curve = (paired.groupby(['length_group_id', 'position'], as_index=False)
         .agg(n_targets=('target_id', 'nunique'), distance_mean_A=('distance_A', 'mean'),
              distance_sd_targets_A=('distance_A', 'std'), delta_mean_A=('delta_from_L20_A', 'mean'),
              delta_sd_targets_A=('delta_from_L20_A', 'std')))
curve.to_csv(outdir / 'distance_and_delta_plot_data_pos1_20.tsv', sep='\t', index=False)

if len(manifest) == 36 and manifest.target_id.nunique() == 12:
    fig, axes = plt.subplots(1, 2, figsize=(12.4, 4.4), constrained_layout=True)
    for group, label, color in [('L16_17', '16 nt', '#c66c9a'),
                                ('L18_19', '18 nt', '#e39700'),
                                ('L20', '20 nt', '#006fa7')]:
        data = curve[curve.length_group_id.eq(group)].sort_values('position')
        if len(data) != 20 or not data.n_targets.eq(12).all():
            raise ValueError(f'{group}: expected all 12 targets at each of 20 positions')
        x = data.position.to_numpy()
        for axis, mean_col, sd_col in [(axes[0], 'distance_mean_A', 'distance_sd_targets_A'),
                                       (axes[1], 'delta_mean_A', 'delta_sd_targets_A')]:
            y = data[mean_col].to_numpy()
            sd = data[sd_col].fillna(0).to_numpy()
            axis.plot(x, y, '-o', ms=3.2, lw=1.8, color=color, label=label)
            axis.fill_between(x, y - sd, y + sd, color=color, alpha=0.13, lw=0)
    axes[0].set_title('DNA interstrand distance')
    axes[1].set_title('Distance change from 20 nt')
    axes[0].set_ylabel("C1'-C1' distance (Å)")
    axes[1].set_ylabel('Δdistance (Å)')
    for axis in axes:
        axis.set_xlabel('Protospacer position (20 nt reference)')
        axis.set_xlim(1, 20)
        axis.set_xticks(range(1, 21))
        axis.spines[['top', 'right']].set_visible(False)
    handles, labels = axes[0].get_legend_handles_labels()
    fig.legend(handles, labels, loc='lower center', ncol=3, bbox_to_anchor=(0.5, -0.08), frameon=False)
    fig.suptitle('12-target mean; shading = ±1 SD across targets', fontsize=11)
    fig.savefig(outdir / 'distance_and_delta_pos1_20.png', dpi=250, bbox_inches='tight')
    fig.savefig(outdir / 'distance_and_delta_pos1_20.pdf', bbox_inches='tight')
    plt.close(fig)
    print('Wrote full 12-target plot and TSV data to', outdir)
else:
    print('Wrote available-condition TSV data; plot requires all 36 conditions')
