use crate::error::PoaError;
use std::collections::{HashMap, HashSet};
#[derive(Debug, Clone)]
pub enum SeedSelection {
Auto,
Explicit(usize),
Shortest,
}
const TERM_K: usize = 15;
const TERM_LEN: usize = 50;
const ANCHOR_FRAC: f64 = 0.3;
pub fn select_seed(reads: &[&[u8]], selection: &SeedSelection) -> Result<usize, PoaError> {
if reads.is_empty() {
return Err(PoaError::EmptyInput);
}
match selection {
SeedSelection::Explicit(idx) => {
if *idx >= reads.len() {
Err(PoaError::SeedOutOfBounds {
index: *idx,
len: reads.len(),
})
} else {
Ok(*idx)
}
}
SeedSelection::Shortest => Ok(shortest(reads)),
SeedSelection::Auto => auto_select(reads),
}
}
fn shortest(reads: &[&[u8]]) -> usize {
reads
.iter()
.enumerate()
.min_by_key(|(_, r)| r.len())
.map(|(i, _)| i)
.unwrap()
}
fn longest(reads: &[&[u8]]) -> usize {
reads
.iter()
.enumerate()
.max_by_key(|(_, r)| r.len())
.map(|(i, _)| i)
.unwrap()
}
fn auto_select(reads: &[&[u8]]) -> Result<usize, PoaError> {
let n = reads.len();
let threshold = ((n as f64 * ANCHOR_FRAC) as usize).max(2).min(n);
let left_sets: Vec<HashSet<u64>> = reads.iter().map(|r| terminal_kmers(r, true)).collect();
let right_sets: Vec<HashSet<u64>> = reads.iter().map(|r| terminal_kmers(r, false)).collect();
let mut left_freq: HashMap<u64, usize> = HashMap::new();
let mut right_freq: HashMap<u64, usize> = HashMap::new();
for s in &left_sets {
for &h in s {
*left_freq.entry(h).or_insert(0) += 1;
}
}
for s in &right_sets {
for &h in s {
*right_freq.entry(h).or_insert(0) += 1;
}
}
let is_left_anchor = |h: &u64| -> bool {
*left_freq.get(h).unwrap_or(&0) >= threshold && *right_freq.get(h).unwrap_or(&0) < threshold
};
let is_right_anchor = |h: &u64| -> bool {
*right_freq.get(h).unwrap_or(&0) >= threshold && *left_freq.get(h).unwrap_or(&0) < threshold
};
let left_ok: Vec<bool> = left_sets
.iter()
.map(|s| s.iter().any(is_left_anchor))
.collect();
let right_ok: Vec<bool> = right_sets
.iter()
.map(|s| s.iter().any(is_right_anchor))
.collect();
let candidates: Vec<usize> = (0..n).filter(|&i| left_ok[i] && right_ok[i]).collect();
if !candidates.is_empty() {
return Ok(*candidates.iter().min_by_key(|&&i| reads[i].len()).unwrap());
}
let left_only = (0..n).filter(|&i| left_ok[i] && !right_ok[i]).count();
let right_only = (0..n).filter(|&i| !left_ok[i] && right_ok[i]).count();
if left_only > 0 && right_only > 0 {
return Err(PoaError::NoSpanningReads {
left_depth: left_only,
right_depth: right_only,
});
}
Ok(longest(reads))
}
fn terminal_kmers(read: &[u8], left: bool) -> HashSet<u64> {
if read.len() < TERM_K {
return HashSet::new();
}
let term_len = TERM_LEN.min(read.len());
let slice = if left {
&read[..term_len]
} else {
&read[read.len() - term_len..]
};
kmers_of(slice)
}
fn kmers_of(seq: &[u8]) -> HashSet<u64> {
let mut out = HashSet::new();
if seq.len() < TERM_K {
return out;
}
let mask = (1u64 << (2 * TERM_K)) - 1;
let mut h: u64 = 0;
let mut valid: usize = 0;
for &b in seq {
let bits = match b {
b'A' | b'a' => {
valid += 1;
0u64
}
b'C' | b'c' => {
valid += 1;
1
}
b'G' | b'g' => {
valid += 1;
2
}
b'T' | b't' => {
valid += 1;
3
}
_ => {
valid = 0;
h = 0;
continue;
}
};
h = ((h << 2) | bits) & mask;
if valid >= TERM_K {
out.insert(h);
}
}
out
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn explicit_valid() {
let reads: &[&[u8]] = &[b"AAA", b"CCC", b"GGG"];
assert_eq!(select_seed(reads, &SeedSelection::Explicit(2)).unwrap(), 2);
}
#[test]
fn explicit_out_of_bounds() {
let reads: &[&[u8]] = &[b"AAA", b"CCC"];
assert!(matches!(
select_seed(reads, &SeedSelection::Explicit(5)),
Err(PoaError::SeedOutOfBounds { index: 5, len: 2 })
));
}
#[test]
fn shortest_picks_min_len() {
let reads: &[&[u8]] = &[b"AAAAAAA", b"AAA", b"AAAAA"];
assert_eq!(select_seed(reads, &SeedSelection::Shortest).unwrap(), 1);
}
#[test]
fn empty_input_errors() {
let reads: &[&[u8]] = &[];
assert!(matches!(
select_seed(reads, &SeedSelection::Auto),
Err(PoaError::EmptyInput)
));
assert!(matches!(
select_seed(reads, &SeedSelection::Shortest),
Err(PoaError::EmptyInput)
));
assert!(matches!(
select_seed(reads, &SeedSelection::Explicit(0)),
Err(PoaError::EmptyInput)
));
}
const LEFT_FLANK: &[u8] = b"ACGTACGTACGTACGTACGTACGTACGTAC"; const RIGHT_FLANK: &[u8] = b"TGCATGCATGCATGCATGCATGCATGCATG"; const REPEAT: &[u8] = b"CAT";
fn make_spanning(n_repeat: usize) -> Vec<u8> {
let mut v = LEFT_FLANK.to_vec();
for _ in 0..n_repeat {
v.extend_from_slice(REPEAT);
}
v.extend_from_slice(RIGHT_FLANK);
v
}
fn make_left_only(n_repeat: usize) -> Vec<u8> {
let mut v = LEFT_FLANK.to_vec();
for _ in 0..n_repeat {
v.extend_from_slice(REPEAT);
}
v
}
fn make_right_only(n_repeat: usize) -> Vec<u8> {
let mut v: Vec<u8> = Vec::new();
for _ in 0..n_repeat {
v.extend_from_slice(REPEAT);
}
v.extend_from_slice(RIGHT_FLANK);
v
}
#[test]
fn auto_picks_shortest_spanning() {
let r_short = make_spanning(15); let r_medium = make_spanning(20); let r_long = make_spanning(25); let reads: Vec<&[u8]> = vec![&r_medium, &r_long, &r_short];
let idx = select_seed(&reads, &SeedSelection::Auto).unwrap();
assert_eq!(idx, 2, "expected shortest spanning read at index 2");
}
#[test]
fn auto_ignores_partials_picks_spanning() {
let s_short = make_spanning(15); let s_long = make_spanning(20); let left1 = make_left_only(30);
let left2 = make_left_only(30);
let left3 = make_left_only(30);
let right1 = make_right_only(30);
let right2 = make_right_only(30);
let reads: Vec<&[u8]> = vec![&s_short, &s_long, &left1, &left2, &left3, &right1, &right2];
let idx = select_seed(&reads, &SeedSelection::Auto).unwrap();
assert_eq!(
reads[idx].len(),
s_short.len(),
"expected the shortest spanning read"
);
}
#[test]
fn auto_two_cluster_errors() {
let left1 = make_left_only(30);
let left2 = make_left_only(30);
let left3 = make_left_only(30);
let right1 = make_right_only(30);
let right2 = make_right_only(30);
let right3 = make_right_only(30);
let reads: Vec<&[u8]> = vec![&left1, &left2, &left3, &right1, &right2, &right3];
match select_seed(&reads, &SeedSelection::Auto) {
Err(PoaError::NoSpanningReads {
left_depth,
right_depth,
}) => {
assert_eq!(left_depth, 3);
assert_eq!(right_depth, 3);
}
other => panic!("expected NoSpanningReads, got {other:?}"),
}
}
#[test]
fn auto_fallback_to_longest_when_no_cluster() {
let reads: Vec<&[u8]> = vec![b"ACGTACG", b"ACGT", b"ACGTACGTACG"];
let idx = select_seed(&reads, &SeedSelection::Auto).unwrap();
assert_eq!(idx, 2, "expected longest read as fallback");
}
}