#!/usr/bin/env python3
"""
Protenix Fine-tuned RNA Inference - OFFLINE VERSION
Stanford RNA 3D Folding Part 2
Based on 7th place Part 1 solution approach
Internet OFF for code competition submission
"""
import os
import sys
import json
import time
import subprocess
import warnings
import gc
import glob
import shutil
warnings.filterwarnings('ignore')

# ============================================================
# PATHS
# ============================================================
COMPETITION_DATA = "/kaggle/input/stanford-rna-3d-folding-2"
if not os.path.exists(COMPETITION_DATA):
    COMPETITION_DATA = "/kaggle/input/competitions/stanford-rna-3d-folding-2"
OUTPUT_DIR = "/kaggle/working"
PROTENIX_SRC = "/kaggle/working/Protenix-RNA-Kaggle"

# ============================================================
# STEP 1: Extract and install Protenix (7th place fork)
# ============================================================
print("=" * 60)
print("STEP 1: Installing Protenix (7th place fork) - OFFLINE")
print("=" * 60)

# Debug: list all mounted datasets FIRST
print("Mounted datasets in /kaggle/input/:")
if os.path.exists("/kaggle/input"):
    for d in sorted(os.listdir("/kaggle/input")):
        dpath = os.path.join("/kaggle/input", d)
        if os.path.isdir(dpath):
            contents = os.listdir(dpath)
            print(f"  {d}/ ({len(contents)} items): {contents[:15]}")
        else:
            print(f"  {d} ({os.path.getsize(dpath)} bytes)")
else:
    print("  /kaggle/input does not exist!")

# Auto-detect dataset paths (Kaggle can mount at different paths)
def find_dataset(slug, expected_file=None):
    """Find a dataset by trying common Kaggle mount patterns."""
    candidates = [
        f"/kaggle/input/{slug}",
        f"/kaggle/input/datasets/artembredikhin/{slug}",
    ]
    # Also check if slug matches any existing dir
    if os.path.exists("/kaggle/input"):
        for d in os.listdir("/kaggle/input"):
            if slug in d.lower() or d.lower() in slug:
                candidates.append(f"/kaggle/input/{d}")
    
    for path in candidates:
        if os.path.exists(path):
            if expected_file:
                if os.path.exists(os.path.join(path, expected_file)):
                    return path
            else:
                return path
    return candidates[0]  # Return default even if not found

CHECKPOINT_DATASET = find_dataset("protenix-rna-finetuned-v2")
CCD_CACHE_DATASET = find_dataset("protenix-ccd-cache")
PROTENIX_CODE_DATASET = find_dataset("protenix-rna-kaggle-code")
PIP_DEPS_DATASET = find_dataset("protenix-pip-deps")

print(f"\nResolved paths:")
print(f"  Checkpoint: {CHECKPOINT_DATASET} (exists: {os.path.exists(CHECKPOINT_DATASET)})")
print(f"  CCD cache:  {CCD_CACHE_DATASET} (exists: {os.path.exists(CCD_CACHE_DATASET)})")
print(f"  Code:       {PROTENIX_CODE_DATASET} (exists: {os.path.exists(PROTENIX_CODE_DATASET)})")
print(f"  Pip deps:   {PIP_DEPS_DATASET} (exists: {os.path.exists(PIP_DEPS_DATASET)})")

