Documentation
use crate::{
    ab1::ab1_file::{
        self, build_cmnt_1, build_fwo_1, build_lane, build_pbas_1, build_pdmf_1, build_pdmf_2,
        build_ploc1, build_smpl_1, build_snr,
    },
    pileup_counter::{BASE2IDX, PlpInfo},
};
use ndarray::{Array1, s};
use std::f32::consts::PI;

pub const MAX_PULSE: f32 = 2400.;

/// peak width . for example 10
/// 1 2 3 4 5 p 6 7 8 9 10 11  12 13 14 15 16 p
/// mu means timestep
pub fn get_mu_list(length: usize, peak_width: usize) -> Vec<usize> {
    let mut mu_list = vec![];

    for idx in 0..length {
        if idx == 0 {
            mu_list.push(peak_width / 2);
        } else {
            let cur_mu = *mu_list.last().unwrap() + peak_width + 1;
            mu_list.push(cur_mu);
        }
    }

    mu_list
}

pub fn get_mu_list_with_deletion_shrink(
    length: usize,
    major: &Vec<usize>,
    minor: &Vec<usize>,
    peak_width: usize,
) -> Vec<usize> {
    let mut major_base_cnts = vec![];
    major.iter().zip(minor.iter()).for_each(|(&maj, &mio)| {
        if major_base_cnts.is_empty() {
            major_base_cnts.push((maj, mio + 1));
        } else {
            if major_base_cnts.last().unwrap().0 == maj {
                major_base_cnts.last_mut().unwrap().1 = mio + 1;
            } else {
                major_base_cnts.push((maj, mio + 1));
            }
        }
    });

    let mut peak_position_cursor = peak_width / 2;

    let mut mu_list = vec![];

    major_base_cnts.into_iter().for_each(|(_, cnt)| {
        mu_list.push(peak_position_cursor);

        if cnt > 1 {
            let inner_interval = peak_width / (cnt - 1 + 1);
            for idx in 1..cnt {
                mu_list.push(peak_position_cursor + idx * inner_interval);
            }
        }

        peak_position_cursor += peak_width + 1;
    });

    assert_eq!(mu_list.len(), length);
    mu_list
}

pub fn get_sigma_list(mu_list: &Vec<usize>, peak_width: usize) -> Vec<f32> {
    let width = peak_width as f32 / 2.;

    let k = 0.00005;
    let c = width / 4.0;
    let thr = width / 3.0;
    mu_list
        .iter()
        .map(|&mu| k * (mu as f32) + c)
        .map(|v| v.min(thr))
        .collect()
}

pub struct Gaussian {
    mu: f32,
    sigma: f32,
    pulse: f32,

    three_sigma_law_low: f32,
    three_sigma_law_high: f32,
}

impl Gaussian {
    pub fn new(mu: f32, sigma: f32, pulse: f32) -> Self {
        let three_sigma_law_low = mu - 3.0 * sigma;
        let three_sigma_law_high = mu + 3.0 * sigma;

        Self {
            mu,
            sigma,
            pulse,
            three_sigma_law_low,
            three_sigma_law_high,
        }
    }

    pub fn compute(&self, x: f32) -> f32 {
        if x <= self.three_sigma_law_low || x >= self.three_sigma_law_high {
            return 0.0;
        }
        let v1 = self.pulse / ((2.0 * PI).sqrt() * self.sigma);
        let v2 = (-(x - self.mu).powi(2) / (2.0 * self.sigma.powi(2))).exp();
        v1 * v2
    }
}

