#!/usr/bin/env python3
"""
Step 2: Data Acquisition & QC for GSE244574
Download gene expression data, process metadata, map probes to genes, and perform QC.
"""

import os
import sys
import json
import time
import numpy as np
import pandas as pd
import matplotlib
matplotlib.use('Agg')  # Non-interactive backend
import matplotlib.pyplot as plt
import seaborn as sns
from sklearn.decomposition import PCA
from sklearn.preprocessing import StandardScaler
import warnings
warnings.filterwarnings('ignore')

# Set random seed for reproducibility
np.random.seed(42)

# Paths
BASE_DIR = "/app/sandbox/session_20251217_123457_77eda5efa279"
WORKFLOW_DIR = f"{BASE_DIR}/workflow"
DATA_DIR = f"{BASE_DIR}/data"
FIGURES_DIR = f"{BASE_DIR}/figures"

# Ensure directories exist
os.makedirs(DATA_DIR, exist_ok=True)
os.makedirs(FIGURES_DIR, exist_ok=True)

print("=" * 80)
print("STEP 2: DATA ACQUISITION & QC FOR GSE244574")
print("=" * 80)

# Load dataset info from Step 1
print("\n[1/8] Loading dataset information...")
with open(f"{WORKFLOW_DIR}/dataset_info.json", 'r') as f:
    dataset_info = json.load(f)

gse_id = dataset_info['gse_id']
platform_id = dataset_info['platform']
print(f"  → GSE ID: {gse_id}")
print(f"  → Platform: GPL{platform_id}")
print(f"  → Samples: {dataset_info['n_samples']}")

# Step 1: Download Series Matrix File
print(f"\n[2/8] Downloading Series Matrix File for {gse_id}...")
print("  → Attempting download from NCBI FTP site...")

import urllib.request
import gzip
import io

series_matrix_url = f"https://ftp.ncbi.nlm.nih.gov/geo/series/GSE244nnn/{gse_id}/matrix/{gse_id}_series_matrix.txt.gz"
series_matrix_file = f"{DATA_DIR}/{gse_id}_series_matrix.txt"

try:
    # Download compressed file
    print(f"  → Downloading from: {series_matrix_url}")
    response = urllib.request.urlopen(series_matrix_url, timeout=120)
    compressed_data = response.read()
    print(f"  → Downloaded {len(compressed_data)} bytes")

    # Decompress
    print("  → Decompressing...")
    decompressed_data = gzip.decompress(compressed_data)

    # Save decompressed file
    with open(series_matrix_file, 'wb') as f:
        f.write(decompressed_data)

    print(f"  → Saved to: {series_matrix_file}")
    print(f"  → File size: {len(decompressed_data)} bytes")

except Exception as e:
    print(f"  ✗ Error downloading data: {e}")
    sys.exit(1)

# Step 2: Parse Series Matrix File
print(f"\n[3/8] Parsing Series Matrix File...")

def parse_series_matrix(filepath):
    """Parse GEO Series Matrix file to extract metadata and expression data."""
    metadata = {}
    data_start_line = None

    with open(filepath, 'r') as f:
        lines = f.readlines()

    # Extract metadata
    print("  → Extracting metadata...")
    for i, line in enumerate(lines):
        if line.startswith('!Sample_'):
            key = line.split('\t')[0].replace('!Sample_', '').strip('"')
            values = [v.strip().strip('"') for v in line.split('\t')[1:]]
            metadata[key] = values
        elif line.startswith('!series_matrix_table_begin'):
            data_start_line = i + 1
            break

    # Find data end
    data_end_line = None
    for i in range(data_start_line, len(lines)):
        if lines[i].startswith('!series_matrix_table_end'):
            data_end_line = i
            break

    # Extract expression data
    print("  → Extracting expression matrix...")
    data_lines = lines[data_start_line:data_end_line]

    # Parse header (sample IDs)
    header = data_lines[0].strip().split('\t')
    sample_ids = header[1:]  # Skip first column (ID_REF)

    # Parse expression values
    probe_ids = []
    expression_values = []

    for i, line in enumerate(data_lines[1:]):
        if i % 5000 == 0:
            print(f"    Processing row {i}/{len(data_lines)-1}...")
        parts = line.strip().split('\t')
        probe_ids.append(parts[0])
        expression_values.append([float(x) if x != 'null' else np.nan for x in parts[1:]])

    # Create expression dataframe
    expr_df = pd.DataFrame(expression_values, index=probe_ids, columns=sample_ids)

    print(f"  → Expression matrix shape: {expr_df.shape}")
    print(f"  → Probes: {len(probe_ids)}, Samples: {len(sample_ids)}")

    return metadata, expr_df

