From cf7023a6daea2180409a58bdb57726c61107b50c Mon Sep 17 00:00:00 2001 From: Johanna Date: Wed, 27 May 2026 14:54:42 +0100 Subject: [PATCH 1/3] Add cross-query kNN distance mode --- src/cli.rs | 17 +- src/distances/distance_matrix.rs | 70 ++++-- src/distances/jaccard.rs | 7 +- src/distances/mod.rs | 123 +++++++++- src/io.rs | 64 +++--- src/lib.rs | 110 +++++---- tests/completeness.rs | 12 +- tests/distance.rs | 360 +++++++++++++++++++++++++++++- tests/test_files_in/qfile.txt | 2 + tests/test_files_in/rfile_ref.txt | 2 + 10 files changed, 657 insertions(+), 110 deletions(-) create mode 100644 tests/test_files_in/qfile.txt create mode 100644 tests/test_files_in/rfile_ref.txt diff --git a/src/cli.rs b/src/cli.rs index 67fa7a1..5b8f88c 100644 --- a/src/cli.rs +++ b/src/cli.rs @@ -178,8 +178,8 @@ pub enum Commands { #[arg(short)] output: Option, - /// Calculate sparse distances with k nearest-neighbours - #[arg(long, group = "query")] + /// Calculate sparse distances with k nearest-neighbours (ref-vs-ref or ref-vs-query) + #[arg(long)] knn: Option, /// Sample names to analyse @@ -198,9 +198,14 @@ pub enum Commands { #[arg(long, value_parser = valid_cpus, default_value_t = 1)] threads: usize, - /// File listing sample and completeness estimate 0.0-1.0 (tab separated) + /// File listing reference sample completeness estimates 0.0-1.0 (tab separated: genome_idcompleteness) + #[arg(long)] + ref_completeness_file: Option, + + /// File listing query sample completeness estimates 0.0-1.0 (tab separated: genome_idcompleteness). + /// Only used in cross-query mode (when a query database is provided). #[arg(long)] - completeness_file: Option, + query_completeness_file: Option, /// minimum completeness for a sample to be corrected but the completeness correction #[arg(long, default_value_t = 0.64)] @@ -419,9 +424,9 @@ pub enum InvertedCommands { #[arg(long, value_parser = valid_cpus, default_value_t = 1)] threads: usize, - /// Completeness file + /// File listing sample completeness estimates 0.0-1.0 (tab separated: genome_idcompleteness) #[arg(long)] - completeness_file: Option, + ref_completeness_file: Option, /// minimum completeness for a sample to be corrected but the completeness correction #[arg(long, default_value_t = 0.64)] diff --git a/src/distances/distance_matrix.rs b/src/distances/distance_matrix.rs index 6c6e742..ec39dd9 100644 --- a/src/distances/distance_matrix.rs +++ b/src/distances/distance_matrix.rs @@ -274,7 +274,14 @@ pub enum DistVec { CoreAcc(Vec), } -/// A sparse distance matrix with a maximum of `knn` distances for each sample +/// A sparse distance matrix with a maximum of `knn` distances for each sample. +/// +/// In self-query mode (one database), `query_names` is `None` and `ref_names` is +/// used for both row iteration and column index lookup. +/// +/// In cross-query mode (two databases), `query_names` holds the query genome names +/// (one per row) and `ref_names` holds the reference genome names (indexed by +/// the stored neighbour index inside each distance item). pub struct SparseDistanceMatrix<'a> { /// Total number of distances pub n_distances: usize, @@ -282,16 +289,17 @@ pub struct SparseDistanceMatrix<'a> { pub knn: usize, jaccard: DistType, distances: DistVec, + /// Reference genome names — used as column labels (neighbour index lookup) ref_names: Vec<&'a str>, + /// Query genome names — used as row labels in cross-query mode (`None` in self-query mode) + query_names: Option>, } impl<'a> SparseDistanceMatrix<'a> { - /// Create a new sparse distance matrix for a [`MultiSketch`] keeping the - /// minimum `knn` distances with distance options set by `jaccard` + /// Self-query constructor: one database, rows and columns are the same set. pub fn new(ref_sketches: &'a MultiSketch, knn: usize, jaccard: DistType) -> Self { let n_distances = ref_sketches.number_samples_loaded() * knn; - // Pre-allocate distances let distances = match jaccard { DistType::CoreAcc => DistVec::CoreAcc(vec![SparseCoreAcc(0, 0.0, 0.0); n_distances]), DistType::Jaccard(_, _, _) => { @@ -305,6 +313,35 @@ impl<'a> SparseDistanceMatrix<'a> { jaccard, distances, ref_names: Self::sketch_names(ref_sketches), + query_names: None, + } + } + + /// Cross-query constructor: two databases. + /// Rows are query genomes; column indices index into the reference genome list. + pub fn new_cross_query( + ref_sketches: &'a MultiSketch, + query_sketches: &'a MultiSketch, + knn: usize, + jaccard: DistType, + ) -> Self { + let n_query = query_sketches.number_samples_loaded(); + let n_distances = n_query * knn; + + let distances = match jaccard { + DistType::CoreAcc => DistVec::CoreAcc(vec![SparseCoreAcc(0, 0.0, 0.0); n_distances]), + DistType::Jaccard(_, _, _) => { + DistVec::Jaccard(vec![SparseJaccard(0, 0.0); n_distances]) + } + }; + + Self { + n_distances, + knn, + jaccard, + distances, + ref_names: Self::sketch_names(ref_sketches), + query_names: Some(Self::sketch_names(query_sketches)), } } @@ -322,24 +359,25 @@ impl<'a> Distances<'a> for SparseDistanceMatrix<'a> { impl fmt::Display for SparseDistanceMatrix<'_> { fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result { - let mut ref_name_iter = self.ref_names.iter(); - let mut ref_name = ref_name_iter.next().unwrap(); + // In cross-query mode use query genome names as row labels; + // in self-query mode fall back to ref_names. + let query_names = self.query_names.as_deref().unwrap_or(&self.ref_names); + let mut row_name_iter = query_names.iter(); + let mut row_name = row_name_iter.next().unwrap(); let mut k = 0; match &self.distances { DistVec::Jaccard(dists) => { for dist_item in dists { k += 1; if k > self.knn { - ref_name = ref_name_iter.next().unwrap(); + row_name = row_name_iter.next().unwrap(); k = 1; } - let query_name = self.ref_names[dist_item.0]; - // If fewer items than knn, padded with ref=query and dist=1.0 values - // which are skipped here - // TODO: more rust-like way of doing this would be to have - // SparseJaccard as an enum with an empty value - if dist_item.1 < 1.0_f32 || query_name != *ref_name { - writeln!(f, "{ref_name}\t{query_name}\t{}", dist_item.1)?; + // dist_item.0 is always an index into ref_names + let col_name = self.ref_names[dist_item.0]; + // Padding entries (dist == 1.0, col == row) are skipped + if dist_item.1 < 1.0_f32 || col_name != *row_name { + writeln!(f, "{row_name}\t{col_name}\t{}", dist_item.1)?; } } } @@ -347,12 +385,12 @@ impl fmt::Display for SparseDistanceMatrix<'_> { for dist_item in dists { k += 1; if k > self.knn { - ref_name = ref_name_iter.next().unwrap(); + row_name = row_name_iter.next().unwrap(); k = 1; } writeln!( f, - "{ref_name}\t{}\t{}\t{}", + "{row_name}\t{}\t{}\t{}", self.ref_names[dist_item.0], dist_item.1, dist_item.2, )?; } diff --git a/src/distances/jaccard.rs b/src/distances/jaccard.rs index 2ce8ae6..91f82d6 100644 --- a/src/distances/jaccard.rs +++ b/src/distances/jaccard.rs @@ -63,7 +63,8 @@ pub fn core_acc_dist( query_sketches: &MultiSketch, ref_sketch_idx: usize, query_sketch_idx: usize, - completeness_vec: Option<&Vec>, + ref_completeness_vec: Option<&Vec>, + query_completeness_vec: Option<&Vec>, completeness_cutoff: f64, ) -> (f32, f32) { if ref_sketches.kmer_lengths().len() < 2 { @@ -74,8 +75,8 @@ pub fn core_acc_dist( let tolerance = (2.0_f64 / ((ref_sketches.sketch_size * u64::BITS as u64) as f64)).ln(); //let tolerance = -100.0_f32; for (k_idx, k) in ref_sketches.kmer_lengths().iter().enumerate() { - let c1 = completeness_vec.map(|cv| cv[ref_sketch_idx]); - let c2 = completeness_vec.map(|cv| cv[query_sketch_idx]); + let c1 = ref_completeness_vec.map(|cv| cv[ref_sketch_idx]); + let c2 = query_completeness_vec.map(|cv| cv[query_sketch_idx]); let y = jaccard_index( ref_sketches.get_sketch_slice(ref_sketch_idx, k_idx), query_sketches.get_sketch_slice(query_sketch_idx, k_idx), diff --git a/src/distances/mod.rs b/src/distances/mod.rs index f2ac1b8..327fade 100644 --- a/src/distances/mod.rs +++ b/src/distances/mod.rs @@ -106,6 +106,7 @@ pub fn self_dists_all<'a>( i, j, completeness_vec, + completeness_vec, completeness_cutoff, ); dist_slice[dist_idx * 2] = dist.0; @@ -207,6 +208,7 @@ pub fn self_dists_knn<'a>( i, j, completeness_vec, + completeness_vec, completeness_cutoff, ); let dist_item = SparseCoreAcc(j, dists.0, dists.1); @@ -220,15 +222,16 @@ pub fn self_dists_knn<'a>( sp_distances } -/// Self query mode (dense, all distances) -pub fn self_query_dists_all<'a>( +/// Cross-query mode (dense, all distances) +pub fn cross_dists_all<'a>( ref_sketches: &'a MultiSketch, query_sketches: &'a MultiSketch, n: usize, - nq: usize, + n_query: usize, dist_type: DistType, quiet: bool, - completeness_vec: Option<&Vec>, + ref_completeness_vec: Option<&Vec>, + query_completeness_vec: Option<&Vec>, completeness_cutoff: f64, ) -> DistanceMatrix<'a> { let mut distances = DistanceMatrix::new(ref_sketches, Some(query_sketches), dist_type); @@ -244,14 +247,11 @@ pub fn self_query_dists_all<'a>( .for_each(|(chunk_idx, dist_slice)| { // Get first i, j index for the chunk let start_dist_idx = chunk_idx * CHUNK_SIZE; - let (mut i, mut j) = calc_query_indices(start_dist_idx, nq); + let (mut i, mut j) = calc_query_indices(start_dist_idx, n_query); for dist_idx in 0..CHUNK_SIZE { if let Some((k_idx, k_f64)) = k_vals { - // If completeness_vec is Some, extract the value at index i (or j) from the inner vector. - // If completeness_vec is None, the result will also be None. - // This uses Option::map to safely access the completeness value for each sample. - let c1 = completeness_vec.map(|cv| cv[i]); - let c2 = completeness_vec.map(|cv| cv[j]); + let c1 = ref_completeness_vec.map(|cv| cv[i]); + let c2 = query_completeness_vec.map(|cv| cv[j]); let j_index = jaccard_index( ref_sketches.get_sketch_slice(i, k_idx), query_sketches.get_sketch_slice(j, k_idx), @@ -272,7 +272,8 @@ pub fn self_query_dists_all<'a>( query_sketches, i, j, - completeness_vec, + ref_completeness_vec, + query_completeness_vec, completeness_cutoff, ); dist_slice[dist_idx * 2] = dist.0; @@ -281,7 +282,7 @@ pub fn self_query_dists_all<'a>( // Move to next index j += 1; - if j >= nq { + if j >= n_query { i += 1; j = 0; // End of all dists reached (final chunk) @@ -294,6 +295,104 @@ pub fn self_query_dists_all<'a>( distances } +/// Cross-query mode with kNN filtering. +/// +/// For each query genome, computes distances to all `n` reference genomes and +/// retains only the `knn` nearest neighbours using a priority queue (max-heap +/// capped at `knn`). Output has `n_query × knn` entries — one row per query genome. +/// +/// This is the cross-database analogue of [`self_dists_knn`]. +pub fn cross_dists_knn<'a>( + ref_sketches: &'a MultiSketch, + query_sketches: &'a MultiSketch, + n: usize, + n_query: usize, + knn: usize, + dist_type: DistType, + quiet: bool, + ref_completeness_vec: Option<&Vec>, + query_completeness_vec: Option<&Vec>, + completeness_cutoff: f64, +) -> SparseDistanceMatrix<'a> { + if n == 0 { + panic!("Reference database has no loaded samples"); + } + if n_query == 0 { + panic!("Query database has no loaded samples"); + } + // Can't have more neighbours than there are reference genomes + let knn = knn.min(n); + let mut sp_distances = + SparseDistanceMatrix::new_cross_query(ref_sketches, query_sketches, knn, dist_type); + let k_vals = sp_distances.k_vals(); + let ani = sp_distances.ani(); + let progress_bar = get_progress_bar(n_query, BAR_PERCENT, quiet); + match sp_distances.dists_mut() { + DistVec::Jaccard(distances) => { + let (k_idx, k_f64) = k_vals.unwrap(); + distances + .par_chunks_mut(knn) + .progress_with(progress_bar) + .enumerate() + .for_each(|(qi, row_dist_slice)| { + let mut heap = BinaryHeap::with_capacity(knn + 1); + let qi_sketch = query_sketches.get_sketch_slice(qi, k_idx); + for ri in 0..n { + let c1 = query_completeness_vec.map(|cv| cv[qi]); + let c2 = ref_completeness_vec.map(|cv| cv[ri]); + let dist = jaccard_index( + qi_sketch, + ref_sketches.get_sketch_slice(ri, k_idx), + ref_sketches.sketchsize64, + c1, + c2, + completeness_cutoff, + ); + let dist_f32 = if ani { + (1.0_f64 - ani_pois(dist, k_f64)) as f32 + } else { + (1.0_f64 - dist) as f32 + }; + push_heap(&mut heap, SparseJaccard(ri, dist_f32), knn); + } + if ani { + heap.into_sorted_vec().iter().zip(row_dist_slice).for_each( + |(inverse_ani, output_ani)| { + *output_ani = + SparseJaccard(inverse_ani.0, 1.0_f32 - inverse_ani.1); + }, + ); + } else { + row_dist_slice.clone_from_slice(&heap.into_sorted_vec()); + } + }); + } + DistVec::CoreAcc(distances) => { + distances + .par_chunks_mut(knn) + .progress_with(progress_bar) + .enumerate() + .for_each(|(qi, row_dist_slice)| { + let mut heap = BinaryHeap::with_capacity(knn + 1); + for ri in 0..n { + let dists = core_acc_dist( + ref_sketches, + query_sketches, + ri, + qi, + ref_completeness_vec, + query_completeness_vec, + completeness_cutoff, + ); + push_heap(&mut heap, SparseCoreAcc(ri, dists.0, dists.1), knn); + } + row_dist_slice.clone_from_slice(&heap.into_sorted_vec()); + }); + } + } + sp_distances +} + /// Same as [`self_dists_knn`], but also using an inverted_index to precluster /// to reduce the number of comparisons pub fn self_dists_knn_precluster<'a>( diff --git a/src/io.rs b/src/io.rs index cc5030b..426748c 100644 --- a/src/io.rs +++ b/src/io.rs @@ -210,60 +210,72 @@ pub fn read_subset_names(subset_file: &str) -> Vec { } /// Read completeness values from a file and create a vector for each genome in the provided sketches. -/// If a genome is not found in the file, a default value of 1.0 is used. +/// Completeness values must be in [0.0, 1.0] — not percentages. Genomes not in the file default to 1.0. pub fn read_completeness_file( completeness_file: &str, sketches: &MultiSketch, ) -> Result, crate::Error> { - // Pre-allocate vector with default values (1.0 for missing genomes) let mut completeness_vec = vec![1.0_f64; sketches.number_samples_loaded()]; - let missing_genomes = Mutex::new(Vec::new()); + let not_in_sketch = Mutex::new(Vec::new()); + let out_of_range = Mutex::new(Vec::new()); - // Open file with buffered reader for streaming let f = File::open(completeness_file) .with_context(|| format!("Failed to open completeness file: {completeness_file}"))?; let f = BufReader::new(f); - // Read lines and collect completeness values let lines: Vec = f.lines().collect::, _>>().with_context(|| { format!("Failed to read lines from completeness file: {completeness_file}") })?; - // Parse all lines in parallel and collect completeness values let updates: Vec<(usize, f64)> = lines .par_iter() .filter_map(|line| { - if let Some((genome_id, completeness_str)) = line.split_once('\t') { - if let Ok(completeness) = completeness_str.parse::() { - // Use MultiSketch's method to get the logical index - if let Some(index) = sketches.get_sample_index(genome_id) { - Some((index, completeness)) - } else { - // Track missing genomes - missing_genomes.lock().unwrap().push(genome_id.to_string()); - None - } - } else { - None - } + let Some((genome_id, completeness_str)) = line.split_once('\t') else { + return None; + }; + let Ok(completeness) = completeness_str.trim().parse::() else { + log::warn!( + "Could not parse completeness value for '{genome_id}': '{completeness_str}' — skipping" + ); + return None; + }; + if !(0.0..=1.0).contains(&completeness) { + out_of_range + .lock() + .unwrap() + .push(format!("{genome_id}: {completeness}")); + return None; + } + if let Some(index) = sketches.get_sample_index(genome_id) { + Some((index, completeness)) } else { + not_in_sketch.lock().unwrap().push(genome_id.to_string()); None } }) .collect(); - // Apply updates to completeness_vec (thread safe) + // Out-of-range values are most likely percentages (0–100) instead of fractions (0–1.0) + let invalid = out_of_range.into_inner().unwrap(); + if !invalid.is_empty() { + anyhow::bail!( + "Completeness values must be in [0.0, 1.0], not percentages. \ + Found {} out-of-range value(s) in {completeness_file}:\n {}", + invalid.len(), + invalid.join("\n ") + ); + } + for (index, completeness) in updates { completeness_vec[index] = completeness; } - // Report missing genomes - let missing = missing_genomes.into_inner().unwrap(); - if !missing.is_empty() { + let not_found = not_in_sketch.into_inner().unwrap(); + if !not_found.is_empty() { log::warn!( - "Found {} genomes not in completeness file, using default 1.0: {}", - missing.len(), - missing.join(", ") + "{} genome(s) in completeness file not found in sketch database (ignored): {}", + not_found.len(), + not_found.join(", ") ); } diff --git a/src/lib.rs b/src/lib.rs index ce1c884..a861f8b 100644 --- a/src/lib.rs +++ b/src/lib.rs @@ -257,12 +257,13 @@ pub fn main() -> Result<(), Error> { ref_db, query_db, output, - mut knn, + knn, subset, kmer, ani, threads, - completeness_file, + ref_completeness_file, + query_completeness_file, completeness_cutoff, } => { check_and_set_threads(*threads); @@ -283,18 +284,12 @@ pub fn main() -> Result<(), Error> { } log::info!("Read reference sketches:\n{references:?}"); let n = references.number_samples_loaded(); - if let Some(nn) = knn { - if nn >= n { - log::warn!("knn={nn} is higher than number of samples={n}"); - knn = Some(n - 1); - } - } - // Read in completeness (parallel implementation) - let completeness_vec: Option> = if let Some(file_path) = completeness_file { - Some(read_completeness_file(file_path, &references)?) - } else { - None - }; + let ref_completeness_vec: Option> = + if let Some(file_path) = ref_completeness_file { + Some(read_completeness_file(file_path, &references)?) + } else { + None + }; let dist_type = set_k(&references, *kmer, *ani).unwrap_or_else(|e| { panic!("Error setting k size: {e}"); @@ -325,15 +320,19 @@ pub fn main() -> Result<(), Error> { n, dist_type, args.quiet, - completeness_vec.as_ref(), + ref_completeness_vec.as_ref(), *completeness_cutoff, ); log::info!("Writing out in long matrix form"); write!(output_file, "{distances}") .expect("Error writing output distances"); } - Some(nn) => { - // Self mode (sparse) + Some(mut nn) => { + // Self mode (sparse): a genome cannot be its own neighbour + if nn >= n { + log::warn!("knn={nn} is higher than number of samples={n}"); + nn = n - 1; + } log::info!("Calculating sparse ref vs ref distances with {nn} nearest neighbours"); let distances = self_dists_knn( &references, @@ -341,7 +340,7 @@ pub fn main() -> Result<(), Error> { nn, dist_type, args.quiet, - completeness_vec.as_ref(), + ref_completeness_vec.as_ref(), *completeness_cutoff, ); @@ -352,23 +351,55 @@ pub fn main() -> Result<(), Error> { } } Some(query_db) => { - // Ref v query mode - log::info!("Calculating all ref vs query distances"); - - let nq = query_db.number_samples_loaded(); - let distances = self_query_dists_all( - &references, - &query_db, - n, - nq, - dist_type, - args.quiet, - completeness_vec.as_ref(), - *completeness_cutoff, - ); - - log::info!("Writing out in long matrix form"); - write!(output_file, "{distances}").expect("Error writing output distances"); + let query_completeness_vec: Option> = + if let Some(file_path) = query_completeness_file { + Some(read_completeness_file(file_path, &query_db)?) + } else { + None + }; + let n_query = query_db.number_samples_loaded(); + match knn { + Some(mut nn) => { + // Cross-query mode: query genomes never overlap ref genomes, so knn=n is valid + if nn > n { + log::warn!("knn={nn} is higher than number of reference samples={n}"); + nn = n; + } + // Cross-query mode (sparse kNN) + log::info!("Calculating sparse ref vs query distances with {nn} nearest neighbours"); + let distances = cross_dists_knn( + &references, + &query_db, + n, + n_query, + nn, + dist_type, + args.quiet, + ref_completeness_vec.as_ref(), + query_completeness_vec.as_ref(), + *completeness_cutoff, + ); + log::info!("Writing out in sparse matrix form"); + write!(output_file, "{distances}").expect("Error writing output distances"); + } + None => { + // Cross-query mode (dense, all pairs) + log::info!("Calculating all ref vs query distances"); + let distances = cross_dists_all( + &references, + &query_db, + n, + n_query, + dist_type, + args.quiet, + ref_completeness_vec.as_ref(), + query_completeness_vec.as_ref(), + *completeness_cutoff, + ); + log::info!("Writing out in long matrix form"); + write!(output_file, "{distances}").expect("Error writing output distances"); + } + } } } Ok(()) @@ -552,7 +583,7 @@ pub fn main() -> Result<(), Error> { mut knn, ani, threads, - completeness_file, + ref_completeness_file, completeness_cutoff, } => { check_and_set_threads(*threads); @@ -613,9 +644,8 @@ pub fn main() -> Result<(), Error> { panic!("K-mer size {kmer} used for .ski not found in .skd: {e}"); }); - // Read in completeness (parallel implementation) - let completeness_vec: Option> = - if let Some(file_path) = completeness_file { + let ref_completeness_vec: Option> = + if let Some(file_path) = ref_completeness_file { Some(read_completeness_file(file_path, &references)?) } else { None @@ -638,7 +668,7 @@ pub fn main() -> Result<(), Error> { knn, dist_type, args.quiet, - completeness_vec.as_ref(), + ref_completeness_vec.as_ref(), *completeness_cutoff, ); diff --git a/tests/completeness.rs b/tests/completeness.rs index 8ca7a0f..c219413 100644 --- a/tests/completeness.rs +++ b/tests/completeness.rs @@ -93,7 +93,7 @@ mod tests { .arg("dist") .arg("test_genomes") .args(["-o", "distances_with_correction_default_cutoff"]) - .arg("--completeness-file") + .arg("--ref-completeness-file") .arg("completeness_cutoff_test.txt") .args(["-k", "31"]) .arg("-v") @@ -107,7 +107,7 @@ mod tests { .arg("dist") .arg("test_genomes") .args(["-o", "distances_with_correction_high_cutoff"]) - .arg("--completeness-file") + .arg("--ref-completeness-file") .arg("completeness_cutoff_test.txt") .arg("--completeness-cutoff") .arg("0.8") @@ -279,7 +279,7 @@ mod tests { .arg("dist") .arg("test_missing") .args(["-k", "21", "-o", "distances_missing"]) - .arg("--completeness-file") + .arg("--ref-completeness-file") .arg("completeness_missing.txt") .arg("-v") .assert() @@ -348,7 +348,7 @@ mod tests { .arg("dist") .arg("test_extra") .args(["-k", "21", "-o", "distances_extra"]) - .arg("--completeness-file") + .arg("--ref-completeness-file") .arg("completeness_extra.txt") .arg("-v") .assert() @@ -433,7 +433,7 @@ mod tests { .arg("--skd") .arg("precluster_sketches") .args(["-v", "--knn", "2"]) - .arg("--completeness-file") + .arg("--ref-completeness-file") .arg("completeness_precluster.txt") .args(["-o", "precluster_distances"]) .assert() @@ -512,7 +512,7 @@ mod tests { .arg("dist") .arg("formula_test") .args(["-o", "distances_corrected"]) - .arg("--completeness-file") + .arg("--ref-completeness-file") .arg("completeness_formula.txt") .args(["-k", "31"]) .arg("-v") diff --git a/tests/distance.rs b/tests/distance.rs index 047a9b1..9499e74 100644 --- a/tests/distance.rs +++ b/tests/distance.rs @@ -1,10 +1,10 @@ use snapbox::cmd::{cargo_bin, Command}; +use std::collections::{HashMap, HashSet}; use std::path::Path; pub mod common; use crate::common::*; -use std::collections::HashMap; use std::fs::File; use std::io::{BufRead, BufReader}; @@ -328,6 +328,364 @@ mod tests { .stdout_eq(sandbox.snapbox_file("dists_knn_ani.stdout", TestDir::Correct)); } + /// Helper: parse ANI distance output lines into (query, reference, ani) triples. + fn parse_dist_output(stdout: &str) -> Vec<(String, String, f64)> { + stdout + .lines() + .filter(|l| !l.is_empty()) + .map(|line| { + let parts: Vec<&str> = line.split_whitespace().collect(); + assert!(parts.len() >= 3, "Unexpected dist output line: {line}"); + ( + parts[0].to_string(), + parts[1].to_string(), + parts[2].parse::().expect("Could not parse ANI"), + ) + }) + .collect() + } + + /// Sketch databases into the sandbox: + /// - `bact_db`: 14412_3#82 + 14412_3#84 (disjoint from query genomes) + /// - `query_db`: R6 + TIGR4 + /// - `ref_db`: all 4 genomes (used for self-kNN consistency test) + fn sketch_ref_and_query(sandbox: &TestSetup) { + for f in &[ + "14412_3#82.contigs_velvet.fa.gz", + "14412_3#84.contigs_velvet.fa.gz", + "R6.fa.gz", + "TIGR4.fa.gz", + "rfile.txt", + "rfile_ref.txt", + "qfile.txt", + ] { + sandbox.copy_input_file_to_wd(f); + } + + // Disjoint reference: 14412_3#82 + 14412_3#84 only + Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args(["sketch", "-f", "rfile_ref.txt", "--k-seq", "17,31,4", "-s", "1000", "-o", "bact_db"]) + .assert() + .success(); + + // Query: R6 + TIGR4 + Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args(["sketch", "-f", "qfile.txt", "--k-seq", "17,31,4", "-s", "1000", "-o", "query_db"]) + .assert() + .success(); + + // All 4 genomes — used for self-kNN consistency test only + Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args(["sketch", "-f", "rfile.txt", "--k-seq", "17,31,4", "-s", "1000", "-o", "ref_db"]) + .assert() + .success(); + } + + /// Test 1: output has exactly nq × knn rows. + /// + /// 2 query genomes × knn=2 → 4 rows. + #[test] + fn knn_cross_query_row_count() { + let sandbox = TestSetup::setup(); + sketch_ref_and_query(&sandbox); + + let output = std::process::Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args(["dist", "bact_db", "query_db", "--knn", "1", "-k", "21", "--ani"]) + .output() + .expect("Failed to run sketchlib"); + + let stdout = String::from_utf8(output.stdout).unwrap(); + let n_lines = stdout.lines().filter(|l| !l.is_empty()).count(); + assert_eq!(n_lines, 2, "Expected 2 queries × 1 neighbour = 2 rows, got {n_lines}"); + } + + /// Test 2: kNN output contains the same top-k neighbours as the dense output sorted by ANI. + /// + /// Runs both dense and kNN cross-query, then verifies that for each query genome + /// the kNN output matches the top-2 hits from the dense output ranked by ANI. + #[test] + fn knn_cross_query_matches_dense_top_k() { + let sandbox = TestSetup::setup(); + sketch_ref_and_query(&sandbox); + + // Dense: all bact_ref × query pairs (format: ref \t query \t ani) + let dense_out = std::process::Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args(["dist", "bact_db", "query_db", "-k", "21", "--ani"]) + .output() + .expect("Failed to run dense dist"); + let dense_stdout = String::from_utf8(dense_out.stdout).unwrap(); + + // Dense output columns: ref(0) \t query(1) \t ani(2) — group by query genome + let mut dense_by_query: HashMap> = HashMap::new(); + for line in dense_stdout.lines().filter(|l| !l.is_empty()) { + let parts: Vec<&str> = line.split_whitespace().collect(); + assert!(parts.len() >= 3, "Unexpected dense output line: {line}"); + let reference = parts[0].to_string(); + let query = parts[1].to_string(); + let ani: f64 = parts[2].parse().expect("Could not parse ANI"); + dense_by_query.entry(query).or_default().push((reference, ani)); + } + + // For each query, sort by ANI descending and keep top-1 + let mut dense_top1: HashMap = HashMap::new(); + for (query, mut hits) in dense_by_query { + hits.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal)); + dense_top1.insert(query, hits.into_iter().next().unwrap()); + } + + // kNN output columns: query(0) \t ref(1) \t ani(2) + let knn_out = std::process::Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args(["dist", "bact_db", "query_db", "--knn", "1", "-k", "21", "--ani"]) + .output() + .expect("Failed to run kNN dist"); + let knn_triples = parse_dist_output(&String::from_utf8(knn_out.stdout).unwrap()); + + // Assert kNN top-1 matches dense top-1 for every query genome + for (query, reference, knn_ani) in &knn_triples { + let (dense_ref, dense_ani) = dense_top1.get(query) + .unwrap_or_else(|| panic!("Query {query} not found in dense output")); + assert_eq!(reference, dense_ref, + "Top neighbour mismatch for {query}: knn={reference}, dense={dense_ref}"); + assert!((knn_ani - dense_ani).abs() < 1e-5, + "ANI mismatch for {query}/{reference}: knn={knn_ani}, dense={dense_ani}"); + } + } + + /// Test 3: every row name is a query genome and every column name is a reference genome. + /// + /// Verifies that the Display impl uses row_names (query) for rows and ref_names + /// (reference) for the neighbour column — not ref_names for both. + #[test] + fn knn_cross_query_correct_name_columns() { + let sandbox = TestSetup::setup(); + sketch_ref_and_query(&sandbox); + + let output = std::process::Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args(["dist", "bact_db", "query_db", "--knn", "1", "-k", "21", "--ani"]) + .output() + .expect("Failed to run sketchlib"); + + let stdout = String::from_utf8(output.stdout).unwrap(); + let triples = parse_dist_output(&stdout); + + let query_names: HashSet<&str> = ["R6.fa.gz", "TIGR4.fa.gz"].iter().cloned().collect(); + let ref_names: HashSet<&str> = [ + "14412_3#82.contigs_velvet.fa.gz", + "14412_3#84.contigs_velvet.fa.gz", + ].iter().cloned().collect(); + + for (query, reference, _) in &triples { + assert!(query_names.contains(query.as_str()), + "Row '{query}' is not a query genome"); + assert!(ref_names.contains(reference.as_str()), + "Column '{reference}' is not a reference genome"); + } + } + + /// Test 4: cross-query kNN distances are consistent with self-query kNN. + /// + /// Runs self-kNN on all 4 genomes. For R6 and TIGR4, extracts their nearest + /// neighbours from the self-kNN output restricted to the 2 bacterial reference + /// genomes (14412_3#82, 14412_3#84). Then runs cross-query kNN with those 2 + /// genomes as reference and R6/TIGR4 as query. Both should give the same distances. + #[test] + fn knn_cross_query_consistent_with_self_knn() { + let sandbox = TestSetup::setup(); + sketch_ref_and_query(&sandbox); + + // Self-kNN on all 4 genomes with knn=3 + let self_out = std::process::Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args(["dist", "ref_db", "--knn", "3", "-k", "21", "--ani"]) + .output() + .expect("Failed to run self kNN dist"); + let self_triples = parse_dist_output(&String::from_utf8(self_out.stdout).unwrap()); + + // Cross-query kNN=1: query=R6+TIGR4 against ref=14412_3#82+14412_3#84 + // kNN output columns: query(0) \t ref(1) \t ani(2) + let cross_out = std::process::Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args(["dist", "bact_db", "query_db", "--knn", "1", "-k", "21", "--ani"]) + .output() + .expect("Failed to run cross-query kNN dist"); + let cross_triples = parse_dist_output(&String::from_utf8(cross_out.stdout).unwrap()); + + // From self-kNN output, extract the best bacterial hit for R6 and TIGR4. + // Self-kNN output columns: query(0) \t neighbour(1) \t ani(2) + let bact_names: HashSet<&str> = [ + "14412_3#82.contigs_velvet.fa.gz", + "14412_3#84.contigs_velvet.fa.gz", + ].iter().cloned().collect(); + + let mut self_best_bact: HashMap = HashMap::new(); + for query in &["R6.fa.gz", "TIGR4.fa.gz"] { + let best = self_triples.iter() + .filter(|(q, r, _)| q == query && bact_names.contains(r.as_str())) + .max_by(|a, b| a.2.partial_cmp(&b.2).unwrap_or(std::cmp::Ordering::Equal)) + .unwrap_or_else(|| panic!("No bacterial hits for {query} in self-kNN")); + self_best_bact.insert(query.to_string(), (best.1.clone(), best.2)); + } + + // Compare: cross-query kNN top-1 should match self-kNN top-1 from bacterial genomes + for (query, cross_ref, cross_ani) in &cross_triples { + let (self_ref, self_ani) = self_best_bact.get(query) + .unwrap_or_else(|| panic!("Query {query} not found in self-kNN")); + assert_eq!(cross_ref, self_ref, + "Nearest bacterial genome mismatch for {query}: cross={cross_ref}, self={self_ref}"); + assert!((cross_ani - self_ani).abs() < 1e-5, + "ANI mismatch for {query}/{cross_ref}: cross={cross_ani}, self={self_ani}"); + } + } + + /// Test 5: cross-query kNN with ref and query completeness files. + /// + /// Runs cross-query kNN twice — without and with completeness correction. + /// Verifies the command succeeds with both completeness flags, output has + /// the correct row count, ANI values are in [0.0, 1.0], and that out-of-range + /// completeness values (percentages instead of fractions) are rejected. + #[test] + fn knn_cross_query_completeness() { + let sandbox = TestSetup::setup(); + sketch_ref_and_query(&sandbox); + + TestSetup::create_completeness_file( + &sandbox, + "ref_completeness.txt", + &[ + ("14412_3#82.contigs_velvet.fa.gz", 0.8), + ("14412_3#84.contigs_velvet.fa.gz", 0.85), + ], + ); + TestSetup::create_completeness_file( + &sandbox, + "query_completeness.txt", + &[("R6.fa.gz", 0.9), ("TIGR4.fa.gz", 0.75)], + ); + + // Both completeness flags accepted; output has correct row count and valid ANI range + let comp_out = std::process::Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args([ + "dist", "bact_db", "query_db", + "--knn", "1", "-k", "21", "--ani", + "--ref-completeness-file", "ref_completeness.txt", + "--query-completeness-file", "query_completeness.txt", + ]) + .output() + .expect("Failed to run with completeness"); + assert!( + comp_out.status.success(), + "Command failed: {}", + String::from_utf8_lossy(&comp_out.stderr) + ); + let comp_triples = parse_dist_output(&String::from_utf8(comp_out.stdout).unwrap()); + assert_eq!(comp_triples.len(), 2, "Expected 2 rows (2 queries × knn=1)"); + for (query, ref_name, ani) in &comp_triples { + assert!( + (0.0..=1.0).contains(ani), + "ANI out of range for query={query} ref={ref_name}: {ani}" + ); + } + + // Out-of-range completeness values (percentages) must be rejected + TestSetup::create_completeness_file( + &sandbox, + "bad_completeness.txt", + &[ + ("14412_3#82.contigs_velvet.fa.gz", 80.0), + ("14412_3#84.contigs_velvet.fa.gz", 85.0), + ], + ); + let bad_out = std::process::Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args([ + "dist", "bact_db", "query_db", + "--knn", "1", "-k", "21", "--ani", + "--ref-completeness-file", "bad_completeness.txt", + ]) + .output() + .expect("Failed to run with bad completeness"); + assert!( + !bad_out.status.success(), + "Expected failure for out-of-range completeness values" + ); + let stderr = String::from_utf8_lossy(&bad_out.stderr); + assert!( + stderr.contains("[0.0, 1.0]"), + "Error message should mention [0.0, 1.0] range, got: {stderr}" + ); + } + + /// Test 6: cross-query kNN in CoreAcc mode (no -k flag). + /// + /// Verifies that cross-query kNN works without a k-mer length, producing + /// 4-column output (query, ref, core, acc). + #[test] + fn knn_cross_query_core_acc() { + let sandbox = TestSetup::setup(); + sketch_ref_and_query(&sandbox); + + let output = std::process::Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args(["dist", "bact_db", "query_db", "--knn", "1"]) + .output() + .expect("Failed to run CoreAcc cross-query kNN"); + + assert!( + output.status.success(), + "CoreAcc cross-query failed: {}", + String::from_utf8_lossy(&output.stderr) + ); + + let stdout = String::from_utf8(output.stdout).unwrap(); + let lines: Vec<&str> = stdout.lines().filter(|l| !l.is_empty()).collect(); + assert_eq!(lines.len(), 2, "Expected 2 rows (2 queries × knn=1), got {}", lines.len()); + + for line in &lines { + let parts: Vec<&str> = line.split_whitespace().collect(); + assert_eq!(parts.len(), 4, "Expected 4 columns (query, ref, core, acc): {line}"); + parts[2].parse::().expect("Core distance not a float"); + parts[3].parse::().expect("Acc distance not a float"); + } + } + + /// Test 7: cross-query kNN with knn equal to the number of reference genomes. + /// + /// bact_db has n=2 reference genomes. With knn=2 every query genome should + /// get all 2 reference genomes as neighbours (2 queries × 2 = 4 rows). + /// Previously Bug 2 silently clamped knn=n to knn=n-1, giving only 2 rows. + #[test] + fn knn_cross_query_knn_equals_n_ref() { + let sandbox = TestSetup::setup(); + sketch_ref_and_query(&sandbox); + + let output = std::process::Command::new(cargo_bin("sketchlib")) + .current_dir(sandbox.get_wd()) + .args(["dist", "bact_db", "query_db", "--knn", "2", "-k", "21", "--ani"]) + .output() + .expect("Failed to run knn=n_ref cross-query"); + + assert!( + output.status.success(), + "Command failed: {}", + String::from_utf8_lossy(&output.stderr) + ); + + let stdout = String::from_utf8(output.stdout).unwrap(); + let n_lines = stdout.lines().filter(|l| !l.is_empty()).count(); + assert_eq!( + n_lines, 4, + "Expected 2 queries × 2 neighbours = 4 rows, got {n_lines}" + ); + } + #[test] fn subset_dists() { let sandbox = TestSetup::setup(); diff --git a/tests/test_files_in/qfile.txt b/tests/test_files_in/qfile.txt new file mode 100644 index 0000000..0a07145 --- /dev/null +++ b/tests/test_files_in/qfile.txt @@ -0,0 +1,2 @@ +R6.fa.gz R6.fa.gz +TIGR4.fa.gz TIGR4.fa.gz diff --git a/tests/test_files_in/rfile_ref.txt b/tests/test_files_in/rfile_ref.txt new file mode 100644 index 0000000..e039d21 --- /dev/null +++ b/tests/test_files_in/rfile_ref.txt @@ -0,0 +1,2 @@ +14412_3#82.contigs_velvet.fa.gz 14412_3#82.contigs_velvet.fa.gz +14412_3#84.contigs_velvet.fa.gz 14412_3#84.contigs_velvet.fa.gz From 40c37d4a90914d885f1727b7f1406f87ae9226a2 Mon Sep 17 00:00:00 2001 From: Johanna Date: Fri, 29 May 2026 16:26:29 +0100 Subject: [PATCH 2/3] fixed clippy warning --- src/io.rs | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/src/io.rs b/src/io.rs index 426748c..5c42fcf 100644 --- a/src/io.rs +++ b/src/io.rs @@ -230,9 +230,7 @@ pub fn read_completeness_file( let updates: Vec<(usize, f64)> = lines .par_iter() .filter_map(|line| { - let Some((genome_id, completeness_str)) = line.split_once('\t') else { - return None; - }; + let (genome_id, completeness_str) = line.split_once('\t')?; let Ok(completeness) = completeness_str.trim().parse::() else { log::warn!( "Could not parse completeness value for '{genome_id}': '{completeness_str}' — skipping" From f7ffb94cfcb192078e6afaf4a22782e47ff6aba7 Mon Sep 17 00:00:00 2001 From: Johanna Date: Tue, 16 Jun 2026 12:16:07 +0100 Subject: [PATCH 3/3] Updated codecov version --- .github/workflows/codecov.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/codecov.yml b/.github/workflows/codecov.yml index 808bb34..aefdccb 100644 --- a/.github/workflows/codecov.yml +++ b/.github/workflows/codecov.yml @@ -32,7 +32,7 @@ jobs: - name: Codecov # You may pin to the exact commit or the version. - uses: codecov/codecov-action@v5.1.2 + uses: codecov/codecov-action@v4.6.0 with: # Repository upload token - get it from codecov.io. Required only for private repositories token: ${{ secrets.CODECOV_TOKEN }}