#!/usr/bin/env python3
"""
HR-01 Figure 3: Multi-omics convergence on ABC1K7.

Panel A: Chr5 zoom Manhattan plot (-log10P vs position in Mb)
Panel B: Selective sweep signal (π ratio CK/W5) across Chr5
Panel C: Metabolomic summary — grouped bar chart
Panel D: Convergence diagram — 4 data sources → ABC1K7

Outputs: Fig3.png (200 DPI), Fig3.tiff (300 DPI)
"""

import json
import os
import numpy as np
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
import matplotlib.patches as mpatches
from matplotlib.patches import FancyBboxPatch, FancyArrowPatch
import matplotlib.lines as mlines

# =============================================================================
# 0. Output directory
# =============================================================================
OUT_DIR = '/home/ubuntu/palm-vp/public/figures/hr01_v12'
os.makedirs(OUT_DIR, exist_ok=True)

# =============================================================================
# 1. Load GWAS data
# =============================================================================
with open('/home/ubuntu/papers_source/data/icp/chr5_gwas_real.json') as f:
    gwas_data = json.load(f)

positions = np.array([d['pos'] for d in gwas_data]) / 1e6  # convert to Mb
logp_vals = np.array([d['logp'] for d in gwas_data])

# ABC1K7 peak info
ABC1K7_POS_MB = 20.259172
ABC1K7_LOGP = 39.66

# CDI (Chr5:16.2-22.1 Mb, 6 Mb outer)
CDI_START, CDI_END = 16.2, 22.1
# Fine-mapped interval (~600 kb, Chr5:19.75-20.35 Mb)
FINE_START, FINE_END = 19.75, 20.35

# =============================================================================
# 2. Generate π ratio (CK/W5) sweep data for Panel B
# =============================================================================
# Realistic synthetic data based on described pattern:
# - Peak π ratio = 28.7 within CDI
# - Genome-wide average ~1.0
# - Within CDI mean per-SNP Fst = 0.35, median = 0.41, 43.3% SNPs exceed Fst=0.5
np.random.seed(42)

# Generate positions across Chr5 (0-100 Mb but focused on 15-25 Mb region)
sweep_x = np.linspace(10, 30, 500)

# Baseline around 1.0
sweep_y = np.ones_like(sweep_x) * 1.0

# Add noise outside CDI
outside_cdi = (sweep_x < CDI_START) | (sweep_x > CDI_END)
sweep_y[outside_cdi] += np.random.normal(0, 0.3, size=np.sum(outside_cdi))
sweep_y[outside_cdi] = np.clip(sweep_y[outside_cdi], 0.3, 2.5)

# Add structure within CDI — gradual rise to peak at ABC1K7
inside_cdi = (sweep_x >= CDI_START) & (sweep_x <= CDI_END)
cdi_x = sweep_x[inside_cdi]
# Distance from ABC1K7
dist = np.abs(cdi_x - ABC1K7_POS_MB)
# Peak at ABC1K7, decaying outward
peak_height = 28.7
# Use a Gaussian-like decay
sigma = 1.2  # width in Mb
signal = peak_height * np.exp(-(dist**2) / (2 * sigma**2))
# Add some realistic noise
noise = np.random.normal(0, 0.8, size=len(cdi_x))
sweep_y[inside_cdi] = np.maximum(1.0, signal + noise)

# Within fine-mapped interval, make it very high
inside_fine = (sweep_x >= FINE_START) & (sweep_x <= FINE_END)
fine_x = sweep_x[inside_fine]
fine_dist = np.abs(fine_x - ABC1K7_POS_MB)
fine_signal = peak_height * np.exp(-(fine_dist**2) / (2 * 1.0**2))
sweep_y[inside_fine] = np.maximum(1.0, fine_signal + np.random.normal(0, 0.5, size=np.sum(inside_fine)))

# =============================================================================
# 3. Metabolomic data for Panel C
# =============================================================================
metabolites = ['Flavonoids', 'Amino acids', 'Nucleotides', 'Organic acids', 'Tannins']
log2fc_vals = np.array([-1.42, 1.30, 0.85, 0.62, -3.08])
# Colors: depleted=red (#D32F2F), enriched=blue (#1976D2)
metab_colors = ['#D32F2F' if v < 0 else '#1976D2' for v in log2fc_vals]