pub fn transform_plp_info_2_ab1_data(
    plp_info: &PlpInfo,
    target_seq: &str,
    peak_width: Option<usize>,
    sample_name: Option<String>,
) -> ab1_file::AbiFile {
    let peak_width = peak_width.unwrap_or(20);
    let mu_list = get_mu_list(plp_info.major.len(), peak_width);
    let sigma_list = get_sigma_list(&mu_list, peak_width);

    let last_mu = (*mu_list.last().unwrap()) as f32;
    let last_sigma = *sigma_list.last().unwrap();
    let last_tt = (last_mu + 3.0 * last_sigma).ceil() as usize;

    let g_pulse = plp_info
        .normed_count
        .slice(s![BASE2IDX['G' as usize] as usize - 1, ..])
        .mapv(|v| v * MAX_PULSE);

    let a_pulse = plp_info
        .normed_count
        .slice(s![BASE2IDX['A' as usize] as usize - 1, ..])
        .mapv(|v| v * MAX_PULSE);

    let t_pulse = plp_info
        .normed_count
        .slice(s![BASE2IDX['T' as usize] as usize - 1, ..])
        .mapv(|v| v * MAX_PULSE);

    let c_pulse = plp_info
        .normed_count
        .slice(s![BASE2IDX['C' as usize] as usize - 1, ..])
        .mapv(|v| v * MAX_PULSE);

    let g_data = build_data(g_pulse, &mu_list, &sigma_list, last_tt);
    let a_data = build_data(a_pulse, &mu_list, &sigma_list, last_tt);
    let t_data = build_data(t_pulse, &mu_list, &sigma_list, last_tt);
    let c_data = build_data(c_pulse, &mu_list, &sigma_list, last_tt);

    let peek_loc = mu_list
        .iter()
        .zip(plp_info.minor.iter())
        .filter(|(_, mino)| **mino == 0)
        .map(|(loc, _)| loc)
        .copied()
        .map(|v| v as u16)
        .collect::<Vec<_>>();

    let ploc1 = build_ploc1(peek_loc);
    let fwo_1 = build_fwo_1();
    let lane = build_lane();
    let snr = build_snr();
    let sample_name = build_smpl_1(sample_name.unwrap_or(format!("sample")));
    let bases = build_pbas_1(target_seq.to_string());
    let cmnt1 = build_cmnt_1(format!("Generated by bam2ab1"));
    let pdmf_1 = build_pdmf_1(format!("none"));
    let pdmf_2 = build_pdmf_2(format!("none"));
    let data9 = ab1_file::build_data(9, g_data);
    let data10 = ab1_file::build_data(10, a_data);
    let data11 = ab1_file::build_data(11, t_data);
    let data12 = ab1_file::build_data(12, c_data);

    let mut ab1_file = ab1_file::AbiFile::new(101);
    ab1_file.push_tagged_data(ploc1);
    ab1_file.push_tagged_data(fwo_1);
    ab1_file.push_tagged_data(lane);
    ab1_file.push_tagged_data(snr);
    ab1_file.push_tagged_data(sample_name);
    ab1_file.push_tagged_data(bases);
    ab1_file.push_tagged_data(cmnt1);
    ab1_file.push_tagged_data(pdmf_1);
    ab1_file.push_tagged_data(pdmf_2);
    ab1_file.push_tagged_data(data9);
    ab1_file.push_tagged_data(data10);
    ab1_file.push_tagged_data(data11);
    ab1_file.push_tagged_data(data12);

    ab1_file
}

