use hyalite::{Mode, Scoring, SearchType, align_pair};
#[derive(Clone, Copy, PartialEq)]
enum Move {
Start,
Right, Down, }
struct Prob<'a> {
q: &'a [u8],
t: &'a [u8],
mat: &'a [i32],
al: usize,
go: i32,
ge: i32,
}
fn brute_nw(q: &[u8], t: &[u8], mat: &[i32], al: usize, go: i32, ge: i32) -> i32 {
fn rec(p: &Prob, i: usize, j: usize, last: Move) -> i32 {
let (m, n) = (p.q.len(), p.t.len());
if i == m && j == n {
return 0;
}
let mut best = i32::MIN;
if i < m && j < n {
let s = p.mat[p.q[i] as usize * p.al + p.t[j] as usize];
best = best.max(s.saturating_add(rec(p, i + 1, j + 1, Move::Start)));
}
if j < n {
let cost = if last == Move::Right { p.ge } else { p.go };
best = best.max((-cost).saturating_add(rec(p, i, j + 1, Move::Right)));
}
if i < m {
let cost = if last == Move::Down { p.ge } else { p.go };
best = best.max((-cost).saturating_add(rec(p, i + 1, j, Move::Down)));
}
best
}
rec(
&Prob {
q,
t,
mat,
al,
go,
ge,
},
0,
0,
Move::Start,
)
}
fn brute_sw(q: &[u8], t: &[u8], mat: &[i32], al: usize, go: i32, ge: i32) -> i32 {
let (m, n) = (q.len(), t.len());
let mut best = 0;
for a in 0..=m {
for b in a..=m {
for c in 0..=n {
for d in c..=n {
best = best.max(brute_nw(&q[a..b], &t[c..d], mat, al, go, ge));
}
}
}
}
best
}
fn brute_hw(q: &[u8], t: &[u8], mat: &[i32], al: usize, go: i32, ge: i32) -> i32 {
let n = t.len();
let mut best = i32::MIN;
for c in 0..=n {
for d in c..=n {
best = best.max(brute_nw(q, &t[c..d], mat, al, go, ge));
}
}
best
}
fn brute_ov(q: &[u8], t: &[u8], mat: &[i32], al: usize, go: i32, ge: i32) -> i32 {
let (m, n) = (q.len(), t.len());
let mut best = 0;
for a in 0..=m {
for b in a..=m {
for c in 0..=n {
for d in c..=n {
let touches_start = a == 0 || c == 0;
let touches_end = b == m || d == n;
if touches_start && touches_end {
best = best.max(brute_nw(&q[a..b], &t[c..d], mat, al, go, ge));
}
}
}
}
}
best
}
fn brute(mode: Mode, q: &[u8], t: &[u8], mat: &[i32], al: usize, go: i32, ge: i32) -> i32 {
match mode {
Mode::Nw => brute_nw(q, t, mat, al, go, ge),
Mode::Sw => brute_sw(q, t, mat, al, go, ge),
Mode::Hw => brute_hw(q, t, mat, al, go, ge),
Mode::Ov => brute_ov(q, t, mat, al, go, ge),
_ => unreachable!("ALL_MODES covers every mode this test exercises"),
}
}
fn all_sequences(alphabet: u8, max_len: usize) -> Vec<Vec<u8>> {
let mut out = vec![vec![]];
let mut frontier = vec![vec![]];
for _ in 0..max_len {
let mut next = Vec::new();
for seq in &frontier {
for sym in 0..alphabet {
let mut s = seq.clone();
s.push(sym);
next.push(s);
}
}
out.extend(next.iter().cloned());
frontier = next;
}
out
}
const ALL_MODES: [Mode; 4] = [Mode::Sw, Mode::Nw, Mode::Hw, Mode::Ov];
fn identity_matrix(al: usize, m: i32, x: i32) -> Vec<i32> {
let mut v = vec![x; al * al];
for i in 0..al {
v[i * al + i] = m;
}
v
}
#[test]
fn scalar_matches_brute_force_over_all_short_pairs() {
let schemes = [
(2usize, 3usize, 2, -1, 2, 1),
(2, 3, 1, -1, 3, 3),
(2, 3, 3, -2, 4, 2),
(3, 2, 2, -1, 2, 1),
];
for (al, max_len, m, x, go, ge) in schemes {
let mat = identity_matrix(al, m, x);
let scoring = Scoring::new(al, mat.clone(), go, ge).unwrap();
let seqs = all_sequences(al as u8, max_len);
for q in &seqs {
for t in &seqs {
for mode in ALL_MODES {
let expected = brute(mode, q, t, &mat, al, go, ge);
let hit = align_pair(q, t, &scoring, mode, SearchType::Score).unwrap();
assert_eq!(
hit.score, expected,
"mode {mode}, scheme (al={al}, m={m}, x={x}, go={go}, ge={ge}), \
q={q:?}, t={t:?}"
);
}
}
}
}
}
fn dna() -> Scoring {
Scoring::new(4, identity_matrix(4, 2, -1), 2, 1).unwrap()
}
fn score(mode: Mode, q: &[u8], t: &[u8]) -> i32 {
align_pair(q, t, &dna(), mode, SearchType::Score)
.unwrap()
.score
}
#[test]
fn perfect_match_scores_full_match_in_every_mode() {
let seq = [0u8, 1, 2, 3];
for mode in ALL_MODES {
assert_eq!(score(mode, &seq, &seq), 8, "{mode}");
let hit = align_pair(&seq, &seq, &dna(), mode, SearchType::ScoreEnd).unwrap();
assert_eq!(
(hit.query_end, hit.target_end),
(Some(3), Some(3)),
"{mode} ends"
);
}
}
#[test]
fn single_internal_mismatch() {
let q = [0u8, 1, 2, 3];
let t = [0u8, 1, 0, 3];
assert_eq!(score(Mode::Nw, &q, &t), 5);
assert_eq!(score(Mode::Sw, &q, &t), 5);
}
#[test]
fn single_gap_costs_gap_open_only() {
let q = [0u8, 1, 2, 3];
let t = [0u8, 2, 3];
assert_eq!(score(Mode::Nw, &q, &t), 4);
}
#[test]
fn two_base_gap_uses_open_plus_extend() {
let q = [0u8, 1, 2, 3];
let t = [0u8, 3];
assert_eq!(score(Mode::Nw, &q, &t), 1);
}
#[test]
fn hw_fits_query_into_longer_target_ignoring_target_flanks() {
let q = [1u8, 2];
let t = [0u8, 1, 2, 3];
assert_eq!(score(Mode::Hw, &q, &t), 4);
let hit = align_pair(&q, &t, &dna(), Mode::Hw, SearchType::ScoreEnd).unwrap();
assert_eq!((hit.query_end, hit.target_end), (Some(1), Some(2)));
assert_eq!(score(Mode::Nw, &q, &t), 0);
}
#[test]
fn ov_scores_suffix_prefix_overlap_for_free_ends() {
let q = [0u8, 0, 2, 3]; let t = [2u8, 3, 1, 1]; assert_eq!(score(Mode::Ov, &q, &t), 4);
}
#[test]
fn mode_score_ordering_holds_for_all_short_pairs() {
let scoring = dna();
let mat = identity_matrix(4, 2, -1);
let seqs = all_sequences(4, 3);
for q in &seqs {
for t in &seqs {
let sw = score_with(&scoring, Mode::Sw, q, t);
let ov = score_with(&scoring, Mode::Ov, q, t);
let hw = score_with(&scoring, Mode::Hw, q, t);
let nw = score_with(&scoring, Mode::Nw, q, t);
let _ = &mat;
assert!(sw >= 0, "SW negative: q={q:?} t={t:?} -> {sw}");
assert!(sw >= ov, "SW<OV: q={q:?} t={t:?} ({sw}<{ov})");
assert!(ov >= hw, "OV<HW: q={q:?} t={t:?} ({ov}<{hw})");
assert!(hw >= nw, "HW<NW: q={q:?} t={t:?} ({hw}<{nw})");
}
}
}
fn score_with(s: &Scoring, mode: Mode, q: &[u8], t: &[u8]) -> i32 {
align_pair(q, t, s, mode, SearchType::Score).unwrap().score
}
#[test]
fn tie_break_prefers_smallest_target_then_query_end() {
let q = [0u8];
let t = [0u8, 0, 0];
let hit = align_pair(&q, &t, &dna(), Mode::Sw, SearchType::ScoreEnd).unwrap();
assert_eq!(hit.score, 2);
assert_eq!(
hit.target_end,
Some(0),
"should pick the leftmost target match"
);
assert_eq!(hit.query_end, Some(0));
let hit = align_pair(&q, &t, &dna(), Mode::Hw, SearchType::ScoreEnd).unwrap();
assert_eq!((hit.score, hit.target_end), (2, Some(0)));
}
#[test]
fn tie_break_is_stable_regardless_of_target_length() {
for run in 1..=6 {
let t = vec![0u8; run];
let hit = align_pair(&[0u8], &t, &dna(), Mode::Sw, SearchType::ScoreEnd).unwrap();
assert_eq!(hit.target_end, Some(0), "run={run}");
}
}
#[test]
fn empty_inputs_never_panic_and_score_correctly() {
let s = dna();
for mode in ALL_MODES {
let hit = align_pair(&[], &[], &s, mode, SearchType::ScoreEnd).unwrap();
assert_eq!(hit.score, 0, "{mode} empty/empty");
assert_eq!((hit.query_end, hit.target_end), (None, None), "{mode} ends");
}
let hit = align_pair(&[], &[0, 1, 2, 3], &s, Mode::Nw, SearchType::ScoreEnd).unwrap();
assert_eq!(hit.score, -5);
assert_eq!(hit.query_end, None);
assert_eq!(hit.target_end, Some(3));
assert_eq!(score(Mode::Sw, &[0, 1, 2], &[]), 0);
assert_eq!(score(Mode::Ov, &[], &[0, 1, 2]), 0);
assert_eq!(score(Mode::Hw, &[], &[0, 1, 2]), 0);
}
#[test]
fn score_is_symmetric_for_symmetric_modes_and_matrix() {
let scoring = dna();
let seqs = all_sequences(4, 3);
for q in &seqs {
for t in &seqs {
for mode in [Mode::Sw, Mode::Nw, Mode::Ov] {
let a = score_with(&scoring, mode, q, t);
let b = score_with(&scoring, mode, t, q);
assert_eq!(a, b, "{mode} not symmetric: q={q:?} t={t:?} ({a} != {b})");
}
}
}
}
#[test]
fn hw_is_generally_not_symmetric() {
let short = [1u8, 2];
let long = [0u8, 1, 2, 3];
assert_ne!(
score(Mode::Hw, &short, &long),
score(Mode::Hw, &long, &short)
);
}