1use std::cmp::min;
10
11type IndexAndAlignments = (usize, Vec<(Option<usize>, Option<usize>)>);
12
13#[derive(PartialEq, Eq, Clone, Copy, Debug)]
14pub enum Trace {
15 Stop,
16 Left,
17 Up,
18 MatchDiagonal,
19 MisMatchDiagonal,
20}
21
22#[derive(PartialEq, Eq, Debug)]
23pub enum NeedlemanWunschError {
24 MatrixTooLarge { cells: usize, max_cells: usize },
25}
26
27fn simple_score<T: PartialEq>(a: T, b: T) -> i32 {
28 if a == b { 1 } else { -5 }
29}
30
31pub struct NeedlemanWunsch<T: PartialEq + Clone> {
32 pub first_sequence: Vec<T>,
33 pub second_sequence: Vec<T>,
34 pub first_extra: Vec<T>,
35 pub second_extra: Vec<T>,
36 pub mismatches: Vec<(usize, usize)>,
37 pub num_mismatches: usize,
38}
39
40impl<T: PartialEq + Clone> Default for NeedlemanWunsch<T> {
41 fn default() -> NeedlemanWunsch<T> {
42 NeedlemanWunsch {
43 first_sequence: vec![],
44 second_sequence: vec![],
45 first_extra: vec![],
46 second_extra: vec![],
47 mismatches: vec![],
48 num_mismatches: 0,
49 }
50 }
51}
52
53impl<T: PartialEq + Clone> NeedlemanWunsch<T> {
54 pub fn match_sequences_base(&mut self) -> IndexAndAlignments {
55 self.match_sequences(None, None, None)
56 }
57
58 pub fn generate_midpoint_schedule(
61 &mut self,
62 start_index: usize,
63 alignment_difference: Vec<(Option<usize>, Option<usize>)>,
64 ) -> Vec<T> {
65 let mut midpoint_schedule: Vec<T> = vec![];
66 let total_swaps = self.num_mismatches / 2;
67 let available_full_swaps = self.mismatches.len();
68 let mut current_swaps = 0;
69 let mut full_swaps = 0;
70 {
71 let mut i = 0;
72 while i < start_index {
73 midpoint_schedule.push(self.first_sequence[i].clone());
74 i += 1;
75 }
76 }
77 {
78 let mut i = 0;
79 while i < alignment_difference.len() {
80 if alignment_difference[i].0.is_none() && alignment_difference[i].1.is_none() {
81 if full_swaps < available_full_swaps && current_swaps < total_swaps {
82 midpoint_schedule
83 .push(self.first_sequence[self.mismatches[full_swaps].0].clone());
84 } else if full_swaps < available_full_swaps {
85 midpoint_schedule
86 .push(self.second_sequence[self.mismatches[full_swaps].1].clone());
87 }
88 current_swaps += 1;
89 full_swaps += 1;
90 } else if alignment_difference[i].0.is_some() && current_swaps < total_swaps {
91 midpoint_schedule
92 .push(self.first_sequence[alignment_difference[i].0.unwrap()].clone());
93 current_swaps += 1;
94 } else if alignment_difference[i].1.is_some() && current_swaps >= total_swaps {
95 midpoint_schedule
96 .push(self.second_sequence[alignment_difference[i].1.unwrap()].clone());
97 current_swaps += 1;
98 }
99 i += 1;
100 }
101 }
102 midpoint_schedule
103 }
104
105 pub fn match_sequences(
111 &mut self,
112 scoring_function: Option<fn(T, T) -> i32>,
113 gap_penalty: Option<i32>,
114 mismatch_penalty: Option<i32>,
115 ) -> IndexAndAlignments {
116 self.match_sequences_bounded(scoring_function, gap_penalty, mismatch_penalty, usize::MAX)
117 .expect("Needleman-Wunsch matrix dimensions overflowed")
118 }
119
120 pub fn match_sequences_bounded(
123 &mut self,
124 scoring_function: Option<fn(T, T) -> i32>,
125 gap_penalty: Option<i32>,
126 mismatch_penalty: Option<i32>,
127 max_matrix_cells: usize,
128 ) -> Result<IndexAndAlignments, NeedlemanWunschError> {
129 let gap_penalty: i32 = gap_penalty.unwrap_or(-1);
130 let mismatch_penalty: i32 = mismatch_penalty.unwrap_or(-2);
131 let scoring_function = scoring_function.unwrap_or(simple_score);
132
133 self.mismatches.clear();
134 self.num_mismatches = 0;
135
136 let mut start_index = 0;
137 let max_length: usize = min(self.first_sequence.len(), self.second_sequence.len());
138 while start_index < max_length
139 && scoring_function(
140 self.first_sequence[start_index].clone(),
141 self.second_sequence[start_index].clone(),
142 ) > 0
143 {
144 start_index += 1;
145 }
146 let row = self.first_sequence.len() + 1 - start_index;
147 let col = self.second_sequence.len() + 1 - start_index;
148 let cells = row
149 .checked_mul(col)
150 .ok_or(NeedlemanWunschError::MatrixTooLarge {
151 cells: usize::MAX,
152 max_cells: max_matrix_cells,
153 })?;
154 if cells > max_matrix_cells {
155 return Err(NeedlemanWunschError::MatrixTooLarge {
156 cells,
157 max_cells: max_matrix_cells,
158 });
159 }
160
161 let mut previous_scores = (0..col).map(|j| gap_penalty * j as i32).collect::<Vec<_>>();
162 let mut current_scores = vec![0; col];
163 let mut tracing_matrix = vec![Trace::Stop; cells];
164 for trace in tracing_matrix.iter_mut().take(col).skip(1) {
165 *trace = Trace::Left;
166 }
167 for i in 1..row {
168 current_scores[0] = gap_penalty * i as i32;
169 tracing_matrix[i * col] = Trace::Up;
170 for j in 1..col {
171 let match_value: i32 = scoring_function(
172 self.first_sequence[i - 1 + start_index].clone(),
173 self.second_sequence[j - 1 + start_index].clone(),
174 );
175
176 let diagonal_score = previous_scores[j - 1]
177 + if match_value > 0 {
178 match_value
179 } else {
180 mismatch_penalty
181 };
182 let horizontal_score = current_scores[j - 1] + gap_penalty;
183 let vertical_score = previous_scores[j] + gap_penalty;
184
185 let (score, trace) =
186 if diagonal_score >= horizontal_score && diagonal_score >= vertical_score {
187 (
188 diagonal_score,
189 if match_value > 0 {
190 Trace::MatchDiagonal
191 } else {
192 Trace::MisMatchDiagonal
193 },
194 )
195 } else if horizontal_score >= vertical_score {
196 (horizontal_score, Trace::Left)
197 } else {
198 (vertical_score, Trace::Up)
199 };
200 current_scores[j] = score;
201 tracing_matrix[i * col + j] = trace;
202 }
203 std::mem::swap(&mut previous_scores, &mut current_scores);
204 }
205
206 let mut aligned_seq: Vec<(Option<usize>, Option<usize>)> = vec![];
207 let (mut i, mut j) = (row - 1, col - 1);
208
209 while i > 0 || j > 0 {
210 match tracing_matrix[i * col + j] {
211 Trace::MatchDiagonal => {
212 aligned_seq.push((Some(i - 1 + start_index), Some(j - 1 + start_index)));
213 i -= 1;
214 j -= 1;
215 }
216 Trace::Up => {
217 aligned_seq.push((Some(i - 1 + start_index), None));
218 i -= 1;
219 self.num_mismatches += 1;
220 }
221 Trace::Left => {
222 aligned_seq.push((None, Some(j - 1 + start_index)));
223 j -= 1;
224 self.num_mismatches += 1;
225 }
226 Trace::MisMatchDiagonal => {
227 aligned_seq.push((None, None));
228 self.num_mismatches += 2;
229 self.mismatches
230 .push((i - 1 + start_index, j - 1 + start_index));
231 i -= 1;
232 j -= 1;
233 }
234 Trace::Stop => unreachable!("traceback stopped before reaching matrix origin"),
235 }
236 }
237 self.mismatches.reverse();
238 aligned_seq.reverse();
239
240 Ok((start_index, aligned_seq))
241 }
242}
243
244#[cfg(test)]
245mod tests {
246 use super::*;
247
248 #[test]
249 fn test_single_complete_match() {
250 let mut sw_object = NeedlemanWunsch {
251 first_sequence: vec![10],
252 second_sequence: vec![10],
253 ..Default::default()
254 };
255 assert_eq!(sw_object.match_sequences_base(), (1, vec![]));
256 }
257
258 #[test]
259 fn test_single_complete_mismatch() {
260 let mut sw_object = NeedlemanWunsch {
261 first_sequence: vec![0],
262 second_sequence: vec![1],
263 ..Default::default()
264 };
265 assert_eq!(sw_object.match_sequences_base(), (0, vec![(None, None)],));
266 }
267
268 #[test]
269 fn test_partial_one_off_match() {
270 let mut sw_object = NeedlemanWunsch {
271 first_sequence: vec![3, 3, 4, 4, 3, 1, 2, 4, 1],
272 second_sequence: vec![4, 3, 4, 4, 1, 2, 3, 3],
273 ..Default::default()
274 };
275 assert_eq!(
276 sw_object.match_sequences_base(),
277 (
278 0,
279 vec![
280 (None, None),
281 (Some(1), Some(1)),
282 (Some(2), Some(2)),
283 (Some(3), Some(3)),
284 (Some(4), None),
285 (Some(5), Some(4)),
286 (Some(6), Some(5)),
287 (None, None),
288 (None, None)
289 ]
290 )
291 );
292 }
293
294 #[test]
295 fn test_gap_match() {
296 let mut sw_object = NeedlemanWunsch {
297 first_sequence: vec![1, 2, 3],
298 second_sequence: vec![1, 4, 3],
299 ..Default::default()
300 };
301 assert_eq!(
302 sw_object.match_sequences_base(),
303 (1, vec![(None, None), (Some(2), Some(2))])
304 );
305 }
306
307 #[test]
308 fn test_gap_start_match() {
309 let mut sw_object = NeedlemanWunsch {
310 first_sequence: vec![2, 4, 3],
311 second_sequence: vec![1, 4, 3],
312 ..Default::default()
313 };
314 assert_eq!(
315 sw_object.match_sequences_base(),
316 (
317 0,
318 vec![(None, None), (Some(1), Some(1)), (Some(2), Some(2))],
319 )
320 );
321 }
322
323 #[test]
324 fn bounded_alignment_rejects_large_matrices() {
325 let mut sw_object = NeedlemanWunsch {
326 first_sequence: vec![1; 100],
327 second_sequence: vec![2; 100],
328 ..Default::default()
329 };
330 assert_eq!(
331 sw_object.match_sequences_bounded(None, None, None, 10_000),
332 Err(NeedlemanWunschError::MatrixTooLarge {
333 cells: 10_201,
334 max_cells: 10_000,
335 })
336 );
337 }
338
339 #[test]
340 fn alignment_includes_unmatched_suffix() {
341 let mut sw_object = NeedlemanWunsch {
342 first_sequence: vec![1, 2],
343 second_sequence: vec![1, 2, 3],
344 ..Default::default()
345 };
346 assert_eq!(sw_object.match_sequences_base(), (2, vec![(None, Some(2))]));
347 }
348}