# Copy the source code from the dataset
if not os.path.exists(PROTENIX_SRC):
    print("\nCopying Protenix source code...")
    
    # Try multiple strategies
    found = False
    
    # Strategy 1: configs/ at root of dataset
    if os.path.exists(os.path.join(PROTENIX_CODE_DATASET, "configs")):
        print(f"Copying from: {PROTENIX_CODE_DATASET} (configs at root)")
        shutil.copytree(PROTENIX_CODE_DATASET, PROTENIX_SRC, symlinks=True)
        found = True
    
    # Strategy 2: Protenix-RNA-Kaggle subdirectory
    if not found:
        sub = os.path.join(PROTENIX_CODE_DATASET, "Protenix-RNA-Kaggle")
        if os.path.exists(sub) and os.path.isdir(sub):
            print(f"Copying from: {sub}")
            shutil.copytree(sub, PROTENIX_SRC, symlinks=True)
            found = True
    
    # Strategy 3: tar.gz file in dataset
    if not found:
        for f in glob.glob(os.path.join(PROTENIX_CODE_DATASET, "*.tar*")) + \
                 glob.glob(os.path.join(PROTENIX_CODE_DATASET, "*.tgz")):
            print(f"Found archive: {f}, extracting...")
            os.makedirs(PROTENIX_SRC, exist_ok=True)
            subprocess.run(["tar", "xf", f, "-C", PROTENIX_SRC, "--strip-components=1"], check=True)
            found = True
            break
    
    # Strategy 4: Any directory in dataset with configs/
    if not found and os.path.exists(PROTENIX_CODE_DATASET):
        for item in os.listdir(PROTENIX_CODE_DATASET):
            item_path = os.path.join(PROTENIX_CODE_DATASET, item)
            if os.path.isdir(item_path) and os.path.exists(os.path.join(item_path, "configs")):
                print(f"Copying from: {item_path}")
                shutil.copytree(item_path, PROTENIX_SRC, symlinks=True)
                found = True
                break
    
    if not found:
        # Last resort: copy whatever is there
        if os.path.exists(PROTENIX_CODE_DATASET) and os.path.isdir(PROTENIX_CODE_DATASET):
            print(f"WARNING: No standard structure found, copying entire dataset")
            print(f"Dataset contents: {os.listdir(PROTENIX_CODE_DATASET)}")
            shutil.copytree(PROTENIX_CODE_DATASET, PROTENIX_SRC, symlinks=True)
            found = True
        else:
            raise FileNotFoundError(
                f"Could not find Protenix source. "
                f"PROTENIX_CODE_DATASET={PROTENIX_CODE_DATASET} "
                f"exists={os.path.exists(PROTENIX_CODE_DATASET)}"
            )
    
    print("Source code copied")
else:
    print("Protenix-RNA-Kaggle already exists")

# Install dependencies from local wheels (OFFLINE)
print("Installing dependencies from local wheels...")
print(f"Wheels directory: {PIP_DEPS_DATASET}")
if os.path.exists(PIP_DEPS_DATASET):
    print(f"Contents: {os.listdir(PIP_DEPS_DATASET)[:20]}")
else:
    print("WARNING: PIP_DEPS_DATASET not found!")

# Find ALL .whl files (search recursively in case of subdirectories)
whl_files = sorted(glob.glob(os.path.join(PIP_DEPS_DATASET, "**", "*.whl"), recursive=True))
if not whl_files:
    whl_files = sorted(glob.glob(os.path.join(PIP_DEPS_DATASET, "*.whl")))
print(f"Found {len(whl_files)} wheel files")
for w in whl_files:
    print(f"  {os.path.basename(w)}")

if whl_files:
    subprocess.run(
        [sys.executable, "-m", "pip", "install", "--quiet", "--no-deps"] + whl_files,
        check=True
    )
else:
    print("WARNING: No wheel files found, trying --find-links fallback...")
    # Fallback: try --find-links approach
    subprocess.run([
        sys.executable, "-m", "pip", "install", "--quiet",
        "--no-index", "--find-links=" + PIP_DEPS_DATASET,
        "rdkit", "ml_collections", "optree", "modelcif==0.7", "biotite==1.0.1",
        "scikit-learn-extra", "biopython==1.83", "gemmi",
        "protobuf==3.20.2"
    ], check=True)

# Verify rdkit
subprocess.run([sys.executable, "-c", "from rdkit import Chem; print('rdkit OK')"], check=True)

# Install protenix from source (editable)
print("Installing Protenix from source...")
subprocess.run([
    sys.executable, "-m", "pip", "install", "--quiet", "--no-deps", "-e", PROTENIX_SRC
], check=True)

print("Protenix installed (offline mode)")

