import argparse
import csv
import pandas as pd
import matplotlib.pyplot as plt

def plot_damage_by_taxid(df, taxid_to_plot):
    # Ensure type matching (convert column to string or int as needed)
    taxon_df = df[df['taxid'].astype(str) == str(taxid_to_plot)]

    if taxon_df.empty:
        print(f"No data found for taxid: {taxid_to_plot}")
        return

    taxon_name = taxon_df['name'].iloc[0] if 'name' in taxon_df.columns else f"TaxID {taxid_to_plot}"

    fw_cols = [col for col in df.columns if col.startswith('fwf')]
    bw_cols = [col for col in df.columns if col.startswith('bwf')]
    
    if not fw_cols or not bw_cols:
        print("Error: Could not find 'fwf' or 'bwf' columns in the metaDMG output file.")
        return

    positions = [int(col[3:]) for col in fw_cols]

    plot_data = pd.DataFrame({
        'Position': positions,
        '5\'_end': taxon_df[fw_cols].iloc[0].values,
        '3\'_end': taxon_df[bw_cols].iloc[0].values
    })
    print(f"--- Data for TaxID {taxid_to_plot} ({taxon_name}) ---")
    print(plot_data)

    fig, ax = plt.subplots(figsize=(10, 6))

    ax.plot(plot_data['Position'], plot_data['5\'_end'], marker='o', label='5\' end', color='blue')
    ax.plot(plot_data['Position'], plot_data['3\'_end'], marker='o', label='3\' end', color='red')

    ax.set_title(f"DNA Damage for {taxon_name} (TaxID: {taxid_to_plot})")
    ax.set_xlabel('Position from read end')
    ax.set_ylabel('Damage Fraction (C -> T or G -> A)')
    ax.grid(True, linestyle='--', alpha=0.6)
    ax.legend()
    ax.set_xticks(range(0, max(positions) + 1, 5))

    plot_filename = f"damage_plot_{taxid_to_plot}.png"
    plt.savefig(plot_filename)
    plt.close(fig)  # Close figure to free memory when looping through many plots
    print(f"Plot saved to {plot_filename}\n")

def main():
    parser = argparse.ArgumentParser(description="Generate metaDMG smile plots for a list of target species.")
    parser.add_argument("-c", "--csv", required=True, help="Path to the CSV file containing target species (taxid in 4th column)")
    parser.add_argument("-s", "--stats", required=True, help="Path to the metaDMG aggregate statistics file (e.g., h1.agg.stat.gz)")
    args = parser.parse_args()

    # 1. Load target taxids from the fourth column (index 3) of the CSV file
    target_taxids = set()
    with open(args.csv, mode='r', encoding='utf-8') as f:
        reader = csv.reader(f)
        for row_idx, row in enumerate(reader):
            if len(row) > 3:
                taxid = row[3].strip()
                # Skip header row if it contains text instead of numbers
                if row_idx == 0 and not taxid.isdigit():
                    continue
                if taxid:
                    target_taxids.add(taxid)

    print(f"Loaded {len(target_taxids)} unique taxids to process from {args.csv}.")

    # 2. Load metaDMG aggregate statistics table from command-line path
    print(f"Loading metaDMG data from {args.stats}...")
    df = pd.read_csv(args.stats, sep='\t')

    # 3. Loop through each target taxid and generate plots
    for taxid in target_taxids:
        plot_damage_by_taxid(df, taxid)

if __name__ == '__main__':
    main()