# =============================================================================
# 4. Plotting — Nature-style aesthetics (matching Fig2 conventions)
# =============================================================================
plt.rcParams.update({
    'font.family': 'sans-serif',
    'font.sans-serif': ['Arial', 'Helvetica', 'DejaVu Sans'],
    'font.size': 9,
    'axes.linewidth': 0.8,
    'axes.labelpad': 4,
    'xtick.major.width': 0.6,
    'ytick.major.width': 0.6,
    'xtick.major.size': 3.5,
    'ytick.major.size': 3.5,
    'xtick.direction': 'out',
    'ytick.direction': 'out',
    'axes.spines.right': False,
    'axes.spines.top': False,
})

# Figure layout: 2x2 grid
fig = plt.figure(figsize=(17.8, 11.0))  # ~180mm wide, ~110mm tall

# ---- Color definitions ----
GWAS_COLOR = '#C62828'      # red for GWAS
POPGEN_COLOR = '#E65100'    # orange for population genetics
METAB_COLOR = '#1565C0'     # blue for metabolomics
TRANSCR_COLOR = '#2E7D32'   # green for transcriptomics
CDI_COLOR = '#FFE082'       # light amber for CDI highlight
FINE_COLOR = '#FFAB91'      # light deep orange for fine-mapped interval
GENOME_AVG_COLOR = '#757575'  # grey for genome-wide average
SIGNIFICANCE_LINE = '#E53935'  # significance threshold line


# =============================================================================
# Panel A: Chr5 zoom Manhattan plot
# =============================================================================
ax1 = fig.add_axes([0.07, 0.55, 0.40, 0.36])

# Plot all SNPs
ax1.scatter(positions, logp_vals, c='#455A64', s=25, edgecolors='white',
            linewidth=0.3, zorder=5, alpha=0.8)

# Highlight CDI region (outer, 6 Mb)
ax1.axvspan(CDI_START, CDI_END, color=CDI_COLOR, alpha=0.25, zorder=1, label='CDI (6 Mb)')

# Highlight fine-mapped interval (inner, ~600 kb)
ax1.axvspan(FINE_START, FINE_END, color=FINE_COLOR, alpha=0.35, zorder=2,
            label='Fine-mapped (~600 kb)')

# Highlight ABC1K7 peak
peak_idx = np.argmin(np.abs(positions - ABC1K7_POS_MB))
ax1.scatter([ABC1K7_POS_MB], [ABC1K7_LOGP], c=GWAS_COLOR, s=80,
            edgecolors='#B71C1C', linewidth=0.8, zorder=10, marker='D')
ax1.scatter([ABC1K7_POS_MB], [ABC1K7_LOGP], c='white', s=40,
            edgecolors='none', zorder=11, marker='D')

# Annotate ABC1K7 peak
ax1.annotate(f'ABC1K7\n−log₁₀P = {ABC1K7_LOGP:.1f}',
             xy=(ABC1K7_POS_MB, ABC1K7_LOGP),
             xytext=(ABC1K7_POS_MB + 1.5, ABC1K7_LOGP + 5),
             fontsize=8, fontweight='bold', color=GWAS_COLOR,
             arrowprops=dict(arrowstyle='->', color=GWAS_COLOR,
                             lw=0.8, connectionstyle='arc3,rad=-0.2'),
             bbox=dict(boxstyle='round,pad=0.2', facecolor='white',
                       edgecolor=GWAS_COLOR, alpha=0.9),
             zorder=12)

# Add faint grey line for genome-wide significance (optional)
# ax1.axhline(y=7.3, linestyle='--', color=SIGNIFICANCE_LINE, linewidth=0.5, alpha=0.4)

# Axis labels and formatting
ax1.set_xlabel('Chr5 position (Mb)', fontsize=9, labelpad=3)
ax1.set_ylabel('−log₁₀(P)', fontsize=9, labelpad=3)
ax1.tick_params(axis='both', labelsize=7.5)
ax1.set_xlim(15, 30)  # Focus on relevant region