# ============================================================
# STEP 2: Find checkpoint
# ============================================================
print("\n" + "=" * 60)
print("STEP 2: Finding checkpoint")
print("=" * 60)

print("Checkpoint dataset contents:")
if os.path.exists(CHECKPOINT_DATASET):
    for root, dirs, files in os.walk(CHECKPOINT_DATASET):
        for f in files:
            fpath = os.path.join(root, f)
            print(f"  {fpath} ({os.path.getsize(fpath)/1e9:.2f} GB)")
else:
    print(f"  WARNING: {CHECKPOINT_DATASET} not found!")
    # Try alternative names
    for d in os.listdir("/kaggle/input"):
        print(f"  /kaggle/input/{d}")

# Find the checkpoint
checkpoint_path = None
if os.path.exists(CHECKPOINT_DATASET):
    for root, dirs, files in os.walk(CHECKPOINT_DATASET):
        for f in files:
            if f.endswith('_ema_0.995.pt'):
                checkpoint_path = os.path.join(root, f)
                break
            elif f.endswith('.pt') and checkpoint_path is None:
                checkpoint_path = os.path.join(root, f)
        if checkpoint_path and '_ema_' in checkpoint_path:
            break

if checkpoint_path is None:
    # Check all input directories for .pt files
    for d in os.listdir("/kaggle/input"):
        for root, dirs, files in os.walk(f"/kaggle/input/{d}"):
            for f in files:
                if f.endswith('.pt'):
                    checkpoint_path = os.path.join(root, f)
                    print(f"  Found checkpoint at: {checkpoint_path}")
                    break
            if checkpoint_path:
                break
        if checkpoint_path:
            break

print(f"Checkpoint: {checkpoint_path}")
assert checkpoint_path is not None, "No checkpoint found!"

# ============================================================
# STEP 3: Fix hardcoded paths in inference.py
# ============================================================
print("\n" + "=" * 60)
print("STEP 3: Patching source code")
print("=" * 60)

# Fix inference.py - remove hardcoded checkpoint paths
inference_py = os.path.join(PROTENIX_SRC, "runner", "inference.py")
with open(inference_py, 'r') as f:
    content = f.read()

# Replace the entire load_checkpoint method to use config path properly
old_load = """    def load_checkpoint(self) -> None:
        checkpoint_path = self.configs.load_checkpoint_path
        checkpoint_path = '/home/lhw/work/rna2025/Protenix/output/protenix_finetune_20250409_185733/checkpoints/499_ema_0.995.pt'
        #checkpoint_path = '/home/lhw/work/rna2025/Protenix/output/protenix_finetune_20250409_180051/checkpoints/99_ema_0.999.pt'
        #checkpoint_path = '/home/lhw/work/rna2025/Protenix/output/protenix_finetune_20250409_180051/checkpoints/99.pt'
        #checkpoint_path = '/home/lhw/work/rna2025/Protenix/output2/protenix_finetune_20250410_081032/checkpoints/999_ema_0.995.pt'
        #checkpoint_path = '/home/lhw/work/rna2025/Protenix/output_top5/protenix_finetune_20250410_142904/checkpoints/999_ema_0.995.pt'
        checkpoint_path = '/home/lhw/work/rna2025/Protenix/output_comp_no_msa/protenix_finetune_20250410_204706/checkpoints/3999_ema_0.995.pt'
        print(checkpoint_path)
        #exit(0)"""

new_load = """    def load_checkpoint(self) -> None:
        checkpoint_path = self.configs.load_checkpoint_path
        print(f"Using checkpoint: {checkpoint_path}")"""

if old_load in content:
    content = content.replace(old_load, new_load)
    print("Patched load_checkpoint (exact match)")