metadata, expr_df = parse_series_matrix(series_matrix_file)

# Step 3: Process Metadata
print(f"\n[4/8] Processing sample metadata...")

# Extract relevant metadata fields
sample_ids = metadata.get('geo_accession', [])
titles = metadata.get('title', [])
sources = metadata.get('source_name_ch1', [])
characteristics = metadata.get('characteristics_ch1', [])

print(f"  → Found {len(sample_ids)} samples")

# Parse treatment information
conditions = []
for i, sample_id in enumerate(sample_ids):
    # Look for treatment information in characteristics or title
    title = titles[i] if i < len(titles) else ""
    char = characteristics[i] if i < len(characteristics) else ""

    # Determine condition
    if 'control' in title.lower() or 'control' in char.lower() or 'untreated' in title.lower():
        condition = 'Control'
    elif 'doxorubicin' in title.lower() or 'dox' in title.lower() or 'doxorubicin' in char.lower():
        condition = 'Treatment'
    else:
        # Default heuristic: first half control, second half treatment
        condition = 'Control' if i < len(sample_ids) // 2 else 'Treatment'

    conditions.append(condition)
    print(f"    {sample_id}: {condition} (from: {title})")

# Create metadata dataframe
metadata_df = pd.DataFrame({
    'sample_id': sample_ids,
    'condition': conditions,
    'title': titles[:len(sample_ids)]
})

print(f"\n  → Control samples: {sum(np.array(conditions) == 'Control')}")
print(f"  → Treatment samples: {sum(np.array(conditions) == 'Treatment')}")

# Save metadata
metadata_output = f"{WORKFLOW_DIR}/sample_metadata.csv"
metadata_df.to_csv(metadata_output, index=False)
print(f"  → Saved metadata to: {metadata_output}")

# Step 4: Map Probes to Gene Symbols
print(f"\n[5/8] Mapping probes to Gene Symbols (GPL{platform_id})...")

# Download platform annotation
platform_url = f"https://ftp.ncbi.nlm.nih.gov/geo/platforms/GPL21nnn/GPL{platform_id}/annot/GPL{platform_id}.annot.gz"
platform_file = f"{DATA_DIR}/GPL{platform_id}.annot"