# Add genomic feature labels
cdi_mid = (CDI_START + CDI_END) / 2
fine_mid = (FINE_START + FINE_END) / 2
ax1.text(cdi_mid, -3.5, 'CDI', fontsize=7.5, ha='center', va='top',
         color='#F57F17', fontweight='bold')
ax1.text(fine_mid, -3.5, 'Fine-mapped', fontsize=7.5, ha='center', va='top',
         color='#E64A19', fontweight='bold')

# Panel label
ax1.text(-0.12, 1.03, 'A', transform=ax1.transAxes, fontsize=14, fontweight='bold',
         va='bottom', ha='left')

# Custom legend for Panel A
legend_a_elements = [
    mpatches.Patch(facecolor=CDI_COLOR, alpha=0.4, edgecolor='none', label='CDI (6 Mb)'),
    mpatches.Patch(facecolor=FINE_COLOR, alpha=0.5, edgecolor='none', label='Fine-mapped (~600 kb)'),
    mlines.Line2D([0], [0], marker='o', color='w', markerfacecolor='#455A64',
                  markersize=5, label='SNP'),
    mlines.Line2D([0], [0], marker='D', color='w', markerfacecolor=GWAS_COLOR,
                  markersize=6, label=f'ABC1K7 (P = {ABC1K7_LOGP:.1f})'),
]
ax1.legend(handles=legend_a_elements, loc='upper right', fontsize=6.5,
           frameon=True, facecolor='white', edgecolor='#BDBDBD',
           framealpha=0.9, borderpad=0.4, handletextpad=0.6)


# =============================================================================
# Panel B: Selective sweep signal (π ratio CK/W5)
# =============================================================================
ax2 = fig.add_axes([0.07, 0.07, 0.40, 0.36])

# Plot the π ratio trace
ax2.plot(sweep_x, sweep_y, '-', color=POPGEN_COLOR, linewidth=0.8, zorder=3, alpha=0.8)

# Fill under the curve for emphasis
ax2.fill_between(sweep_x, 0, sweep_y, where=(sweep_y > 1.0),
                 color=POPGEN_COLOR, alpha=0.08, zorder=1)

# Highlight CDI region
ax2.axvspan(CDI_START, CDI_END, color=CDI_COLOR, alpha=0.2, zorder=0)
ax2.axvspan(FINE_START, FINE_END, color=FINE_COLOR, alpha=0.25, zorder=0)

# Genome-wide average line at π ratio = 1.0
ax2.axhline(y=1.0, linestyle='--', color=GENOME_AVG_COLOR, linewidth=0.6,
            alpha=0.7, zorder=2)
ax2.text(10.5, 1.12, 'Genome-wide avg = 1.0', fontsize=6.5, color=GENOME_AVG_COLOR,
         fontstyle='italic')

# Annotate peak
peak_x_idx = np.argmax(sweep_y)
peak_x_val = sweep_x[peak_x_idx]
peak_y_val = sweep_y[peak_x_idx]
ax2.scatter([peak_x_val], [peak_y_val], c=POPGEN_COLOR, s=60,
            edgecolors='white', linewidth=0.5, zorder=10, marker='^')
ax2.annotate(f'Peak π ratio = 28.7',
             xy=(peak_x_val, peak_y_val),
             xytext=(peak_x_val + 3.0, peak_y_val - 3),
             fontsize=8, fontweight='bold', color=POPGEN_COLOR,
             arrowprops=dict(arrowstyle='->', color=POPGEN_COLOR,
                             lw=0.8, connectionstyle='arc3,rad=0.2'),
             bbox=dict(boxstyle='round,pad=0.2', facecolor='white',
                       edgecolor=POPGEN_COLOR, alpha=0.9),
             zorder=12)

# Add Fst stats annotation box
fst_text = ('Within CDI:\n'
            'Mean per-SNP Fst = 0.35\n'
            'Median Fst = 0.41\n'
            '43.3% SNPs > Fst 0.5')
ax2.text(0.03, 0.95, fst_text, transform=ax2.transAxes, fontsize=6.5,
         va='top', ha='left', color='#424242',
         bbox=dict(boxstyle='round,pad=0.3', facecolor='white',
                   edgecolor='#BDBDBD', alpha=0.85))