else:
    # Try line-by-line approach
    lines = content.split('\n')
    new_lines = []
    in_load_checkpoint = False
    skip_hardcoded = False
    
    for i, line in enumerate(lines):
        if 'def load_checkpoint(self)' in line:
            in_load_checkpoint = True
            new_lines.append(line)
            continue
        
        if in_load_checkpoint and "checkpoint_path = '/home/lhw/" in line:
            if not skip_hardcoded:
                skip_hardcoded = True
            continue
        elif in_load_checkpoint and line.strip().startswith("#checkpoint_path"):
            continue
        elif in_load_checkpoint and line.strip() == "print(checkpoint_path)":
            new_lines.append("        print(f'Using checkpoint: {checkpoint_path}')")
            continue
        elif in_load_checkpoint and line.strip() == "#exit(0)":
            in_load_checkpoint = False
            skip_hardcoded = False
            continue
        else:
            new_lines.append(line)
    
    content = '\n'.join(new_lines)
    print("Patched load_checkpoint (line-by-line)")

with open(inference_py, 'w') as f:
    f.write(content)

# Fix DATA_ROOT_DIR in configs_data.py
configs_data_py = os.path.join(PROTENIX_SRC, "configs", "configs_data.py")
with open(configs_data_py, 'r') as f:
    content = f.read()

content = content.replace(
    'DATA_ROOT_DIR = "/home/lhw/work/rna2025/release_data/"',
    'DATA_ROOT_DIR = "/kaggle/working/release_data/"'
)
content = content.replace(
    'DATA_ROOT_DIR = "/workspace/release_data/"',
    'DATA_ROOT_DIR = "/kaggle/working/release_data/"'
)

with open(configs_data_py, 'w') as f:
    f.write(content)

print("Patched DATA_ROOT_DIR")

# ============================================================
# STEP 4: Setup CCD cache
# ============================================================
print("\n" + "=" * 60)
print("STEP 4: Setting up CCD cache")
print("=" * 60)

release_data_dir = "/kaggle/working/release_data"
os.makedirs(release_data_dir, exist_ok=True)

# Check CCD cache dataset
if os.path.exists(CCD_CACHE_DATASET):
    print("CCD cache dataset found:")
    for f in os.listdir(CCD_CACHE_DATASET):
        src = os.path.join(CCD_CACHE_DATASET, f)
        dst = os.path.join(release_data_dir, f)
        if not os.path.exists(dst):
            os.symlink(src, dst)
            print(f"  Linked: {f}")
        
        # Also create non-versioned names
        if "v20240608" in f:
            short_name = f.replace(".v20240608", "")
            short_dst = os.path.join(release_data_dir, short_name)
            if not os.path.exists(short_dst):
                os.symlink(src, short_dst)
                print(f"  Linked: {short_name} -> {f}")
else:
    print("WARNING: CCD cache dataset not found - inference may fail!")

print(f"Release data contents:")
for f in os.listdir(release_data_dir):
    fpath = os.path.join(release_data_dir, f)
    if os.path.islink(fpath):
        target = os.readlink(fpath)
        print(f"  {f} -> {target}")
    else:
        print(f"  {f} ({os.path.getsize(fpath)/1e6:.0f}MB)")

# ============================================================
# STEP 5: Prepare test data
# ============================================================
print("\n" + "=" * 60)
print("STEP 5: Preparing test data")
print("=" * 60)

import pandas as pd

test_csv = os.path.join(COMPETITION_DATA, "test_sequences.csv")
test_df = pd.read_csv(test_csv)
print(f"Test sequences: {len(test_df)}")
print(f"Lengths: min={test_df['sequence'].str.len().min()}, "
      f"max={test_df['sequence'].str.len().max()}, "
      f"mean={test_df['sequence'].str.len().mean():.0f}")

msa_dir = os.path.join(COMPETITION_DATA, "MSA")
has_msa = os.path.exists(msa_dir)
if has_msa:
    msa_files = [f for f in os.listdir(msa_dir) if f.endswith('.fasta')]
    print(f"MSA files: {len(msa_files)}")

# Create input JSON for ALL sequences at once (batch inference)
all_entries = []
for _, row in test_df.iterrows():
    entry = {
        "sequences": [{
            "rnaSequence": {
                "sequence": row['sequence'],
                "count": 1,
            }
        }],
        "name": row['target_id']
    }
    all_entries.append(entry)

