use mafft_align::{local_align, GapModel};
use mafft_scoring::build_context;
use mafft_tree::parttree_dist::{
common_sextets_p, composition_table, encode_points_dna,
};
use mafft_types::{ScoringModel, Sequence, SeqType, SequenceSet};
const REFERENCE_LIMIT_KMER: usize = 5000;
const REFERENCE_LIMIT_DP: usize = 100;
const TSIZE: usize = 4096;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum AdjustMode {
Kmer,
Dp,
}
fn creverse(c: u8) -> u8 {
match c {
b'A' => b'T', b'C' => b'G', b'G' => b'C', b'T' => b'A', b'U' => b'A',
b'M' => b'K', b'R' => b'Y', b'W' => b'W', b'S' => b'S', b'Y' => b'R',
b'K' => b'M', b'V' => b'B', b'H' => b'D', b'D' => b'H', b'B' => b'V',
b'N' => b'N',
b'a' => b't', b'c' => b'g', b'g' => b'c', b't' => b'a', b'u' => b'a',
b'm' => b'k', b'r' => b'y', b'w' => b'w', b's' => b's', b'y' => b'r',
b'k' => b'm', b'v' => b'b', b'h' => b'd', b'd' => b'h', b'b' => b'v',
b'n' => b'n',
other => other,
}
}
pub fn reverse_complement(seq: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(seq.len());
for &c in seq.iter().rev() {
out.push(creverse(c));
}
let num_t = seq.iter().filter(|&&c| c == b't' || c == b'T').count();
let num_u = seq.iter().filter(|&&c| c == b'u' || c == b'U').count();
if num_u > num_t {
for c in &mut out {
if *c == b't' { *c = b'u'; }
else if *c == b'T' { *c = b'U'; }
}
}
out
}
fn gappick0(seq: &[u8]) -> Vec<u8> {
seq.iter()
.copied()
.filter(|&c| c != b'-' && c != b'.')
.map(|c| c.to_ascii_lowercase())
.collect()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Direction { Forward, Reverse }
pub fn adjust_direction(input: &SequenceSet) -> SequenceSet {
adjust_direction_mode(input, AdjustMode::Kmer)
}
pub fn adjust_direction_mode_add(
input: &SequenceSet,
mode: AdjustMode,
nadd: usize,
) -> SequenceSet {
adjust_direction_with(input, mode, nadd)
}
pub fn adjust_direction_mode(input: &SequenceSet, mode: AdjustMode) -> SequenceSet {
adjust_direction_with(input, mode, 0)
}
fn adjust_direction_with(input: &SequenceSet, mode: AdjustMode, nadd: usize) -> SequenceSet {
if !input.seq_type.is_nucleotide() {
return input.clone();
}
let nseq = input.nseq();
if nseq == 0 {
return input.clone();
}
let forward: Vec<Vec<u8>> = input.sequences.iter()
.map(|s| gappick0(&s.data))
.collect();
let reverse: Vec<Vec<u8>> = forward.iter()
.map(|s| reverse_complement(s))
.collect();
let scoring = if mode == AdjustMode::Dp {
Some(build_context(ScoringModel::Dna, SeqType::Dna))
} else {
None
};
let gap_dp = scoring.as_ref().map(|s| {
GapModel::new(s.gap.open as f64, s.gap.extend as f64)
});
let points_fwd: Vec<Vec<u32>> = forward.iter().map(|s| encode_points_dna(s)).collect();
let points_rev: Vec<Vec<u32>> = reverse.iter().map(|s| encode_points_dna(s)).collect();
let n_anchor = if nadd > 0 && nadd <= nseq { nseq - nadd } else { 0 };
let mut contrast_order: Vec<(usize, f64)> = Vec::with_capacity(nseq);
for i in 0..n_anchor { contrast_order.push((i, 0.0)); }
let mut testable: Vec<(usize, f64)> = match mode {
AdjustMode::Kmer => (n_anchor..nseq).map(|i| {
let p_fwd = &points_fwd[i];
let p_rev = &points_rev[i];
let t_fwd = composition_table(p_fwd, TSIZE);
let t_rev = composition_table(p_rev, TSIZE);
let dif = (common_sextets_p(&t_fwd, p_fwd, TSIZE)
- common_sextets_p(&t_rev, p_fwd, TSIZE)) as f64;
(i, dif)
}).collect(),
AdjustMode::Dp => {
let sc = scoring.as_ref().unwrap();
let gap = gap_dp.as_ref().unwrap();
(n_anchor..nseq).map(|i| {
let fwd_self = local_align(
&forward[i], &forward[i],
&sc.consweight_matrix, &sc.amino_map, gap, 0.0,
).alignment.score;
let rev_self = local_align(
&forward[i], &reverse[i],
&sc.consweight_matrix, &sc.amino_map, gap, 0.0,
).alignment.score;
(i, fwd_self - rev_self)
}).collect()
}
};
testable.sort_unstable_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
contrast_order.extend(testable);
let order: Vec<usize> = contrast_order.iter().map(|(i, _)| *i).collect();
let mut direction = vec![Direction::Forward; nseq];
let mut chosen_points: Vec<&[u32]> = Vec::with_capacity(nseq);
let mut chosen_seqs: Vec<&[u8]> = Vec::with_capacity(nseq);
let pivot_count = n_anchor.max(1);
for step in 0..pivot_count.min(nseq) {
chosen_points.push(&points_fwd[order[step]]);
chosen_seqs.push(&forward[order[step]]);
}
let reflim = match mode { AdjustMode::Kmer => REFERENCE_LIMIT_KMER,
AdjustMode::Dp => REFERENCE_LIMIT_DP };
for step in pivot_count..nseq {
let ic = order[step];
let iend = step.min(reflim);
let (res_forward, res_reverse) = match mode {
AdjustMode::Kmer => {
let table_fwd = composition_table(&points_fwd[ic], TSIZE);
let table_rev = composition_table(&points_rev[ic], TSIZE);
let mut sum_f: f64 = 0.0;
let mut sum_r: f64 = 0.0;
for j in 0..iend {
let ref_points = chosen_points[j];
sum_f += common_sextets_p(&table_fwd, ref_points, TSIZE) as f64;
sum_r += common_sextets_p(&table_rev, ref_points, TSIZE) as f64;
}
(sum_f / iend as f64, sum_r / iend as f64)
}
AdjustMode::Dp => {
let sc = scoring.as_ref().unwrap();
let gap = gap_dp.as_ref().unwrap();
let mut sum_f: f64 = 0.0;
let mut sum_r: f64 = 0.0;
for j in 0..iend {
let r = chosen_seqs[j];
sum_f += local_align(
&forward[ic], r,
&sc.consweight_matrix, &sc.amino_map, gap, 0.0,
).alignment.score;
sum_r += local_align(
&reverse[ic], r,
&sc.consweight_matrix, &sc.amino_map, gap, 0.0,
).alignment.score;
}
(sum_f / iend as f64, sum_r / iend as f64)
}
};
if res_reverse > res_forward {
direction[ic] = Direction::Reverse;
chosen_points.push(&points_rev[ic]);
chosen_seqs.push(&reverse[ic]);
} else {
direction[ic] = Direction::Forward;
chosen_points.push(&points_fwd[ic]);
chosen_seqs.push(&forward[ic]);
}
}
if direction[0] == Direction::Reverse {
for d in direction.iter_mut() {
*d = match *d { Direction::Forward => Direction::Reverse, Direction::Reverse => Direction::Forward };
}
}
let mut adjusted = input.clone();
for (i, dir) in direction.iter().enumerate() {
if *dir == Direction::Reverse {
let rc = reverse_complement(&input.sequences[i].data);
adjusted.sequences[i] = Sequence {
name: format!("_R_{}", input.sequences[i].name),
data: rc,
};
}
}
adjusted
}
#[cfg(test)]
mod tests {
use super::*;
use mafft_types::{Sequence, SeqType};
fn dna_set(seqs: &[(&str, &str)]) -> SequenceSet {
SequenceSet {
sequences: seqs.iter().map(|(n, s)| Sequence {
name: (*n).to_string(),
data: s.as_bytes().to_vec(),
}).collect(),
seq_type: SeqType::Dna,
}
}
#[test]
fn reverse_complement_basic() {
assert_eq!(reverse_complement(b"ACGT"), b"ACGT");
assert_eq!(reverse_complement(b"AAAA"), b"TTTT");
assert_eq!(reverse_complement(b"AcGt"), b"aCgT");
}
#[test]
fn rna_t_to_u_when_majority_u() {
assert_eq!(reverse_complement(b"uuuu"), b"aaaa");
assert_eq!(reverse_complement(b"auug"), b"caau");
}
#[test]
fn protein_passthrough() {
let set = SequenceSet {
sequences: vec![Sequence { name: "p1".into(), data: b"MKLVN".to_vec() }],
seq_type: SeqType::Protein,
};
let adjusted = adjust_direction(&set);
assert_eq!(adjusted.sequences[0].data, b"MKLVN");
assert_eq!(adjusted.sequences[0].name, "p1");
}
#[test]
fn detects_reversed_sequence() {
let fwd = "atggcaattcgcatggcaattcgcatggcaattcgc";
let rc: String = fwd.chars().rev().map(|c| match c {
'a' => 't', 'c' => 'g', 'g' => 'c', 't' => 'a',
_ => c,
}).collect();
let set = dna_set(&[
("s1", fwd),
("s2", fwd),
("s3_rc", &rc),
]);
let adjusted = adjust_direction(&set);
assert!(adjusted.sequences[2].name.starts_with("_R_"),
"expected _R_ prefix, got {}", adjusted.sequences[2].name);
assert_eq!(adjusted.sequences[2].data, fwd.as_bytes(),
"reversed sequence should now equal forward");
assert_eq!(adjusted.sequences[0].name, "s1");
assert_eq!(adjusted.sequences[1].name, "s2");
}
}