try:
    print(f"  → Downloading platform annotation from: {platform_url}")
    response = urllib.request.urlopen(platform_url, timeout=120)
    compressed_data = response.read()
    decompressed_data = gzip.decompress(compressed_data)

    with open(platform_file, 'wb') as f:
        f.write(decompressed_data)

    print(f"  → Downloaded and saved platform annotation")

    # Parse platform file
    print("  → Parsing platform annotation...")
    with open(platform_file, 'r', encoding='latin-1') as f:
        lines = f.readlines()

    # Find header line
    header_idx = None
    for i, line in enumerate(lines):
        if line.startswith('ID\t'):
            header_idx = i
            break

    if header_idx is None:
        raise ValueError("Could not find header in platform file")

    # Parse annotation data
    annot_lines = lines[header_idx:]
    annot_df = pd.read_csv(io.StringIO(''.join(annot_lines)), sep='\t', comment='#', low_memory=False)

    print(f"  → Platform annotation shape: {annot_df.shape}")
    print(f"  → Columns: {list(annot_df.columns[:10])}")

    # Look for gene symbol column
    gene_col = None
    for col in ['Gene Symbol', 'GENE_SYMBOL', 'Gene symbol', 'Symbol', 'gene_assignment']:
        if col in annot_df.columns:
            gene_col = col
            break

    if gene_col is None:
        print("  ! Warning: Could not find gene symbol column, using probe IDs")
        expr_df_genes = expr_df.copy()
        expr_df_genes.index.name = 'Gene_Symbol'
    else:
        print(f"  → Using column '{gene_col}' for gene symbols")

        # Create mapping dictionary
        probe_to_gene = dict(zip(annot_df['ID'], annot_df[gene_col]))

        # Map probes to genes
        print("  → Mapping probes to genes...")
        gene_symbols = []
        for probe in expr_df.index:
            gene = probe_to_gene.get(probe, probe)
            # Clean gene symbol
            if pd.isna(gene) or gene == '' or gene == '---':
                gene = probe  # Keep probe ID if no gene symbol
            else:
                # Take first gene if multiple (separated by ///)
                gene = str(gene).split('///')[0].strip()
            gene_symbols.append(gene)

        expr_df['Gene_Symbol'] = gene_symbols

        # Handle duplicates by averaging
        print("  → Handling duplicate gene symbols (averaging)...")
        expr_df_genes = expr_df.groupby('Gene_Symbol').mean()

        print(f"  → After mapping: {expr_df_genes.shape[0]} unique genes")
        print(f"  → Reduced from {expr_df.shape[0]} probes to {expr_df_genes.shape[0]} genes")

except Exception as e:
    print(f"  ! Warning: Could not download platform annotation: {e}")
    print("  → Using probe IDs as identifiers")
    expr_df_genes = expr_df.copy()
    expr_df_genes.index.name = 'Gene_Symbol'

# Step 5: Log Transformation Check
print(f"\n[6/8] Checking log transformation status...")

# Check value range
max_val = expr_df_genes.max().max()
median_val = expr_df_genes.median().median()
min_val = expr_df_genes.min().min()

print(f"  → Value range: [{min_val:.2f}, {max_val:.2f}]")
print(f"  → Median value: {median_val:.2f}")

if max_val > 100:
    print("  → Data appears to be non-log transformed (values > 100)")
    print("  → Applying log2(x+1) transformation...")
    expr_df_genes = np.log2(expr_df_genes + 1)
    print(f"  → After transformation: [{expr_df_genes.min().min():.2f}, {expr_df_genes.max().max():.2f}]")
else:
    print("  → Data appears to be already log-transformed (values typically < 20)")

# Save expression matrix
expr_output = f"{WORKFLOW_DIR}/expression_matrix.csv"
expr_df_genes.to_csv(expr_output)
print(f"\n  → Saved expression matrix to: {expr_output}")
print(f"  → Shape: {expr_df_genes.shape}")

# Step 6: Quality Control - Boxplot
print(f"\n[7/8] Generating QC plots...")
print("  → Creating boxplot of expression distributions...")

plt.figure(figsize=(12, 6))
expr_df_genes.boxplot(rot=90, figsize=(14, 6))
plt.title(f'{gse_id}: Expression Value Distribution Across Samples', fontsize=14, weight='bold')
plt.xlabel('Sample ID', fontsize=12)
plt.ylabel('Log2 Expression', fontsize=12)
plt.xticks(fontsize=8)
plt.tight_layout()

boxplot_file = f"{FIGURES_DIR}/qc_boxplot.png"
plt.savefig(boxplot_file, dpi=300, bbox_inches='tight')
plt.close()
print(f"  → Saved boxplot to: {boxplot_file}")