input_json_path = os.path.join(OUTPUT_DIR, "test_input.json")
with open(input_json_path, 'w') as f:
    json.dump(all_entries, f)
print(f"Input JSON: {input_json_path} ({len(all_entries)} entries)")

# ============================================================
# STEP 6: Run inference (per-sequence to handle OOM gracefully)
# ============================================================
print("\n" + "=" * 60)
print("STEP 6: Running inference")
print("=" * 60)

predictions_dir = os.path.join(OUTPUT_DIR, "predictions")
os.makedirs(predictions_dir, exist_ok=True)

SEEDS = [101, 102]
N_STEP = 200
N_CYCLE = 10

# Sort sequences by length (shortest first -- more likely to succeed)
seq_lengths = [(row['target_id'], len(row['sequence'])) for _, row in test_df.iterrows()]
seq_lengths.sort(key=lambda x: x[1])
print(f"Sequence lengths: {[(t, l) for t, l in seq_lengths]}")

# T4 has 16GB VRAM. Adjust N_sample based on sequence length to avoid OOM.
def get_n_sample(seq_len):
    if seq_len <= 100:
        return 5
    elif seq_len <= 200:
        return 3
    elif seq_len <= 350:
        return 2
    else:
        return 1

succeeded = set()
failed = set()

MAX_SEQ_LEN = 1500  # Skip sequences longer than this to avoid system OOM (kills kernel)

for target_id, seq_len in seq_lengths:
    if seq_len > MAX_SEQ_LEN:
        print(f"\n  {target_id} (len={seq_len}): SKIPPED (>{MAX_SEQ_LEN}nt, would OOM)")
        failed.add(target_id)
        continue
    
    n_sample = get_n_sample(seq_len)
    
    # Create individual input JSON for this sequence
    row = test_df[test_df['target_id'] == target_id].iloc[0]
    entry = [{
        "sequences": [{"rnaSequence": {"sequence": row['sequence'], "count": 1}}],
        "name": target_id
    }]
    single_json = os.path.join(OUTPUT_DIR, f"input_{target_id}.json")
    with open(single_json, 'w') as f:
        json.dump(entry, f)
    
    target_ok = False
    for seed in SEEDS:
        # Skip if already have predictions from previous seed
        existing_cifs = glob.glob(os.path.join(predictions_dir, "**", f"*{target_id}*seed*{seed}*.cif"), recursive=True)
        if existing_cifs:
            print(f"  {target_id} seed {seed}: already have {len(existing_cifs)} CIFs, skipping")
            target_ok = True
            continue
        
        cmd = [
            sys.executable, "-m", "runner.inference",
            "--input_json_path", single_json,
            "--dump_dir", predictions_dir,
            "--load_checkpoint_path", checkpoint_path,
            "--dtype", "bf16",
            "--use_msa", "false",
            "--num_workers", "0",
            "--seeds", str(seed),
            "--sample_diffusion.N_step", str(N_STEP),
            "--sample_diffusion.N_sample", str(n_sample),
            "--model.N_cycle", str(N_CYCLE),
        ]
        
        t0 = time.time()
        print(f"\n  {target_id} (len={seq_len}, samples={n_sample}, seed={seed})...", end=" ", flush=True)
        
        try:
            result = subprocess.run(
                cmd,
                capture_output=True,
                text=True,
                timeout=3600,  # 1 hour max per sequence
                cwd=PROTENIX_SRC,
                env={**os.environ, "PYTHONPATH": PROTENIX_SRC}
            )
            elapsed = time.time() - t0
            
            if result.returncode == 0:
                print(f"OK ({elapsed:.0f}s)")
                target_ok = True
            else:
                print(f"FAIL ({elapsed:.0f}s)")
                # Check if it's OOM
                if "CUDA out of memory" in result.stderr or result.returncode == -9:
                    print(f"    GPU OOM on {target_id} (len={seq_len})")
                    if n_sample > 1:
                        # Retry with N_sample=1
                        print(f"    Retrying with N_sample=1...")
                        cmd_retry = cmd.copy()
                        idx = cmd_retry.index("--sample_diffusion.N_sample") + 1
                        cmd_retry[idx] = "1"
                        try:
                            result2 = subprocess.run(
                                cmd_retry, capture_output=True, text=True,
                                timeout=3600, cwd=PROTENIX_SRC,
                                env={**os.environ, "PYTHONPATH": PROTENIX_SRC}
                            )
                            if result2.returncode == 0:
                                print(f"    Retry OK")
                                target_ok = True
                            else:
                                print(f"    Retry also failed")
                        except:
                            print(f"    Retry exception")
                else:
                    # Print last few lines of stderr for debugging
                    stderr_tail = result.stderr.strip().split('\n')[-5:]
                    for line in stderr_tail:
                        print(f"    {line}")
        except subprocess.TimeoutExpired:
            print(f"TIMEOUT (>3600s)")
        except Exception as e:
            print(f"EXCEPTION: {e}")
        
        gc.collect()
    
    if target_ok:
        succeeded.add(target_id)
    else:
        failed.add(target_id)