pub fn transform_plp_info_2_ab1_data_with_deletion_shrink(
    plp_info: &PlpInfo,
    target_seq: &str,
    peak_width: Option<usize>,
    sample_name: Option<String>,
) -> ab1_file::AbiFile {
    let peak_width = peak_width.unwrap_or(20);
    let mu_list = get_mu_list_with_deletion_shrink(
        plp_info.major.len(),
        &plp_info.major,
        &plp_info.minor,
        peak_width,
    );
    let sigma_list = get_sigma_list(&mu_list, peak_width);

    let last_mu = (*mu_list.last().unwrap()) as f32;
    let last_sigma = *sigma_list.last().unwrap();
    let last_tt = (last_mu + 3.0 * last_sigma).ceil() as usize;

    let g_pulse = plp_info
        .normed_count
        .slice(s![BASE2IDX['G' as usize] as usize - 1, ..])
        .mapv(|v| v * MAX_PULSE);

    let a_pulse = plp_info
        .normed_count
        .slice(s![BASE2IDX['A' as usize] as usize - 1, ..])
        .mapv(|v| v * MAX_PULSE);

    let t_pulse = plp_info
        .normed_count
        .slice(s![BASE2IDX['T' as usize] as usize - 1, ..])
        .mapv(|v| v * MAX_PULSE);

    let c_pulse = plp_info
        .normed_count
        .slice(s![BASE2IDX['C' as usize] as usize - 1, ..])
        .mapv(|v| v * MAX_PULSE);

    let g_data = build_data(g_pulse, &mu_list, &sigma_list, last_tt);
    let a_data = build_data(a_pulse, &mu_list, &sigma_list, last_tt);
    let t_data = build_data(t_pulse, &mu_list, &sigma_list, last_tt);
    let c_data = build_data(c_pulse, &mu_list, &sigma_list, last_tt);

    let peek_loc = mu_list
        .iter()
        .zip(plp_info.minor.iter())
        .filter(|(_, mino)| **mino == 0)
        .map(|(loc, _)| loc)
        .copied()
        .map(|v| v as u16)
        .collect::<Vec<_>>();

    let ploc1 = build_ploc1(peek_loc);
    let fwo_1 = build_fwo_1();
    let lane = build_lane();
    let snr = build_snr();
    let sample_name = build_smpl_1(sample_name.unwrap_or(format!("sample")));
    let bases = build_pbas_1(target_seq.to_string());
    let cmnt1 = build_cmnt_1(format!("Generated by bam2ab1"));
    let pdmf_1 = build_pdmf_1(format!("none"));
    let pdmf_2 = build_pdmf_2(format!("none"));
    let data9 = ab1_file::build_data(9, g_data);
    let data10 = ab1_file::build_data(10, a_data);
    let data11 = ab1_file::build_data(11, t_data);
    let data12 = ab1_file::build_data(12, c_data);
    let pcon1 = ab1_file::build_pcon1(vec![40; target_seq.len()]);

    let mut ab1_file = ab1_file::AbiFile::new(101);
    ab1_file.push_tagged_data(ploc1);
    ab1_file.push_tagged_data(fwo_1);
    ab1_file.push_tagged_data(lane);
    ab1_file.push_tagged_data(snr);
    ab1_file.push_tagged_data(sample_name);
    ab1_file.push_tagged_data(bases);
    ab1_file.push_tagged_data(cmnt1);
    ab1_file.push_tagged_data(pdmf_1);
    ab1_file.push_tagged_data(pdmf_2);
    ab1_file.push_tagged_data(data9);
    ab1_file.push_tagged_data(data10);
    ab1_file.push_tagged_data(data11);
    ab1_file.push_tagged_data(data12);
    // ab1_file.push_tagged_data(pcon1);

    ab1_file
}

pub fn build_data(
    peak_pulse: Array1<f32>,
    mu_list: &Vec<usize>,
    sigma_list: &Vec<f32>,
    last_tt: usize,
) -> Vec<u16> {
    let gaussians = peak_pulse
        .iter()
        .zip(mu_list.iter())
        .zip(sigma_list.iter())
        .map(|((&pulse, &mu), &sigma)| Gaussian::new(mu as f32, sigma, pulse))
        .collect::<Vec<_>>();

    (1..last_tt)
        .into_iter()
        .map(|tt| {
            gaussians
                .iter()
                .map(|gauss| {
                    let left = gauss.compute((tt - 1) as f32);
                    let cur = gauss.compute(tt as f32);
                    // let right = gauss.compute((tt + 1).min(last_tt - 1) as f32);

                    // (left + cur + right) / 3.0
                    (left + cur) / 2.0
                })
                .sum::<f32>()
                .round() as u16
        })
        .collect()
}