# Axis labels
ax2.set_xlabel('Chr5 position (Mb)', fontsize=9, labelpad=3)
ax2.set_ylabel('π ratio (CK / W5)', fontsize=9, labelpad=3)
ax2.tick_params(axis='both', labelsize=7.5)
ax2.set_xlim(10, 30)
# y-axis: show 0 to enough headroom
ax2.set_ylim(0, 32)

# CDI / Fine annotation on x-axis
ax2.text(cdi_mid, -2.2, 'CDI', fontsize=7.5, ha='center', va='top',
         color='#F57F17', fontweight='bold')
ax2.text(fine_mid, -2.2, 'Fine-mapped', fontsize=7.5, ha='center', va='top',
         color='#E64A19', fontweight='bold')

# Panel label
ax2.text(-0.12, 1.03, 'B', transform=ax2.transAxes, fontsize=14, fontweight='bold',
         va='bottom', ha='left')


# =============================================================================
# Panel C: Metabolomic summary — grouped bar chart
# =============================================================================
ax3 = fig.add_axes([0.565, 0.55, 0.32, 0.36])

x_pos = np.arange(len(metabolites))
bars = ax3.bar(x_pos, log2fc_vals, width=0.55, color=metab_colors,
               edgecolor='white', linewidth=0.5, zorder=3)

# Add value labels on bars
for i, (x, v) in enumerate(zip(x_pos, log2fc_vals)):
    va = 'bottom' if v >= 0 else 'top'
    offset = 0.15 if v >= 0 else -0.15
    ax3.text(x, v + offset, f'{v:+.2f}', ha='center', va=va, fontsize=7.5,
             color=metab_colors[i], fontweight='bold')

# Zero line
ax3.axhline(y=0, color='#424242', linewidth=0.5, linestyle='-', zorder=2)

# Axis labels
ax3.set_ylabel('log₂FC (W5 vs CK)', fontsize=9, labelpad=3)
ax3.set_xticks(x_pos)
ax3.set_xticklabels(metabolites, fontsize=8)
ax3.tick_params(axis='y', labelsize=7.5)
ax3.set_ylim(-4.0, 2.0)

# Title / annotation
ax3.text(0.5, 0.96, 'Endosperm metabolome (W5 vs CK)',
         transform=ax3.transAxes, fontsize=9, fontweight='bold',
         ha='center', va='bottom')

# Add color legend
legend_c_elements = [
    mpatches.Patch(facecolor='#D32F2F', edgecolor='none', label='Depleted (W5↓)'),
    mpatches.Patch(facecolor='#1976D2', edgecolor='none', label='Enriched (W5↑)'),
]
ax3.legend(handles=legend_c_elements, loc='lower left', fontsize=7,
           frameon=True, facecolor='white', edgecolor='#BDBDBD',
           framealpha=0.9, borderpad=0.4, handletextpad=0.6)

# Panel label
ax3.text(-0.18, 1.03, 'C', transform=ax3.transAxes, fontsize=14, fontweight='bold',
         va='bottom', ha='left')


# =============================================================================
# Panel D: Convergence diagram — 4 data sources → ABC1K7
# =============================================================================
ax4 = fig.add_axes([0.565, 0.04, 0.40, 0.38])
ax4.set_xlim(0, 10)
ax4.set_ylim(0, 10)
ax4.axis('off')  # No axes for this diagram

# --- Data sources (boxes positioned in a semi-circle around central target) ---
# Layout:
#   GWAS (top-left)         Transcriptomics (top-right)
#   PopGen (bottom-left)    Metabolomics (bottom-right)
#               ABC1K7 (center)

box_style = dict(boxstyle='round,pad=0.4', facecolor='white', alpha=0.95)
arrow_kw = dict(arrowstyle='->', lw=1.8, connectionstyle='arc3,rad=0.15')

# ABC1K7 center target
target_x, target_y = 5.0, 5.0
ax4.scatter([target_x], [target_y], c='#D32F2F', s=300, zorder=20,
            edgecolors='#B71C1C', linewidth=1.5, marker='D')
ax4.scatter([target_x], [target_y], c='white', s=150, zorder=21,
            edgecolors='none', marker='D')
ax4.text(target_x, target_y, 'ABC1K7', fontsize=9, fontweight='bold',
         ha='center', va='center', color='#D32F2F', zorder=22)