print(f"\n\nInference summary: {len(succeeded)} succeeded, {len(failed)} failed")
if failed:
    print(f"Failed targets: {failed}")

# ============================================================
# STEP 7: Collect CIF outputs and extract coordinates
# ============================================================
print("\n" + "=" * 60)
print("STEP 7: Processing predictions")
print("=" * 60)

import numpy as np
from Bio.PDB import MMCIFParser

def extract_c1_prime_coords(cif_path):
    """Extract C1' atom coordinates from a CIF file."""
    parser = MMCIFParser(QUIET=True)
    structure = parser.get_structure('pred', cif_path)
    coords = []
    for model in structure:
        for chain in model:
            for residue in chain:
                for atom in residue:
                    if atom.get_name() == "C1'":
                        coords.append(atom.get_coord().tolist())
    return coords

def simplified_tm_score(coords1, coords2):
    """Simplified TM-score using C1' atoms."""
    coords1, coords2 = np.array(coords1), np.array(coords2)
    if len(coords1) != len(coords2) or len(coords1) == 0:
        return 0.0
    L = len(coords1)
    d0 = max(1.24 * (L - 15) ** (1.0/3.0) - 1.8, 0.5)
    c1 = coords1 - coords1.mean(axis=0)
    c2 = coords2 - coords2.mean(axis=0)
    H = c1.T @ c2
    U, S, Vt = np.linalg.svd(H)
    d = np.sign(np.linalg.det(Vt.T @ U.T))
    R = Vt.T @ np.diag([1, 1, d]) @ U.T
    c2a = (R @ c2.T).T
    dists = np.sqrt(np.sum((c1 - c2a) ** 2, axis=1))
    return np.sum(1.0 / (1.0 + (dists / d0) ** 2)) / L

def select_diverse_k(all_coords, k=5):
    """Select k diverse structures using greedy farthest-point."""
    n = len(all_coords)
    if n <= k:
        return list(range(n))
    dist = np.zeros((n, n))
    for i in range(n):
        for j in range(i+1, n):
            tm = simplified_tm_score(all_coords[i], all_coords[j])
            dist[i,j] = dist[j,i] = 1 - tm
    try:
        from sklearn_extra.cluster import KMedoids
        km = KMedoids(n_clusters=k, metric='precomputed', random_state=42)
        km.fit(dist)
        return km.medoid_indices_.tolist()
    except:
        selected = [0]
        for _ in range(k - 1):
            md = np.min(dist[selected], axis=0)
            md[selected] = -1
            selected.append(np.argmax(md))
        return selected

# List predictions directory structure for debugging
print("Predictions directory structure:")
for root, dirs, files in os.walk(predictions_dir):
    depth = root.replace(predictions_dir, '').count(os.sep)
    indent = '  ' * depth
    print(f"{indent}{os.path.basename(root)}/")
    if depth < 3:  # Don't go too deep
        for f in files[:10]:
            print(f"{indent}  {f}")
        if len(files) > 10:
            print(f"{indent}  ... and {len(files)-10} more files")