# Step 7: Quality Control - PCA
print("  → Performing PCA analysis...")

# Remove genes with any missing values for PCA
expr_clean = expr_df_genes.dropna()
print(f"    Removed {expr_df_genes.shape[0] - expr_clean.shape[0]} genes with missing values")
print(f"    Using {expr_clean.shape[0]} genes for PCA")

# Transpose: samples as rows, genes as columns
expr_transposed = expr_clean.T

# Standardize features
scaler = StandardScaler()
expr_scaled = scaler.fit_transform(expr_transposed)

# Perform PCA
pca = PCA(n_components=2)
pca_result = pca.fit_transform(expr_scaled)

# Create PCA dataframe
pca_df = pd.DataFrame(
    pca_result,
    columns=['PC1', 'PC2'],
    index=expr_transposed.index
)

# Add condition labels - create a mapping dict for safer access
condition_map = dict(zip(metadata_df['sample_id'], metadata_df['condition']))
pca_df['condition'] = [condition_map.get(sid, 'Unknown') for sid in pca_df.index]

# Verify all samples have condition labels
if 'Unknown' in pca_df['condition'].values:
    print(f"    Warning: Some samples could not be matched to conditions")
    print(f"    Expression sample IDs: {list(pca_df.index[:3])}")
    print(f"    Metadata sample IDs: {list(metadata_df['sample_id'][:3])}")

print(f"    PC1 variance explained: {pca.explained_variance_ratio_[0]:.2%}")
print(f"    PC2 variance explained: {pca.explained_variance_ratio_[1]:.2%}")
print(f"    Total variance explained: {pca.explained_variance_ratio_.sum():.2%}")

# Plot PCA
plt.figure(figsize=(10, 8))
colors = {'Control': '#3498db', 'Treatment': '#e74c3c'}

for condition in pca_df['condition'].unique():
    mask = pca_df['condition'] == condition
    plt.scatter(
        pca_df[mask]['PC1'],
        pca_df[mask]['PC2'],
        c=colors.get(condition, 'gray'),
        label=condition,
        s=100,
        alpha=0.7,
        edgecolors='black',
        linewidth=1
    )

plt.xlabel(f'PC1 ({pca.explained_variance_ratio_[0]:.1%} variance)', fontsize=12)
plt.ylabel(f'PC2 ({pca.explained_variance_ratio_[1]:.1%} variance)', fontsize=12)
plt.title(f'{gse_id}: PCA of Gene Expression', fontsize=14, weight='bold')
plt.legend(fontsize=11, frameon=True, shadow=True)
plt.grid(True, alpha=0.3)
plt.tight_layout()

pca_file = f"{FIGURES_DIR}/qc_pca.png"
plt.savefig(pca_file, dpi=300, bbox_inches='tight')
plt.close()
print(f"  → Saved PCA plot to: {pca_file}")

# Final Summary
print("\n" + "=" * 80)
print("STEP 2 COMPLETE - DATA ACQUISITION & QC SUMMARY")
print("=" * 80)
print(f"\n✓ Dataset: {gse_id} (GPL{platform_id})")
print(f"✓ Samples: {expr_df_genes.shape[1]} ({sum(metadata_df['condition'] == 'Control')} Control, {sum(metadata_df['condition'] == 'Treatment')} Treatment)")
print(f"✓ Genes: {expr_df_genes.shape[0]}")
print(f"✓ Expression range: [{expr_df_genes.min().min():.2f}, {expr_df_genes.max().max():.2f}]")
print(f"✓ PCA separation: {pca.explained_variance_ratio_[0]:.1%} variance on PC1")

print("\n📁 Output Files:")
print(f"  • {expr_output}")
print(f"  • {metadata_output}")
print(f"  • {boxplot_file}")
print(f"  • {pca_file}")

print("\n✓ SUCCESS: Data acquisition and QC completed successfully!")
print("=" * 80)