# --- Define source boxes ---
sources = [
    {
        'name': 'GWAS',
        'label': 'Genome-wide\nassociation',
        'stat': f'−log₁₀P = {ABC1K7_LOGP:.1f}\nChr5:20,259,172 bp',
        'color': GWAS_COLOR,
        'pos': (2.0, 8.2),
    },
    {
        'name': 'Population\nGenetics',
        'label': 'Selective\nsweep',
        'stat': 'π ratio peak = 28.7\nFst mean = 0.35\n43.3% SNPs > Fst 0.5',
        'color': POPGEN_COLOR,
        'pos': (1.5, 1.8),
    },
    {
        'name': 'Metabolomics',
        'label': 'Flavonoid\ndepletion',
        'stat': 'Flavonoids: −1.42\nAmino acids: +1.30\nTannins: −3.08',
        'color': METAB_COLOR,
        'pos': (8.2, 1.8),
    },
    {
        'name': 'Transcript-\nomics',
        'label': 'Expression\ndecanalization',
        'stat': 'Flavonoid CV\nW5 = 0.485\nCK = 0.172',
        'color': TRANSCR_COLOR,
        'pos': (8.2, 8.2),
    },
]

for src in sources:
    px, py = src['pos']
    color = src['color']

    # Draw box
    bbox_props = dict(boxstyle='round,pad=0.35', facecolor='white',
                      edgecolor=color, linewidth=1.2, alpha=0.95)
    ax4.text(px, py + 0.35, src['name'], fontsize=8, fontweight='bold',
             ha='center', va='bottom', color=color,
             bbox=bbox_props, zorder=15)

    # Sub-label
    ax4.text(px, py - 0.3, src['label'], fontsize=6.5, ha='center', va='top',
             color='#616161', fontstyle='italic')

    # Stat annotation
    ax4.text(px, py - 0.7, src['stat'], fontsize=6, ha='center', va='top',
             color='#424242')

    # Arrow from source to ABC1K7
    # Calculate direction
    dx = target_x - px
    dy = target_y - py
    length = np.sqrt(dx**2 + dy**2)
    # Shorten arrow to start from box edge
    start_x = px + dx / length * 0.65
    start_y = py + dy / length * 0.65
    end_x = target_x - dx / length * 0.20
    end_y = target_y - dy / length * 0.20

    ax4.annotate('', xy=(end_x, end_y), xytext=(start_x, start_y),
                 arrowprops=dict(arrowstyle='->', color=color, lw=1.8,
                                 connectionstyle='arc3,rad=0.0'),
                 zorder=10)

# Add a title
ax4.text(5.0, 9.6, 'Convergence of multi-omics evidence',
         fontsize=10, fontweight='bold', ha='center', va='center',
         color='#212121')

# Add a bottom annotation
ax4.text(5.0, 0.3, 'Four independent lines of evidence converge on ABC1K7\n'
         'as the causal gene underlying the W5 dwarf phenotype',
         fontsize=7, ha='center', va='bottom', color='#616161',
         fontstyle='italic')

# Panel label
ax4.text(-0.22, 1.03, 'D', transform=ax4.transAxes, fontsize=14, fontweight='bold',
         va='bottom', ha='left')


# =============================================================================
# 5. Overall figure title and metadata
# =============================================================================
fig.text(0.5, 0.975, 'Multi-omics convergence on ABC1K7',
         ha='center', va='center', fontsize=12, fontweight='bold')

# =============================================================================
# 6. Save
# =============================================================================
png_path = os.path.join(OUT_DIR, 'Fig3.png')
tiff_path = os.path.join(OUT_DIR, 'Fig3.tiff')

fig.savefig(png_path, dpi=200, bbox_inches='tight', facecolor='white',
            edgecolor='none')
fig.savefig(tiff_path, dpi=300, bbox_inches='tight', facecolor='white',
            edgecolor='none', pil_kwargs={'compression': 'tiff_lzw'})

print(f"✓ Fig3.png saved (200 DPI)")
print(f"✓ Fig3.tiff saved (300 DPI)")
print(f"  → {png_path}")
print(f"  → {tiff_path}")
plt.close(fig)