# Find all CIF files
all_cifs = glob.glob(os.path.join(predictions_dir, "**", "*.cif"), recursive=True)
print(f"\nTotal CIF files found: {len(all_cifs)}")
for c in all_cifs[:20]:
    print(f"  {c}")

# Group by target
target_cifs = {}
for cif_path in all_cifs:
    basename = os.path.basename(cif_path)
    for _, row in test_df.iterrows():
        tid = row['target_id']
        if tid in basename or tid in cif_path:
            if tid not in target_cifs:
                target_cifs[tid] = []
            target_cifs[tid].append(cif_path)
            break

print(f"Targets with predictions: {len(target_cifs)}/{len(test_df)}")

# Extract coords and select diverse structures
all_predictions = {}
for target_id in test_df['target_id']:
    cifs = target_cifs.get(target_id, [])
    if not cifs:
        print(f"  {target_id}: NO predictions")
        continue
    
    all_coords = []
    for cif_path in sorted(cifs):
        try:
            coords = extract_c1_prime_coords(cif_path)
            if coords:
                all_coords.append(coords)
        except Exception as e:
            print(f"  Error: {os.path.basename(cif_path)}: {e}")
    
    if all_coords:
        if len(all_coords) > 5:
            idx = select_diverse_k(all_coords, k=5)
            selected = [all_coords[i] for i in idx]
        else:
            selected = all_coords
        all_predictions[target_id] = selected
        print(f"  {target_id}: {len(all_coords)} -> {len(selected)} selected")

# Clean up cloned repo from /kaggle/working to reduce output file count
# (Kaggle saves everything in /kaggle/working as kernel output)
print("Cleaning up source repo from output...")
if os.path.exists(PROTENIX_SRC):
    shutil.rmtree(PROTENIX_SRC, ignore_errors=True)
    print("Removed Protenix-RNA-Kaggle from working dir")

# Also clean up individual input JSONs
for f in glob.glob(os.path.join(OUTPUT_DIR, "input_*.json")):
    os.remove(f)

# ============================================================
# STEP 8: Generate submission.csv
# ============================================================
print("\n" + "=" * 60)
print("STEP 8: Generating submission")
print("=" * 60)

rows = []
for _, row in test_df.iterrows():
    target_id = row['target_id']
    sequence = row['sequence']
    seq_len = len(sequence)
    pred_coords = all_predictions.get(target_id, None)
    
    for resid in range(1, seq_len + 1):
        row_data = {
            "ID": f"{target_id}_{resid}",
            "resname": sequence[resid - 1],
            "resid": resid,
        }
        for i in range(5):
            if pred_coords is not None:
                src = min(i, len(pred_coords) - 1)
                cidx = resid - 1
                if cidx < len(pred_coords[src]):
                    row_data[f"x_{i+1}"] = pred_coords[src][cidx][0]
                    row_data[f"y_{i+1}"] = pred_coords[src][cidx][1]
                    row_data[f"z_{i+1}"] = pred_coords[src][cidx][2]
                else:
                    last = pred_coords[src][-1]
                    row_data[f"x_{i+1}"] = last[0]
                    row_data[f"y_{i+1}"] = last[1]
                    row_data[f"z_{i+1}"] = last[2]
            else:
                row_data[f"x_{i+1}"] = 0.0
                row_data[f"y_{i+1}"] = 0.0
                row_data[f"z_{i+1}"] = 0.0
        rows.append(row_data)

submission = pd.DataFrame(rows)
submission_path = os.path.join(OUTPUT_DIR, "submission.csv")
submission.to_csv(submission_path, index=False)

print(f"Submission: {submission_path}")
print(f"Shape: {submission.shape}")
print(f"Targets predicted: {len(all_predictions)}/{len(test_df)}")
missing = set(test_df['target_id']) - set(all_predictions.keys())
if missing:
    print(f"Missing: {missing}")

print("\nDone!")
