Push zpwxxpnpktps #67
@@ -1,578 +0,0 @@
|
||||
use std::fs::{self, File};
|
||||
use std::io::{self, BufWriter, Read as _, Write as _};
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use memmap2::Mmap;
|
||||
use ndarray::{Array1, Array2};
|
||||
use rayon::prelude::*;
|
||||
|
||||
use crate::bitvec::{PersistentBitVec, PersistentBitVecBuilder};
|
||||
use crate::colgroup::{ColGroup, MatrixGroupOps};
|
||||
use crate::layer_meta::LayerMeta;
|
||||
use crate::meta::MatrixMeta;
|
||||
use crate::tempbitvec::{TempBitVec, TempBitVecBuilder};
|
||||
use crate::tempintvec::{TempCompactIntVec, TempCompactIntVecBuilder};
|
||||
use crate::views::BitSliceView;
|
||||
|
||||
fn col_path(dir: &Path, col: usize) -> PathBuf {
|
||||
dir.join(format!("col_{col:06}.pbiv"))
|
||||
}
|
||||
|
||||
// ── ColumnarBitMatrix ─────────────────────────────────────────────────────────
|
||||
|
||||
/// Per-column file layout (original format).
|
||||
pub struct ColumnarBitMatrix {
|
||||
cols: Vec<PersistentBitVec>,
|
||||
n: usize,
|
||||
}
|
||||
|
||||
impl ColumnarBitMatrix {
|
||||
pub(crate) fn open(dir: &Path) -> io::Result<Self> {
|
||||
let meta = MatrixMeta::load(dir)?;
|
||||
let cols = (0..meta.n_cols)
|
||||
.map(|c| PersistentBitVec::open(&col_path(dir, c)))
|
||||
.collect::<io::Result<Vec<_>>>()?;
|
||||
Ok(Self { cols, n: meta.n })
|
||||
}
|
||||
|
||||
pub(crate) fn n(&self) -> usize { self.n }
|
||||
pub(crate) fn n_cols(&self) -> usize { self.cols.len() }
|
||||
pub(crate) fn col(&self, c: usize) -> &PersistentBitVec { &self.cols[c] }
|
||||
|
||||
pub(crate) fn row(&self, slot: usize) -> Box<[bool]> {
|
||||
self.cols.iter().map(|c| c.get(slot)).collect()
|
||||
}
|
||||
|
||||
pub(crate) fn fill_row(&self, slot: usize, buf: &mut [u32]) {
|
||||
for (c, col) in self.cols.iter().enumerate() {
|
||||
buf[c] = col.get(slot) as u32;
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn count_ones(&self) -> Array1<u64> {
|
||||
let counts: Vec<u64> = (0..self.n_cols())
|
||||
.into_par_iter()
|
||||
.map(|c| self.col(c).count_ones())
|
||||
.collect();
|
||||
Array1::from_vec(counts)
|
||||
}
|
||||
|
||||
pub(crate) fn partial_jaccard_dist_matrix(&self) -> (Array2<u64>, Array2<u64>) {
|
||||
pairwise2_matrix(self.n_cols(), |i, j| self.col(i).partial_jaccard_dist(self.col(j)))
|
||||
}
|
||||
|
||||
pub(crate) fn partial_hamming_dist_matrix(&self) -> Array2<u64> {
|
||||
pairwise_matrix(self.n_cols(), |i, j| self.col(i).hamming_dist(self.col(j)))
|
||||
}
|
||||
|
||||
pub(crate) fn append_column(dir: &Path, value_of: impl Fn(usize) -> bool) -> io::Result<()> {
|
||||
let mut meta = MatrixMeta::load(dir)?;
|
||||
let mut b = PersistentBitVecBuilder::new(meta.n, &col_path(dir, meta.n_cols))?;
|
||||
for slot in 0..meta.n {
|
||||
b.set(slot, value_of(slot));
|
||||
}
|
||||
b.close()?;
|
||||
meta.n_cols += 1;
|
||||
meta.save(dir)
|
||||
}
|
||||
}
|
||||
|
||||
// ── PackedBitMatrix ───────────────────────────────────────────────────────────
|
||||
|
||||
const PBMX_MAGIC: [u8; 4] = *b"PBMX";
|
||||
const PBMX_HEADER: usize = 24; // magic(4) + pad(4) + n_rows(8) + n_cols(8)
|
||||
const PBIV_HEADER: usize = 16; // magic(4) + pad(4) + n(8)
|
||||
|
||||
/// Single-file packed layout: all columns concatenated behind a header.
|
||||
pub struct PackedBitMatrix {
|
||||
mmap: Mmap,
|
||||
n_rows: usize,
|
||||
n_cols: usize,
|
||||
/// Absolute byte offset to the start of each column's bit data
|
||||
/// (= file offset of the PBIV blob + PBIV_HEADER).
|
||||
data_offsets: Vec<usize>,
|
||||
}
|
||||
|
||||
impl PackedBitMatrix {
|
||||
pub(crate) fn open(path: &Path) -> io::Result<Self> {
|
||||
let mmap = unsafe { Mmap::map(&File::open(path)?)? };
|
||||
if mmap.len() < PBMX_HEADER {
|
||||
return Err(io::Error::new(io::ErrorKind::InvalidData, "PBMX file too short"));
|
||||
}
|
||||
if &mmap[0..4] != &PBMX_MAGIC {
|
||||
return Err(io::Error::new(io::ErrorKind::InvalidData, "bad PBMX magic"));
|
||||
}
|
||||
let n_rows = u64::from_le_bytes(mmap[8..16].try_into().unwrap()) as usize;
|
||||
let n_cols = u64::from_le_bytes(mmap[16..24].try_into().unwrap()) as usize;
|
||||
|
||||
let mut data_offsets = Vec::with_capacity(n_cols);
|
||||
for c in 0..n_cols {
|
||||
let off_pos = PBMX_HEADER + c * 8;
|
||||
let col_file_off = u64::from_le_bytes(mmap[off_pos..off_pos+8].try_into().unwrap()) as usize;
|
||||
data_offsets.push(col_file_off + PBIV_HEADER);
|
||||
}
|
||||
|
||||
Ok(Self { mmap, n_rows, n_cols, data_offsets })
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub(crate) fn fill_row(&self, slot: usize, buf: &mut [u32]) {
|
||||
for (c, &data_off) in self.data_offsets.iter().enumerate() {
|
||||
buf[c] = ((self.mmap[data_off + (slot >> 3)] >> (slot & 7)) & 1) as u32;
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn row(&self, slot: usize) -> Box<[bool]> {
|
||||
(0..self.n_cols).map(|c| {
|
||||
(self.mmap[self.data_offsets[c] + (slot >> 3)] >> (slot & 7)) & 1 != 0
|
||||
}).collect()
|
||||
}
|
||||
|
||||
fn col_bytes(&self, c: usize) -> &[u8] {
|
||||
let start = self.data_offsets[c];
|
||||
&self.mmap[start..start + self.n_rows.div_ceil(8)]
|
||||
}
|
||||
|
||||
fn col_words(&self, c: usize) -> &[u64] {
|
||||
let nw = self.n_rows.div_ceil(64);
|
||||
// SAFETY: data_offsets[c] is always 8-byte aligned.
|
||||
// PBMX header = 24 + n_cols×8 (multiple of 8); each PBIV blob =
|
||||
// 16 + nwords×8 (multiple of 8); mmap base is page-aligned.
|
||||
let ptr = self.mmap[self.data_offsets[c]..].as_ptr() as *const u64;
|
||||
unsafe { std::slice::from_raw_parts(ptr, nw) }
|
||||
}
|
||||
|
||||
pub(crate) fn col_slice(&self, c: usize) -> BitSliceView<'_> {
|
||||
BitSliceView::new(self.col_words(c), self.n_rows)
|
||||
}
|
||||
|
||||
pub(crate) fn col_persist(&self, c: usize, path: &Path) -> io::Result<PersistentBitVecBuilder> {
|
||||
PersistentBitVecBuilder::from_raw_bytes(self.col_bytes(c), self.n_rows, path)
|
||||
}
|
||||
|
||||
pub(crate) fn count_ones(&self) -> Array1<u64> {
|
||||
Array1::from_vec(
|
||||
(0..self.n_cols).into_par_iter()
|
||||
.map(|c| self.col_slice(c).count_ones())
|
||||
.collect()
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn partial_jaccard_dist_matrix(&self) -> (Array2<u64>, Array2<u64>) {
|
||||
pairwise2_matrix(self.n_cols, |i, j| {
|
||||
self.col_slice(i).partial_jaccard_dist(self.col_slice(j))
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn partial_hamming_dist_matrix(&self) -> Array2<u64> {
|
||||
pairwise_matrix(self.n_cols, |i, j| {
|
||||
self.col_slice(i).hamming_dist(self.col_slice(j))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Reads just the `n_cols` field from an existing packed matrix's header,
|
||||
/// without mapping the file. Used by `pack_bit_matrix` to tell a genuinely
|
||||
/// complete pack from a stale one that predates a later column-widening.
|
||||
fn packed_bit_matrix_n_cols(path: &Path) -> io::Result<usize> {
|
||||
let mut f = File::open(path)?;
|
||||
let mut header = [0u8; PBMX_HEADER];
|
||||
f.read_exact(&mut header)?;
|
||||
Ok(u64::from_le_bytes(header[16..24].try_into().unwrap()) as usize)
|
||||
}
|
||||
|
||||
/// Build `presence/matrix.pbmx` from existing `col_*.pbiv` files.
|
||||
pub fn pack_bit_matrix(dir: &Path) -> io::Result<()> {
|
||||
let packed_path = dir.join("matrix.pbmx");
|
||||
|
||||
let meta = match MatrixMeta::load(dir) {
|
||||
Ok(meta) => meta,
|
||||
Err(e) => {
|
||||
// No columnar data pending: either this layer was already
|
||||
// packed and cleaned up (matrix.pbmx complete, nothing left to
|
||||
// do), or genuinely nothing was ever written here.
|
||||
return if packed_path.exists() { Ok(()) } else { Err(e) };
|
||||
}
|
||||
};
|
||||
|
||||
// A `matrix.pbmx` can already exist here even though columnar data is
|
||||
// still pending — e.g. copied verbatim from a merge's base source
|
||||
// before this layer was widened with more genome columns (see
|
||||
// `obikpartitionner::merge_partition`). Only skip (re-)packing if the
|
||||
// existing file already reflects the current column count; otherwise
|
||||
// the columnar files are newer and must be (re-)packed, overwriting the
|
||||
// stale one — never silently discarded as "leftover cleanup".
|
||||
if packed_bit_matrix_n_cols(&packed_path).ok() == Some(meta.n_cols) {
|
||||
for c in 0..meta.n_cols { let _ = fs::remove_file(col_path(dir, c)); }
|
||||
let _ = fs::remove_file(dir.join("meta.json"));
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let n_cols = meta.n_cols;
|
||||
|
||||
// Compute offsets from file sizes — no column data loaded into RAM.
|
||||
let col_sizes: Vec<u64> = (0..n_cols)
|
||||
.map(|c| fs::metadata(col_path(dir, c)).map(|m| m.len()))
|
||||
.collect::<io::Result<_>>()?;
|
||||
|
||||
let header_size = (PBMX_HEADER + n_cols * 8) as u64;
|
||||
let mut col_offset = header_size;
|
||||
let mut offsets = Vec::with_capacity(n_cols);
|
||||
for &size in &col_sizes {
|
||||
offsets.push(col_offset);
|
||||
col_offset += size;
|
||||
}
|
||||
|
||||
// Write to a temp file; rename atomically so a killed process never leaves
|
||||
// a truncated matrix.pbmx that would be mistaken for a complete file.
|
||||
let tmp_path = dir.join("matrix.pbmx.tmp");
|
||||
let mut out = BufWriter::new(File::create(&tmp_path)?);
|
||||
out.write_all(&PBMX_MAGIC)?;
|
||||
out.write_all(&[0u8; 4])?;
|
||||
out.write_all(&(meta.n as u64).to_le_bytes())?;
|
||||
out.write_all(&(n_cols as u64).to_le_bytes())?;
|
||||
for &off in &offsets { out.write_all(&off.to_le_bytes())?; }
|
||||
for c in 0..n_cols {
|
||||
io::copy(&mut File::open(col_path(dir, c))?, &mut out)?;
|
||||
}
|
||||
out.flush()?;
|
||||
drop(out);
|
||||
fs::rename(&tmp_path, &packed_path)?;
|
||||
|
||||
for c in 0..n_cols { fs::remove_file(col_path(dir, c))?; }
|
||||
fs::remove_file(dir.join("meta.json"))?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ── PersistentBitMatrix — public enum ────────────────────────────────────────
|
||||
|
||||
/// Bit matrix that transparently handles columnar, packed, and implicit formats.
|
||||
///
|
||||
/// - `Columnar`: per-column `.pbiv` files (original format, used during build)
|
||||
/// - `Packed`: single `matrix.pbmx` file (optimised for query — one `mmap`)
|
||||
/// - `Implicit`: no file — all values are 1 (mono-genome presence/absence)
|
||||
pub enum PersistentBitMatrix {
|
||||
Columnar(ColumnarBitMatrix),
|
||||
Packed(PackedBitMatrix),
|
||||
Implicit { n_rows: usize, n_cols: usize },
|
||||
}
|
||||
|
||||
impl PersistentBitMatrix {
|
||||
/// Open from `layer_dir`, auto-detecting the format.
|
||||
///
|
||||
/// Checks (in order):
|
||||
/// 1. `layer_dir/presence/matrix.pbmx` → Packed
|
||||
/// 2. `layer_dir/presence/meta.json` → Columnar
|
||||
/// 3. `layer_dir/layer_meta.json` → Implicit (new index)
|
||||
/// 4. `layer_dir/unitigs.bin` → Implicit with warning (old index)
|
||||
pub fn open(layer_dir: &Path) -> io::Result<Self> {
|
||||
let presence_dir = layer_dir.join("presence");
|
||||
|
||||
if presence_dir.join("matrix.pbmx").exists() {
|
||||
return Ok(Self::Packed(PackedBitMatrix::open(&presence_dir.join("matrix.pbmx"))?));
|
||||
}
|
||||
|
||||
if MatrixMeta::load(&presence_dir).is_ok() {
|
||||
return Ok(Self::Columnar(ColumnarBitMatrix::open(&presence_dir)?));
|
||||
}
|
||||
|
||||
// No presence matrix → Implicit; requires layer_meta.json
|
||||
let meta = LayerMeta::load(layer_dir).map_err(|_| io::Error::new(
|
||||
io::ErrorKind::NotFound,
|
||||
format!(
|
||||
"no presence matrix and no layer_meta.json in {} — run 'obikmer upgrade'",
|
||||
layer_dir.display()
|
||||
),
|
||||
))?;
|
||||
Ok(Self::Implicit { n_rows: meta.n, n_cols: 1 })
|
||||
}
|
||||
|
||||
pub fn n(&self) -> usize {
|
||||
match self {
|
||||
Self::Columnar(m) => m.n(),
|
||||
Self::Packed(m) => m.n_rows,
|
||||
Self::Implicit { n_rows, .. } => *n_rows,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn n_cols(&self) -> usize {
|
||||
match self {
|
||||
Self::Columnar(m) => m.n_cols(),
|
||||
Self::Packed(m) => m.n_cols,
|
||||
Self::Implicit { n_cols, .. } => *n_cols,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn col(&self, c: usize) -> &PersistentBitVec {
|
||||
match self {
|
||||
Self::Columnar(m) => m.col(c),
|
||||
_ => panic!("col() only available on Columnar PersistentBitMatrix"),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn col_view(&self, c: usize) -> BitSliceView<'_> {
|
||||
match self {
|
||||
Self::Columnar(m) => m.col(c).view(),
|
||||
Self::Packed(m) => m.col_slice(c),
|
||||
Self::Implicit { .. } => panic!("col_view() not available on Implicit PersistentBitMatrix"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Column-major point lookup: value at column `c`, slot `slot`, as 0/1.
|
||||
///
|
||||
/// Unlike [`col_view`](Self::col_view), this never panics on `Implicit`
|
||||
/// (every column reads as present, per the mono-genome fast path) — safe
|
||||
/// to call for any `c < self.n_cols()`.
|
||||
pub fn get(&self, c: usize, slot: usize) -> u32 {
|
||||
match self {
|
||||
Self::Columnar(m) => m.col(c).get(slot) as u32,
|
||||
Self::Packed(m) => m.col_slice(c).get(slot) as u32,
|
||||
Self::Implicit { .. } => 1,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn col_persist(&self, c: usize, path: &Path) -> io::Result<PersistentBitVecBuilder> {
|
||||
match self {
|
||||
Self::Columnar(m) => PersistentBitVecBuilder::build_from(m.col(c), path),
|
||||
Self::Packed(m) => m.col_persist(c, path),
|
||||
Self::Implicit { n_rows, .. } => {
|
||||
PersistentBitVecBuilder::new_ones(*n_rows, path)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn row(&self, slot: usize) -> Box<[bool]> {
|
||||
match self {
|
||||
Self::Columnar(m) => m.row(slot),
|
||||
Self::Packed(m) => m.row(slot),
|
||||
Self::Implicit { n_cols, .. } => vec![true; *n_cols].into_boxed_slice(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Fill `buf[i]` with `col_i[slot]` as 0/1 u32, without allocating.
|
||||
pub fn fill_row(&self, slot: usize, buf: &mut [u32]) {
|
||||
match self {
|
||||
Self::Columnar(m) => m.fill_row(slot, buf),
|
||||
Self::Packed(m) => m.fill_row(slot, buf),
|
||||
Self::Implicit { n_cols, .. } => buf[..*n_cols].fill(1),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn count_ones(&self) -> Array1<u64> {
|
||||
match self {
|
||||
Self::Columnar(m) => m.count_ones(),
|
||||
Self::Packed(m) => m.count_ones(),
|
||||
Self::Implicit { n_rows, n_cols } => Array1::from_elem(*n_cols, *n_rows as u64),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn partial_jaccard_dist_matrix(&self) -> (Array2<u64>, Array2<u64>) {
|
||||
match self {
|
||||
Self::Columnar(m) => m.partial_jaccard_dist_matrix(),
|
||||
Self::Packed(m) => m.partial_jaccard_dist_matrix(),
|
||||
Self::Implicit { n_rows, n_cols } => {
|
||||
let v = *n_rows as u64;
|
||||
let n = *n_cols;
|
||||
let mut inter = Array2::zeros((n, n));
|
||||
let mut union = Array2::zeros((n, n));
|
||||
for i in 0..n { for j in 0..n {
|
||||
inter[[i, j]] = v; union[[i, j]] = v;
|
||||
}}
|
||||
(inter, union)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn partial_hamming_dist_matrix(&self) -> Array2<u64> {
|
||||
match self {
|
||||
Self::Columnar(m) => m.partial_hamming_dist_matrix(),
|
||||
Self::Packed(m) => m.partial_hamming_dist_matrix(),
|
||||
Self::Implicit { n_cols, .. } => Array2::zeros((*n_cols, *n_cols)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Append a new column to an on-disk Columnar matrix.
|
||||
pub fn append_column(dir: &Path, value_of: impl Fn(usize) -> bool) -> io::Result<()> {
|
||||
ColumnarBitMatrix::append_column(dir, value_of)
|
||||
}
|
||||
}
|
||||
|
||||
// ── Trait impls ───────────────────────────────────────────────────────────────
|
||||
|
||||
use crate::traits::{BitPartials, ColumnWeights};
|
||||
|
||||
impl ColumnWeights for PersistentBitMatrix {
|
||||
fn col_weights(&self) -> Array1<u64> { self.count_ones() }
|
||||
}
|
||||
|
||||
impl BitPartials for PersistentBitMatrix {
|
||||
fn partial_jaccard(&self) -> (Array2<u64>, Array2<u64>) {
|
||||
self.partial_jaccard_dist_matrix()
|
||||
}
|
||||
fn partial_hamming(&self) -> Array2<u64> {
|
||||
self.partial_hamming_dist_matrix()
|
||||
}
|
||||
}
|
||||
|
||||
// ── Builder (unchanged — always builds Columnar) ──────────────────────────────
|
||||
|
||||
pub struct PersistentBitMatrixBuilder {
|
||||
dir: PathBuf,
|
||||
n: usize,
|
||||
n_cols: usize,
|
||||
}
|
||||
|
||||
impl PersistentBitMatrixBuilder {
|
||||
pub fn new(n: usize, dir: &Path) -> io::Result<Self> {
|
||||
fs::create_dir_all(dir)?;
|
||||
Ok(Self { dir: dir.to_path_buf(), n, n_cols: 0 })
|
||||
}
|
||||
|
||||
pub fn n(&self) -> usize { self.n }
|
||||
pub fn n_cols(&self) -> usize { self.n_cols }
|
||||
|
||||
pub fn add_col(&mut self) -> io::Result<PersistentBitVecBuilder> {
|
||||
let path = col_path(&self.dir, self.n_cols);
|
||||
self.n_cols += 1;
|
||||
PersistentBitVecBuilder::new(self.n, &path)
|
||||
}
|
||||
|
||||
pub fn add_col_ones(&mut self) -> io::Result<PersistentBitVecBuilder> {
|
||||
let path = col_path(&self.dir, self.n_cols);
|
||||
self.n_cols += 1;
|
||||
PersistentBitVecBuilder::new_ones(self.n, &path)
|
||||
}
|
||||
|
||||
pub fn add_col_from(&mut self, src: &TempBitVec) -> io::Result<()> {
|
||||
src.make_persistent(&col_path(&self.dir, self.n_cols))?;
|
||||
self.n_cols += 1;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn add_col_from_int(&mut self, src: &TempCompactIntVec) -> io::Result<()> {
|
||||
let path = col_path(&self.dir, self.n_cols);
|
||||
self.n_cols += 1;
|
||||
let mut b = PersistentBitVecBuilder::new(self.n, &path)?;
|
||||
b.or_where(src.view(), |v| v > 0);
|
||||
b.close()
|
||||
}
|
||||
|
||||
pub fn close(self) -> io::Result<()> {
|
||||
MatrixMeta { n: self.n, n_cols: self.n_cols }.save(&self.dir)
|
||||
}
|
||||
}
|
||||
|
||||
// ── MatrixGroupOps ────────────────────────────────────────────────────────────
|
||||
|
||||
impl MatrixGroupOps for PersistentBitMatrix {
|
||||
fn partial_group_presence_count(&self, g: &ColGroup, _threshold: u32) -> io::Result<TempCompactIntVec> {
|
||||
// Bit matrices store 0/1 — threshold is structurally always 1.
|
||||
let n = self.n();
|
||||
if g.indices.len() < 255 {
|
||||
let mut builder = TempCompactIntVecBuilder::new(n)?;
|
||||
for &c in &g.indices {
|
||||
builder.inc_present_fast(self.col_view(c));
|
||||
}
|
||||
builder.freeze()
|
||||
} else {
|
||||
let mut result = TempCompactIntVecBuilder::new(n)?;
|
||||
for chunk in g.indices.chunks(254) {
|
||||
let mut chunk_b = TempCompactIntVecBuilder::new(n)?;
|
||||
for &c in chunk {
|
||||
chunk_b.inc_present_fast(self.col_view(c));
|
||||
}
|
||||
let frozen = chunk_b.freeze()?;
|
||||
result.add(frozen.view());
|
||||
}
|
||||
result.freeze()
|
||||
}
|
||||
}
|
||||
|
||||
fn partial_group_sum(&self, g: &ColGroup) -> io::Result<TempCompactIntVec> {
|
||||
// For bit matrices, sum = count of 1-bits — identical to presence_count.
|
||||
self.partial_group_presence_count(g, 1)
|
||||
}
|
||||
|
||||
fn partial_group_any(&self, g: &ColGroup, _threshold: u32) -> io::Result<TempBitVec> {
|
||||
let n = self.n();
|
||||
let mut result = TempBitVecBuilder::new(n)?;
|
||||
for &c in &g.indices {
|
||||
result.or(self.col_view(c));
|
||||
}
|
||||
result.freeze()
|
||||
}
|
||||
|
||||
fn partial_group_min(&self, g: &ColGroup) -> io::Result<TempCompactIntVec> {
|
||||
// min of 0/1 values = AND: 1 only if ALL columns are 1
|
||||
let n = self.n();
|
||||
let mut result = TempCompactIntVecBuilder::new(n)?;
|
||||
if let Some((&first, rest)) = g.indices.split_first() {
|
||||
result.inc_present_fast(self.col_view(first));
|
||||
for &c in rest { result.mask_with(self.col_view(c)); }
|
||||
}
|
||||
result.freeze()
|
||||
}
|
||||
|
||||
fn partial_group_max(&self, g: &ColGroup) -> io::Result<TempCompactIntVec> {
|
||||
// max of 0/1 values = OR: 1 if any column is 1
|
||||
let any = self.partial_group_any(g, 1)?;
|
||||
let n = any.len();
|
||||
let mut result = TempCompactIntVecBuilder::new(n)?;
|
||||
result.inc_present(any.view());
|
||||
result.freeze()
|
||||
}
|
||||
}
|
||||
|
||||
// ── Shared matrix helpers (also used by intmatrix.rs) ─────────────────────────
|
||||
|
||||
fn upper_pairs(n: usize) -> Vec<(usize, usize)> {
|
||||
(0..n).flat_map(|i| (i + 1..n).map(move |j| (i, j))).collect()
|
||||
}
|
||||
|
||||
fn fill_symmetric<T>(n: usize, vals: impl Iterator<Item = (usize, usize, T, T)>) -> Array2<T>
|
||||
where T: Clone + Default {
|
||||
let mut m = Array2::from_elem((n, n), T::default());
|
||||
for (i, j, vij, vji) in vals { m[[i, j]] = vij; m[[j, i]] = vji; }
|
||||
m
|
||||
}
|
||||
|
||||
/// Compute a symmetric `n×n` matrix in parallel by evaluating `f(i,j)` for
|
||||
/// all upper-triangle pairs, plus `f(i,i)` for the diagonal. `T: Copy` avoids
|
||||
/// the `.clone()` needed for the lower-triangle mirror.
|
||||
///
|
||||
/// The diagonal is *not* generally `T::default()`: for a self-comparison,
|
||||
/// `f(i,i)` is often the column's own weight (e.g. intersection-with-self —
|
||||
/// see `pairwise2_matrix`), not zero. Distance finalisations that need a
|
||||
/// zero diagonal (self-distance) already overwrite it explicitly.
|
||||
pub(crate) fn pairwise_matrix<T>(n: usize, f: impl Fn(usize, usize) -> T + Sync) -> Array2<T>
|
||||
where T: Copy + Default + Send {
|
||||
let results: Vec<(usize, usize, T)> = upper_pairs(n)
|
||||
.into_par_iter().map(|(i, j)| (i, j, f(i, j))).collect();
|
||||
let mut m = fill_symmetric(n, results.into_iter().map(|(i, j, v)| (i, j, v, v)));
|
||||
for i in 0..n { m[[i, i]] = f(i, i); }
|
||||
m
|
||||
}
|
||||
|
||||
/// Same as `pairwise_matrix` but `f` returns two values that fill two
|
||||
/// symmetric matrices simultaneously (e.g. intersection + union for Jaccard).
|
||||
/// The diagonal is `f(i,i)` (e.g. a genome's kmer count intersected with
|
||||
/// itself), not `T::default()` — see `pairwise_matrix` for why that matters.
|
||||
pub(crate) fn pairwise2_matrix<T>(n: usize, f: impl Fn(usize, usize) -> (T, T) + Sync) -> (Array2<T>, Array2<T>)
|
||||
where T: Copy + Default + Send {
|
||||
let results: Vec<(usize, usize, T, T)> = upper_pairs(n)
|
||||
.into_par_iter()
|
||||
.map(|(i, j)| { let (a, b) = f(i, j); (i, j, a, b) })
|
||||
.collect();
|
||||
let mut m0 = Array2::from_elem((n, n), T::default());
|
||||
let mut m1 = Array2::from_elem((n, n), T::default());
|
||||
for (i, j, a, b) in results {
|
||||
m0[[i, j]] = a; m0[[j, i]] = a;
|
||||
m1[[i, j]] = b; m1[[j, i]] = b;
|
||||
}
|
||||
for i in 0..n {
|
||||
let (a, b) = f(i, i);
|
||||
m0[[i, i]] = a;
|
||||
m1[[i, i]] = b;
|
||||
}
|
||||
(m0, m1)
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
use std::fs;
|
||||
use std::io;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use crate::bitvec::PersistentBitVecBuilder;
|
||||
use crate::meta::MatrixMeta;
|
||||
use crate::tempbitvec::TempBitVec;
|
||||
use crate::tempintvec::TempCompactIntVec;
|
||||
|
||||
use super::col_path;
|
||||
|
||||
// ── Builder (unchanged — always builds Columnar) ──────────────────────────────
|
||||
|
||||
pub struct PersistentBitMatrixBuilder {
|
||||
dir: PathBuf,
|
||||
n: usize,
|
||||
n_cols: usize,
|
||||
}
|
||||
|
||||
impl PersistentBitMatrixBuilder {
|
||||
pub fn new(n: usize, dir: &Path) -> io::Result<Self> {
|
||||
fs::create_dir_all(dir)?;
|
||||
Ok(Self { dir: dir.to_path_buf(), n, n_cols: 0 })
|
||||
}
|
||||
|
||||
pub fn n(&self) -> usize { self.n }
|
||||
pub fn n_cols(&self) -> usize { self.n_cols }
|
||||
|
||||
pub fn add_col(&mut self) -> io::Result<PersistentBitVecBuilder> {
|
||||
let path = col_path(&self.dir, self.n_cols);
|
||||
self.n_cols += 1;
|
||||
PersistentBitVecBuilder::new(self.n, &path)
|
||||
}
|
||||
|
||||
pub fn add_col_ones(&mut self) -> io::Result<PersistentBitVecBuilder> {
|
||||
let path = col_path(&self.dir, self.n_cols);
|
||||
self.n_cols += 1;
|
||||
PersistentBitVecBuilder::new_ones(self.n, &path)
|
||||
}
|
||||
|
||||
pub fn add_col_from(&mut self, src: &TempBitVec) -> io::Result<()> {
|
||||
src.make_persistent(&col_path(&self.dir, self.n_cols))?;
|
||||
self.n_cols += 1;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn add_col_from_int(&mut self, src: &TempCompactIntVec) -> io::Result<()> {
|
||||
let path = col_path(&self.dir, self.n_cols);
|
||||
self.n_cols += 1;
|
||||
let mut b = PersistentBitVecBuilder::new(self.n, &path)?;
|
||||
b.or_where(src.view(), |v| v > 0);
|
||||
b.close()
|
||||
}
|
||||
|
||||
pub fn close(self) -> io::Result<()> {
|
||||
MatrixMeta { n: self.n, n_cols: self.n_cols }.save(&self.dir)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,70 @@
|
||||
use std::io;
|
||||
use std::path::Path;
|
||||
|
||||
use ndarray::{Array1, Array2};
|
||||
use rayon::prelude::*;
|
||||
|
||||
use crate::bitvec::{PersistentBitVec, PersistentBitVecBuilder};
|
||||
use crate::meta::MatrixMeta;
|
||||
|
||||
use super::col_path;
|
||||
use super::pairwise::{pairwise_matrix, pairwise2_matrix};
|
||||
|
||||
// ── ColumnarBitMatrix ─────────────────────────────────────────────────────────
|
||||
|
||||
/// Per-column file layout (original format).
|
||||
pub struct ColumnarBitMatrix {
|
||||
cols: Vec<PersistentBitVec>,
|
||||
n: usize,
|
||||
}
|
||||
|
||||
impl ColumnarBitMatrix {
|
||||
pub(crate) fn open(dir: &Path) -> io::Result<Self> {
|
||||
let meta = MatrixMeta::load(dir)?;
|
||||
let cols = (0..meta.n_cols)
|
||||
.map(|c| PersistentBitVec::open(&col_path(dir, c)))
|
||||
.collect::<io::Result<Vec<_>>>()?;
|
||||
Ok(Self { cols, n: meta.n })
|
||||
}
|
||||
|
||||
pub(crate) fn n(&self) -> usize { self.n }
|
||||
pub(crate) fn n_cols(&self) -> usize { self.cols.len() }
|
||||
pub(crate) fn col(&self, c: usize) -> &PersistentBitVec { &self.cols[c] }
|
||||
|
||||
pub(crate) fn row(&self, slot: usize) -> Box<[bool]> {
|
||||
self.cols.iter().map(|c| c.get(slot)).collect()
|
||||
}
|
||||
|
||||
pub(crate) fn fill_row(&self, slot: usize, buf: &mut [u32]) {
|
||||
for (c, col) in self.cols.iter().enumerate() {
|
||||
buf[c] = col.get(slot) as u32;
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn count_ones(&self) -> Array1<u64> {
|
||||
let counts: Vec<u64> = (0..self.n_cols())
|
||||
.into_par_iter()
|
||||
.map(|c| self.col(c).count_ones())
|
||||
.collect();
|
||||
Array1::from_vec(counts)
|
||||
}
|
||||
|
||||
pub(crate) fn partial_jaccard_dist_matrix(&self) -> (Array2<u64>, Array2<u64>) {
|
||||
pairwise2_matrix(self.n_cols(), |i, j| self.col(i).partial_jaccard_dist(self.col(j)))
|
||||
}
|
||||
|
||||
pub(crate) fn partial_hamming_dist_matrix(&self) -> Array2<u64> {
|
||||
pairwise_matrix(self.n_cols(), |i, j| self.col(i).hamming_dist(self.col(j)))
|
||||
}
|
||||
|
||||
pub(crate) fn append_column(dir: &Path, value_of: impl Fn(usize) -> bool) -> io::Result<()> {
|
||||
let mut meta = MatrixMeta::load(dir)?;
|
||||
let mut b = PersistentBitVecBuilder::new(meta.n, &col_path(dir, meta.n_cols))?;
|
||||
for slot in 0..meta.n {
|
||||
b.set(slot, value_of(slot));
|
||||
}
|
||||
b.close()?;
|
||||
meta.n_cols += 1;
|
||||
meta.save(dir)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
use std::io;
|
||||
|
||||
use crate::colgroup::{ColGroup, MatrixGroupOps};
|
||||
use crate::tempbitvec::{TempBitVec, TempBitVecBuilder};
|
||||
use crate::tempintvec::{TempCompactIntVec, TempCompactIntVecBuilder};
|
||||
|
||||
use super::persistent::PersistentBitMatrix;
|
||||
|
||||
// ── MatrixGroupOps ────────────────────────────────────────────────────────────
|
||||
|
||||
impl MatrixGroupOps for PersistentBitMatrix {
|
||||
fn partial_group_presence_count(&self, g: &ColGroup, _threshold: u32) -> io::Result<TempCompactIntVec> {
|
||||
// Bit matrices store 0/1 — threshold is structurally always 1.
|
||||
let n = self.n();
|
||||
if g.indices.len() < 255 {
|
||||
let mut builder = TempCompactIntVecBuilder::new(n)?;
|
||||
for &c in &g.indices {
|
||||
builder.inc_present_fast(self.col_view(c));
|
||||
}
|
||||
builder.freeze()
|
||||
} else {
|
||||
let mut result = TempCompactIntVecBuilder::new(n)?;
|
||||
for chunk in g.indices.chunks(254) {
|
||||
let mut chunk_b = TempCompactIntVecBuilder::new(n)?;
|
||||
for &c in chunk {
|
||||
chunk_b.inc_present_fast(self.col_view(c));
|
||||
}
|
||||
let frozen = chunk_b.freeze()?;
|
||||
result.add(frozen.view());
|
||||
}
|
||||
result.freeze()
|
||||
}
|
||||
}
|
||||
|
||||
fn partial_group_sum(&self, g: &ColGroup) -> io::Result<TempCompactIntVec> {
|
||||
// For bit matrices, sum = count of 1-bits — identical to presence_count.
|
||||
self.partial_group_presence_count(g, 1)
|
||||
}
|
||||
|
||||
fn partial_group_any(&self, g: &ColGroup, _threshold: u32) -> io::Result<TempBitVec> {
|
||||
let n = self.n();
|
||||
let mut result = TempBitVecBuilder::new(n)?;
|
||||
for &c in &g.indices {
|
||||
result.or(self.col_view(c));
|
||||
}
|
||||
result.freeze()
|
||||
}
|
||||
|
||||
fn partial_group_min(&self, g: &ColGroup) -> io::Result<TempCompactIntVec> {
|
||||
// min of 0/1 values = AND: 1 only if ALL columns are 1
|
||||
let n = self.n();
|
||||
let mut result = TempCompactIntVecBuilder::new(n)?;
|
||||
if let Some((&first, rest)) = g.indices.split_first() {
|
||||
result.inc_present_fast(self.col_view(first));
|
||||
for &c in rest { result.mask_with(self.col_view(c)); }
|
||||
}
|
||||
result.freeze()
|
||||
}
|
||||
|
||||
fn partial_group_max(&self, g: &ColGroup) -> io::Result<TempCompactIntVec> {
|
||||
// max of 0/1 values = OR: 1 if any column is 1
|
||||
let any = self.partial_group_any(g, 1)?;
|
||||
let n = any.len();
|
||||
let mut result = TempCompactIntVecBuilder::new(n)?;
|
||||
result.inc_present(any.view());
|
||||
result.freeze()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
//! Bit matrices (presence/absence), in three on-disk formats transparently
|
||||
//! handled by [`PersistentBitMatrix`]: per-column `Columnar`, single-file
|
||||
//! mmap'd `Packed`, and implicit (mono-genome, no file at all).
|
||||
//!
|
||||
//! Submodules: [`columnar`] (build-time per-column format), [`packed`]
|
||||
//! (query-optimised single-mmap format + [`pack_bit_matrix`]),
|
||||
//! [`persistent`] (the format-dispatching [`PersistentBitMatrix`] enum),
|
||||
//! [`builder`] ([`PersistentBitMatrixBuilder`], always builds Columnar),
|
||||
//! [`group_ops`] (`MatrixGroupOps` impl), [`pairwise`] (shared symmetric
|
||||
//! pairwise-matrix helpers, also used by `intmatrix.rs`).
|
||||
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
mod builder;
|
||||
mod columnar;
|
||||
mod group_ops;
|
||||
mod packed;
|
||||
mod pairwise;
|
||||
mod persistent;
|
||||
|
||||
pub use builder::PersistentBitMatrixBuilder;
|
||||
pub use packed::pack_bit_matrix;
|
||||
pub use persistent::PersistentBitMatrix;
|
||||
|
||||
pub(crate) use pairwise::{pairwise_matrix, pairwise2_matrix};
|
||||
|
||||
fn col_path(dir: &Path, col: usize) -> PathBuf {
|
||||
dir.join(format!("col_{col:06}.pbiv"))
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
use std::fs::{self, File};
|
||||
use std::io::{self, BufWriter, Read as _, Write as _};
|
||||
use std::path::Path;
|
||||
|
||||
use memmap2::Mmap;
|
||||
use ndarray::{Array1, Array2};
|
||||
use rayon::prelude::*;
|
||||
|
||||
use crate::bitvec::PersistentBitVecBuilder;
|
||||
use crate::meta::MatrixMeta;
|
||||
use crate::views::BitSliceView;
|
||||
|
||||
use super::col_path;
|
||||
use super::pairwise::{pairwise_matrix, pairwise2_matrix};
|
||||
|
||||
// ── PackedBitMatrix ───────────────────────────────────────────────────────────
|
||||
|
||||
const PBMX_MAGIC: [u8; 4] = *b"PBMX";
|
||||
const PBMX_HEADER: usize = 24; // magic(4) + pad(4) + n_rows(8) + n_cols(8)
|
||||
const PBIV_HEADER: usize = 16; // magic(4) + pad(4) + n(8)
|
||||
|
||||
/// Single-file packed layout: all columns concatenated behind a header.
|
||||
pub struct PackedBitMatrix {
|
||||
mmap: Mmap,
|
||||
pub(super) n_rows: usize,
|
||||
pub(super) n_cols: usize,
|
||||
/// Absolute byte offset to the start of each column's bit data
|
||||
/// (= file offset of the PBIV blob + PBIV_HEADER).
|
||||
data_offsets: Vec<usize>,
|
||||
}
|
||||
|
||||
impl PackedBitMatrix {
|
||||
pub(crate) fn open(path: &Path) -> io::Result<Self> {
|
||||
let mmap = unsafe { Mmap::map(&File::open(path)?)? };
|
||||
if mmap.len() < PBMX_HEADER {
|
||||
return Err(io::Error::new(io::ErrorKind::InvalidData, "PBMX file too short"));
|
||||
}
|
||||
if &mmap[0..4] != &PBMX_MAGIC {
|
||||
return Err(io::Error::new(io::ErrorKind::InvalidData, "bad PBMX magic"));
|
||||
}
|
||||
let n_rows = u64::from_le_bytes(mmap[8..16].try_into().unwrap()) as usize;
|
||||
let n_cols = u64::from_le_bytes(mmap[16..24].try_into().unwrap()) as usize;
|
||||
|
||||
let mut data_offsets = Vec::with_capacity(n_cols);
|
||||
for c in 0..n_cols {
|
||||
let off_pos = PBMX_HEADER + c * 8;
|
||||
let col_file_off = u64::from_le_bytes(mmap[off_pos..off_pos+8].try_into().unwrap()) as usize;
|
||||
data_offsets.push(col_file_off + PBIV_HEADER);
|
||||
}
|
||||
|
||||
Ok(Self { mmap, n_rows, n_cols, data_offsets })
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub(crate) fn fill_row(&self, slot: usize, buf: &mut [u32]) {
|
||||
for (c, &data_off) in self.data_offsets.iter().enumerate() {
|
||||
buf[c] = ((self.mmap[data_off + (slot >> 3)] >> (slot & 7)) & 1) as u32;
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn row(&self, slot: usize) -> Box<[bool]> {
|
||||
(0..self.n_cols).map(|c| {
|
||||
(self.mmap[self.data_offsets[c] + (slot >> 3)] >> (slot & 7)) & 1 != 0
|
||||
}).collect()
|
||||
}
|
||||
|
||||
fn col_bytes(&self, c: usize) -> &[u8] {
|
||||
let start = self.data_offsets[c];
|
||||
&self.mmap[start..start + self.n_rows.div_ceil(8)]
|
||||
}
|
||||
|
||||
fn col_words(&self, c: usize) -> &[u64] {
|
||||
let nw = self.n_rows.div_ceil(64);
|
||||
// SAFETY: data_offsets[c] is always 8-byte aligned.
|
||||
// PBMX header = 24 + n_cols×8 (multiple of 8); each PBIV blob =
|
||||
// 16 + nwords×8 (multiple of 8); mmap base is page-aligned.
|
||||
let ptr = self.mmap[self.data_offsets[c]..].as_ptr() as *const u64;
|
||||
unsafe { std::slice::from_raw_parts(ptr, nw) }
|
||||
}
|
||||
|
||||
pub(crate) fn col_slice(&self, c: usize) -> BitSliceView<'_> {
|
||||
BitSliceView::new(self.col_words(c), self.n_rows)
|
||||
}
|
||||
|
||||
pub(crate) fn col_persist(&self, c: usize, path: &Path) -> io::Result<PersistentBitVecBuilder> {
|
||||
PersistentBitVecBuilder::from_raw_bytes(self.col_bytes(c), self.n_rows, path)
|
||||
}
|
||||
|
||||
pub(crate) fn count_ones(&self) -> Array1<u64> {
|
||||
Array1::from_vec(
|
||||
(0..self.n_cols).into_par_iter()
|
||||
.map(|c| self.col_slice(c).count_ones())
|
||||
.collect()
|
||||
)
|
||||
}
|
||||
|
||||
pub(crate) fn partial_jaccard_dist_matrix(&self) -> (Array2<u64>, Array2<u64>) {
|
||||
pairwise2_matrix(self.n_cols, |i, j| {
|
||||
self.col_slice(i).partial_jaccard_dist(self.col_slice(j))
|
||||
})
|
||||
}
|
||||
|
||||
pub(crate) fn partial_hamming_dist_matrix(&self) -> Array2<u64> {
|
||||
pairwise_matrix(self.n_cols, |i, j| {
|
||||
self.col_slice(i).hamming_dist(self.col_slice(j))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/// Reads just the `n_cols` field from an existing packed matrix's header,
|
||||
/// without mapping the file. Used by `pack_bit_matrix` to tell a genuinely
|
||||
/// complete pack from a stale one that predates a later column-widening.
|
||||
fn packed_bit_matrix_n_cols(path: &Path) -> io::Result<usize> {
|
||||
let mut f = File::open(path)?;
|
||||
let mut header = [0u8; PBMX_HEADER];
|
||||
f.read_exact(&mut header)?;
|
||||
Ok(u64::from_le_bytes(header[16..24].try_into().unwrap()) as usize)
|
||||
}
|
||||
|
||||
/// Build `presence/matrix.pbmx` from existing `col_*.pbiv` files.
|
||||
pub fn pack_bit_matrix(dir: &Path) -> io::Result<()> {
|
||||
let packed_path = dir.join("matrix.pbmx");
|
||||
|
||||
let meta = match MatrixMeta::load(dir) {
|
||||
Ok(meta) => meta,
|
||||
Err(e) => {
|
||||
// No columnar data pending: either this layer was already
|
||||
// packed and cleaned up (matrix.pbmx complete, nothing left to
|
||||
// do), or genuinely nothing was ever written here.
|
||||
return if packed_path.exists() { Ok(()) } else { Err(e) };
|
||||
}
|
||||
};
|
||||
|
||||
// A `matrix.pbmx` can already exist here even though columnar data is
|
||||
// still pending — e.g. copied verbatim from a merge's base source
|
||||
// before this layer was widened with more genome columns (see
|
||||
// `obikpartitionner::merge_partition`). Only skip (re-)packing if the
|
||||
// existing file already reflects the current column count; otherwise
|
||||
// the columnar files are newer and must be (re-)packed, overwriting the
|
||||
// stale one — never silently discarded as "leftover cleanup".
|
||||
if packed_bit_matrix_n_cols(&packed_path).ok() == Some(meta.n_cols) {
|
||||
for c in 0..meta.n_cols { let _ = fs::remove_file(col_path(dir, c)); }
|
||||
let _ = fs::remove_file(dir.join("meta.json"));
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let n_cols = meta.n_cols;
|
||||
|
||||
// Compute offsets from file sizes — no column data loaded into RAM.
|
||||
let col_sizes: Vec<u64> = (0..n_cols)
|
||||
.map(|c| fs::metadata(col_path(dir, c)).map(|m| m.len()))
|
||||
.collect::<io::Result<_>>()?;
|
||||
|
||||
let header_size = (PBMX_HEADER + n_cols * 8) as u64;
|
||||
let mut col_offset = header_size;
|
||||
let mut offsets = Vec::with_capacity(n_cols);
|
||||
for &size in &col_sizes {
|
||||
offsets.push(col_offset);
|
||||
col_offset += size;
|
||||
}
|
||||
|
||||
// Write to a temp file; rename atomically so a killed process never leaves
|
||||
// a truncated matrix.pbmx that would be mistaken for a complete file.
|
||||
let tmp_path = dir.join("matrix.pbmx.tmp");
|
||||
let mut out = BufWriter::new(File::create(&tmp_path)?);
|
||||
out.write_all(&PBMX_MAGIC)?;
|
||||
out.write_all(&[0u8; 4])?;
|
||||
out.write_all(&(meta.n as u64).to_le_bytes())?;
|
||||
out.write_all(&(n_cols as u64).to_le_bytes())?;
|
||||
for &off in &offsets { out.write_all(&off.to_le_bytes())?; }
|
||||
for c in 0..n_cols {
|
||||
io::copy(&mut File::open(col_path(dir, c))?, &mut out)?;
|
||||
}
|
||||
out.flush()?;
|
||||
drop(out);
|
||||
fs::rename(&tmp_path, &packed_path)?;
|
||||
|
||||
for c in 0..n_cols { fs::remove_file(col_path(dir, c))?; }
|
||||
fs::remove_file(dir.join("meta.json"))?;
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,56 @@
|
||||
use ndarray::Array2;
|
||||
use rayon::prelude::*;
|
||||
|
||||
// ── Shared matrix helpers (also used by intmatrix.rs) ─────────────────────────
|
||||
|
||||
fn upper_pairs(n: usize) -> Vec<(usize, usize)> {
|
||||
(0..n).flat_map(|i| (i + 1..n).map(move |j| (i, j))).collect()
|
||||
}
|
||||
|
||||
fn fill_symmetric<T>(n: usize, vals: impl Iterator<Item = (usize, usize, T, T)>) -> Array2<T>
|
||||
where T: Clone + Default {
|
||||
let mut m = Array2::from_elem((n, n), T::default());
|
||||
for (i, j, vij, vji) in vals { m[[i, j]] = vij; m[[j, i]] = vji; }
|
||||
m
|
||||
}
|
||||
|
||||
/// Compute a symmetric `n×n` matrix in parallel by evaluating `f(i,j)` for
|
||||
/// all upper-triangle pairs, plus `f(i,i)` for the diagonal. `T: Copy` avoids
|
||||
/// the `.clone()` needed for the lower-triangle mirror.
|
||||
///
|
||||
/// The diagonal is *not* generally `T::default()`: for a self-comparison,
|
||||
/// `f(i,i)` is often the column's own weight (e.g. intersection-with-self —
|
||||
/// see `pairwise2_matrix`), not zero. Distance finalisations that need a
|
||||
/// zero diagonal (self-distance) already overwrite it explicitly.
|
||||
pub(crate) fn pairwise_matrix<T>(n: usize, f: impl Fn(usize, usize) -> T + Sync) -> Array2<T>
|
||||
where T: Copy + Default + Send {
|
||||
let results: Vec<(usize, usize, T)> = upper_pairs(n)
|
||||
.into_par_iter().map(|(i, j)| (i, j, f(i, j))).collect();
|
||||
let mut m = fill_symmetric(n, results.into_iter().map(|(i, j, v)| (i, j, v, v)));
|
||||
for i in 0..n { m[[i, i]] = f(i, i); }
|
||||
m
|
||||
}
|
||||
|
||||
/// Same as `pairwise_matrix` but `f` returns two values that fill two
|
||||
/// symmetric matrices simultaneously (e.g. intersection + union for Jaccard).
|
||||
/// The diagonal is `f(i,i)` (e.g. a genome's kmer count intersected with
|
||||
/// itself), not `T::default()` — see `pairwise_matrix` for why that matters.
|
||||
pub(crate) fn pairwise2_matrix<T>(n: usize, f: impl Fn(usize, usize) -> (T, T) + Sync) -> (Array2<T>, Array2<T>)
|
||||
where T: Copy + Default + Send {
|
||||
let results: Vec<(usize, usize, T, T)> = upper_pairs(n)
|
||||
.into_par_iter()
|
||||
.map(|(i, j)| { let (a, b) = f(i, j); (i, j, a, b) })
|
||||
.collect();
|
||||
let mut m0 = Array2::from_elem((n, n), T::default());
|
||||
let mut m1 = Array2::from_elem((n, n), T::default());
|
||||
for (i, j, a, b) in results {
|
||||
m0[[i, j]] = a; m0[[j, i]] = a;
|
||||
m1[[i, j]] = b; m1[[j, i]] = b;
|
||||
}
|
||||
for i in 0..n {
|
||||
let (a, b) = f(i, i);
|
||||
m0[[i, i]] = a;
|
||||
m1[[i, i]] = b;
|
||||
}
|
||||
(m0, m1)
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
use std::io;
|
||||
use std::path::Path;
|
||||
|
||||
use ndarray::{Array1, Array2};
|
||||
|
||||
use crate::bitvec::{PersistentBitVec, PersistentBitVecBuilder};
|
||||
use crate::layer_meta::LayerMeta;
|
||||
use crate::meta::MatrixMeta;
|
||||
use crate::traits::{BitPartials, ColumnWeights};
|
||||
use crate::views::BitSliceView;
|
||||
|
||||
use super::columnar::ColumnarBitMatrix;
|
||||
use super::packed::PackedBitMatrix;
|
||||
|
||||
// ── PersistentBitMatrix — public enum ────────────────────────────────────────
|
||||
|
||||
/// Bit matrix that transparently handles columnar, packed, and implicit formats.
|
||||
///
|
||||
/// - `Columnar`: per-column `.pbiv` files (original format, used during build)
|
||||
/// - `Packed`: single `matrix.pbmx` file (optimised for query — one `mmap`)
|
||||
/// - `Implicit`: no file — all values are 1 (mono-genome presence/absence)
|
||||
pub enum PersistentBitMatrix {
|
||||
Columnar(ColumnarBitMatrix),
|
||||
Packed(PackedBitMatrix),
|
||||
Implicit { n_rows: usize, n_cols: usize },
|
||||
}
|
||||
|
||||
impl PersistentBitMatrix {
|
||||
/// Open from `layer_dir`, auto-detecting the format.
|
||||
///
|
||||
/// Checks (in order):
|
||||
/// 1. `layer_dir/presence/matrix.pbmx` → Packed
|
||||
/// 2. `layer_dir/presence/meta.json` → Columnar
|
||||
/// 3. `layer_dir/layer_meta.json` → Implicit (new index)
|
||||
/// 4. `layer_dir/unitigs.bin` → Implicit with warning (old index)
|
||||
pub fn open(layer_dir: &Path) -> io::Result<Self> {
|
||||
let presence_dir = layer_dir.join("presence");
|
||||
|
||||
if presence_dir.join("matrix.pbmx").exists() {
|
||||
return Ok(Self::Packed(PackedBitMatrix::open(&presence_dir.join("matrix.pbmx"))?));
|
||||
}
|
||||
|
||||
if MatrixMeta::load(&presence_dir).is_ok() {
|
||||
return Ok(Self::Columnar(ColumnarBitMatrix::open(&presence_dir)?));
|
||||
}
|
||||
|
||||
// No presence matrix → Implicit; requires layer_meta.json
|
||||
let meta = LayerMeta::load(layer_dir).map_err(|_| io::Error::new(
|
||||
io::ErrorKind::NotFound,
|
||||
format!(
|
||||
"no presence matrix and no layer_meta.json in {} — run 'obikmer upgrade'",
|
||||
layer_dir.display()
|
||||
),
|
||||
))?;
|
||||
Ok(Self::Implicit { n_rows: meta.n, n_cols: 1 })
|
||||
}
|
||||
|
||||
pub fn n(&self) -> usize {
|
||||
match self {
|
||||
Self::Columnar(m) => m.n(),
|
||||
Self::Packed(m) => m.n_rows,
|
||||
Self::Implicit { n_rows, .. } => *n_rows,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn n_cols(&self) -> usize {
|
||||
match self {
|
||||
Self::Columnar(m) => m.n_cols(),
|
||||
Self::Packed(m) => m.n_cols,
|
||||
Self::Implicit { n_cols, .. } => *n_cols,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn col(&self, c: usize) -> &PersistentBitVec {
|
||||
match self {
|
||||
Self::Columnar(m) => m.col(c),
|
||||
_ => panic!("col() only available on Columnar PersistentBitMatrix"),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn col_view(&self, c: usize) -> BitSliceView<'_> {
|
||||
match self {
|
||||
Self::Columnar(m) => m.col(c).view(),
|
||||
Self::Packed(m) => m.col_slice(c),
|
||||
Self::Implicit { .. } => panic!("col_view() not available on Implicit PersistentBitMatrix"),
|
||||
}
|
||||
}
|
||||
|
||||
/// Column-major point lookup: value at column `c`, slot `slot`, as 0/1.
|
||||
///
|
||||
/// Unlike [`col_view`](Self::col_view), this never panics on `Implicit`
|
||||
/// (every column reads as present, per the mono-genome fast path) — safe
|
||||
/// to call for any `c < self.n_cols()`.
|
||||
pub fn get(&self, c: usize, slot: usize) -> u32 {
|
||||
match self {
|
||||
Self::Columnar(m) => m.col(c).get(slot) as u32,
|
||||
Self::Packed(m) => m.col_slice(c).get(slot) as u32,
|
||||
Self::Implicit { .. } => 1,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn col_persist(&self, c: usize, path: &Path) -> io::Result<PersistentBitVecBuilder> {
|
||||
match self {
|
||||
Self::Columnar(m) => PersistentBitVecBuilder::build_from(m.col(c), path),
|
||||
Self::Packed(m) => m.col_persist(c, path),
|
||||
Self::Implicit { n_rows, .. } => {
|
||||
PersistentBitVecBuilder::new_ones(*n_rows, path)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn row(&self, slot: usize) -> Box<[bool]> {
|
||||
match self {
|
||||
Self::Columnar(m) => m.row(slot),
|
||||
Self::Packed(m) => m.row(slot),
|
||||
Self::Implicit { n_cols, .. } => vec![true; *n_cols].into_boxed_slice(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Fill `buf[i]` with `col_i[slot]` as 0/1 u32, without allocating.
|
||||
pub fn fill_row(&self, slot: usize, buf: &mut [u32]) {
|
||||
match self {
|
||||
Self::Columnar(m) => m.fill_row(slot, buf),
|
||||
Self::Packed(m) => m.fill_row(slot, buf),
|
||||
Self::Implicit { n_cols, .. } => buf[..*n_cols].fill(1),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn count_ones(&self) -> Array1<u64> {
|
||||
match self {
|
||||
Self::Columnar(m) => m.count_ones(),
|
||||
Self::Packed(m) => m.count_ones(),
|
||||
Self::Implicit { n_rows, n_cols } => Array1::from_elem(*n_cols, *n_rows as u64),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn partial_jaccard_dist_matrix(&self) -> (Array2<u64>, Array2<u64>) {
|
||||
match self {
|
||||
Self::Columnar(m) => m.partial_jaccard_dist_matrix(),
|
||||
Self::Packed(m) => m.partial_jaccard_dist_matrix(),
|
||||
Self::Implicit { n_rows, n_cols } => {
|
||||
let v = *n_rows as u64;
|
||||
let n = *n_cols;
|
||||
let mut inter = Array2::zeros((n, n));
|
||||
let mut union = Array2::zeros((n, n));
|
||||
for i in 0..n { for j in 0..n {
|
||||
inter[[i, j]] = v; union[[i, j]] = v;
|
||||
}}
|
||||
(inter, union)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn partial_hamming_dist_matrix(&self) -> Array2<u64> {
|
||||
match self {
|
||||
Self::Columnar(m) => m.partial_hamming_dist_matrix(),
|
||||
Self::Packed(m) => m.partial_hamming_dist_matrix(),
|
||||
Self::Implicit { n_cols, .. } => Array2::zeros((*n_cols, *n_cols)),
|
||||
}
|
||||
}
|
||||
|
||||
/// Append a new column to an on-disk Columnar matrix.
|
||||
pub fn append_column(dir: &Path, value_of: impl Fn(usize) -> bool) -> io::Result<()> {
|
||||
ColumnarBitMatrix::append_column(dir, value_of)
|
||||
}
|
||||
}
|
||||
|
||||
// ── Trait impls ───────────────────────────────────────────────────────────────
|
||||
|
||||
impl ColumnWeights for PersistentBitMatrix {
|
||||
fn col_weights(&self) -> Array1<u64> { self.count_ones() }
|
||||
}
|
||||
|
||||
impl BitPartials for PersistentBitMatrix {
|
||||
fn partial_jaccard(&self) -> (Array2<u64>, Array2<u64>) {
|
||||
self.partial_jaccard_dist_matrix()
|
||||
}
|
||||
fn partial_hamming(&self) -> Array2<u64> {
|
||||
self.partial_hamming_dist_matrix()
|
||||
}
|
||||
}
|
||||
@@ -1,250 +1,24 @@
|
||||
//use ahash::RandomState;
|
||||
use crossbeam_channel;
|
||||
use hashbrown::HashMap;
|
||||
use obikseq::k;
|
||||
use obikseq::{CanonicalKmer, Sequence, Unitig};
|
||||
use obikseq::{CanonicalKmer, Unitig};
|
||||
#[cfg(not(any(test, feature = "test-utils")))]
|
||||
use rayon::iter::{IntoParallelRefIterator, ParallelIterator};
|
||||
use std::cell::RefCell;
|
||||
use std::fmt;
|
||||
use std::sync::atomic::{AtomicU8, Ordering};
|
||||
use xxhash_rust::xxh3::Xxh3Builder;
|
||||
|
||||
use super::node::{IS_VISITED_MASK, Node};
|
||||
use super::unitig_iter::UnitigNucIter;
|
||||
use super::walk::WalkState;
|
||||
|
||||
// ── Types ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
type FastHashMap<K, V> = HashMap<K, V, Xxh3Builder>;
|
||||
|
||||
// ── Node ──────────────────────────────────────────────────────────────────────
|
||||
//
|
||||
// bit layout (LSB first):
|
||||
// bit 0 : can_extend_right — exactly one right canonical neighbour exists
|
||||
// bit 1 : can_extend_left — exactly one left canonical neighbour exists
|
||||
// bit 2 : visited
|
||||
// bits 3–4 : right_nuc — index 0–3 (A/C/G/T) of that neighbour; valid iff bit 0 = 1
|
||||
// bits 5–6 : left_nuc — index 0–3 (A/C/G/T) of that neighbour; valid iff bit 1 = 1
|
||||
// bit 7 : marked as start node (1)
|
||||
//
|
||||
// "can_extend" = false covers both 0 neighbours and ≥2 neighbours; the only
|
||||
// information needed for traversal is "exactly one".
|
||||
|
||||
#[repr(transparent)]
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
pub struct Node(u8);
|
||||
|
||||
const CAN_EXTEND_RIGHT_MASK: u8 = 0b0000_0001; // bit 0: can_extend_right — exactly one right canonical neighbour exists
|
||||
const CAN_EXTEND_LEFT_MASK: u8 = 0b0000_0010; // bit 1: can_extend_left — exactly one left canonical neighbour exists
|
||||
const IS_VISITED_MASK: u8 = 0b0000_0100; // bit 2: visited
|
||||
const RIGHT_NUC_MASK: u8 = 0b0001_1000; // bits 3–4: right_nuc — index 0–3 (A/C/G/T) of that neighbour; valid iff bit 0 = 1
|
||||
const LEFT_NUC_MASK: u8 = 0b0110_0000; // bits 5–6: left_nuc — index 0–3 (A/C/G/T) of that neighbour; valid iff bit 1 = 1
|
||||
const IS_START_MASK: u8 = 0b1000_0000; // bit 7: marked as start node
|
||||
|
||||
impl Node {
|
||||
/// Returns `true` if the node can be extended to the right.
|
||||
///
|
||||
/// A single right neighbour exists.
|
||||
#[inline]
|
||||
pub fn can_extend_right(self) -> bool {
|
||||
self.0 & CAN_EXTEND_RIGHT_MASK != 0
|
||||
}
|
||||
|
||||
/// Returns `true` if the node can be extended to the left.
|
||||
///
|
||||
/// A single left neighbour exists.
|
||||
#[inline]
|
||||
pub fn can_extend_left(self) -> bool {
|
||||
self.0 & CAN_EXTEND_LEFT_MASK != 0
|
||||
}
|
||||
|
||||
/// Returns `true` if the node has been visited.
|
||||
#[inline]
|
||||
pub fn is_visited(self) -> bool {
|
||||
self.0 & IS_VISITED_MASK != 0
|
||||
}
|
||||
|
||||
/// Returns `true` if the node is a start node.
|
||||
#[inline]
|
||||
pub fn is_start(self) -> bool {
|
||||
self.0 & IS_START_MASK != 0
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub fn set_start(&mut self) {
|
||||
self.0 |= IS_START_MASK;
|
||||
}
|
||||
|
||||
pub fn unset_start(&mut self) {
|
||||
self.0 &= !IS_START_MASK;
|
||||
}
|
||||
|
||||
/// Index of the unique right neighbour (0=A, 1=C, 2=G, 3=T).
|
||||
/// Only meaningful when `can_extend_right()` is true.
|
||||
#[inline]
|
||||
pub fn right_nuc(self) -> u8 {
|
||||
debug_assert!(
|
||||
self.can_extend_right(),
|
||||
"from: right_nuc -> The node cannot be extended to the right"
|
||||
);
|
||||
(self.0 >> 3) & 0b11
|
||||
}
|
||||
|
||||
/// Index of the unique left neighbour (0=A, 1=C, 2=G, 3=T).
|
||||
/// Only meaningful when `can_extend_left()` is true.
|
||||
#[inline]
|
||||
pub fn left_nuc(self) -> u8 {
|
||||
debug_assert!(
|
||||
self.can_extend_left(),
|
||||
"from: left_nuc -> The node cannot be extended to the left"
|
||||
);
|
||||
(self.0 >> 5) & 0b11
|
||||
}
|
||||
|
||||
/// Marks the node as visited.
|
||||
#[inline]
|
||||
pub fn set_visited(&mut self) {
|
||||
debug_assert!(
|
||||
!self.is_visited(),
|
||||
"from: is_visited -> The node has already been visited"
|
||||
);
|
||||
self.0 |= IS_VISITED_MASK;
|
||||
}
|
||||
|
||||
/// `nuc` = Some(i) → exactly one neighbour (bit 0 set, bits 3–4 = nucleotide index).
|
||||
/// `nuc` = None → 0 or ≥2 neighbours; `count` encoded in bits 3–4 as count.sat_sub(1).
|
||||
pub fn set_right(&mut self, count: u8, nuc: Option<u8>) {
|
||||
self.0 &= !(CAN_EXTEND_RIGHT_MASK | RIGHT_NUC_MASK);
|
||||
if count == 1 {
|
||||
self.0 |= CAN_EXTEND_RIGHT_MASK;
|
||||
if let Some(n) = nuc {
|
||||
self.0 |= (n & 0b11) << 3;
|
||||
return;
|
||||
}
|
||||
unreachable!("nuc must be Some when count is 1");
|
||||
}
|
||||
self.0 |= (count.saturating_sub(1).min(3)) << 3;
|
||||
}
|
||||
|
||||
/// `nuc` = Some(i) → exactly one neighbour (bit 0 set, bits 3–4 = nucleotide index).
|
||||
/// `nuc` = None → 0 or ≥2 neighbours; `count` encoded in bits 3–4 as count.sat_sub(1).
|
||||
pub fn set_left(&mut self, count: u8, nuc: Option<u8>) {
|
||||
self.0 &= !(CAN_EXTEND_LEFT_MASK | LEFT_NUC_MASK);
|
||||
if count == 1 {
|
||||
self.0 |= CAN_EXTEND_LEFT_MASK;
|
||||
if let Some(n) = nuc {
|
||||
self.0 |= (n & 0b11) << 5;
|
||||
return;
|
||||
}
|
||||
unreachable!("nuc must be Some when count is 1");
|
||||
}
|
||||
self.0 |= (count.saturating_sub(1).min(3)) << 5;
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for Node {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
const NUC: [char; 4] = ['A', 'C', 'G', 'T'];
|
||||
let r = if self.can_extend_right() {
|
||||
format!("→{}", NUC[self.right_nuc() as usize])
|
||||
} else if (self.0 >> 3) & 0b11 == 0 {
|
||||
"→0".to_string()
|
||||
} else {
|
||||
"→≥2".to_string()
|
||||
};
|
||||
let l = if self.can_extend_left() {
|
||||
format!("←{}", NUC[self.left_nuc() as usize])
|
||||
} else if (self.0 >> 5) & 0b11 == 0 {
|
||||
"←0".to_string()
|
||||
} else {
|
||||
"←≥2".to_string()
|
||||
};
|
||||
let v = if self.is_visited() { "V" } else { "." };
|
||||
write!(f, "Node({r} {l} {v})")
|
||||
}
|
||||
}
|
||||
|
||||
pub struct WalkState {
|
||||
kmer: CanonicalKmer,
|
||||
node: Node,
|
||||
direct: bool,
|
||||
}
|
||||
|
||||
impl WalkState {
|
||||
pub fn new(kmer: CanonicalKmer, node: Node, direct: bool) -> Self {
|
||||
debug_assert!(!node.is_visited(), "Cannot walk over a visited node");
|
||||
Self { kmer, node, direct }
|
||||
}
|
||||
|
||||
pub fn leavable(&self, graph: &GraphDeBruijn) -> bool {
|
||||
self.walk(graph).is_some()
|
||||
}
|
||||
|
||||
pub fn reachable(&self, graph: &GraphDeBruijn) -> bool {
|
||||
WalkState {
|
||||
kmer: self.kmer,
|
||||
node: self.node,
|
||||
direct: !self.direct,
|
||||
}
|
||||
.leavable(graph)
|
||||
}
|
||||
|
||||
pub fn walk(&self, graph: &GraphDeBruijn) -> Option<(WalkState, u8)> {
|
||||
if self.direct {
|
||||
if !self.node.can_extend_right() {
|
||||
return None;
|
||||
}
|
||||
let nuc = self.node.right_nuc();
|
||||
let next = self.kmer.into_kmer().push_right(nuc);
|
||||
let cnext = next.canonical();
|
||||
let dnext = next.raw() == cnext.raw();
|
||||
let next_node = Node(graph.nodes.get(&cnext).unwrap().load(Ordering::Relaxed));
|
||||
if next_node.is_visited() {
|
||||
return None;
|
||||
}
|
||||
let reachable = if dnext {
|
||||
next_node.can_extend_left()
|
||||
} else {
|
||||
next_node.can_extend_right()
|
||||
};
|
||||
reachable.then_some((
|
||||
WalkState {
|
||||
kmer: cnext,
|
||||
node: next_node,
|
||||
direct: dnext,
|
||||
},
|
||||
nuc,
|
||||
))
|
||||
} else {
|
||||
if !self.node.can_extend_left() {
|
||||
return None;
|
||||
}
|
||||
let nuc = self.node.left_nuc();
|
||||
let next = self.kmer.into_kmer().push_left(nuc);
|
||||
let cnext = next.canonical();
|
||||
let dnext = next.raw() != cnext.raw();
|
||||
let next_node = Node(graph.nodes.get(&cnext).unwrap().load(Ordering::Relaxed));
|
||||
if next_node.is_visited() {
|
||||
return None;
|
||||
}
|
||||
let reachable = if dnext {
|
||||
next_node.can_extend_right()
|
||||
} else {
|
||||
next_node.can_extend_left()
|
||||
};
|
||||
reachable.then_some((
|
||||
WalkState {
|
||||
kmer: cnext,
|
||||
node: next_node,
|
||||
direct: dnext,
|
||||
},
|
||||
3 - nuc,
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
pub(super) type FastHashMap<K, V> = HashMap<K, V, Xxh3Builder>;
|
||||
|
||||
// ── GraphDeBruijn ─────────────────────────────────────────────────────────────
|
||||
|
||||
pub struct GraphDeBruijn {
|
||||
nodes: FastHashMap<CanonicalKmer, AtomicU8>,
|
||||
pub(super) nodes: FastHashMap<CanonicalKmer, AtomicU8>,
|
||||
}
|
||||
|
||||
impl GraphDeBruijn {
|
||||
@@ -346,7 +120,7 @@ impl GraphDeBruijn {
|
||||
Some(WalkState::new(kmer, node, true))
|
||||
}
|
||||
|
||||
fn unitig_nucleotides(&self, kmer: CanonicalKmer, k: usize) -> Option<UnitigNucIter<'_>> {
|
||||
pub(super) fn unitig_nucleotides(&self, kmer: CanonicalKmer, k: usize) -> Option<UnitigNucIter<'_>> {
|
||||
let old = self
|
||||
.nodes
|
||||
.get(&kmer)?
|
||||
@@ -362,13 +136,7 @@ impl GraphDeBruijn {
|
||||
.fetch_or(IS_VISITED_MASK, Ordering::AcqRel);
|
||||
(ext_old & IS_VISITED_MASK == 0).then_some((next_state, nuc))
|
||||
});
|
||||
Some(UnitigNucIter {
|
||||
graph: self,
|
||||
start: kmer,
|
||||
pos: 0,
|
||||
k,
|
||||
next_step,
|
||||
})
|
||||
Some(UnitigNucIter::new(self, kmer, k, next_step))
|
||||
}
|
||||
|
||||
pub fn for_each_unitig(&self, f: impl Fn(UnitigNucIter<'_>) + Sync) {
|
||||
@@ -467,12 +235,7 @@ impl GraphDeBruijn {
|
||||
}
|
||||
|
||||
fn is_start(&self, query: CanonicalKmer, node: Node) -> bool {
|
||||
!WalkState {
|
||||
kmer: query,
|
||||
node,
|
||||
direct: true,
|
||||
}
|
||||
.reachable(self)
|
||||
!WalkState::new(query, node, true).reachable(self)
|
||||
}
|
||||
|
||||
pub fn try_for_each_unitig<E, F>(&self, f: F) -> Result<(), E>
|
||||
@@ -514,44 +277,6 @@ impl GraphDeBruijn {
|
||||
}
|
||||
}
|
||||
|
||||
// ── UnitigNucIter ─────────────────────────────────────────────────────────────
|
||||
|
||||
pub struct UnitigNucIter<'a> {
|
||||
graph: &'a GraphDeBruijn,
|
||||
start: CanonicalKmer,
|
||||
pos: usize,
|
||||
k: usize,
|
||||
next_step: Option<(WalkState, u8)>,
|
||||
}
|
||||
|
||||
impl Iterator for UnitigNucIter<'_> {
|
||||
type Item = u8;
|
||||
|
||||
fn next(&mut self) -> Option<u8> {
|
||||
if self.pos < self.k {
|
||||
let nuc = self.start.nucleotide(self.pos);
|
||||
self.pos += 1;
|
||||
Some(nuc)
|
||||
} else if let Some((state, nuc)) = self.next_step.take() {
|
||||
self.next_step = state.walk(self.graph).and_then(|(next_state, next_nuc)| {
|
||||
let old = self
|
||||
.graph
|
||||
.nodes
|
||||
.get(&next_state.kmer)?
|
||||
.fetch_or(IS_VISITED_MASK, Ordering::AcqRel);
|
||||
(old & IS_VISITED_MASK == 0).then_some((next_state, next_nuc))
|
||||
});
|
||||
Some(nuc)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
fn size_hint(&self) -> (usize, Option<usize>) {
|
||||
(self.k - self.pos.min(self.k), None)
|
||||
}
|
||||
}
|
||||
|
||||
/// Returns the count of neighbors and the index of the first
|
||||
/// neighbor if exactly one of the four canonical neighbours exists in
|
||||
/// the graph, where `i` is its index (0=A, 1=C, 2=G, 3=T).
|
||||
@@ -580,8 +305,3 @@ fn count_neighbors(
|
||||
(0, None)
|
||||
}
|
||||
}
|
||||
|
||||
// ── tests ─────────────────────────────────────────────────────────────────────
|
||||
#[cfg(test)]
|
||||
#[path = "tests/debruijn.rs"]
|
||||
mod tests;
|
||||
@@ -0,0 +1,20 @@
|
||||
//! De Bruijn graph over canonical k-mers, built for unitig extraction.
|
||||
//!
|
||||
//! Submodules: [`node`] (packed per-kmer neighbour/visited/start flags),
|
||||
//! [`walk`] (single-step traversal), [`graph`] ([`GraphDeBruijn`] itself),
|
||||
//! [`unitig_iter`] (nucleotide-by-nucleotide unitig walk iterator).
|
||||
|
||||
mod graph;
|
||||
mod node;
|
||||
mod unitig_iter;
|
||||
mod walk;
|
||||
|
||||
pub use graph::GraphDeBruijn;
|
||||
|
||||
// Only used by `tests/debruijn.rs` (`use super::*`) below.
|
||||
#[cfg(test)]
|
||||
use obikseq::{CanonicalKmer, Sequence};
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "../tests/debruijn.rs"]
|
||||
mod tests;
|
||||
@@ -0,0 +1,148 @@
|
||||
use std::fmt;
|
||||
|
||||
// ── Node ──────────────────────────────────────────────────────────────────────
|
||||
//
|
||||
// bit layout (LSB first):
|
||||
// bit 0 : can_extend_right — exactly one right canonical neighbour exists
|
||||
// bit 1 : can_extend_left — exactly one left canonical neighbour exists
|
||||
// bit 2 : visited
|
||||
// bits 3–4 : right_nuc — index 0–3 (A/C/G/T) of that neighbour; valid iff bit 0 = 1
|
||||
// bits 5–6 : left_nuc — index 0–3 (A/C/G/T) of that neighbour; valid iff bit 1 = 1
|
||||
// bit 7 : marked as start node (1)
|
||||
//
|
||||
// "can_extend" = false covers both 0 neighbours and ≥2 neighbours; the only
|
||||
// information needed for traversal is "exactly one".
|
||||
|
||||
#[repr(transparent)]
|
||||
#[derive(Debug, Clone, Copy, Default)]
|
||||
pub struct Node(pub(super) u8);
|
||||
|
||||
const CAN_EXTEND_RIGHT_MASK: u8 = 0b0000_0001; // bit 0: can_extend_right — exactly one right canonical neighbour exists
|
||||
const CAN_EXTEND_LEFT_MASK: u8 = 0b0000_0010; // bit 1: can_extend_left — exactly one left canonical neighbour exists
|
||||
pub(super) const IS_VISITED_MASK: u8 = 0b0000_0100; // bit 2: visited
|
||||
const RIGHT_NUC_MASK: u8 = 0b0001_1000; // bits 3–4: right_nuc — index 0–3 (A/C/G/T) of that neighbour; valid iff bit 0 = 1
|
||||
const LEFT_NUC_MASK: u8 = 0b0110_0000; // bits 5–6: left_nuc — index 0–3 (A/C/G/T) of that neighbour; valid iff bit 1 = 1
|
||||
const IS_START_MASK: u8 = 0b1000_0000; // bit 7: marked as start node
|
||||
|
||||
impl Node {
|
||||
/// Returns `true` if the node can be extended to the right.
|
||||
///
|
||||
/// A single right neighbour exists.
|
||||
#[inline]
|
||||
pub fn can_extend_right(self) -> bool {
|
||||
self.0 & CAN_EXTEND_RIGHT_MASK != 0
|
||||
}
|
||||
|
||||
/// Returns `true` if the node can be extended to the left.
|
||||
///
|
||||
/// A single left neighbour exists.
|
||||
#[inline]
|
||||
pub fn can_extend_left(self) -> bool {
|
||||
self.0 & CAN_EXTEND_LEFT_MASK != 0
|
||||
}
|
||||
|
||||
/// Returns `true` if the node has been visited.
|
||||
#[inline]
|
||||
pub fn is_visited(self) -> bool {
|
||||
self.0 & IS_VISITED_MASK != 0
|
||||
}
|
||||
|
||||
/// Returns `true` if the node is a start node.
|
||||
#[inline]
|
||||
pub fn is_start(self) -> bool {
|
||||
self.0 & IS_START_MASK != 0
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub fn set_start(&mut self) {
|
||||
self.0 |= IS_START_MASK;
|
||||
}
|
||||
|
||||
pub fn unset_start(&mut self) {
|
||||
self.0 &= !IS_START_MASK;
|
||||
}
|
||||
|
||||
/// Index of the unique right neighbour (0=A, 1=C, 2=G, 3=T).
|
||||
/// Only meaningful when `can_extend_right()` is true.
|
||||
#[inline]
|
||||
pub fn right_nuc(self) -> u8 {
|
||||
debug_assert!(
|
||||
self.can_extend_right(),
|
||||
"from: right_nuc -> The node cannot be extended to the right"
|
||||
);
|
||||
(self.0 >> 3) & 0b11
|
||||
}
|
||||
|
||||
/// Index of the unique left neighbour (0=A, 1=C, 2=G, 3=T).
|
||||
/// Only meaningful when `can_extend_left()` is true.
|
||||
#[inline]
|
||||
pub fn left_nuc(self) -> u8 {
|
||||
debug_assert!(
|
||||
self.can_extend_left(),
|
||||
"from: left_nuc -> The node cannot be extended to the left"
|
||||
);
|
||||
(self.0 >> 5) & 0b11
|
||||
}
|
||||
|
||||
/// Marks the node as visited.
|
||||
#[inline]
|
||||
pub fn set_visited(&mut self) {
|
||||
debug_assert!(
|
||||
!self.is_visited(),
|
||||
"from: is_visited -> The node has already been visited"
|
||||
);
|
||||
self.0 |= IS_VISITED_MASK;
|
||||
}
|
||||
|
||||
/// `nuc` = Some(i) → exactly one neighbour (bit 0 set, bits 3–4 = nucleotide index).
|
||||
/// `nuc` = None → 0 or ≥2 neighbours; `count` encoded in bits 3–4 as count.sat_sub(1).
|
||||
pub fn set_right(&mut self, count: u8, nuc: Option<u8>) {
|
||||
self.0 &= !(CAN_EXTEND_RIGHT_MASK | RIGHT_NUC_MASK);
|
||||
if count == 1 {
|
||||
self.0 |= CAN_EXTEND_RIGHT_MASK;
|
||||
if let Some(n) = nuc {
|
||||
self.0 |= (n & 0b11) << 3;
|
||||
return;
|
||||
}
|
||||
unreachable!("nuc must be Some when count is 1");
|
||||
}
|
||||
self.0 |= (count.saturating_sub(1).min(3)) << 3;
|
||||
}
|
||||
|
||||
/// `nuc` = Some(i) → exactly one neighbour (bit 0 set, bits 3–4 = nucleotide index).
|
||||
/// `nuc` = None → 0 or ≥2 neighbours; `count` encoded in bits 3–4 as count.sat_sub(1).
|
||||
pub fn set_left(&mut self, count: u8, nuc: Option<u8>) {
|
||||
self.0 &= !(CAN_EXTEND_LEFT_MASK | LEFT_NUC_MASK);
|
||||
if count == 1 {
|
||||
self.0 |= CAN_EXTEND_LEFT_MASK;
|
||||
if let Some(n) = nuc {
|
||||
self.0 |= (n & 0b11) << 5;
|
||||
return;
|
||||
}
|
||||
unreachable!("nuc must be Some when count is 1");
|
||||
}
|
||||
self.0 |= (count.saturating_sub(1).min(3)) << 5;
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for Node {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
const NUC: [char; 4] = ['A', 'C', 'G', 'T'];
|
||||
let r = if self.can_extend_right() {
|
||||
format!("→{}", NUC[self.right_nuc() as usize])
|
||||
} else if (self.0 >> 3) & 0b11 == 0 {
|
||||
"→0".to_string()
|
||||
} else {
|
||||
"→≥2".to_string()
|
||||
};
|
||||
let l = if self.can_extend_left() {
|
||||
format!("←{}", NUC[self.left_nuc() as usize])
|
||||
} else if (self.0 >> 5) & 0b11 == 0 {
|
||||
"←0".to_string()
|
||||
} else {
|
||||
"←≥2".to_string()
|
||||
};
|
||||
let v = if self.is_visited() { "V" } else { "." };
|
||||
write!(f, "Node({r} {l} {v})")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
use obikseq::CanonicalKmer;
|
||||
use std::sync::atomic::Ordering;
|
||||
|
||||
use super::graph::GraphDeBruijn;
|
||||
use super::node::IS_VISITED_MASK;
|
||||
use super::walk::WalkState;
|
||||
|
||||
// ── UnitigNucIter ─────────────────────────────────────────────────────────────
|
||||
|
||||
pub struct UnitigNucIter<'a> {
|
||||
graph: &'a GraphDeBruijn,
|
||||
start: CanonicalKmer,
|
||||
pos: usize,
|
||||
k: usize,
|
||||
next_step: Option<(WalkState, u8)>,
|
||||
}
|
||||
|
||||
impl<'a> UnitigNucIter<'a> {
|
||||
pub(super) fn new(
|
||||
graph: &'a GraphDeBruijn,
|
||||
start: CanonicalKmer,
|
||||
k: usize,
|
||||
next_step: Option<(WalkState, u8)>,
|
||||
) -> Self {
|
||||
Self {
|
||||
graph,
|
||||
start,
|
||||
pos: 0,
|
||||
k,
|
||||
next_step,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Iterator for UnitigNucIter<'_> {
|
||||
type Item = u8;
|
||||
|
||||
fn next(&mut self) -> Option<u8> {
|
||||
if self.pos < self.k {
|
||||
let nuc = self.start.nucleotide(self.pos);
|
||||
self.pos += 1;
|
||||
Some(nuc)
|
||||
} else if let Some((state, nuc)) = self.next_step.take() {
|
||||
self.next_step = state.walk(self.graph).and_then(|(next_state, next_nuc)| {
|
||||
let old = self
|
||||
.graph
|
||||
.nodes
|
||||
.get(&next_state.kmer)?
|
||||
.fetch_or(IS_VISITED_MASK, Ordering::AcqRel);
|
||||
(old & IS_VISITED_MASK == 0).then_some((next_state, next_nuc))
|
||||
});
|
||||
Some(nuc)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
fn size_hint(&self) -> (usize, Option<usize>) {
|
||||
(self.k - self.pos.min(self.k), None)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
use obikseq::{CanonicalKmer, Sequence};
|
||||
use std::sync::atomic::Ordering;
|
||||
|
||||
use super::graph::GraphDeBruijn;
|
||||
use super::node::Node;
|
||||
|
||||
pub struct WalkState {
|
||||
pub(super) kmer: CanonicalKmer,
|
||||
pub(super) node: Node,
|
||||
pub(super) direct: bool,
|
||||
}
|
||||
|
||||
impl WalkState {
|
||||
pub fn new(kmer: CanonicalKmer, node: Node, direct: bool) -> Self {
|
||||
debug_assert!(!node.is_visited(), "Cannot walk over a visited node");
|
||||
Self { kmer, node, direct }
|
||||
}
|
||||
|
||||
pub fn leavable(&self, graph: &GraphDeBruijn) -> bool {
|
||||
self.walk(graph).is_some()
|
||||
}
|
||||
|
||||
pub fn reachable(&self, graph: &GraphDeBruijn) -> bool {
|
||||
WalkState {
|
||||
kmer: self.kmer,
|
||||
node: self.node,
|
||||
direct: !self.direct,
|
||||
}
|
||||
.leavable(graph)
|
||||
}
|
||||
|
||||
pub fn walk(&self, graph: &GraphDeBruijn) -> Option<(WalkState, u8)> {
|
||||
if self.direct {
|
||||
if !self.node.can_extend_right() {
|
||||
return None;
|
||||
}
|
||||
let nuc = self.node.right_nuc();
|
||||
let next = self.kmer.into_kmer().push_right(nuc);
|
||||
let cnext = next.canonical();
|
||||
let dnext = next.raw() == cnext.raw();
|
||||
let next_node = Node(graph.nodes.get(&cnext).unwrap().load(Ordering::Relaxed));
|
||||
if next_node.is_visited() {
|
||||
return None;
|
||||
}
|
||||
let reachable = if dnext {
|
||||
next_node.can_extend_left()
|
||||
} else {
|
||||
next_node.can_extend_right()
|
||||
};
|
||||
reachable.then_some((
|
||||
WalkState {
|
||||
kmer: cnext,
|
||||
node: next_node,
|
||||
direct: dnext,
|
||||
},
|
||||
nuc,
|
||||
))
|
||||
} else {
|
||||
if !self.node.can_extend_left() {
|
||||
return None;
|
||||
}
|
||||
let nuc = self.node.left_nuc();
|
||||
let next = self.kmer.into_kmer().push_left(nuc);
|
||||
let cnext = next.canonical();
|
||||
let dnext = next.raw() != cnext.raw();
|
||||
let next_node = Node(graph.nodes.get(&cnext).unwrap().load(Ordering::Relaxed));
|
||||
if next_node.is_visited() {
|
||||
return None;
|
||||
}
|
||||
let reachable = if dnext {
|
||||
next_node.can_extend_right()
|
||||
} else {
|
||||
next_node.can_extend_left()
|
||||
};
|
||||
reachable.then_some((
|
||||
WalkState {
|
||||
kmer: cnext,
|
||||
node: next_node,
|
||||
direct: dnext,
|
||||
},
|
||||
3 - nuc,
|
||||
))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
//! NUMA-aware partition runner via hwlocality.
|
||||
//!
|
||||
//! Detects NUMA topology using hwloc (cross-platform: Linux, macOS, etc.) and
|
||||
//! builds one Rayon ThreadPool per NUMA node with threads pinned to that node's
|
||||
//! CPUs. Linux first-touch policy then places graph allocations in local DRAM
|
||||
//! automatically — no explicit memory binding needed.
|
||||
//!
|
||||
//! UMA systems (single socket, Apple Silicon, etc.) are the degenerate case:
|
||||
//! one synthetic node containing all cores, no pool, no pinning.
|
||||
//!
|
||||
//! Submodules: [`topology`] (NUMA detection, per-node pools, thread pinning),
|
||||
//! [`runner`] ([`PartitionRunner`], the adaptive worker-activation scheduler).
|
||||
|
||||
mod runner;
|
||||
mod topology;
|
||||
|
||||
pub use runner::PartitionRunner;
|
||||
@@ -1,137 +1,11 @@
|
||||
// NUMA-aware partition runner via hwlocality.
|
||||
//
|
||||
// Detects NUMA topology using hwloc (cross-platform: Linux, macOS, etc.) and
|
||||
// builds one Rayon ThreadPool per NUMA node with threads pinned to that node's
|
||||
// CPUs. Linux first-touch policy then places graph allocations in local DRAM
|
||||
// automatically — no explicit memory binding needed.
|
||||
//
|
||||
// UMA systems (single socket, Apple Silicon, etc.) are the degenerate case:
|
||||
// one synthetic node containing all cores, no pool, no pinning.
|
||||
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use crossbeam_channel::unbounded;
|
||||
#[cfg(feature = "numa")]
|
||||
use hwlocality::Topology;
|
||||
#[cfg(feature = "numa")]
|
||||
use hwlocality::cpu::binding::CpuBindingFlags;
|
||||
#[cfg(feature = "numa")]
|
||||
use hwlocality::cpu::cpuset::CpuSet;
|
||||
#[cfg(feature = "numa")]
|
||||
use hwlocality::object::types::ObjectType;
|
||||
use obisys::{CpuSample, IoSample};
|
||||
use tracing::debug;
|
||||
|
||||
// ── Public interface ──────────────────────────────────────────────────────────
|
||||
|
||||
pub struct NumaSetup {
|
||||
/// One entry per NUMA node. `None` on UMA systems (no pool, no pinning).
|
||||
pub pools: Vec<Option<Arc<rayon::ThreadPool>>>,
|
||||
/// CPU indices for each NUMA node, in node order.
|
||||
pub cpus_per_node: Vec<Vec<usize>>,
|
||||
}
|
||||
|
||||
impl NumaSetup {
|
||||
/// Maximum worker slots per node (one per physical core in the node).
|
||||
pub fn workers_per_node(&self) -> usize {
|
||||
self.cpus_per_node
|
||||
.first()
|
||||
.map(|c| c.len().max(1))
|
||||
.unwrap_or(1)
|
||||
}
|
||||
}
|
||||
|
||||
/// Detect NUMA topology and build per-node Rayon pools.
|
||||
/// Always succeeds: falls back to a single synthetic UMA node on failure.
|
||||
#[cfg(feature = "numa")]
|
||||
pub fn build() -> NumaSetup {
|
||||
if let Ok(topology) = Topology::new() {
|
||||
let nodes: Vec<Vec<usize>> = topology
|
||||
.objects_with_type(ObjectType::NUMANode)
|
||||
.filter_map(|obj| obj.cpuset())
|
||||
.map(|cpuset| {
|
||||
cpuset
|
||||
.iter_set()
|
||||
.map(|idx| usize::from(idx))
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.filter(|v| !v.is_empty())
|
||||
.collect();
|
||||
|
||||
if nodes.len() > 1 {
|
||||
if let Some(pools) = nodes
|
||||
.iter()
|
||||
.map(|cpus| build_pool(cpus).map(|p| Some(Arc::new(p))))
|
||||
.collect::<Option<Vec<_>>>()
|
||||
{
|
||||
debug!(
|
||||
"NUMA topology: {} node(s), {} core(s)/node",
|
||||
nodes.len(),
|
||||
nodes.first().map_or(0, |v| v.len()),
|
||||
);
|
||||
return NumaSetup {
|
||||
pools,
|
||||
cpus_per_node: nodes,
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// UMA fallback: single synthetic node, all cores, no pool, no pinning.
|
||||
let n_cores = obisys::effective_parallelism();
|
||||
debug!("UMA: single synthetic node, {} core(s)", n_cores);
|
||||
NumaSetup {
|
||||
pools: vec![None],
|
||||
cpus_per_node: vec![(0..n_cores).collect()],
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "numa"))]
|
||||
pub fn build() -> NumaSetup {
|
||||
let n_cores = obisys::effective_parallelism();
|
||||
debug!("UMA: single synthetic node, {} core(s)", n_cores);
|
||||
NumaSetup {
|
||||
pools: vec![None],
|
||||
cpus_per_node: vec![(0..n_cores).collect()],
|
||||
}
|
||||
}
|
||||
|
||||
/// Bind the calling thread to `cpu_indices` using hwloc.
|
||||
/// Silently returns on any error so the thread still runs, just unbound.
|
||||
#[cfg(feature = "numa")]
|
||||
pub fn pin_current_thread(cpu_indices: &[usize]) {
|
||||
let Ok(topology) = Topology::new() else {
|
||||
return;
|
||||
};
|
||||
let mut cpuset = CpuSet::new();
|
||||
for &idx in cpu_indices {
|
||||
cpuset.set(idx);
|
||||
}
|
||||
let _ = topology.bind_cpu(&cpuset, CpuBindingFlags::THREAD);
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "numa"))]
|
||||
pub fn pin_current_thread(_cpu_indices: &[usize]) {}
|
||||
|
||||
// ── Internal helpers ──────────────────────────────────────────────────────────
|
||||
|
||||
#[cfg(feature = "numa")]
|
||||
fn build_pool(cpus: &[usize]) -> Option<rayon::ThreadPool> {
|
||||
let cpus = cpus.to_vec();
|
||||
rayon::ThreadPoolBuilder::new()
|
||||
.num_threads(cpus.len())
|
||||
.spawn_handler(move |thread| {
|
||||
let cpus = cpus.clone();
|
||||
std::thread::Builder::new().spawn(move || {
|
||||
pin_current_thread(&cpus);
|
||||
thread.run();
|
||||
})?;
|
||||
Ok(())
|
||||
})
|
||||
.build()
|
||||
.ok()
|
||||
}
|
||||
use super::topology::{build, pin_current_thread};
|
||||
|
||||
// ── PartitionRunner ─────────────────────────────────────────────────────────
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
#[cfg(feature = "numa")]
|
||||
use hwlocality::Topology;
|
||||
#[cfg(feature = "numa")]
|
||||
use hwlocality::cpu::binding::CpuBindingFlags;
|
||||
#[cfg(feature = "numa")]
|
||||
use hwlocality::cpu::cpuset::CpuSet;
|
||||
#[cfg(feature = "numa")]
|
||||
use hwlocality::object::types::ObjectType;
|
||||
use tracing::debug;
|
||||
|
||||
// ── Public interface ──────────────────────────────────────────────────────────
|
||||
|
||||
pub struct NumaSetup {
|
||||
/// One entry per NUMA node. `None` on UMA systems (no pool, no pinning).
|
||||
pub pools: Vec<Option<Arc<rayon::ThreadPool>>>,
|
||||
/// CPU indices for each NUMA node, in node order.
|
||||
pub cpus_per_node: Vec<Vec<usize>>,
|
||||
}
|
||||
|
||||
impl NumaSetup {
|
||||
/// Maximum worker slots per node (one per physical core in the node).
|
||||
pub fn workers_per_node(&self) -> usize {
|
||||
self.cpus_per_node
|
||||
.first()
|
||||
.map(|c| c.len().max(1))
|
||||
.unwrap_or(1)
|
||||
}
|
||||
}
|
||||
|
||||
/// Detect NUMA topology and build per-node Rayon pools.
|
||||
/// Always succeeds: falls back to a single synthetic UMA node on failure.
|
||||
#[cfg(feature = "numa")]
|
||||
pub fn build() -> NumaSetup {
|
||||
if let Ok(topology) = Topology::new() {
|
||||
let nodes: Vec<Vec<usize>> = topology
|
||||
.objects_with_type(ObjectType::NUMANode)
|
||||
.filter_map(|obj| obj.cpuset())
|
||||
.map(|cpuset| {
|
||||
cpuset
|
||||
.iter_set()
|
||||
.map(|idx| usize::from(idx))
|
||||
.collect::<Vec<_>>()
|
||||
})
|
||||
.filter(|v| !v.is_empty())
|
||||
.collect();
|
||||
|
||||
if nodes.len() > 1 {
|
||||
if let Some(pools) = nodes
|
||||
.iter()
|
||||
.map(|cpus| build_pool(cpus).map(|p| Some(Arc::new(p))))
|
||||
.collect::<Option<Vec<_>>>()
|
||||
{
|
||||
debug!(
|
||||
"NUMA topology: {} node(s), {} core(s)/node",
|
||||
nodes.len(),
|
||||
nodes.first().map_or(0, |v| v.len()),
|
||||
);
|
||||
return NumaSetup {
|
||||
pools,
|
||||
cpus_per_node: nodes,
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// UMA fallback: single synthetic node, all cores, no pool, no pinning.
|
||||
let n_cores = obisys::effective_parallelism();
|
||||
debug!("UMA: single synthetic node, {} core(s)", n_cores);
|
||||
NumaSetup {
|
||||
pools: vec![None],
|
||||
cpus_per_node: vec![(0..n_cores).collect()],
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "numa"))]
|
||||
pub fn build() -> NumaSetup {
|
||||
let n_cores = obisys::effective_parallelism();
|
||||
debug!("UMA: single synthetic node, {} core(s)", n_cores);
|
||||
NumaSetup {
|
||||
pools: vec![None],
|
||||
cpus_per_node: vec![(0..n_cores).collect()],
|
||||
}
|
||||
}
|
||||
|
||||
/// Bind the calling thread to `cpu_indices` using hwloc.
|
||||
/// Silently returns on any error so the thread still runs, just unbound.
|
||||
#[cfg(feature = "numa")]
|
||||
pub fn pin_current_thread(cpu_indices: &[usize]) {
|
||||
let Ok(topology) = Topology::new() else {
|
||||
return;
|
||||
};
|
||||
let mut cpuset = CpuSet::new();
|
||||
for &idx in cpu_indices {
|
||||
cpuset.set(idx);
|
||||
}
|
||||
let _ = topology.bind_cpu(&cpuset, CpuBindingFlags::THREAD);
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "numa"))]
|
||||
pub fn pin_current_thread(_cpu_indices: &[usize]) {}
|
||||
|
||||
// ── Internal helpers ──────────────────────────────────────────────────────────
|
||||
|
||||
#[cfg(feature = "numa")]
|
||||
fn build_pool(cpus: &[usize]) -> Option<rayon::ThreadPool> {
|
||||
let cpus = cpus.to_vec();
|
||||
rayon::ThreadPoolBuilder::new()
|
||||
.num_threads(cpus.len())
|
||||
.spawn_handler(move |thread| {
|
||||
let cpus = cpus.clone();
|
||||
std::thread::Builder::new().spawn(move || {
|
||||
pin_current_thread(&cpus);
|
||||
thread.run();
|
||||
})?;
|
||||
Ok(())
|
||||
})
|
||||
.build()
|
||||
.ok()
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,185 @@
|
||||
use rayon::prelude::*;
|
||||
|
||||
use obicompactvec::{PersistentBitMatrix, PersistentCompactIntMatrix, SiblingAnnex};
|
||||
use obikpartitionner::KmerPartition;
|
||||
use obikseq::CanonicalKmer;
|
||||
use obilayeredmap::MphfLayer;
|
||||
use obilayeredmap::meta::PartitionMeta;
|
||||
use obiskio::UnitigFileReader;
|
||||
use obisys::progress_bar;
|
||||
|
||||
use crate::error::{OKIError, OKIResult};
|
||||
use crate::index::KmerIndex;
|
||||
|
||||
use super::cache::{Mat, PartitionCache};
|
||||
use super::helpers::{central_base, is_minorant, partition_of};
|
||||
use super::{olm_to_ok, ANNEX_FILE_NAME, INDEX_SUBDIR};
|
||||
|
||||
/// IUPAC ambiguity code for a per-genome family presence mask (bit `b` set
|
||||
/// iff the genome carries the member whose own central base is `b`):
|
||||
/// single bit -> the plain base; 2 or 3 bits -> the matching IUPAC
|
||||
/// ambiguity code (preserves partial information instead of collapsing to
|
||||
/// `N`, the same convention used for diploid heterozygous VCF/FASTA sites);
|
||||
/// all 4 bits -> `N`; no bits (genome carries none of the family's observed
|
||||
/// members) -> `-` (no data at this locus for this genome).
|
||||
fn iupac_code(mask: u8) -> u8 {
|
||||
match mask & 0b1111 {
|
||||
0b0000 => b'-',
|
||||
0b0001 => b'A',
|
||||
0b0010 => b'C',
|
||||
0b0100 => b'G',
|
||||
0b1000 => b'T',
|
||||
0b0101 => b'R', // A/G
|
||||
0b1010 => b'Y', // C/T
|
||||
0b0110 => b'S', // C/G
|
||||
0b1001 => b'W', // A/T
|
||||
0b1100 => b'K', // G/T
|
||||
0b0011 => b'M', // A/C
|
||||
0b1110 => b'B', // C/G/T
|
||||
0b1101 => b'D', // A/G/T
|
||||
0b1011 => b'H', // A/C/T
|
||||
0b0111 => b'V', // A/C/G
|
||||
0b1111 => b'N',
|
||||
_ => unreachable!("masked to 4 bits"),
|
||||
}
|
||||
}
|
||||
|
||||
/// A SNP-only pseudo-alignment: one row (byte sequence, IUPAC-coded) per
|
||||
/// genome, one column per variable family (`family_size() >= 2` — monomorphic
|
||||
/// families carry no signal and are skipped, unlike `raw_snp_distance`'s
|
||||
/// tally which does count them as `shared`). Column order is the same,
|
||||
/// deterministic sweep order as the annex build (partition, then layer, then
|
||||
/// slot) — arbitrary but stable and identical across genomes, which is all a
|
||||
/// pseudo-alignment needs (there is no natural genomic coordinate to sort by
|
||||
/// once flanks are dropped). See `docmd/theory/evolutionary_distances.md`,
|
||||
/// "Multi-genome framing: family as pseudo-alignment column".
|
||||
pub struct SnpAlignment {
|
||||
/// `sequences[g]` = genome `g`'s IUPAC-coded row, same length for every
|
||||
/// genome (`sequences.len()` columns).
|
||||
pub sequences: Vec<Vec<u8>>,
|
||||
}
|
||||
|
||||
impl KmerIndex {
|
||||
/// Build the SNP-only pseudo-alignment from an already-built sibling
|
||||
/// annex (run [`build_sibling_annex`](Self::build_sibling_annex) first).
|
||||
pub fn snp_pseudo_alignment(&self) -> OKIResult<SnpAlignment> {
|
||||
let n_parts = self.n_partitions();
|
||||
let n_genomes = self.meta.genomes.len();
|
||||
let with_counts = self.meta.config.with_counts;
|
||||
let k = self.kmer_size();
|
||||
let n_bits = n_parts.trailing_zeros() as usize;
|
||||
|
||||
let partition = KmerPartition::open_with_config(
|
||||
&self.root_path,
|
||||
self.kmer_size(),
|
||||
self.minimizer_size(),
|
||||
n_bits,
|
||||
)
|
||||
.map_err(OKIError::Partition)?;
|
||||
let cache = PartitionCache::build(&partition, n_parts, with_counts)?;
|
||||
|
||||
let mut layer_dirs = Vec::new();
|
||||
for part in 0..n_parts {
|
||||
let index_dir = self.partition().part_dir(part).join(INDEX_SUBDIR);
|
||||
if !index_dir.exists() {
|
||||
continue;
|
||||
}
|
||||
let meta = PartitionMeta::load(&index_dir).map_err(olm_to_ok)?;
|
||||
for l in 0..meta.n_layers {
|
||||
let layer_dir = index_dir.join(format!("layer_{l}"));
|
||||
let annex_path = layer_dir.join(ANNEX_FILE_NAME);
|
||||
if !annex_path.exists() {
|
||||
return Err(OKIError::InvalidInput(format!(
|
||||
"no sibling annex at {} — run build_sibling_annex first",
|
||||
annex_path.display()
|
||||
)));
|
||||
}
|
||||
layer_dirs.push(layer_dir);
|
||||
}
|
||||
}
|
||||
|
||||
let pb = progress_bar("snp_pseudo_alignment", layer_dirs.len() as u64, "layers");
|
||||
// `Vec<Vec<u8>>` per layer, one entry (column) per variable family;
|
||||
// `par_iter().map(...).collect()` on this indexed source preserves
|
||||
// input order, so concatenating the results below in order gives a
|
||||
// single deterministic column order across the whole index.
|
||||
let partials: Vec<Vec<Vec<u8>>> = layer_dirs
|
||||
.par_iter()
|
||||
.map(|layer_dir| -> OKIResult<Vec<Vec<u8>>> {
|
||||
let index_dir = layer_dir.parent().expect("layer_dir has a parent index dir");
|
||||
let meta = PartitionMeta::load(index_dir).map_err(olm_to_ok)?;
|
||||
let annex = SiblingAnnex::open(&layer_dir.join(ANNEX_FILE_NAME))?;
|
||||
let mphf = MphfLayer::open(layer_dir, &meta.mode).map_err(olm_to_ok)?;
|
||||
|
||||
let mut slot_kmer: Vec<Option<CanonicalKmer>> = vec![None; annex.len()];
|
||||
let reader = UnitigFileReader::open_sequential(&layer_dir.join("unitigs.bin"))
|
||||
.map_err(OKIError::Partition)?;
|
||||
for (kmer, _, _) in reader.iter_indexed_canonical_kmers() {
|
||||
if let Some(slot) = mphf.find(kmer) {
|
||||
slot_kmer[slot] = Some(kmer);
|
||||
}
|
||||
}
|
||||
|
||||
let use_counts = with_counts && layer_dir.join("counts").exists();
|
||||
let mat = if use_counts {
|
||||
Mat::Count(PersistentCompactIntMatrix::open(layer_dir)?)
|
||||
} else {
|
||||
Mat::Presence(PersistentBitMatrix::open(layer_dir)?)
|
||||
};
|
||||
let n_cols = mat.n_cols().min(n_genomes);
|
||||
|
||||
let mut columns: Vec<Vec<u8>> = Vec::new();
|
||||
let mut genome_mask: Vec<u8> = Vec::with_capacity(n_genomes);
|
||||
|
||||
for slot in 0..annex.len() {
|
||||
let Some(mask) = annex.get(slot) else { continue };
|
||||
let Some(kmer) = slot_kmer[slot] else { continue };
|
||||
if !is_minorant(kmer, mask, k) {
|
||||
continue; // family tallied once, at its minorant
|
||||
}
|
||||
if mask.family_size() < 2 {
|
||||
continue; // monomorphic family — no signal, skip
|
||||
}
|
||||
|
||||
genome_mask.clear();
|
||||
genome_mask.resize(n_genomes, 0);
|
||||
|
||||
for other in kmer.central_canonical_neighbors() {
|
||||
let base = central_base(other, k);
|
||||
if !mask.has(base) {
|
||||
continue;
|
||||
}
|
||||
let presence: Option<Vec<bool>> = if other == kmer {
|
||||
Some((0..n_cols).map(|g| mat.carries(g, slot)).collect())
|
||||
} else {
|
||||
let dest = partition_of(other, n_parts);
|
||||
cache.find_presence(dest, other, n_genomes)
|
||||
};
|
||||
let Some(presence) = presence else { continue };
|
||||
for (g, &present) in presence.iter().enumerate() {
|
||||
if present {
|
||||
genome_mask[g] |= 1 << base;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
columns.push(genome_mask.iter().map(|&m| iupac_code(m)).collect());
|
||||
}
|
||||
|
||||
pb.inc(1);
|
||||
Ok(columns)
|
||||
})
|
||||
.collect::<OKIResult<Vec<_>>>()?;
|
||||
pb.finish_and_clear();
|
||||
|
||||
let mut sequences: Vec<Vec<u8>> = vec![Vec::new(); n_genomes];
|
||||
for layer_columns in partials {
|
||||
for column in layer_columns {
|
||||
for (g, &code) in column.iter().enumerate() {
|
||||
sequences[g].push(code);
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(SnpAlignment { sequences })
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,255 @@
|
||||
use std::path::Path;
|
||||
use std::sync::atomic::{AtomicU8, Ordering};
|
||||
use std::sync::Arc;
|
||||
|
||||
use rayon::prelude::*;
|
||||
|
||||
use obicompactvec::{FamilyMask, SiblingAnnexBuilder};
|
||||
use obikpartitionner::KmerPartition;
|
||||
use obipipeline::ThrottleGuard;
|
||||
use obikseq::CanonicalKmer;
|
||||
use obilayeredmap::MphfLayer;
|
||||
use obilayeredmap::meta::PartitionMeta;
|
||||
use obiskio::UnitigFileReader;
|
||||
use obisys::progress_bar;
|
||||
|
||||
use crate::error::{OKIError, OKIResult};
|
||||
use crate::index::KmerIndex;
|
||||
|
||||
use super::cache::PartitionCache;
|
||||
use super::helpers::{central_base, partition_of};
|
||||
use super::{olm_to_ok, ANNEX_FILE_NAME, INDEX_SUBDIR};
|
||||
|
||||
// ── obipipeline data types ─────────────────────────────────────────────────
|
||||
|
||||
/// A batch of this layer's distinct k-mers (local MPHF slot + k-mer), the
|
||||
/// pipeline's source item — batched, not one k-mer per item, so that
|
||||
/// pipeline messages and their synchronisation cost stay amortised over
|
||||
/// thousands of lookups (see `build_layer_sibling_annex`'s comment on
|
||||
/// `BATCH_SIZE`). Carries the throttle permit for the whole batch, moved
|
||||
/// (not cloned) into the corresponding `VariantBatch` — a 1-to-1 transform,
|
||||
/// unlike the fan-out `Flat` stage this replaced, needs no `Arc` sharing.
|
||||
struct SourceBatch {
|
||||
items: Vec<(usize, CanonicalKmer)>,
|
||||
_permit: ThrottleGuard,
|
||||
}
|
||||
|
||||
/// One batch's worth of central-substitution variants (up to 3 per source
|
||||
/// k-mer), each already routed to its destination partition and carrying
|
||||
/// its own central base (0=A/1=C/2=G/3=T) — the mask bit it will set on a
|
||||
/// hit. `(dest_partition, variant, source_slot, base)` per entry.
|
||||
struct VariantBatch {
|
||||
items: Vec<(usize, CanonicalKmer, usize, u8)>,
|
||||
_permit: ThrottleGuard,
|
||||
}
|
||||
|
||||
enum SibData {
|
||||
Batch(SourceBatch),
|
||||
Variants(VariantBatch),
|
||||
}
|
||||
|
||||
impl KmerIndex {
|
||||
/// Build the sibling-count/minorant annex for every layer of every
|
||||
/// partition of this (already built) index, writing one annex file per
|
||||
/// layer alongside its existing index files. Safe to call again later
|
||||
/// (e.g. after a fresh `merge`) — each run simply overwrites the annex
|
||||
/// files of the index it is called on.
|
||||
///
|
||||
/// Construction only — no statistics gathered here on purpose: this is
|
||||
/// meant to run routinely (it is the artefact the SNP-family distances
|
||||
/// will consume), while the sibling-count distribution
|
||||
/// ([`sibling_annex_stats`](Self::sibling_annex_stats)) is a separate,
|
||||
/// occasional diagnostic pass over the result, not run every time.
|
||||
///
|
||||
/// Cross-partition/cross-layer lookups are required (a k-mer's siblings
|
||||
/// can live in any partition), but the layer loop itself — and thus the
|
||||
/// annex file this produces — stays local to one layer at a time.
|
||||
pub fn build_sibling_annex(&self) -> OKIResult<()> {
|
||||
let n_parts = self.n_partitions();
|
||||
let n_bits = n_parts.trailing_zeros() as usize;
|
||||
|
||||
let partition = KmerPartition::open_with_config(
|
||||
&self.root_path,
|
||||
self.kmer_size(),
|
||||
self.minimizer_size(),
|
||||
n_bits,
|
||||
)
|
||||
.map_err(OKIError::Partition)?;
|
||||
|
||||
tracing::info!("opening {n_parts} partition(s) for the sibling-annex sweep");
|
||||
let cache = Arc::new(PartitionCache::build(&partition, n_parts, self.meta.config.with_counts)?);
|
||||
|
||||
let pb = progress_bar("sibling_annex", n_parts as u64, "partitions");
|
||||
let mut total_slots: u64 = 0;
|
||||
for part in 0..n_parts {
|
||||
let index_dir = self.partition().part_dir(part).join(INDEX_SUBDIR);
|
||||
if !index_dir.exists() {
|
||||
pb.inc(1);
|
||||
continue;
|
||||
}
|
||||
let meta = PartitionMeta::load(&index_dir).map_err(olm_to_ok)?;
|
||||
|
||||
let mut part_slots: u64 = 0;
|
||||
for l in 0..meta.n_layers {
|
||||
let layer_dir = index_dir.join(format!("layer_{l}"));
|
||||
part_slots += self.build_layer_sibling_annex(&layer_dir, n_parts, &cache)?;
|
||||
}
|
||||
total_slots += part_slots;
|
||||
pb.inc(1);
|
||||
pb.set_message(format!("partition {part}: {part_slots} kmers ({total_slots} total)"));
|
||||
}
|
||||
pb.finish_and_clear();
|
||||
tracing::info!("sibling annex built — {total_slots} kmers across {n_parts} partitions");
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Returns the number of distinct k-mers (annex slots) processed, for
|
||||
/// progress reporting.
|
||||
fn build_layer_sibling_annex(
|
||||
&self,
|
||||
layer_dir: &Path,
|
||||
n_parts: usize,
|
||||
cache: &Arc<PartitionCache>,
|
||||
) -> OKIResult<u64> {
|
||||
let index_dir = layer_dir.parent().expect("layer_dir has a parent index dir");
|
||||
let meta = PartitionMeta::load(index_dir).map_err(olm_to_ok)?;
|
||||
let mphf = MphfLayer::open(layer_dir, &meta.mode).map_err(olm_to_ok)?;
|
||||
let n_slots = mphf.n();
|
||||
|
||||
// ── Enumerate this layer's distinct k-mers, one per slot ────────────
|
||||
let mut slot_kmer: Vec<Option<CanonicalKmer>> = vec![None; n_slots];
|
||||
let reader = UnitigFileReader::open_sequential(&layer_dir.join("unitigs.bin"))
|
||||
.map_err(OKIError::Partition)?;
|
||||
for (kmer, _, _) in reader.iter_indexed_canonical_kmers() {
|
||||
if let Some(slot) = mphf.find(kmer) {
|
||||
slot_kmer[slot] = Some(kmer);
|
||||
}
|
||||
}
|
||||
|
||||
let k = self.kmer_size();
|
||||
|
||||
// ── Reconciliation state, initialised with each slot's own base —
|
||||
// that member is trivially present, no lookup needed. Built before
|
||||
// the pipeline runs, from the same enumeration, since `sources`
|
||||
// below is consumed as a throttled iterator, not collected.
|
||||
// `AtomicU8`, not `FamilyMask`, because the gather phase below
|
||||
// parallelises across destination partitions (independent
|
||||
// `query_partition_with` calls, safe to run concurrently) and their
|
||||
// `Found` hits can land on arbitrary, possibly-shared slots — a
|
||||
// lock-free `fetch_or` avoids needing any synchronisation beyond
|
||||
// that. ─────────────────────────────────────────────────────────
|
||||
let mask: Vec<AtomicU8> = (0..n_slots).map(|_| AtomicU8::new(0)).collect();
|
||||
for (slot, kmer) in slot_kmer.iter().enumerate().filter_map(|(s, k)| k.map(|k| (s, k))) {
|
||||
mask[slot].fetch_or(1 << central_base(kmer, k), Ordering::Relaxed);
|
||||
}
|
||||
|
||||
// ── obipipeline: a *batch* transform, not a per-k-mer `Flat` one —
|
||||
// the actual cross-partition lookup reuses
|
||||
// `KmerPartition::query_partition_with` (the same batching mechanism
|
||||
// `obikmer query` already uses: open a partition's files once,
|
||||
// answer a whole batch of queries against it) instead of one lookup
|
||||
// per pipeline item. A per-item lookup (tried first) reopened/
|
||||
// re-mmap'd every target partition's files on every single variant
|
||||
// — fine at toy scale, but ~90% system time against a real index,
|
||||
// observed in practice. A *later* attempt still pushed one pipeline
|
||||
// message per generated variant (a `Flat` stage, `SourceItem` =>
|
||||
// `VariantQuery`, one k-mer in => up to 3 variants out as separate
|
||||
// messages) — cheaper than reopening files, but sampling a real run
|
||||
// showed most wall-clock time going into per-message channel
|
||||
// send/notify syscalls instead of the lookup itself: the pipeline's
|
||||
// whole point is amortising synchronisation over a batch, and a
|
||||
// single k-mer's ≤3 variants is far too fine a granularity for
|
||||
// that. Batching `BATCH_SIZE` source k-mers into one pipeline item
|
||||
// — a plain 1-to-1 (`|`, not `||`) transform, batch in, batch of
|
||||
// variants out, one message either way — keeps the per-message
|
||||
// synchronisation cost amortised over thousands of lookups instead
|
||||
// of one to three. ──────────────────────────────────────────────
|
||||
const BATCH_SIZE: usize = 4096;
|
||||
let n_workers = obisys::effective_parallelism();
|
||||
let capacity = 256;
|
||||
|
||||
// Throttling limits how many *batches* are in flight at once — the
|
||||
// permit is acquired per batch (not per k-mer) in the source
|
||||
// thread, and released once its `VariantBatch` has been read out of
|
||||
// the pipeline by the accumulation loop below. See
|
||||
// `obipipeline::throttle`'s docs for why this is required, not
|
||||
// optional, once a `Flat`-style stage sits in the pipeline.
|
||||
let sources: Vec<(usize, CanonicalKmer)> = slot_kmer
|
||||
.iter()
|
||||
.enumerate()
|
||||
.filter_map(|(slot, maybe_kmer)| maybe_kmer.map(|kmer| (slot, kmer)))
|
||||
.collect();
|
||||
let batches: Vec<Vec<(usize, CanonicalKmer)>> = sources
|
||||
.chunks(BATCH_SIZE)
|
||||
.map(|chunk| chunk.to_vec())
|
||||
.collect();
|
||||
let throttled = obipipeline::throttle(batches.into_iter(), n_workers).map(|t| SourceBatch {
|
||||
items: t.item,
|
||||
_permit: t.guard,
|
||||
});
|
||||
|
||||
let pipe = obipipeline::make_pipe! {
|
||||
SibData : SourceBatch => VariantBatch,
|
||||
| {
|
||||
move |batch: SourceBatch| -> VariantBatch {
|
||||
let mut items = Vec::with_capacity(batch.items.len() * 3);
|
||||
for (slot, kmer) in batch.items {
|
||||
for variant in kmer.central_canonical_neighbors() {
|
||||
if variant == kmer {
|
||||
continue;
|
||||
}
|
||||
items.push((
|
||||
partition_of(variant, n_parts),
|
||||
variant,
|
||||
slot,
|
||||
central_base(variant, k),
|
||||
));
|
||||
}
|
||||
}
|
||||
VariantBatch { items, _permit: batch._permit }
|
||||
}
|
||||
} : Batch => Variants,
|
||||
};
|
||||
|
||||
// ── Group generated variants by destination partition. `cache`
|
||||
// holds every partition already mmap'd (no more `open()` cost), but
|
||||
// `mmap` pages are still loaded on demand and can be evicted — a
|
||||
// lookup is not free just because the file isn't reopened. Grouping
|
||||
// keeps one partition's pages hot while its whole batch is resolved,
|
||||
// instead of faulting pages in and out as lookups jump between
|
||||
// partitions in whatever order the pipeline happens to produce
|
||||
// them. Each batch's throttle permit drops here, once accumulated.
|
||||
let mut outgoing: Vec<Vec<(CanonicalKmer, usize, u8)>> = (0..n_parts).map(|_| Vec::new()).collect();
|
||||
for vb in pipe.apply(throttled, n_workers, capacity) {
|
||||
for (dest_partition, variant, source_slot, base) in vb.items {
|
||||
outgoing[dest_partition].push((variant, source_slot, base));
|
||||
}
|
||||
}
|
||||
|
||||
// ── Resolve each partition's batch against the cache in one
|
||||
// contiguous pass; parallelised across partitions (independent,
|
||||
// read-only) so this keeps using multiple cores without giving up
|
||||
// the per-partition locality above. ─────────────────────────────
|
||||
outgoing.par_iter().enumerate().filter(|(_, q)| !q.is_empty()).for_each(|(dest, queries)| {
|
||||
for &(variant, source_slot, base) in queries {
|
||||
if cache.find(dest, variant) {
|
||||
mask[source_slot].fetch_or(1 << base, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
// ── Write the layer's annex file ─────────────────────────────────────
|
||||
let annex_path = layer_dir.join(ANNEX_FILE_NAME);
|
||||
let mut builder = SiblingAnnexBuilder::new(n_slots, &annex_path)?;
|
||||
for (slot, m) in mask.iter().enumerate() {
|
||||
if slot_kmer[slot].is_none() {
|
||||
continue; // unused MPHF slot, if any — leave at the sentinel
|
||||
}
|
||||
builder.set(slot, FamilyMask::from_bits(m.load(Ordering::Relaxed)));
|
||||
}
|
||||
builder.close()?;
|
||||
|
||||
Ok(n_slots as u64)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,121 @@
|
||||
use rayon::prelude::*;
|
||||
|
||||
use obicompactvec::{PersistentBitMatrix, PersistentCompactIntMatrix};
|
||||
use obikpartitionner::KmerPartition;
|
||||
use obikseq::CanonicalKmer;
|
||||
use obilayeredmap::MphfLayer;
|
||||
use obilayeredmap::meta::PartitionMeta;
|
||||
use obisys::progress_bar;
|
||||
|
||||
use crate::error::OKIResult;
|
||||
|
||||
use super::{olm_to_ok, INDEX_SUBDIR};
|
||||
|
||||
/// Every partition's already-open MPHF layers, built **once** for the whole
|
||||
/// `build_sibling_annex` run and shared (read-only) across every lookup, in
|
||||
/// every source layer, for the rest of the run — not reopened/re-mmap'd per
|
||||
/// query, nor per source layer.
|
||||
///
|
||||
/// Confirmed necessary by sampling a real run: routing lookups through
|
||||
/// `KmerPartition::query_partition_with` (the same batching `obikmer query`
|
||||
/// uses) still reopens+re-mmaps every target partition's files on every
|
||||
/// call, and it is called once per destination partition **per source
|
||||
/// layer** — for an index with many layers this repeats the same
|
||||
/// `MphfLayer::open`/`Evidence::open`/`PersistentBitMatrix::open` work over
|
||||
/// and over. Parallelising those calls (see the gather step below) spread
|
||||
/// the redundant work across more cores but did not reduce it: sampling
|
||||
/// showed Rayon workers spending their time inside repeated `open()`
|
||||
/// syscalls, not computation. This cache amortises that cost to once per
|
||||
/// partition for the entire run, regardless of how many source layers or
|
||||
/// lookups follow.
|
||||
/// A cached layer's opened presence/count matrix, alongside its `MphfLayer`.
|
||||
pub(super) enum Mat {
|
||||
Count(PersistentCompactIntMatrix),
|
||||
Presence(PersistentBitMatrix),
|
||||
}
|
||||
|
||||
impl Mat {
|
||||
pub(super) fn n_cols(&self) -> usize {
|
||||
match self {
|
||||
Mat::Count(m) => m.n_cols(),
|
||||
Mat::Presence(m) => m.n_cols(),
|
||||
}
|
||||
}
|
||||
pub(super) fn carries(&self, g: usize, slot: usize) -> bool {
|
||||
match self {
|
||||
Mat::Count(m) => m.col_view(g).get(slot) != 0,
|
||||
Mat::Presence(m) => m.get(g, slot) != 0,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) struct PartitionCache {
|
||||
/// `layers[partition][layer]` = that partition's opened MPHF layers,
|
||||
/// paired 1:1 with `mats[partition][layer]`; empty if the partition
|
||||
/// directory doesn't exist. Used by both
|
||||
/// [`crate::index::KmerIndex::build_sibling_annex`] (`layers` only) and
|
||||
/// [`crate::index::KmerIndex::sibling_annex_stats`] (both).
|
||||
layers: Vec<Vec<MphfLayer>>,
|
||||
mats: Vec<Vec<Mat>>,
|
||||
}
|
||||
|
||||
impl PartitionCache {
|
||||
pub(super) fn build(partition: &KmerPartition, n_parts: usize, with_counts: bool) -> OKIResult<Self> {
|
||||
let pb = progress_bar("open_partitions", n_parts as u64, "partitions");
|
||||
let built: Vec<(Vec<MphfLayer>, Vec<Mat>)> = (0..n_parts)
|
||||
.into_par_iter()
|
||||
.map(|part| -> OKIResult<(Vec<MphfLayer>, Vec<Mat>)> {
|
||||
let index_dir = partition.part_dir(part).join(INDEX_SUBDIR);
|
||||
if !index_dir.exists() {
|
||||
pb.inc(1);
|
||||
return Ok((Vec::new(), Vec::new()));
|
||||
}
|
||||
let meta = PartitionMeta::load(&index_dir).map_err(olm_to_ok)?;
|
||||
let mut layers = Vec::with_capacity(meta.n_layers);
|
||||
let mut mats = Vec::with_capacity(meta.n_layers);
|
||||
for l in 0..meta.n_layers {
|
||||
let layer_dir = index_dir.join(format!("layer_{l}"));
|
||||
let Ok(mphf) = MphfLayer::open(&layer_dir, &meta.mode) else { continue };
|
||||
let use_counts = with_counts && layer_dir.join("counts").exists();
|
||||
let mat = if use_counts {
|
||||
PersistentCompactIntMatrix::open(&layer_dir).ok().map(Mat::Count)
|
||||
} else {
|
||||
PersistentBitMatrix::open(&layer_dir).ok().map(Mat::Presence)
|
||||
};
|
||||
let Some(mat) = mat else { continue };
|
||||
layers.push(mphf);
|
||||
mats.push(mat);
|
||||
}
|
||||
pb.inc(1);
|
||||
Ok((layers, mats))
|
||||
})
|
||||
.collect::<OKIResult<Vec<_>>>()?;
|
||||
pb.finish_and_clear();
|
||||
let (layers, mats) = built.into_iter().unzip();
|
||||
Ok(Self { layers, mats })
|
||||
}
|
||||
|
||||
/// Existence-only lookup of `variant` in partition `dest_partition`:
|
||||
/// tries each of the partition's already-open layers in turn, stopping
|
||||
/// at the first hit.
|
||||
pub(super) fn find(&self, dest_partition: usize, variant: CanonicalKmer) -> bool {
|
||||
self.layers
|
||||
.get(dest_partition)
|
||||
.is_some_and(|layers| layers.iter().any(|mphf| mphf.find(variant).is_some()))
|
||||
}
|
||||
|
||||
/// Per-genome presence vector for `variant` in partition `dest_partition`
|
||||
/// (`true` iff that genome carries it), `None` on a miss. Same shape as
|
||||
/// `find`, but also reads the cached matrix instead of just the MPHF.
|
||||
pub(super) fn find_presence(&self, dest_partition: usize, variant: CanonicalKmer, n_genomes: usize) -> Option<Vec<bool>> {
|
||||
let layers = self.layers.get(dest_partition)?;
|
||||
let mats = self.mats.get(dest_partition)?;
|
||||
for (mphf, mat) in layers.iter().zip(mats.iter()) {
|
||||
if let Some(slot) = mphf.find(variant) {
|
||||
let n_cols = mat.n_cols().min(n_genomes);
|
||||
return Some((0..n_cols).map(|g| mat.carries(g, slot)).collect());
|
||||
}
|
||||
}
|
||||
None
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,204 @@
|
||||
use ndarray::Array2;
|
||||
use rayon::prelude::*;
|
||||
|
||||
use obicompactvec::{PersistentBitMatrix, PersistentCompactIntMatrix, SiblingAnnex};
|
||||
use obikpartitionner::KmerPartition;
|
||||
use obikseq::CanonicalKmer;
|
||||
use obilayeredmap::MphfLayer;
|
||||
use obilayeredmap::meta::PartitionMeta;
|
||||
use obiskio::UnitigFileReader;
|
||||
use obisys::progress_bar;
|
||||
|
||||
use crate::error::{OKIError, OKIResult};
|
||||
use crate::index::KmerIndex;
|
||||
|
||||
use super::cache::{Mat, PartitionCache};
|
||||
use super::distance::RawSnpDistanceOutput;
|
||||
use super::helpers::{central_base, is_minorant, partition_of};
|
||||
use super::{olm_to_ok, ANNEX_FILE_NAME, INDEX_SUBDIR};
|
||||
|
||||
/// See [`KmerIndex::cardinality_tally`].
|
||||
pub struct CardinalityTally {
|
||||
/// `counts[a][b] == counts[b][a]` = number of family sites, pooled over
|
||||
/// included genome pairs, where one genome's family cardinality
|
||||
/// (popcount of its presence mask, `0..=4`) is `a` and the other's is
|
||||
/// `b`. Diagonal is real data here (both genomes at the same
|
||||
/// cardinality), unlike [`super::distance::BasePairTally::counts`].
|
||||
pub counts: [[u64; 5]; 5],
|
||||
}
|
||||
|
||||
impl KmerIndex {
|
||||
/// Cardinality co-occurrence, pooled only over genome pairs whose
|
||||
/// overall SNP ratio in `raw` is at or below `ratio_ceiling` — same
|
||||
/// saturation/no-data exclusion discipline as
|
||||
/// [`base_pair_tally`](Self::base_pair_tally). Unlike
|
||||
/// `scan_family_pairs` (which resolves each genome to a single form and
|
||||
/// silently drops any genome carrying more than one member of the
|
||||
/// family), this needs the *full* per-genome presence mask — a family
|
||||
/// member count of 2, 3 or 4 is exactly the signal being tallied, not
|
||||
/// noise to discard — so it re-implements the traversal rather than
|
||||
/// reusing that helper.
|
||||
///
|
||||
/// Restricted to variable families (`family_size() >= 2`), matching
|
||||
/// `snp_pseudo_alignment`'s own scope — briefly removed, then
|
||||
/// reinstated: without it, the diagonal is dominated by genome-wide
|
||||
/// invariant background (family_size()<2 loci vastly outnumber the
|
||||
/// ones that ever vary anywhere), which is inconsistent with the
|
||||
/// `+ASC`-corrected alignment this matrix is ultimately used with —
|
||||
/// `+ASC` exists specifically because the likelihood only ever sees
|
||||
/// variable sites, so a rate model calibrated mostly from invariant
|
||||
/// background sites doesn't describe the population it's applied to.
|
||||
/// Verified empirically: removing the filter measurably worsened a
|
||||
/// real IQ-TREE run (log-likelihood dropped, `NNI search needs
|
||||
/// unusual large number of steps to converge` warnings appeared) — see
|
||||
/// `docmd/theory/evolutionary_distances.md` for the full account.
|
||||
/// [`base_pair_tally`](Self::base_pair_tally)'s own diagonal (`same`)
|
||||
/// gets the matching restriction via `scan_family_pairs`'s new
|
||||
/// `variable` flag, rather than a `family_size()` check of its own (it
|
||||
/// doesn't have direct access to the family's mask).
|
||||
pub fn cardinality_tally(&self, raw: &RawSnpDistanceOutput, ratio_ceiling: f64) -> OKIResult<CardinalityTally> {
|
||||
let n_parts = self.n_partitions();
|
||||
let n_genomes = self.meta.genomes.len();
|
||||
let with_counts = self.meta.config.with_counts;
|
||||
let k = self.kmer_size();
|
||||
let n_bits = n_parts.trailing_zeros() as usize;
|
||||
|
||||
let included = Array2::from_shape_fn((n_genomes, n_genomes), |(i, j)| {
|
||||
if i == j {
|
||||
return false;
|
||||
}
|
||||
let snp = raw.snp[[i, j]];
|
||||
let total = snp + raw.shared[[i, j]];
|
||||
total > 0 && (snp as f64 / total as f64) <= ratio_ceiling
|
||||
});
|
||||
|
||||
let partition = KmerPartition::open_with_config(
|
||||
&self.root_path,
|
||||
self.kmer_size(),
|
||||
self.minimizer_size(),
|
||||
n_bits,
|
||||
)
|
||||
.map_err(OKIError::Partition)?;
|
||||
let cache = PartitionCache::build(&partition, n_parts, with_counts)?;
|
||||
|
||||
let mut layer_dirs = Vec::new();
|
||||
for part in 0..n_parts {
|
||||
let index_dir = self.partition().part_dir(part).join(INDEX_SUBDIR);
|
||||
if !index_dir.exists() {
|
||||
continue;
|
||||
}
|
||||
let meta = PartitionMeta::load(&index_dir).map_err(olm_to_ok)?;
|
||||
for l in 0..meta.n_layers {
|
||||
let layer_dir = index_dir.join(format!("layer_{l}"));
|
||||
let annex_path = layer_dir.join(ANNEX_FILE_NAME);
|
||||
if !annex_path.exists() {
|
||||
return Err(OKIError::InvalidInput(format!(
|
||||
"no sibling annex at {} — run build_sibling_annex first",
|
||||
annex_path.display()
|
||||
)));
|
||||
}
|
||||
layer_dirs.push(layer_dir);
|
||||
}
|
||||
}
|
||||
|
||||
let pb = progress_bar("cardinality_tally", layer_dirs.len() as u64, "layers");
|
||||
let partials: Vec<[[u64; 5]; 5]> = layer_dirs
|
||||
.par_iter()
|
||||
.map(|layer_dir| -> OKIResult<[[u64; 5]; 5]> {
|
||||
let index_dir = layer_dir.parent().expect("layer_dir has a parent index dir");
|
||||
let meta = PartitionMeta::load(index_dir).map_err(olm_to_ok)?;
|
||||
let annex = SiblingAnnex::open(&layer_dir.join(ANNEX_FILE_NAME))?;
|
||||
let mphf = MphfLayer::open(layer_dir, &meta.mode).map_err(olm_to_ok)?;
|
||||
|
||||
let mut slot_kmer: Vec<Option<CanonicalKmer>> = vec![None; annex.len()];
|
||||
let reader = UnitigFileReader::open_sequential(&layer_dir.join("unitigs.bin"))
|
||||
.map_err(OKIError::Partition)?;
|
||||
for (kmer, _, _) in reader.iter_indexed_canonical_kmers() {
|
||||
if let Some(slot) = mphf.find(kmer) {
|
||||
slot_kmer[slot] = Some(kmer);
|
||||
}
|
||||
}
|
||||
|
||||
let use_counts = with_counts && layer_dir.join("counts").exists();
|
||||
let mat = if use_counts {
|
||||
Mat::Count(PersistentCompactIntMatrix::open(layer_dir)?)
|
||||
} else {
|
||||
Mat::Presence(PersistentBitMatrix::open(layer_dir)?)
|
||||
};
|
||||
let n_cols = mat.n_cols().min(n_genomes);
|
||||
|
||||
let mut counts = [[0u64; 5]; 5];
|
||||
let mut genome_mask: Vec<u8> = Vec::with_capacity(n_genomes);
|
||||
|
||||
for slot in 0..annex.len() {
|
||||
let Some(mask) = annex.get(slot) else { continue };
|
||||
let Some(kmer) = slot_kmer[slot] else { continue };
|
||||
if !is_minorant(kmer, mask, k) {
|
||||
continue; // family tallied once, at its minorant
|
||||
}
|
||||
if mask.family_size() < 2 {
|
||||
// Fully invariant family (never varies anywhere in
|
||||
// the index) — genome-wide background, not
|
||||
// SNP-adjacent signal; would otherwise swamp the
|
||||
// diagonal (`c=1/c=1` etc.), which needs to reflect
|
||||
// the same variable-families-only population the
|
||||
// `+ASC`-corrected alignment/likelihood actually
|
||||
// models. See `base_pair_tally`'s `variable` gate
|
||||
// on its own `same` diagonal for the matching fix.
|
||||
continue;
|
||||
}
|
||||
|
||||
genome_mask.clear();
|
||||
genome_mask.resize(n_genomes, 0);
|
||||
|
||||
for other in kmer.central_canonical_neighbors() {
|
||||
let base = central_base(other, k);
|
||||
if !mask.has(base) {
|
||||
continue;
|
||||
}
|
||||
let presence: Option<Vec<bool>> = if other == kmer {
|
||||
Some((0..n_cols).map(|g| mat.carries(g, slot)).collect())
|
||||
} else {
|
||||
let dest = partition_of(other, n_parts);
|
||||
cache.find_presence(dest, other, n_genomes)
|
||||
};
|
||||
let Some(presence) = presence else { continue };
|
||||
for (g, &present) in presence.iter().enumerate() {
|
||||
if present {
|
||||
genome_mask[g] |= 1 << base;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for i in 0..n_genomes {
|
||||
let card_i = genome_mask[i].count_ones() as usize;
|
||||
for j in (i + 1)..n_genomes {
|
||||
if !included[[i, j]] {
|
||||
continue;
|
||||
}
|
||||
let card_j = genome_mask[j].count_ones() as usize;
|
||||
counts[card_i][card_j] += 1;
|
||||
if card_i != card_j {
|
||||
counts[card_j][card_i] += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pb.inc(1);
|
||||
Ok(counts)
|
||||
})
|
||||
.collect::<OKIResult<Vec<_>>>()?;
|
||||
pb.finish_and_clear();
|
||||
|
||||
let mut total = [[0u64; 5]; 5];
|
||||
for partial in partials {
|
||||
for a in 0..5 {
|
||||
for b in 0..5 {
|
||||
total[a][b] += partial[a][b];
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(CardinalityTally { counts: total })
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,305 @@
|
||||
use ndarray::Array2;
|
||||
use rayon::prelude::*;
|
||||
|
||||
use obicompactvec::{PersistentBitMatrix, PersistentCompactIntMatrix, SiblingAnnex};
|
||||
use obikpartitionner::KmerPartition;
|
||||
use obikseq::CanonicalKmer;
|
||||
use obilayeredmap::MphfLayer;
|
||||
use obilayeredmap::meta::PartitionMeta;
|
||||
use obiskio::UnitigFileReader;
|
||||
use obisys::progress_bar;
|
||||
|
||||
use crate::error::{OKIError, OKIResult};
|
||||
use crate::index::KmerIndex;
|
||||
|
||||
use super::cache::{Mat, PartitionCache};
|
||||
use super::helpers::{central_base, is_minorant, partition_of};
|
||||
use super::{olm_to_ok, ANNEX_FILE_NAME, INDEX_SUBDIR};
|
||||
|
||||
/// Raw p-distance restricted to loci that are single-copy in **both**
|
||||
/// genomes of a pair — the "stringent / paralogy-aware" locus eligibility
|
||||
/// rule (`docmd/theory/evolutionary_distances.md`, "Locus eligibility"),
|
||||
/// without the JC/K2P/LogDet finalisation math: `snp[i,j]` + `shared[i,j]`
|
||||
/// is the number of eligible loci, `snp[i,j] / (snp[i,j] + shared[i,j])` is
|
||||
/// `p_hat`. A quick, self-contained way to sanity-check the estimator
|
||||
/// against a real index before the full `SnpTally` design is built.
|
||||
///
|
||||
/// A locus (family, tallied once at its minorant) is eligible for pair
|
||||
/// `(i, j)` iff genome `i` carries exactly one of the family's observed
|
||||
/// forms **and** genome `j` carries exactly one (possibly a different one)
|
||||
/// — presence-only: a genome carrying the same form twice (a same-allele
|
||||
/// duplicate) is indistinguishable from carrying it once when only a
|
||||
/// presence matrix is available, so such cases are not excluded here even
|
||||
/// when a count index exists. See "Locus eligibility", stringent rule, for
|
||||
/// why this matters and how a count index would close the gap — left as a
|
||||
/// follow-up, not applied here.
|
||||
pub struct RawSnpDistanceOutput {
|
||||
/// n×n count of eligible loci where the two genomes' single forms differ.
|
||||
pub snp: Array2<u64>,
|
||||
/// n×n count of eligible loci where the two genomes' single forms agree.
|
||||
pub shared: Array2<u64>,
|
||||
}
|
||||
|
||||
impl KmerIndex {
|
||||
/// Shared traversal behind [`raw_snp_distance`](Self::raw_snp_distance)
|
||||
/// and [`base_pair_tally`](Self::base_pair_tally): for every family
|
||||
/// (tallied once, at its minorant) of every layer of the already-built
|
||||
/// sibling annex, resolves each genome's single observed form (`None`
|
||||
/// if absent or ambiguous/multi-copy), then calls `on_pair(acc, i, j,
|
||||
/// bi, bj, variable)` for every genome pair `(i, j)` where both are
|
||||
/// unambiguous and single-copy (`bi == bj` means shared at that locus,
|
||||
/// `bi != bj` means a SNP). `variable` is the family's own
|
||||
/// `family_size() >= 2` (true if more than one member is observed
|
||||
/// *anywhere* in the family, i.e. it isn't fully invariant across the
|
||||
/// whole index) — `raw_snp_distance` ignores it (a fully-invariant
|
||||
/// family is still legitimately "shared"), but callers whose diagonal
|
||||
/// should only reflect genuine SNP-adjacent agreement, not the
|
||||
/// genome-wide invariant background, need it (see
|
||||
/// [`base_pair_tally`](Self::base_pair_tally)'s `same` field). Layers
|
||||
/// are processed in parallel (rayon); each gets its own accumulator
|
||||
/// from `zero()`, combined pairwise via `combine`.
|
||||
fn scan_family_pairs<Acc, F, C>(
|
||||
&self,
|
||||
label: &str,
|
||||
zero: impl Fn() -> Acc + Sync,
|
||||
on_pair: F,
|
||||
combine: C,
|
||||
) -> OKIResult<Acc>
|
||||
where
|
||||
Acc: Send,
|
||||
F: Fn(&mut Acc, usize, usize, u8, u8, bool) + Sync,
|
||||
C: Fn(Acc, Acc) -> Acc,
|
||||
{
|
||||
let n_parts = self.n_partitions();
|
||||
let n_genomes = self.meta.genomes.len();
|
||||
let with_counts = self.meta.config.with_counts;
|
||||
let k = self.kmer_size();
|
||||
let n_bits = n_parts.trailing_zeros() as usize;
|
||||
|
||||
let partition = KmerPartition::open_with_config(
|
||||
&self.root_path,
|
||||
self.kmer_size(),
|
||||
self.minimizer_size(),
|
||||
n_bits,
|
||||
)
|
||||
.map_err(OKIError::Partition)?;
|
||||
let cache = PartitionCache::build(&partition, n_parts, with_counts)?;
|
||||
|
||||
let mut layer_dirs = Vec::new();
|
||||
for part in 0..n_parts {
|
||||
let index_dir = self.partition().part_dir(part).join(INDEX_SUBDIR);
|
||||
if !index_dir.exists() {
|
||||
continue;
|
||||
}
|
||||
let meta = PartitionMeta::load(&index_dir).map_err(olm_to_ok)?;
|
||||
for l in 0..meta.n_layers {
|
||||
let layer_dir = index_dir.join(format!("layer_{l}"));
|
||||
let annex_path = layer_dir.join(ANNEX_FILE_NAME);
|
||||
if !annex_path.exists() {
|
||||
return Err(OKIError::InvalidInput(format!(
|
||||
"no sibling annex at {} — run build_sibling_annex first",
|
||||
annex_path.display()
|
||||
)));
|
||||
}
|
||||
layer_dirs.push(layer_dir);
|
||||
}
|
||||
}
|
||||
|
||||
let pb = progress_bar(label, layer_dirs.len() as u64, "layers");
|
||||
let partials: Vec<Acc> = layer_dirs
|
||||
.par_iter()
|
||||
.map(|layer_dir| -> OKIResult<Acc> {
|
||||
let mut acc = zero();
|
||||
|
||||
let index_dir = layer_dir.parent().expect("layer_dir has a parent index dir");
|
||||
let meta = PartitionMeta::load(index_dir).map_err(olm_to_ok)?;
|
||||
let annex = SiblingAnnex::open(&layer_dir.join(ANNEX_FILE_NAME))?;
|
||||
let mphf = MphfLayer::open(layer_dir, &meta.mode).map_err(olm_to_ok)?;
|
||||
|
||||
let mut slot_kmer: Vec<Option<CanonicalKmer>> = vec![None; annex.len()];
|
||||
let reader = UnitigFileReader::open_sequential(&layer_dir.join("unitigs.bin"))
|
||||
.map_err(OKIError::Partition)?;
|
||||
for (kmer, _, _) in reader.iter_indexed_canonical_kmers() {
|
||||
if let Some(slot) = mphf.find(kmer) {
|
||||
slot_kmer[slot] = Some(kmer);
|
||||
}
|
||||
}
|
||||
|
||||
let use_counts = with_counts && layer_dir.join("counts").exists();
|
||||
let mat = if use_counts {
|
||||
Mat::Count(PersistentCompactIntMatrix::open(layer_dir)?)
|
||||
} else {
|
||||
Mat::Presence(PersistentBitMatrix::open(layer_dir)?)
|
||||
};
|
||||
let n_cols = mat.n_cols().min(n_genomes);
|
||||
|
||||
// Per family: which single form (if exactly one) each genome
|
||||
// carries — `None` once a second form is seen (ambiguous,
|
||||
// not single-copy, ineligible for either side of a pair).
|
||||
let mut single_form: Vec<Option<u8>> = Vec::with_capacity(n_cols);
|
||||
let mut ambiguous: Vec<bool> = Vec::with_capacity(n_cols);
|
||||
|
||||
for slot in 0..annex.len() {
|
||||
let Some(mask) = annex.get(slot) else { continue };
|
||||
let Some(kmer) = slot_kmer[slot] else { continue };
|
||||
if !is_minorant(kmer, mask, k) {
|
||||
continue; // family tallied once, at its minorant
|
||||
}
|
||||
let variable = mask.family_size() >= 2;
|
||||
|
||||
single_form.clear();
|
||||
single_form.resize(n_cols, None);
|
||||
ambiguous.clear();
|
||||
ambiguous.resize(n_cols, false);
|
||||
|
||||
for other in kmer.central_canonical_neighbors() {
|
||||
let base = central_base(other, k);
|
||||
if !mask.has(base) {
|
||||
continue;
|
||||
}
|
||||
let presence: Option<Vec<bool>> = if other == kmer {
|
||||
Some((0..n_cols).map(|g| mat.carries(g, slot)).collect())
|
||||
} else {
|
||||
let dest = partition_of(other, n_parts);
|
||||
cache.find_presence(dest, other, n_genomes)
|
||||
};
|
||||
let Some(presence) = presence else { continue };
|
||||
for (g, &present) in presence.iter().enumerate() {
|
||||
if !present {
|
||||
continue;
|
||||
}
|
||||
if single_form[g].is_some() {
|
||||
ambiguous[g] = true;
|
||||
} else {
|
||||
single_form[g] = Some(base);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for i in 0..n_cols {
|
||||
if ambiguous[i] {
|
||||
continue;
|
||||
}
|
||||
let Some(bi) = single_form[i] else { continue };
|
||||
for j in (i + 1)..n_cols {
|
||||
if ambiguous[j] {
|
||||
continue;
|
||||
}
|
||||
let Some(bj) = single_form[j] else { continue };
|
||||
on_pair(&mut acc, i, j, bi, bj, variable);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pb.inc(1);
|
||||
Ok(acc)
|
||||
})
|
||||
.collect::<OKIResult<Vec<_>>>()?;
|
||||
pb.finish_and_clear();
|
||||
|
||||
let mut total = zero();
|
||||
for partial in partials {
|
||||
total = combine(total, partial);
|
||||
}
|
||||
Ok(total)
|
||||
}
|
||||
|
||||
/// Compute [`RawSnpDistanceOutput`] from an already-built sibling annex
|
||||
/// (run [`build_sibling_annex`](Self::build_sibling_annex) first).
|
||||
pub fn raw_snp_distance(&self) -> OKIResult<RawSnpDistanceOutput> {
|
||||
let n_genomes = self.meta.genomes.len();
|
||||
let (snp, shared) = self.scan_family_pairs(
|
||||
"raw_snp_distance",
|
||||
|| (Array2::<u64>::zeros((n_genomes, n_genomes)), Array2::<u64>::zeros((n_genomes, n_genomes))),
|
||||
|(snp, shared), i, j, bi, bj, _variable| {
|
||||
if bi == bj {
|
||||
shared[[i, j]] += 1;
|
||||
shared[[j, i]] += 1;
|
||||
} else {
|
||||
snp[[i, j]] += 1;
|
||||
snp[[j, i]] += 1;
|
||||
}
|
||||
},
|
||||
|(mut snp, mut shared), (s, sh)| {
|
||||
snp += &s;
|
||||
shared += &sh;
|
||||
(snp, shared)
|
||||
},
|
||||
)?;
|
||||
Ok(RawSnpDistanceOutput { snp, shared })
|
||||
}
|
||||
|
||||
/// Symmetric 6-category base-pair substitution tally (AC, AG, AT, CG,
|
||||
/// CT, GT — indexed `0=A,1=C,2=G,3=T`), pooled only over genome pairs
|
||||
/// whose overall SNP ratio in `raw` is at or below `ratio_ceiling` —
|
||||
/// same saturation-exclusion discipline as
|
||||
/// [`cardinality_tally`](Self::cardinality_tally), for the same reason:
|
||||
/// a saturated pair's observed base-pair mix trends toward neutral base
|
||||
/// composition, not the true point-mutation spectrum.
|
||||
///
|
||||
/// A second full pass over the annex, sharing
|
||||
/// [`raw_snp_distance`](Self::raw_snp_distance)'s traversal (guided by
|
||||
/// it, not a blind re-scan) — needed because `raw_snp_distance` only
|
||||
/// keeps aggregate SNP/shared counts per genome pair, not which bases
|
||||
/// were actually involved at each locus, and the ratio-ceiling filter
|
||||
/// can only be evaluated once the aggregate counts are known.
|
||||
pub fn base_pair_tally(&self, raw: &RawSnpDistanceOutput, ratio_ceiling: f64) -> OKIResult<BasePairTally> {
|
||||
let n_genomes = self.meta.genomes.len();
|
||||
let included = Array2::from_shape_fn((n_genomes, n_genomes), |(i, j)| {
|
||||
if i == j {
|
||||
return false;
|
||||
}
|
||||
let snp = raw.snp[[i, j]];
|
||||
let total = snp + raw.shared[[i, j]];
|
||||
total > 0 && (snp as f64 / total as f64) <= ratio_ceiling
|
||||
});
|
||||
|
||||
let (counts, same) = self.scan_family_pairs(
|
||||
"base_pair_tally",
|
||||
|| ([[0u64; 4]; 4], [0u64; 4]),
|
||||
|(counts, same), i, j, bi, bj, variable| {
|
||||
if !included[[i, j]] {
|
||||
return;
|
||||
}
|
||||
if bi != bj {
|
||||
counts[bi as usize][bj as usize] += 1;
|
||||
counts[bj as usize][bi as usize] += 1;
|
||||
} else if variable {
|
||||
// Only count "stayed the same" from families that vary
|
||||
// *somewhere* in the index — a fully invariant family
|
||||
// (never varies anywhere) isn't a SNP-adjacent
|
||||
// agreement, it's genome-wide background, and would
|
||||
// otherwise swamp the diagonal (see
|
||||
// `docmd/theory/evolutionary_distances.md`, the
|
||||
// ascertainment-bias regression this was reverting).
|
||||
same[bi as usize] += 1;
|
||||
}
|
||||
},
|
||||
|(mut counts, mut same), (partial_counts, partial_same)| {
|
||||
for a in 0..4 {
|
||||
same[a] += partial_same[a];
|
||||
for b in 0..4 {
|
||||
counts[a][b] += partial_counts[a][b];
|
||||
}
|
||||
}
|
||||
(counts, same)
|
||||
},
|
||||
)?;
|
||||
Ok(BasePairTally { counts, same })
|
||||
}
|
||||
}
|
||||
|
||||
/// See [`KmerIndex::base_pair_tally`].
|
||||
pub struct BasePairTally {
|
||||
/// `counts[a][b] == counts[b][a]` = number of eligible loci, pooled
|
||||
/// over included genome pairs, where the two genomes' single forms are
|
||||
/// `a` and `b` (0=A, 1=C, 2=G, 3=T). Diagonal always `0` — an `a == b`
|
||||
/// locus is counted in `same`, not here.
|
||||
pub counts: [[u64; 4]; 4],
|
||||
/// `same[a]` = number of eligible loci, pooled over included genome
|
||||
/// pairs, where both genomes' single forms are `a` — the diagonal
|
||||
/// `counts` omits, needed to build a proper row-stochastic composition
|
||||
/// probability matrix (the "stay the same base" entries), not just the
|
||||
/// substitution-cost off-diagonal.
|
||||
pub same: [u64; 4],
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
use obikseq::CanonicalKmer;
|
||||
use obiskbuilder::rolling_stat::RollingStat;
|
||||
|
||||
use obicompactvec::FamilyMask;
|
||||
|
||||
/// Central-position base of a canonical k-mer, in the fixed 0=A/1=C/2=G/3=T
|
||||
/// encoding — the mask's bit index. `k` must be odd (project invariant).
|
||||
#[inline]
|
||||
pub(super) fn central_base(kmer: CanonicalKmer, k: usize) -> u8 {
|
||||
kmer.nucleotide((k - 1) / 2)
|
||||
}
|
||||
|
||||
/// Is `kmer` the minorant of its family, given the family's presence mask?
|
||||
/// Regenerates the family's 4 canonical forms from `kmer` itself (cheap, no
|
||||
/// lookup — see the design doc's "Definitions" section for why this is
|
||||
/// always safe: the set of 4 forms is invariant regardless of which member
|
||||
/// you start from), and compares the raw encodings of whichever are marked
|
||||
/// present in `mask`.
|
||||
pub(super) fn is_minorant(kmer: CanonicalKmer, mask: FamilyMask, k: usize) -> bool {
|
||||
kmer.central_canonical_neighbors().into_iter().all(|other| {
|
||||
other == kmer || !mask.has(central_base(other, k)) || kmer.raw() <= other.raw()
|
||||
})
|
||||
}
|
||||
|
||||
/// Minimiser of a single, isolated canonical k-mer (not part of a streamed
|
||||
/// sequence). `RollingStat` computes minimisers incrementally along a
|
||||
/// sequence; this feeds one k-mer's bases through a fresh instance to get
|
||||
/// the same selection for a single, disconnected k-mer. Not the leanest
|
||||
/// possible primitive (an O(1)-amortised dedicated scan, as originally
|
||||
/// sketched in the design doc's Step 0, would avoid the ASCII round-trip and
|
||||
/// `RollingStat` allocation) but correct and reuses already-tested logic;
|
||||
/// left as a follow-up optimisation.
|
||||
fn lone_kmer_minimizer(kmer: CanonicalKmer) -> obikseq::Minimizer {
|
||||
let ascii = kmer.to_ascii();
|
||||
let mut rs = RollingStat::new(0);
|
||||
for b in ascii {
|
||||
rs.push(b);
|
||||
}
|
||||
rs.canonical_minimizer()
|
||||
.expect("RollingStat must be ready after k bases of a valid k-mer")
|
||||
}
|
||||
|
||||
/// Destination partition for a (possibly synthetic) canonical k-mer, using
|
||||
/// the same routing rule as the rest of the index (`minimiser.seq_hash() &
|
||||
/// mask`, `n_partitions` is a power of two).
|
||||
pub(super) fn partition_of(kmer: CanonicalKmer, n_partitions: usize) -> usize {
|
||||
let mask = (n_partitions as u64) - 1;
|
||||
(lone_kmer_minimizer(kmer).seq_hash() & mask) as usize
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
//! Family presence-mask annex construction.
|
||||
//!
|
||||
//! See `docmd/theory/evolutionary_distances.md`, "Definitions: family, and
|
||||
//! the canonical form of a family" and "Step 2b", for the full design
|
||||
//! discussion this implements.
|
||||
//!
|
||||
//! For each distinct k-mer of each layer of the (already built/merged)
|
||||
//! index, computes a 4-bit presence mask for its "family" (the up to 4
|
||||
//! k-mers sharing its flanks, differing only at the central base —
|
||||
//! well-defined for odd k): bit `b` set iff the family member whose own
|
||||
//! canonical central base is `b` (0=A, 1=C, 2=G, 3=T) is observed anywhere
|
||||
//! in the current multi-genome index — a property of the whole index, not
|
||||
//! of any one genome. Sibling count and minorant are *derived* from the
|
||||
//! mask by callers, not stored (see `FamilyMask` and
|
||||
//! [`sibling_annex_stats`](crate::index::KmerIndex::sibling_annex_stats)
|
||||
//! below).
|
||||
//!
|
||||
//! Per layer, an `obipipeline` batch transform (throttled — see
|
||||
//! `obipipeline::throttle`) generates a whole batch's central variants at
|
||||
//! once (`BATCH_SIZE` source k-mers in, that batch's variants out as one
|
||||
//! pipeline message), interleaved across many in-flight batches by the
|
||||
//! scheduler's shared worker pool rather than processed on a single
|
||||
//! thread. The actual cross-partition lookup reuses a `PartitionCache` of
|
||||
//! every partition's already-open MPHF layers, built once for the whole
|
||||
//! `build_sibling_annex` run, rather than reopening files per lookup or
|
||||
//! per source layer. Two earlier, coarser-grained designs were tried and
|
||||
//! measured (not guessed) to be worse, in order: (1) reopening/re-mmap'ing
|
||||
//! every target partition's files on every single lookup — fine at toy
|
||||
//! scale, ~90% system time against a real index; (2) a `Flat` pipeline
|
||||
//! stage pushing one message per generated *variant* (up to 3 per source
|
||||
//! k-mer) — cheaper than reopening files, but sampling a real run showed
|
||||
//! most wall-clock time going into per-message channel send/notify
|
||||
//! syscalls rather than the lookup itself, because a single k-mer's ≤3
|
||||
//! variants is far too fine a granularity to amortise a pipeline's
|
||||
//! synchronisation cost over. See `docmd/theory/evolutionary_distances.md`,
|
||||
//! Step 2b, "Mechanism".
|
||||
//!
|
||||
//! Submodules, in the order data flows through them: [`cache`] (shared
|
||||
//! whole-run partition cache), [`helpers`] (small pure functions used
|
||||
//! throughout), [`build`] (annex construction), [`stats`] (family-size
|
||||
//! diagnostics), [`distance`] (raw SNP distance + base-pair tally),
|
||||
//! [`cardinality`] (cardinality co-occurrence), [`alignment`] (SNP-only
|
||||
//! pseudo-alignment).
|
||||
|
||||
mod alignment;
|
||||
mod build;
|
||||
mod cache;
|
||||
mod cardinality;
|
||||
mod distance;
|
||||
mod helpers;
|
||||
mod stats;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
pub use alignment::SnpAlignment;
|
||||
pub use cardinality::CardinalityTally;
|
||||
pub use distance::{BasePairTally, RawSnpDistanceOutput};
|
||||
pub use stats::SiblingAnnexStats;
|
||||
|
||||
use obilayeredmap::OLMError;
|
||||
|
||||
use crate::error::OKIError;
|
||||
|
||||
pub(super) const INDEX_SUBDIR: &str = "index";
|
||||
pub(super) const ANNEX_FILE_NAME: &str = "siblings.psib";
|
||||
|
||||
pub(super) fn olm_to_ok(e: OLMError) -> OKIError {
|
||||
match e {
|
||||
OLMError::Io(e) => OKIError::Io(e),
|
||||
other => OKIError::InvalidInput(format!("layered-map error: {other}")),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,192 @@
|
||||
use rayon::prelude::*;
|
||||
|
||||
use obicompactvec::{PersistentBitMatrix, PersistentCompactIntMatrix, SiblingAnnex};
|
||||
use obikpartitionner::KmerPartition;
|
||||
use obikseq::CanonicalKmer;
|
||||
use obilayeredmap::MphfLayer;
|
||||
use obilayeredmap::meta::PartitionMeta;
|
||||
use obiskio::UnitigFileReader;
|
||||
use obisys::progress_bar;
|
||||
|
||||
use crate::error::{OKIError, OKIResult};
|
||||
use crate::index::KmerIndex;
|
||||
|
||||
use super::cache::{Mat, PartitionCache};
|
||||
use super::helpers::{central_base, is_minorant, partition_of};
|
||||
use super::{olm_to_ok, ANNEX_FILE_NAME, INDEX_SUBDIR};
|
||||
|
||||
/// Distribution of family sizes (1-4), read back from an already-built
|
||||
/// annex (see [`KmerIndex::build_sibling_annex`]) plus the index's
|
||||
/// presence/count data — a separate, occasional diagnostic pass, not fused
|
||||
/// into construction.
|
||||
///
|
||||
/// Every count here is **per family, not per slot**: a family with `F`
|
||||
/// members occupies `F` annex slots (one per observed member), all sharing
|
||||
/// the same mask. Counting every slot would count each family up to 4
|
||||
/// times over; only the minorant's slot is tallied (minorant is derived on
|
||||
/// the fly — see `is_minorant` — not stored, but cheap: no lookup, pure
|
||||
/// bit arithmetic on already-in-hand data).
|
||||
#[derive(Debug, Clone, Default)]
|
||||
pub struct SiblingAnnexStats {
|
||||
/// `counts[s]` = number of *families* of size `s + 1` (index 0 = size 1,
|
||||
/// i.e. 0 siblings, ... index 3 = size 4, i.e. 3 siblings).
|
||||
pub counts: [u64; 4],
|
||||
/// `per_genome[g][s]` = number of families of size `s + 1` for which
|
||||
/// genome `g` (index into `KmerIndex::meta().genomes`) carries at least
|
||||
/// one member.
|
||||
pub per_genome: Vec<[u64; 4]>,
|
||||
}
|
||||
|
||||
impl KmerIndex {
|
||||
/// Tally the family-size distribution of an already-built annex
|
||||
/// (globally, and per genome), counting each family once (at its
|
||||
/// minorant slot). Errors if [`build_sibling_annex`] has not been run on
|
||||
/// this index first.
|
||||
///
|
||||
/// [`build_sibling_annex`]: Self::build_sibling_annex
|
||||
pub fn sibling_annex_stats(&self) -> OKIResult<SiblingAnnexStats> {
|
||||
let n_parts = self.n_partitions();
|
||||
let n_genomes = self.meta.genomes.len();
|
||||
let with_counts = self.meta.config.with_counts;
|
||||
let k = self.kmer_size();
|
||||
let n_bits = n_parts.trailing_zeros() as usize;
|
||||
|
||||
// Same whole-run cache as `build_sibling_annex` — see its docs for
|
||||
// why re-opening per lookup (or per call to a batching helper) is
|
||||
// not good enough on a real index.
|
||||
let partition = KmerPartition::open_with_config(
|
||||
&self.root_path,
|
||||
self.kmer_size(),
|
||||
self.minimizer_size(),
|
||||
n_bits,
|
||||
)
|
||||
.map_err(OKIError::Partition)?;
|
||||
let cache = PartitionCache::build(&partition, n_parts, with_counts)?;
|
||||
|
||||
// Gather the (partition, layer) pairs to process — cheap metadata
|
||||
// reads only, checking every annex file exists up front so a
|
||||
// missing one is reported before any real work starts.
|
||||
let mut layer_dirs = Vec::new();
|
||||
for part in 0..n_parts {
|
||||
let index_dir = self.partition().part_dir(part).join(INDEX_SUBDIR);
|
||||
if !index_dir.exists() {
|
||||
continue;
|
||||
}
|
||||
let meta = PartitionMeta::load(&index_dir).map_err(olm_to_ok)?;
|
||||
for l in 0..meta.n_layers {
|
||||
let layer_dir = index_dir.join(format!("layer_{l}"));
|
||||
let annex_path = layer_dir.join(ANNEX_FILE_NAME);
|
||||
if !annex_path.exists() {
|
||||
return Err(OKIError::InvalidInput(format!(
|
||||
"no sibling annex at {} — run build_sibling_annex first",
|
||||
annex_path.display()
|
||||
)));
|
||||
}
|
||||
layer_dirs.push(layer_dir);
|
||||
}
|
||||
}
|
||||
|
||||
// One layer's worth of work, parallelised across layers with Rayon
|
||||
// — independent, read-only, each producing its own partial tally
|
||||
// merged at the end.
|
||||
let pb = progress_bar("sibling_annex_stats", layer_dirs.len() as u64, "layers");
|
||||
let partials: Vec<SiblingAnnexStats> = layer_dirs
|
||||
.par_iter()
|
||||
.map(|layer_dir| -> OKIResult<SiblingAnnexStats> {
|
||||
let mut stats = SiblingAnnexStats {
|
||||
per_genome: vec![[0u64; 4]; n_genomes],
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let index_dir = layer_dir.parent().expect("layer_dir has a parent index dir");
|
||||
let meta = PartitionMeta::load(index_dir).map_err(olm_to_ok)?;
|
||||
let annex = SiblingAnnex::open(&layer_dir.join(ANNEX_FILE_NAME))?;
|
||||
|
||||
// Need each slot's own k-mer to derive minorant — same
|
||||
// enumeration as construction.
|
||||
let mphf = MphfLayer::open(layer_dir, &meta.mode).map_err(olm_to_ok)?;
|
||||
let mut slot_kmer: Vec<Option<CanonicalKmer>> = vec![None; annex.len()];
|
||||
let reader = UnitigFileReader::open_sequential(&layer_dir.join("unitigs.bin"))
|
||||
.map_err(OKIError::Partition)?;
|
||||
for (kmer, _, _) in reader.iter_indexed_canonical_kmers() {
|
||||
if let Some(slot) = mphf.find(kmer) {
|
||||
slot_kmer[slot] = Some(kmer);
|
||||
}
|
||||
}
|
||||
|
||||
let use_counts = with_counts && layer_dir.join("counts").exists();
|
||||
let mat = if use_counts {
|
||||
Mat::Count(PersistentCompactIntMatrix::open(layer_dir)?)
|
||||
} else {
|
||||
Mat::Presence(PersistentBitMatrix::open(layer_dir)?)
|
||||
};
|
||||
let n_cols = mat.n_cols().min(n_genomes);
|
||||
|
||||
for slot in 0..annex.len() {
|
||||
let Some(mask) = annex.get(slot) else { continue };
|
||||
let Some(kmer) = slot_kmer[slot] else { continue };
|
||||
if !is_minorant(kmer, mask, k) {
|
||||
continue; // this family is tallied at its minorant's slot only
|
||||
}
|
||||
let s = mask.siblings() as usize;
|
||||
stats.counts[s] += 1;
|
||||
|
||||
// "Genome g represents this family" means g carries
|
||||
// *any* of its members, not just the minorant's own —
|
||||
// start from the minorant's own presence (already
|
||||
// open, no lookup) and OR in every other present
|
||||
// member's presence vector, resolved against the
|
||||
// whole-run cache (no I/O) — exactly `mask.siblings()`
|
||||
// of them, the mask tells us precisely which to fetch.
|
||||
let mut carries = vec![false; n_cols];
|
||||
for g in 0..n_cols {
|
||||
carries[g] = mat.carries(g, slot);
|
||||
}
|
||||
for other in kmer.central_canonical_neighbors() {
|
||||
if other == kmer {
|
||||
continue;
|
||||
}
|
||||
let base = central_base(other, k);
|
||||
if !mask.has(base) {
|
||||
continue;
|
||||
}
|
||||
let dest = partition_of(other, n_parts);
|
||||
if let Some(other_presence) = cache.find_presence(dest, other, n_genomes) {
|
||||
for (g, &present) in other_presence.iter().enumerate() {
|
||||
if present {
|
||||
carries[g] = true;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
for (g, &carried) in carries.iter().enumerate() {
|
||||
if carried {
|
||||
stats.per_genome[g][s] += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pb.inc(1);
|
||||
Ok(stats)
|
||||
})
|
||||
.collect::<OKIResult<Vec<_>>>()?;
|
||||
pb.finish_and_clear();
|
||||
|
||||
let mut stats = SiblingAnnexStats {
|
||||
per_genome: vec![[0u64; 4]; n_genomes],
|
||||
..Default::default()
|
||||
};
|
||||
for part in partials {
|
||||
for s in 0..4 {
|
||||
stats.counts[s] += part.counts[s];
|
||||
}
|
||||
for g in 0..n_genomes {
|
||||
for s in 0..4 {
|
||||
stats.per_genome[g][s] += part.per_genome[g][s];
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(stats)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,186 @@
|
||||
use std::io::Write;
|
||||
use std::path::Path;
|
||||
|
||||
use obicompactvec::{FamilyMask, SiblingAnnex};
|
||||
use obikseq::{CanonicalKmer, Kmer, Sequence};
|
||||
use obilayeredmap::MphfLayer;
|
||||
use obilayeredmap::meta::PartitionMeta;
|
||||
use obisys::Reporter;
|
||||
use tempfile::tempdir;
|
||||
|
||||
use crate::index::KmerIndex;
|
||||
use crate::meta::{GenomeInfo, IndexConfig};
|
||||
use crate::merge::MergeMode;
|
||||
|
||||
use super::helpers::is_minorant;
|
||||
use super::{ANNEX_FILE_NAME, INDEX_SUBDIR};
|
||||
|
||||
// k must be >= 11 (project constraint, "k ∈ [11,31]"); k=11, level_max=1,
|
||||
// theta=0.0 mirror `obiskbuilder`'s own tests (smaller k/level_max
|
||||
// combinations trip an unrelated pre-existing bug in `obikentropy`'s
|
||||
// sliding-window ring buffer — not this feature's concern).
|
||||
const K: usize = 11;
|
||||
const M: usize = 5;
|
||||
|
||||
/// Build a single-genome index from one in-memory FASTA sequence, driving
|
||||
/// the same primitives `obikmer`'s `scatter` step uses (minus the
|
||||
/// multi-file `obipipeline` wrapper — a single sequence needs none of
|
||||
/// that): normalise -> build superkmers -> route -> write.
|
||||
/// `cargo test` doesn't install a `tracing` subscriber the way `obikmer`'s
|
||||
/// CLI does, so `debug!`/etc. are silent no-ops by default — including the
|
||||
/// `PartitionRunner` instrumentation that would matter most for
|
||||
/// re-diagnosing a hang here. `try_init` is idempotent across concurrently
|
||||
/// running tests (later calls just find a subscriber already installed).
|
||||
fn init_tracing() {
|
||||
let _ = tracing_subscriber::fmt()
|
||||
.with_env_filter(
|
||||
tracing_subscriber::EnvFilter::try_from_default_env()
|
||||
.unwrap_or_else(|_| tracing_subscriber::EnvFilter::new("info")),
|
||||
)
|
||||
.with_writer(std::io::stderr)
|
||||
.try_init();
|
||||
}
|
||||
|
||||
fn build_single_genome_index(dir: &Path, label: &str, seq: &[u8]) -> KmerIndex {
|
||||
init_tracing();
|
||||
let fasta_path = dir.join(format!("{label}.fasta"));
|
||||
let mut f = std::fs::File::create(&fasta_path).unwrap();
|
||||
writeln!(f, ">{label}").unwrap();
|
||||
f.write_all(seq).unwrap();
|
||||
writeln!(f).unwrap();
|
||||
drop(f);
|
||||
|
||||
let index_path = dir.join(format!("{label}.idx"));
|
||||
let config = IndexConfig {
|
||||
kmer_size: K,
|
||||
minimizer_size: M,
|
||||
n_bits: 0, // 1 partition — keeps the test deterministic and simple
|
||||
with_counts: false,
|
||||
evidence: obilayeredmap::IndexMode::Exact,
|
||||
block_bits: 0,
|
||||
};
|
||||
let mut idx = KmerIndex::create(&index_path, config, Some(GenomeInfo::new(label)), false)
|
||||
.expect("create");
|
||||
|
||||
let mut rep = Reporter::new();
|
||||
let stream = obiread::open_nuc_stream(fasta_path.to_str().unwrap(), K).expect("open fasta");
|
||||
for page in stream {
|
||||
let batch = obiskbuilder::build_superkmers_page(page, K, /* level_max */ 1, /* theta */ 0.0);
|
||||
idx.partition_mut().write_batch(batch).expect("write_batch");
|
||||
}
|
||||
idx.partition_mut().close().expect("close partition writers");
|
||||
idx.mark_scattered().expect("mark_scattered");
|
||||
idx.dereplicate_and_count(false, &mut rep).expect("dereplicate_and_count");
|
||||
idx.build_layers(1, None, false, &mut rep).expect("build_layers");
|
||||
idx
|
||||
}
|
||||
|
||||
fn canonical(ascii: &[u8]) -> CanonicalKmer {
|
||||
Kmer::from_ascii(ascii).unwrap().canonical()
|
||||
}
|
||||
|
||||
/// Read back the annex entry for a given canonical k-mer from the merged
|
||||
/// index's (single) partition/layer, asserting it was found at all.
|
||||
fn annex_info_for(idx: &KmerIndex, kmer: CanonicalKmer) -> FamilyMask {
|
||||
let index_dir = idx.partition().part_dir(0).join(INDEX_SUBDIR);
|
||||
let meta = PartitionMeta::load(&index_dir).unwrap();
|
||||
for l in 0..meta.n_layers {
|
||||
let layer_dir = index_dir.join(format!("layer_{l}"));
|
||||
let mphf = MphfLayer::open(&layer_dir, &meta.mode).unwrap();
|
||||
if let Some(slot) = mphf.find(kmer) {
|
||||
let annex = SiblingAnnex::open(&layer_dir.join(ANNEX_FILE_NAME)).unwrap();
|
||||
return annex.get(slot).expect("slot must have a computed annex entry");
|
||||
}
|
||||
}
|
||||
panic!("kmer not found in any layer of partition 0");
|
||||
}
|
||||
|
||||
fn merge_two(dir: &Path, g1: &KmerIndex, g2: &KmerIndex) -> KmerIndex {
|
||||
let mut rep = Reporter::new();
|
||||
KmerIndex::merge(
|
||||
&dir.join("merged.idx"),
|
||||
&[g1, g2],
|
||||
MergeMode::Presence,
|
||||
false,
|
||||
false,
|
||||
1.0,
|
||||
&mut rep,
|
||||
)
|
||||
.expect("merge")
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sibling_annex_one_sibling_each() {
|
||||
// k=11, centre = index 5 (0-based). Two genomes, each exactly one
|
||||
// k-mer, sharing every base except the centre:
|
||||
// g1 = "AACCGCTTAAG" (centre 'C', base index 1)
|
||||
// g2 = "AACCGGTTAAG" (centre 'G', base index 2)
|
||||
// Hand-verified: both stay forward-oriented under canonicalisation
|
||||
// (each is lexicographically smaller than its own reverse
|
||||
// complement, since both start with "AA"), and raw(g1) < raw(g2)
|
||||
// (only differing base: C=0b01 < G=0b10 at the centre) — so g1 is
|
||||
// the minorant, g2 is not. The mask is a family-wide value: both
|
||||
// slots must read back the *same* mask (bits 1 and 2 set).
|
||||
let dir = tempdir().unwrap();
|
||||
let g1 = build_single_genome_index(dir.path(), "g1", b"AACCGCTTAAG");
|
||||
let g2 = build_single_genome_index(dir.path(), "g2", b"AACCGGTTAAG");
|
||||
let merged = merge_two(dir.path(), &g1, &g2);
|
||||
merged.build_sibling_annex().expect("build_sibling_annex");
|
||||
|
||||
let g1_kmer = canonical(b"AACCGCTTAAG");
|
||||
let g2_kmer = canonical(b"AACCGGTTAAG");
|
||||
let expected_mask = FamilyMask::EMPTY.with(1).with(2);
|
||||
|
||||
let a = annex_info_for(&merged, g1_kmer);
|
||||
assert_eq!(a, expected_mask, "AACCGCTTAAG");
|
||||
assert_eq!(a.siblings(), 1);
|
||||
assert!(is_minorant(g1_kmer, a, K), "g1 should be the minorant");
|
||||
|
||||
let b = annex_info_for(&merged, g2_kmer);
|
||||
assert_eq!(b, expected_mask, "AACCGGTTAAG");
|
||||
assert_eq!(b.siblings(), 1);
|
||||
assert!(!is_minorant(g2_kmer, b, K), "g2 should not be the minorant");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sibling_annex_zero_siblings_when_identical_and_no_variant_exists() {
|
||||
// Same k-mer in both genomes, no other genome around to carry a
|
||||
// variant -> 0 siblings, trivially its own minorant.
|
||||
let dir = tempdir().unwrap();
|
||||
let g1 = build_single_genome_index(dir.path(), "g1", b"GATTACAGATC");
|
||||
let g2 = build_single_genome_index(dir.path(), "g2", b"GATTACAGATC");
|
||||
let merged = merge_two(dir.path(), &g1, &g2);
|
||||
merged.build_sibling_annex().expect("build_sibling_annex");
|
||||
|
||||
let kmer = canonical(b"GATTACAGATC");
|
||||
let mask = annex_info_for(&merged, kmer);
|
||||
assert_eq!(mask.siblings(), 0, "GATTACAGATC");
|
||||
assert_eq!(mask.family_size(), 1);
|
||||
assert!(is_minorant(kmer, mask, K));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sibling_annex_stats_counts_each_family_once_and_per_genome() {
|
||||
// Reuses the one-sibling-each fixture: a single family of size 2
|
||||
// (g1's centre-C form + g2's centre-G form), each genome carrying
|
||||
// exactly one of the two members. Stats must report exactly one
|
||||
// family of size 2 (`counts[1] == 1`, since index 1 = size 2), not
|
||||
// two (which naively summing both slots would give), and both
|
||||
// genomes represented at size 2, neither at any other size.
|
||||
let dir = tempdir().unwrap();
|
||||
let g1 = build_single_genome_index(dir.path(), "g1", b"AACCGCTTAAG");
|
||||
let g2 = build_single_genome_index(dir.path(), "g2", b"AACCGGTTAAG");
|
||||
let merged = merge_two(dir.path(), &g1, &g2);
|
||||
merged.build_sibling_annex().expect("build_sibling_annex");
|
||||
|
||||
let stats = merged.sibling_annex_stats().expect("sibling_annex_stats");
|
||||
|
||||
assert_eq!(stats.counts, [0, 1, 0, 0], "one family of size 2, counted once");
|
||||
assert_eq!(stats.per_genome.len(), 2);
|
||||
for g in 0..2 {
|
||||
assert_eq!(
|
||||
stats.per_genome[g], [0, 1, 0, 0],
|
||||
"genome {g} should represent exactly one size-2 family"
|
||||
);
|
||||
}
|
||||
}
|
||||
+14
-87
@@ -1,5 +1,14 @@
|
||||
//! Merging a source partition's new layer into a destination partition:
|
||||
//! de Bruijn graph union (pass 1) then column fill (pass 2).
|
||||
//!
|
||||
//! Submodules: [`src_layer`] (`SrcLayerData`, the opened-source-matrix
|
||||
//! lookup used by pass 2 here and by `rebuild_layer`). The `merge_partition`
|
||||
//! orchestration itself stays in this file — its ~400-line body is one
|
||||
//! tightly threaded pipeline (shared `Arc`/`Mutex` state across pass 1,
|
||||
//! builder setup, and pass 2), not a set of independently callable steps.
|
||||
|
||||
use std::fs;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::path::PathBuf;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use tracing::debug;
|
||||
@@ -10,7 +19,6 @@ use obipipeline::{
|
||||
};
|
||||
|
||||
use obicompactvec::{
|
||||
MatrixGroupOps,
|
||||
PersistentBitMatrix, PersistentBitMatrixBuilder, PersistentBitVecBuilder,
|
||||
PersistentCompactIntMatrix, PersistentCompactIntMatrixBuilder, PersistentCompactIntVecBuilder,
|
||||
};
|
||||
@@ -23,6 +31,10 @@ use crate::common::{ColBuilder, col_path_bit, col_path_int, load_meta, olm_to_sk
|
||||
use crate::graph_pipeline::{build_graph, materialize_layer};
|
||||
use crate::partition::KmerPartition;
|
||||
|
||||
mod src_layer;
|
||||
|
||||
pub(crate) use src_layer::SrcLayerData;
|
||||
|
||||
// ── MergeMode ─────────────────────────────────────────────────────────────────
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
@@ -31,91 +43,6 @@ pub enum MergeMode {
|
||||
Count,
|
||||
}
|
||||
|
||||
// ── SrcLayerData — opened source matrix for pass-2 lookup ─────────────────────
|
||||
|
||||
pub(crate) enum SrcLayerData {
|
||||
Presence(MphfOnly, PersistentBitMatrix),
|
||||
Count(MphfOnly, PersistentCompactIntMatrix),
|
||||
}
|
||||
|
||||
impl SrcLayerData {
|
||||
pub(crate) fn open(layer_dir: &Path, merge_mode: MergeMode) -> SKResult<Self> {
|
||||
let counts_dir = layer_dir.join("counts");
|
||||
match merge_mode {
|
||||
MergeMode::Presence => {
|
||||
if counts_dir.exists() && !layer_dir.join("presence").exists() {
|
||||
let mphf = MphfOnly::open(layer_dir).map_err(|e| olm_to_sk(e, "merge"))?;
|
||||
let mat = PersistentCompactIntMatrix::open(layer_dir).map_err(SKError::Io)?;
|
||||
Ok(SrcLayerData::Count(mphf, mat))
|
||||
} else {
|
||||
// presence dir exists, or neither exists → Implicit handled by open()
|
||||
let mphf = MphfOnly::open(layer_dir).map_err(|e| olm_to_sk(e, "merge"))?;
|
||||
let mat = PersistentBitMatrix::open(layer_dir).map_err(SKError::Io)?;
|
||||
Ok(SrcLayerData::Presence(mphf, mat))
|
||||
}
|
||||
}
|
||||
MergeMode::Count => {
|
||||
let mphf = MphfOnly::open(layer_dir).map_err(|e| olm_to_sk(e, "merge"))?;
|
||||
if counts_dir.exists() {
|
||||
let mat = PersistentCompactIntMatrix::open(layer_dir).map_err(SKError::Io)?;
|
||||
Ok(SrcLayerData::Count(mphf, mat))
|
||||
} else {
|
||||
// No counts → treat as implicit presence (all 1s)
|
||||
let mat = PersistentBitMatrix::open(layer_dir).map_err(SKError::Io)?;
|
||||
Ok(SrcLayerData::Presence(mphf, mat))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Return one value per source genome for `kmer`.
|
||||
/// The caller guarantees `kmer` is in the source MPHF domain.
|
||||
#[inline]
|
||||
pub(crate) fn lookup(&self, kmer: CanonicalKmer, n_genomes: usize) -> Vec<u32> {
|
||||
let mut buf = vec![0u32; n_genomes];
|
||||
match self {
|
||||
SrcLayerData::Presence(mphf, mat) => mat.fill_row(mphf.index(kmer), &mut buf),
|
||||
SrcLayerData::Count(mphf, mat) => mat.fill_row(mphf.index(kmer), &mut buf),
|
||||
}
|
||||
buf
|
||||
}
|
||||
|
||||
pub(crate) fn n_slots(&self) -> usize {
|
||||
match self {
|
||||
SrcLayerData::Presence(_, mat) => mat.n(),
|
||||
SrcLayerData::Count(_, mat) => mat.n(),
|
||||
}
|
||||
}
|
||||
|
||||
/// MPHF lookup: returns the slot index for `kmer` (kmer must be in the domain).
|
||||
#[inline]
|
||||
pub(crate) fn slot(&self, kmer: CanonicalKmer) -> usize {
|
||||
match self {
|
||||
SrcLayerData::Presence(mphf, _) => mphf.index(kmer),
|
||||
SrcLayerData::Count(mphf, _) => mphf.index(kmer),
|
||||
}
|
||||
}
|
||||
|
||||
/// Row lookup by slot index, bypassing the MPHF.
|
||||
#[inline]
|
||||
pub(crate) fn fill_row_by_slot(&self, slot: usize, n_genomes: usize) -> Vec<u32> {
|
||||
let mut buf = vec![0u32; n_genomes];
|
||||
match self {
|
||||
SrcLayerData::Presence(_, mat) => mat.fill_row(slot, &mut buf),
|
||||
SrcLayerData::Count(_, mat) => mat.fill_row(slot, &mut buf),
|
||||
}
|
||||
buf
|
||||
}
|
||||
|
||||
/// Call `f` with a reference to the underlying matrix as `&dyn MatrixGroupOps`.
|
||||
pub(crate) fn with_matrix<R>(&self, f: impl FnOnce(&dyn MatrixGroupOps) -> R) -> R {
|
||||
match self {
|
||||
SrcLayerData::Presence(_, mat) => f(mat),
|
||||
SrcLayerData::Count(_, mat) => f(mat),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── helpers ───────────────────────────────────────────────────────────────────
|
||||
|
||||
const INDEX_SUBDIR: &str = "index";
|
||||
@@ -0,0 +1,95 @@
|
||||
use std::path::Path;
|
||||
|
||||
use obicompactvec::{MatrixGroupOps, PersistentBitMatrix, PersistentCompactIntMatrix};
|
||||
use obikseq::CanonicalKmer;
|
||||
use obilayeredmap::MphfOnly;
|
||||
use obiskio::{SKError, SKResult};
|
||||
|
||||
use crate::common::olm_to_sk;
|
||||
|
||||
use super::MergeMode;
|
||||
|
||||
// ── SrcLayerData — opened source matrix for pass-2 lookup ─────────────────────
|
||||
|
||||
pub(crate) enum SrcLayerData {
|
||||
Presence(MphfOnly, PersistentBitMatrix),
|
||||
Count(MphfOnly, PersistentCompactIntMatrix),
|
||||
}
|
||||
|
||||
impl SrcLayerData {
|
||||
pub(crate) fn open(layer_dir: &Path, merge_mode: MergeMode) -> SKResult<Self> {
|
||||
let counts_dir = layer_dir.join("counts");
|
||||
match merge_mode {
|
||||
MergeMode::Presence => {
|
||||
if counts_dir.exists() && !layer_dir.join("presence").exists() {
|
||||
let mphf = MphfOnly::open(layer_dir).map_err(|e| olm_to_sk(e, "merge"))?;
|
||||
let mat = PersistentCompactIntMatrix::open(layer_dir).map_err(SKError::Io)?;
|
||||
Ok(SrcLayerData::Count(mphf, mat))
|
||||
} else {
|
||||
// presence dir exists, or neither exists → Implicit handled by open()
|
||||
let mphf = MphfOnly::open(layer_dir).map_err(|e| olm_to_sk(e, "merge"))?;
|
||||
let mat = PersistentBitMatrix::open(layer_dir).map_err(SKError::Io)?;
|
||||
Ok(SrcLayerData::Presence(mphf, mat))
|
||||
}
|
||||
}
|
||||
MergeMode::Count => {
|
||||
let mphf = MphfOnly::open(layer_dir).map_err(|e| olm_to_sk(e, "merge"))?;
|
||||
if counts_dir.exists() {
|
||||
let mat = PersistentCompactIntMatrix::open(layer_dir).map_err(SKError::Io)?;
|
||||
Ok(SrcLayerData::Count(mphf, mat))
|
||||
} else {
|
||||
// No counts → treat as implicit presence (all 1s)
|
||||
let mat = PersistentBitMatrix::open(layer_dir).map_err(SKError::Io)?;
|
||||
Ok(SrcLayerData::Presence(mphf, mat))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Return one value per source genome for `kmer`.
|
||||
/// The caller guarantees `kmer` is in the source MPHF domain.
|
||||
#[inline]
|
||||
pub(crate) fn lookup(&self, kmer: CanonicalKmer, n_genomes: usize) -> Vec<u32> {
|
||||
let mut buf = vec![0u32; n_genomes];
|
||||
match self {
|
||||
SrcLayerData::Presence(mphf, mat) => mat.fill_row(mphf.index(kmer), &mut buf),
|
||||
SrcLayerData::Count(mphf, mat) => mat.fill_row(mphf.index(kmer), &mut buf),
|
||||
}
|
||||
buf
|
||||
}
|
||||
|
||||
pub(crate) fn n_slots(&self) -> usize {
|
||||
match self {
|
||||
SrcLayerData::Presence(_, mat) => mat.n(),
|
||||
SrcLayerData::Count(_, mat) => mat.n(),
|
||||
}
|
||||
}
|
||||
|
||||
/// MPHF lookup: returns the slot index for `kmer` (kmer must be in the domain).
|
||||
#[inline]
|
||||
pub(crate) fn slot(&self, kmer: CanonicalKmer) -> usize {
|
||||
match self {
|
||||
SrcLayerData::Presence(mphf, _) => mphf.index(kmer),
|
||||
SrcLayerData::Count(mphf, _) => mphf.index(kmer),
|
||||
}
|
||||
}
|
||||
|
||||
/// Row lookup by slot index, bypassing the MPHF.
|
||||
#[inline]
|
||||
pub(crate) fn fill_row_by_slot(&self, slot: usize, n_genomes: usize) -> Vec<u32> {
|
||||
let mut buf = vec![0u32; n_genomes];
|
||||
match self {
|
||||
SrcLayerData::Presence(_, mat) => mat.fill_row(slot, &mut buf),
|
||||
SrcLayerData::Count(_, mat) => mat.fill_row(slot, &mut buf),
|
||||
}
|
||||
buf
|
||||
}
|
||||
|
||||
/// Call `f` with a reference to the underlying matrix as `&dyn MatrixGroupOps`.
|
||||
pub(crate) fn with_matrix<R>(&self, f: impl FnOnce(&dyn MatrixGroupOps) -> R) -> R {
|
||||
match self {
|
||||
SrcLayerData::Presence(_, mat) => f(mat),
|
||||
SrcLayerData::Count(_, mat) => f(mat),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1,662 +0,0 @@
|
||||
use std::collections::{BTreeMap, HashMap};
|
||||
use std::fs;
|
||||
use std::io;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::time::Instant;
|
||||
use tracing::debug;
|
||||
|
||||
use obisys::progress_bar;
|
||||
|
||||
use cacheline_ef::{CachelineEf, CachelineEfVec};
|
||||
use epserde::ser::Serialize as EpSerialize;
|
||||
use memmap2::Mmap;
|
||||
use obicompactvec::PersistentCompactIntVecBuilder;
|
||||
use obikseq::RoutableSuperKmer;
|
||||
use obikseq::Sequence;
|
||||
use obikseq::superkmer::SuperKmer;
|
||||
use obiskio::{SKFileMeta, SKFileReader, SKFileWriter, SKResult};
|
||||
use ptr_hash::{PtrHash, PtrHashParams, bucket_fn::CubicEps, hash::Xx64};
|
||||
use rayon::prelude::*;
|
||||
use remove_dir_all::remove_dir_all;
|
||||
use sysinfo::System;
|
||||
|
||||
use niffler::Level;
|
||||
use niffler::send::compression::Format;
|
||||
|
||||
use crate::kmer_sort::{chunk_size_from_ram, sort_unique_kmers};
|
||||
|
||||
type Mphf = PtrHash<u64, CubicEps, CachelineEfVec<Vec<CachelineEf>>, Xx64, Vec<u8>>;
|
||||
|
||||
pub struct KmerSpectrum {
|
||||
pub f0: u64,
|
||||
pub f1: u64,
|
||||
pub counts: BTreeMap<u32, u64>,
|
||||
}
|
||||
|
||||
const SK_EXT: &str = "skmer.zst";
|
||||
pub const PARTITIONS_SUBDIR: &str = "partitions";
|
||||
|
||||
pub struct KmerPartition {
|
||||
root_path: PathBuf,
|
||||
n_partitions: usize,
|
||||
partitions_mask: u64,
|
||||
kmer_size: usize,
|
||||
minimizer_size: usize,
|
||||
writers: Vec<Option<SKFileWriter>>,
|
||||
level: Level,
|
||||
closed: bool,
|
||||
}
|
||||
|
||||
impl KmerPartition {
|
||||
pub fn create<P: AsRef<Path>>(
|
||||
path: P,
|
||||
n_bits: usize,
|
||||
kmer_size: usize,
|
||||
minimizer_size: usize,
|
||||
force: bool,
|
||||
) -> SKResult<Self> {
|
||||
Self::create_with(path, n_bits, kmer_size, minimizer_size, Level::One, force)
|
||||
}
|
||||
|
||||
pub fn create_with<P: AsRef<Path>>(
|
||||
path: P,
|
||||
n_bits: usize,
|
||||
kmer_size: usize,
|
||||
minimizer_size: usize,
|
||||
level: Level,
|
||||
force: bool,
|
||||
) -> SKResult<Self> {
|
||||
let root_path = path.as_ref().to_owned();
|
||||
if root_path.exists() {
|
||||
if force {
|
||||
remove_dir_all(&root_path)?;
|
||||
} else {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::AlreadyExists,
|
||||
format!(
|
||||
"{}: partition directory already exists",
|
||||
root_path.display()
|
||||
),
|
||||
)
|
||||
.into());
|
||||
}
|
||||
}
|
||||
fs::create_dir_all(root_path.join(PARTITIONS_SUBDIR))?;
|
||||
let n_partitions = 1usize << n_bits;
|
||||
let writers = (0..n_partitions).map(|_| None).collect();
|
||||
let partition = Self {
|
||||
root_path,
|
||||
n_partitions,
|
||||
partitions_mask: (1u64 << n_bits) - 1,
|
||||
kmer_size,
|
||||
minimizer_size,
|
||||
writers,
|
||||
level,
|
||||
closed: false,
|
||||
};
|
||||
Ok(partition)
|
||||
}
|
||||
|
||||
pub fn open_with_config<P: AsRef<Path>>(
|
||||
path: P,
|
||||
kmer_size: usize,
|
||||
minimizer_size: usize,
|
||||
n_bits: usize,
|
||||
) -> SKResult<Self> {
|
||||
let root_path = path.as_ref().to_owned();
|
||||
if !root_path.exists() {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::NotFound,
|
||||
format!("{}: partition directory not found", root_path.display()),
|
||||
)
|
||||
.into());
|
||||
}
|
||||
let n_partitions = 1usize << n_bits;
|
||||
let writers = (0..n_partitions).map(|_| None).collect();
|
||||
Ok(Self {
|
||||
root_path,
|
||||
n_partitions,
|
||||
partitions_mask: (1u64 << n_bits) - 1,
|
||||
kmer_size,
|
||||
minimizer_size,
|
||||
writers,
|
||||
level: Level::One,
|
||||
closed: true,
|
||||
})
|
||||
}
|
||||
|
||||
/// Route and write one super-kmer to its partition file.
|
||||
pub fn write(&mut self, rsk: RoutableSuperKmer) -> SKResult<()> {
|
||||
self.check_not_closed()?;
|
||||
let partition = (rsk.minimizer().seq_hash() & self.partitions_mask) as usize;
|
||||
let sk = rsk.into_superkmer();
|
||||
self.ensure_writer(partition)?.write(&sk)
|
||||
}
|
||||
|
||||
/// Route and write a batch of super-kmers.
|
||||
pub fn write_batch(&mut self, rsks: Vec<RoutableSuperKmer>) -> SKResult<()> {
|
||||
self.check_not_closed()?;
|
||||
for rsk in rsks {
|
||||
let partition = (rsk.minimizer().seq_hash() & self.partitions_mask) as usize;
|
||||
let sk = rsk.into_superkmer();
|
||||
self.ensure_writer(partition)?.write(&sk)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn flush(&mut self) -> SKResult<()> {
|
||||
self.check_not_closed()?;
|
||||
for writer in self.writers.iter_mut().flatten() {
|
||||
writer.flush()?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn close(&mut self) -> SKResult<()> {
|
||||
if self.closed {
|
||||
return Ok(());
|
||||
}
|
||||
self.closed = true;
|
||||
for writer in self.writers.iter_mut().flatten() {
|
||||
writer.close()?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn is_open(&self) -> bool {
|
||||
!self.closed
|
||||
}
|
||||
|
||||
pub fn path(&self) -> &Path {
|
||||
&self.root_path
|
||||
}
|
||||
|
||||
/// Path of partition `i` directory.
|
||||
pub fn part_dir(&self, i: usize) -> PathBuf {
|
||||
self.root_path.join(PARTITIONS_SUBDIR).join(format!("part_{i:05}"))
|
||||
}
|
||||
|
||||
pub fn kmer_size(&self) -> usize {
|
||||
self.kmer_size
|
||||
}
|
||||
|
||||
pub fn minimizer_size(&self) -> usize {
|
||||
self.minimizer_size
|
||||
}
|
||||
|
||||
pub fn n_partitions(&self) -> usize {
|
||||
self.n_partitions
|
||||
}
|
||||
|
||||
/// Deduplicate all `raw.{ext}` files in parallel, replacing each with a
|
||||
/// `dereplicated.{ext}` file where identical canonical sequences are merged
|
||||
/// and their counts summed.
|
||||
///
|
||||
/// Each partition file is processed in two phases to bound memory use:
|
||||
///
|
||||
/// 1. **Split** — the raw file is scattered into `2^temp_bits` temporary
|
||||
/// files routed by `hash(canonical_seq) & temp_mask`. Because duplicates
|
||||
/// always share the same hash, they always land in the same temp file.
|
||||
/// 2. **Merge** — each temp file is loaded fully into a `HashMap`, counts
|
||||
/// are accumulated in `u64` (no 24-bit overflow risk), and the result is
|
||||
/// appended to `dereplicated.{ext}`.
|
||||
///
|
||||
/// If a merged count exceeds the 24-bit header limit, the sequence is
|
||||
/// emitted as multiple records whose counts sum to the true total.
|
||||
///
|
||||
/// `temp_bits` controls the split fan-out (`2^temp_bits` temp files per
|
||||
/// partition). Higher values reduce per-temp-file memory at the cost of
|
||||
/// more temporary file descriptors — all managed by the global fd pool.
|
||||
pub fn dereplicate(&self) -> SKResult<()> {
|
||||
let level = self.level;
|
||||
let sys = System::new_all();
|
||||
// available_memory() can return 0 on macOS when the compressor page count exceeds
|
||||
// free+inactive+purgeable pages (sysinfo saturating_sub). Fall back to half of total.
|
||||
let available = match sys.available_memory() {
|
||||
0 => sys.total_memory() / 2,
|
||||
n => n,
|
||||
};
|
||||
let n_threads = rayon::current_num_threads().max(1) as u64;
|
||||
let available_per_thread = available / n_threads;
|
||||
|
||||
let pb = progress_bar("dereplication", self.n_partitions as u64, "partitions");
|
||||
|
||||
let results: Vec<SKResult<()>> = (0..self.n_partitions)
|
||||
.into_par_iter()
|
||||
.map(|i| {
|
||||
let dir = self.part_dir(i);
|
||||
if !dir.exists() {
|
||||
pb.inc(1);
|
||||
return Ok(());
|
||||
}
|
||||
let raw_path = dir.join(format!("raw.{SK_EXT}"));
|
||||
let t = Instant::now();
|
||||
let n_buckets = optimal_buckets(&raw_path, available_per_thread);
|
||||
let result = dereplicate_partition(&dir, level, n_buckets);
|
||||
pb.set_message(format!("last {:.0}ms", t.elapsed().as_millis()));
|
||||
pb.inc(1);
|
||||
result
|
||||
})
|
||||
.collect();
|
||||
|
||||
pb.finish_and_clear();
|
||||
for r in results {
|
||||
r?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// For each partition that has a `dereplicated.{ext}` file:
|
||||
/// 1. Enumerates all unique canonical kmers (two passes over the file).
|
||||
/// 2. Builds a provisional MPHF (FMPHGO) over those kmers.
|
||||
/// 3. Writes a flat binary count file (`counts1.bin`, one `u32` per slot,
|
||||
/// memory-mapped) accumulating kmer abundances from the superkmer counts.
|
||||
/// 4. Persists the MPHF to `mphf1.bin` for downstream use.
|
||||
///
|
||||
/// Returns the aggregated `KmerSpectrum`. Per-partition spectrum files are
|
||||
/// deleted after aggregation unless `keep_partial` is true.
|
||||
///
|
||||
/// Partitions are processed in parallel via Rayon (one task per thread).
|
||||
/// Peak memory per partition is ~80 MB, so n_threads partitions run simultaneously.
|
||||
pub fn count_kmer(&self, keep_partial: bool) -> SKResult<KmerSpectrum> {
|
||||
let sys = System::new_all();
|
||||
let available = match sys.available_memory() {
|
||||
0 => sys.total_memory() / 2,
|
||||
n => n,
|
||||
};
|
||||
let n_threads = rayon::current_num_threads().max(1) as u64;
|
||||
let chunk_kmers = chunk_size_from_ram(available / n_threads);
|
||||
|
||||
let pb = progress_bar("counting", self.n_partitions as u64, "partitions");
|
||||
|
||||
let results: Vec<SKResult<()>> = (0..self.n_partitions)
|
||||
.into_par_iter()
|
||||
.map(|i| {
|
||||
let dir = self.part_dir(i);
|
||||
let dedup_path = dir.join(format!("dereplicated.{SK_EXT}"));
|
||||
if !dedup_path.exists() {
|
||||
pb.inc(1);
|
||||
return Ok(());
|
||||
}
|
||||
let t = Instant::now();
|
||||
let result = count_partition(&dir, &dedup_path, chunk_kmers);
|
||||
pb.set_message(format!("last {:.0}ms", t.elapsed().as_millis()));
|
||||
pb.inc(1);
|
||||
result
|
||||
})
|
||||
.collect();
|
||||
|
||||
pb.finish_and_clear();
|
||||
for r in results {
|
||||
r?;
|
||||
}
|
||||
|
||||
// Aggregate per-partition spectra.
|
||||
let mut counts: BTreeMap<u32, u64> = BTreeMap::new();
|
||||
let mut f0: u64 = 0;
|
||||
let mut f1: u64 = 0;
|
||||
|
||||
for i in 0..self.n_partitions {
|
||||
let path = self.part_dir(i).join("kmer_spectrum_raw.json");
|
||||
if !path.exists() {
|
||||
continue;
|
||||
}
|
||||
let v: serde_json::Value =
|
||||
serde_json::from_str(&fs::read_to_string(&path)?).map_err(io::Error::other)?;
|
||||
f0 += v["f0"].as_u64().unwrap_or(0);
|
||||
f1 += v["f1"].as_u64().unwrap_or(0);
|
||||
if let Some(obj) = v["spectrum"].as_object() {
|
||||
for (c_str, freq) in obj {
|
||||
if let (Ok(c), Some(f)) = (c_str.parse::<u32>(), freq.as_u64()) {
|
||||
*counts.entry(c).or_insert(0) += f;
|
||||
}
|
||||
}
|
||||
}
|
||||
if !keep_partial {
|
||||
let _ = fs::remove_file(&path);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(KmerSpectrum { f0, f1, counts })
|
||||
}
|
||||
|
||||
// ── private ───────────────────────────────────────────────────────────────
|
||||
|
||||
fn check_not_closed(&self) -> SKResult<()> {
|
||||
if self.closed {
|
||||
Err(io::Error::new(io::ErrorKind::BrokenPipe, "write to closed KmerPartition").into())
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn ensure_writer(&mut self, partition: usize) -> SKResult<&mut SKFileWriter> {
|
||||
if self.writers[partition].is_none() {
|
||||
let dir = self.root_path.join(PARTITIONS_SUBDIR).join(format!("part_{:05}", partition));
|
||||
fs::create_dir_all(&dir)?;
|
||||
let file_path = dir.join(format!("raw.{SK_EXT}"));
|
||||
let writer = SKFileWriter::create_with(file_path, Format::Zstd, self.level)?;
|
||||
self.writers[partition] = Some(writer);
|
||||
}
|
||||
Ok(self.writers[partition].as_mut().unwrap())
|
||||
}
|
||||
}
|
||||
|
||||
// ── free helpers ─────────────────────────────────────────────────────────────
|
||||
|
||||
/// Estimate the number of in-memory buckets needed to deduplicate the partition
|
||||
/// file at `raw_path` given `available_bytes` of free RAM.
|
||||
///
|
||||
/// Memory per HashMap entry:
|
||||
/// key Box (1 + avg_seq_bytes) + SuperKmer header (4 B) + avg seq bytes + u64 count (8 B),
|
||||
/// multiplied by 1.5 for hashbrown load-factor overhead.
|
||||
///
|
||||
/// Returns 1 if the partition fits comfortably in memory (no split needed).
|
||||
/// Always returns a power of two.
|
||||
/// Remove a SuperKmer file and its sidecar (if present).
|
||||
fn remove_skmer_file(path: &Path) -> SKResult<()> {
|
||||
fs::remove_file(path)?;
|
||||
let sidecar = SKFileMeta::sidecar_path(path);
|
||||
match fs::remove_file(&sidecar) {
|
||||
Ok(()) => {}
|
||||
Err(e) if e.kind() == io::ErrorKind::NotFound => {}
|
||||
Err(e) => return Err(e.into()),
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn optimal_buckets(raw_path: &Path, available_bytes: u64) -> usize {
|
||||
// Use 60 % of available RAM to leave headroom for the rest of the process.
|
||||
let budget = (available_bytes as f64 * 0.60) as u64;
|
||||
|
||||
let meta = match SKFileMeta::read(raw_path) {
|
||||
Ok(Some(m)) if m.instances > 0 => m,
|
||||
_ => return 1,
|
||||
};
|
||||
|
||||
let avg_seq_bytes = ((meta.length_sum + meta.instances - 1) / meta.instances + 3) / 4;
|
||||
// SuperKmer: header (4 B) + Box<[u8]> ptr+len (16 B) + heap seq bytes; value: u64 (8 B); ×1.5 for hashbrown overhead.
|
||||
let bytes_per_entry = ((4 + 16 + avg_seq_bytes + 8) as f64 * 1.5) as u64;
|
||||
let estimated = meta.instances * bytes_per_entry;
|
||||
|
||||
if estimated <= budget {
|
||||
debug!("Dereplication: estimated={estimated} budget={budget} n_temp=1");
|
||||
return 1;
|
||||
}
|
||||
|
||||
// Round up to the next power of two.
|
||||
let n = (estimated + budget - 1) / budget;
|
||||
debug!("Dereplication: estimated={estimated} budget={budget} n_temp={n}");
|
||||
n.next_power_of_two() as usize
|
||||
}
|
||||
|
||||
|
||||
/// Maximum value that fits in the 24-bit COUNT field of a SuperKmer header.
|
||||
const MAX_SK_COUNT: u64 = (1 << 24) - 1;
|
||||
|
||||
/// Deduplicate one partition directory in place (two-phase split + merge).
|
||||
fn dereplicate_partition(dir: &Path, level: Level, n_temp: usize) -> SKResult<()> {
|
||||
let raw_path = dir.join(format!("raw.{SK_EXT}"));
|
||||
if !raw_path.exists() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let out_path = dir.join(format!("dereplicated.{SK_EXT}"));
|
||||
let mut writer = SKFileWriter::create_with(&out_path, Format::Zstd, level)?;
|
||||
|
||||
if n_temp == 1 {
|
||||
// ── Direct path: partition fits in memory, no split needed ────────────
|
||||
let map = load_bucket(&raw_path)?;
|
||||
remove_skmer_file(&raw_path)?;
|
||||
flush_map(map, &mut writer)?;
|
||||
} else {
|
||||
// ── Phase 1: split raw file into temp buckets ─────────────────────────
|
||||
let temp_mask = (n_temp as u64) - 1;
|
||||
let temp_paths: Vec<PathBuf> = (0..n_temp)
|
||||
.map(|j| dir.join(format!("temp_{j:04}.{SK_EXT}")))
|
||||
.collect();
|
||||
|
||||
{
|
||||
let mut writers: Vec<SKFileWriter> = temp_paths
|
||||
.iter()
|
||||
.map(|p| SKFileWriter::create_with(p, Format::Zstd, level))
|
||||
.collect::<SKResult<_>>()?;
|
||||
|
||||
let mut reader = SKFileReader::open(&raw_path)?;
|
||||
while let Some(sk) = reader.read()? {
|
||||
let bucket = (sk.seq_hash() & temp_mask) as usize;
|
||||
writers[bucket].write(&sk)?;
|
||||
}
|
||||
for w in &mut writers {
|
||||
w.close()?;
|
||||
}
|
||||
}
|
||||
remove_skmer_file(&raw_path)?;
|
||||
|
||||
// ── Phase 2: merge each temp bucket into the output ───────────────────
|
||||
for temp_path in &temp_paths {
|
||||
let map = load_bucket(temp_path)?;
|
||||
remove_skmer_file(temp_path)?;
|
||||
flush_map(map, &mut writer)?;
|
||||
}
|
||||
}
|
||||
|
||||
writer.close()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Read a SuperKmer file into a deduplication map (already canonical).
|
||||
fn load_bucket(path: &Path) -> SKResult<HashMap<SuperKmer, u64>> {
|
||||
let capacity = SKFileMeta::read(path)
|
||||
.ok()
|
||||
.flatten()
|
||||
.map(|m| m.instances as usize)
|
||||
.unwrap_or(0);
|
||||
let mut map: HashMap<SuperKmer, u64> = HashMap::with_capacity(capacity);
|
||||
let mut reader = SKFileReader::open(path)?;
|
||||
while let Some(sk) = reader.read()? {
|
||||
let count = sk.count() as u64;
|
||||
*map.entry(sk).or_insert(0) += count;
|
||||
}
|
||||
Ok(map)
|
||||
}
|
||||
|
||||
/// Write all entries of a deduplication map to `writer`, splitting oversized counts.
|
||||
fn flush_map(map: HashMap<SuperKmer, u64>, writer: &mut SKFileWriter) -> SKResult<()> {
|
||||
for (mut sk, mut total) in map {
|
||||
while total > MAX_SK_COUNT {
|
||||
sk.set_count(MAX_SK_COUNT as u32);
|
||||
writer.write(&sk)?;
|
||||
total -= MAX_SK_COUNT;
|
||||
}
|
||||
sk.set_count(total as u32);
|
||||
writer.write(&sk)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn build_mphf(unique_path: &Path, f0: usize) -> io::Result<Mphf> {
|
||||
let file = fs::File::open(unique_path)?;
|
||||
let mmap = unsafe { Mmap::map(&file)? };
|
||||
let kmers: &[u64] = unsafe {
|
||||
std::slice::from_raw_parts(mmap.as_ptr() as *const u64, f0)
|
||||
};
|
||||
// Sequential constructor: the outer par_iter over partitions already saturates
|
||||
// the Rayon pool. new_from_par_iter would get no additional threads and adds
|
||||
// coordination overhead. try_new accesses the same mmap'd pages at zero extra cost.
|
||||
Mphf::try_new(kmers, PtrHashParams::<CubicEps>::default())
|
||||
.ok_or_else(|| io::Error::other("ptr_hash construction failed"))
|
||||
}
|
||||
|
||||
fn count_partition(dir: &Path, dedup_path: &Path, chunk_kmers: usize) -> SKResult<()> {
|
||||
let unique_path = dir.join("sorted_unique.bin");
|
||||
let f0 = sort_unique_kmers(dedup_path, dir, &unique_path, chunk_kmers)?;
|
||||
if f0 == 0 {
|
||||
return Ok(());
|
||||
}
|
||||
debug!("{}: f0={f0} unique kmers sorted", dir.display());
|
||||
|
||||
let mphf = build_mphf(&unique_path, f0)?;
|
||||
fs::remove_file(&unique_path)?;
|
||||
|
||||
let counts_path = dir.join("counts1.bin");
|
||||
let mut builder = PersistentCompactIntVecBuilder::new(f0, &counts_path)?;
|
||||
|
||||
{
|
||||
let mut reader = SKFileReader::open(dedup_path)?;
|
||||
while let Some(sk) = reader.read()? {
|
||||
let sk_count = sk.count();
|
||||
for kmer in sk.iter_canonical_kmers() {
|
||||
let slot = mphf.index(&kmer.raw());
|
||||
builder.set(slot, builder.get(slot).saturating_add(sk_count));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut spectrum: BTreeMap<u32, u64> = BTreeMap::new();
|
||||
for slot in 0..f0 {
|
||||
let c = builder.get(slot);
|
||||
if c > 0 {
|
||||
*spectrum.entry(c).or_insert(0) += 1;
|
||||
}
|
||||
}
|
||||
let f1: u64 = spectrum.iter().map(|(&c, &f)| c as u64 * f).sum();
|
||||
builder.close()?;
|
||||
|
||||
let spectrum_map: BTreeMap<String, u64> = spectrum
|
||||
.iter()
|
||||
.map(|(&c, &f)| (format!("{c:010}"), f))
|
||||
.collect();
|
||||
serde_json::to_writer_pretty(
|
||||
fs::File::create(dir.join("kmer_spectrum_raw.json"))?,
|
||||
&serde_json::json!({ "f0": f0 as u64, "f1": f1, "spectrum": &spectrum_map }),
|
||||
)
|
||||
.map_err(io::Error::other)?;
|
||||
|
||||
EpSerialize::store(&mphf, &dir.join("mphf1.bin"))
|
||||
.map_err(|e| io::Error::other(e.to_string()))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
impl Drop for KmerPartition {
|
||||
fn drop(&mut self) {
|
||||
let _ = self.close();
|
||||
}
|
||||
}
|
||||
|
||||
// ── integration tests ─────────────────────────────────────────────────────────
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::collections::HashMap;
|
||||
|
||||
use obikrope::Rope;
|
||||
use obikseq::SuperKmer;
|
||||
use obiskbuilder::build_superkmers;
|
||||
|
||||
const K: usize = 11;
|
||||
const M: usize = 5;
|
||||
|
||||
fn setup() {
|
||||
obikseq::params::set_k(K);
|
||||
obikseq::params::set_m(M);
|
||||
}
|
||||
|
||||
/// Direct canonical k-mer counts from ASCII sequences — ground truth.
|
||||
fn direct_counts(seqs: &[&[u8]]) -> (u64, u64) {
|
||||
let mut counts: HashMap<Vec<u8>, u64> = HashMap::new();
|
||||
for seq in seqs {
|
||||
for i in 0..seq.len().saturating_sub(K - 1) {
|
||||
let km = SuperKmer::from_ascii(&seq[i..i + K]).to_ascii();
|
||||
*counts.entry(km).or_insert(0) += 1;
|
||||
}
|
||||
}
|
||||
let f0 = counts.len() as u64;
|
||||
let f1: u64 = counts.values().sum();
|
||||
(f0, f1)
|
||||
}
|
||||
|
||||
/// Run the full pipeline on a list of sequences and return (f0, f1) from
|
||||
/// the `kmer_spectrum_raw.json` produced by `count_partition`.
|
||||
fn pipeline_counts(seqs: &[&[u8]]) -> (u64, u64) {
|
||||
setup();
|
||||
|
||||
let mut rope_data: Vec<u8> = Vec::new();
|
||||
for seq in seqs {
|
||||
rope_data.extend_from_slice(seq);
|
||||
rope_data.push(0x00);
|
||||
}
|
||||
let mut rope = Rope::new(None);
|
||||
rope.push(rope_data);
|
||||
|
||||
let superkmers: Vec<_> = build_superkmers(rope, K, 1, 0.0);
|
||||
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let mut kp = KmerPartition::create(dir.path(), 0, K, M, true).unwrap();
|
||||
kp.write_batch(superkmers).unwrap();
|
||||
kp.close().unwrap();
|
||||
kp.dereplicate().unwrap();
|
||||
|
||||
let part_dir = dir.path().join(PARTITIONS_SUBDIR).join("part_00000");
|
||||
let dedup_path = part_dir.join("dereplicated.skmer.zst");
|
||||
if !dedup_path.exists() {
|
||||
return (0, 0);
|
||||
}
|
||||
count_partition(&part_dir, &dedup_path, 1 << 20).unwrap();
|
||||
|
||||
let spec: serde_json::Value = serde_json::from_reader(
|
||||
fs::File::open(part_dir.join("kmer_spectrum_raw.json")).unwrap(),
|
||||
).unwrap();
|
||||
let f0 = spec["f0"].as_u64().unwrap_or(0);
|
||||
let f1 = spec["f1"].as_u64().unwrap_or(0);
|
||||
(f0, f1)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn single_sequence_f0_f1_match() {
|
||||
let seqs: &[&[u8]] = &[b"ACGTACGTACGTACGTACGT"];
|
||||
let (ef0, ef1) = direct_counts(seqs);
|
||||
let (gf0, gf1) = pipeline_counts(seqs);
|
||||
assert_eq!(gf0, ef0, "f0 wrong: expected {ef0}, got {gf0}");
|
||||
assert_eq!(gf1, ef1, "f1 wrong: expected {ef1}, got {gf1}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn two_sequences_f0_f1_match() {
|
||||
let seqs: &[&[u8]] = &[
|
||||
b"ACGTACGTACGTACGTACGT",
|
||||
b"TGCATGCATGCATGCATGCA",
|
||||
];
|
||||
let (ef0, ef1) = direct_counts(seqs);
|
||||
let (gf0, gf1) = pipeline_counts(seqs);
|
||||
assert_eq!(gf0, ef0, "f0 wrong: expected {ef0}, got {gf0}");
|
||||
assert_eq!(gf1, ef1, "f1 wrong: expected {ef1}, got {gf1}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn repeated_sequence_f1_doubles() {
|
||||
let seq = b"ACGTACGTACGTACGTACGT";
|
||||
let seqs: &[&[u8]] = &[seq, seq];
|
||||
let (ef0, ef1) = direct_counts(seqs);
|
||||
let (gf0, gf1) = pipeline_counts(seqs);
|
||||
assert_eq!(gf0, ef0, "f0 wrong: expected {ef0}, got {gf0}");
|
||||
assert_eq!(gf1, ef1, "f1 wrong: expected {ef1}, got {gf1}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn many_sequences_f0_f1_match() {
|
||||
// 20 distinct sequences of length 40 — forces multiple super-kmers and
|
||||
// multiple minimizer boundaries per sequence.
|
||||
let bases = b"ACGT";
|
||||
let seqs: Vec<Vec<u8>> = (0..20u32)
|
||||
.map(|i| (0..40).map(|j| bases[((i * 7 + j * 3) % 4) as usize]).collect())
|
||||
.collect();
|
||||
let seq_refs: Vec<&[u8]> = seqs.iter().map(|v| v.as_slice()).collect();
|
||||
let (ef0, ef1) = direct_counts(&seq_refs);
|
||||
let (gf0, gf1) = pipeline_counts(&seq_refs);
|
||||
assert_eq!(gf0, ef0, "f0 wrong: expected {ef0}, got {gf0}");
|
||||
assert_eq!(gf1, ef1, "f1 wrong: expected {ef1}, got {gf1}");
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::fs;
|
||||
use std::io;
|
||||
use std::path::Path;
|
||||
|
||||
use cacheline_ef::{CachelineEf, CachelineEfVec};
|
||||
use epserde::ser::Serialize as EpSerialize;
|
||||
use memmap2::Mmap;
|
||||
use obicompactvec::PersistentCompactIntVecBuilder;
|
||||
use obiskio::{SKFileReader, SKResult};
|
||||
use ptr_hash::{PtrHash, PtrHashParams, bucket_fn::CubicEps, hash::Xx64};
|
||||
use tracing::debug;
|
||||
|
||||
use crate::kmer_sort::sort_unique_kmers;
|
||||
|
||||
pub(super) type Mphf = PtrHash<u64, CubicEps, CachelineEfVec<Vec<CachelineEf>>, Xx64, Vec<u8>>;
|
||||
|
||||
fn build_mphf(unique_path: &Path, f0: usize) -> io::Result<Mphf> {
|
||||
let file = fs::File::open(unique_path)?;
|
||||
let mmap = unsafe { Mmap::map(&file)? };
|
||||
let kmers: &[u64] = unsafe {
|
||||
std::slice::from_raw_parts(mmap.as_ptr() as *const u64, f0)
|
||||
};
|
||||
// Sequential constructor: the outer par_iter over partitions already saturates
|
||||
// the Rayon pool. new_from_par_iter would get no additional threads and adds
|
||||
// coordination overhead. try_new accesses the same mmap'd pages at zero extra cost.
|
||||
Mphf::try_new(kmers, PtrHashParams::<CubicEps>::default())
|
||||
.ok_or_else(|| io::Error::other("ptr_hash construction failed"))
|
||||
}
|
||||
|
||||
pub(super) fn count_partition(dir: &Path, dedup_path: &Path, chunk_kmers: usize) -> SKResult<()> {
|
||||
let unique_path = dir.join("sorted_unique.bin");
|
||||
let f0 = sort_unique_kmers(dedup_path, dir, &unique_path, chunk_kmers)?;
|
||||
if f0 == 0 {
|
||||
return Ok(());
|
||||
}
|
||||
debug!("{}: f0={f0} unique kmers sorted", dir.display());
|
||||
|
||||
let mphf = build_mphf(&unique_path, f0)?;
|
||||
fs::remove_file(&unique_path)?;
|
||||
|
||||
let counts_path = dir.join("counts1.bin");
|
||||
let mut builder = PersistentCompactIntVecBuilder::new(f0, &counts_path)?;
|
||||
|
||||
{
|
||||
let mut reader = SKFileReader::open(dedup_path)?;
|
||||
while let Some(sk) = reader.read()? {
|
||||
let sk_count = sk.count();
|
||||
for kmer in sk.iter_canonical_kmers() {
|
||||
let slot = mphf.index(&kmer.raw());
|
||||
builder.set(slot, builder.get(slot).saturating_add(sk_count));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut spectrum: BTreeMap<u32, u64> = BTreeMap::new();
|
||||
for slot in 0..f0 {
|
||||
let c = builder.get(slot);
|
||||
if c > 0 {
|
||||
*spectrum.entry(c).or_insert(0) += 1;
|
||||
}
|
||||
}
|
||||
let f1: u64 = spectrum.iter().map(|(&c, &f)| c as u64 * f).sum();
|
||||
builder.close()?;
|
||||
|
||||
let spectrum_map: BTreeMap<String, u64> = spectrum
|
||||
.iter()
|
||||
.map(|(&c, &f)| (format!("{c:010}"), f))
|
||||
.collect();
|
||||
serde_json::to_writer_pretty(
|
||||
fs::File::create(dir.join("kmer_spectrum_raw.json"))?,
|
||||
&serde_json::json!({ "f0": f0 as u64, "f1": f1, "spectrum": &spectrum_map }),
|
||||
)
|
||||
.map_err(io::Error::other)?;
|
||||
|
||||
EpSerialize::store(&mphf, &dir.join("mphf1.bin"))
|
||||
.map_err(|e| io::Error::other(e.to_string()))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
use std::collections::HashMap;
|
||||
use std::fs;
|
||||
use std::io;
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use tracing::debug;
|
||||
|
||||
use niffler::Level;
|
||||
use niffler::send::compression::Format;
|
||||
use obikseq::Sequence;
|
||||
use obikseq::superkmer::SuperKmer;
|
||||
use obiskio::{SKFileMeta, SKFileReader, SKFileWriter, SKResult};
|
||||
|
||||
use super::SK_EXT;
|
||||
|
||||
/// Estimate the number of in-memory buckets needed to deduplicate the partition
|
||||
/// file at `raw_path` given `available_bytes` of free RAM.
|
||||
///
|
||||
/// Memory per HashMap entry:
|
||||
/// key Box (1 + avg_seq_bytes) + SuperKmer header (4 B) + avg seq bytes + u64 count (8 B),
|
||||
/// multiplied by 1.5 for hashbrown load-factor overhead.
|
||||
///
|
||||
/// Returns 1 if the partition fits comfortably in memory (no split needed).
|
||||
/// Always returns a power of two.
|
||||
pub(super) fn optimal_buckets(raw_path: &Path, available_bytes: u64) -> usize {
|
||||
// Use 60 % of available RAM to leave headroom for the rest of the process.
|
||||
let budget = (available_bytes as f64 * 0.60) as u64;
|
||||
|
||||
let meta = match SKFileMeta::read(raw_path) {
|
||||
Ok(Some(m)) if m.instances > 0 => m,
|
||||
_ => return 1,
|
||||
};
|
||||
|
||||
let avg_seq_bytes = ((meta.length_sum + meta.instances - 1) / meta.instances + 3) / 4;
|
||||
// SuperKmer: header (4 B) + Box<[u8]> ptr+len (16 B) + heap seq bytes; value: u64 (8 B); ×1.5 for hashbrown overhead.
|
||||
let bytes_per_entry = ((4 + 16 + avg_seq_bytes + 8) as f64 * 1.5) as u64;
|
||||
let estimated = meta.instances * bytes_per_entry;
|
||||
|
||||
if estimated <= budget {
|
||||
debug!("Dereplication: estimated={estimated} budget={budget} n_temp=1");
|
||||
return 1;
|
||||
}
|
||||
|
||||
// Round up to the next power of two.
|
||||
let n = (estimated + budget - 1) / budget;
|
||||
debug!("Dereplication: estimated={estimated} budget={budget} n_temp={n}");
|
||||
n.next_power_of_two() as usize
|
||||
}
|
||||
|
||||
/// Remove a SuperKmer file and its sidecar (if present).
|
||||
fn remove_skmer_file(path: &Path) -> SKResult<()> {
|
||||
fs::remove_file(path)?;
|
||||
let sidecar = SKFileMeta::sidecar_path(path);
|
||||
match fs::remove_file(&sidecar) {
|
||||
Ok(()) => {}
|
||||
Err(e) if e.kind() == io::ErrorKind::NotFound => {}
|
||||
Err(e) => return Err(e.into()),
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Maximum value that fits in the 24-bit COUNT field of a SuperKmer header.
|
||||
const MAX_SK_COUNT: u64 = (1 << 24) - 1;
|
||||
|
||||
/// Deduplicate one partition directory in place (two-phase split + merge).
|
||||
pub(super) fn dereplicate_partition(dir: &Path, level: Level, n_temp: usize) -> SKResult<()> {
|
||||
let raw_path = dir.join(format!("raw.{SK_EXT}"));
|
||||
if !raw_path.exists() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let out_path = dir.join(format!("dereplicated.{SK_EXT}"));
|
||||
let mut writer = SKFileWriter::create_with(&out_path, Format::Zstd, level)?;
|
||||
|
||||
if n_temp == 1 {
|
||||
// ── Direct path: partition fits in memory, no split needed ────────────
|
||||
let map = load_bucket(&raw_path)?;
|
||||
remove_skmer_file(&raw_path)?;
|
||||
flush_map(map, &mut writer)?;
|
||||
} else {
|
||||
// ── Phase 1: split raw file into temp buckets ─────────────────────────
|
||||
let temp_mask = (n_temp as u64) - 1;
|
||||
let temp_paths: Vec<PathBuf> = (0..n_temp)
|
||||
.map(|j| dir.join(format!("temp_{j:04}.{SK_EXT}")))
|
||||
.collect();
|
||||
|
||||
{
|
||||
let mut writers: Vec<SKFileWriter> = temp_paths
|
||||
.iter()
|
||||
.map(|p| SKFileWriter::create_with(p, Format::Zstd, level))
|
||||
.collect::<SKResult<_>>()?;
|
||||
|
||||
let mut reader = SKFileReader::open(&raw_path)?;
|
||||
while let Some(sk) = reader.read()? {
|
||||
let bucket = (sk.seq_hash() & temp_mask) as usize;
|
||||
writers[bucket].write(&sk)?;
|
||||
}
|
||||
for w in &mut writers {
|
||||
w.close()?;
|
||||
}
|
||||
}
|
||||
remove_skmer_file(&raw_path)?;
|
||||
|
||||
// ── Phase 2: merge each temp bucket into the output ───────────────────
|
||||
for temp_path in &temp_paths {
|
||||
let map = load_bucket(temp_path)?;
|
||||
remove_skmer_file(temp_path)?;
|
||||
flush_map(map, &mut writer)?;
|
||||
}
|
||||
}
|
||||
|
||||
writer.close()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Read a SuperKmer file into a deduplication map (already canonical).
|
||||
fn load_bucket(path: &Path) -> SKResult<HashMap<SuperKmer, u64>> {
|
||||
let capacity = SKFileMeta::read(path)
|
||||
.ok()
|
||||
.flatten()
|
||||
.map(|m| m.instances as usize)
|
||||
.unwrap_or(0);
|
||||
let mut map: HashMap<SuperKmer, u64> = HashMap::with_capacity(capacity);
|
||||
let mut reader = SKFileReader::open(path)?;
|
||||
while let Some(sk) = reader.read()? {
|
||||
let count = sk.count() as u64;
|
||||
*map.entry(sk).or_insert(0) += count;
|
||||
}
|
||||
Ok(map)
|
||||
}
|
||||
|
||||
/// Write all entries of a deduplication map to `writer`, splitting oversized counts.
|
||||
fn flush_map(map: HashMap<SuperKmer, u64>, writer: &mut SKFileWriter) -> SKResult<()> {
|
||||
for (mut sk, mut total) in map {
|
||||
while total > MAX_SK_COUNT {
|
||||
sk.set_count(MAX_SK_COUNT as u32);
|
||||
writer.write(&sk)?;
|
||||
total -= MAX_SK_COUNT;
|
||||
}
|
||||
sk.set_count(total as u32);
|
||||
writer.write(&sk)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -0,0 +1,341 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::fs;
|
||||
use std::io;
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::time::Instant;
|
||||
|
||||
use obisys::progress_bar;
|
||||
|
||||
use obikseq::RoutableSuperKmer;
|
||||
use obiskio::SKResult;
|
||||
use rayon::prelude::*;
|
||||
use remove_dir_all::remove_dir_all;
|
||||
use sysinfo::System;
|
||||
|
||||
use niffler::Level;
|
||||
use niffler::send::compression::Format;
|
||||
use obiskio::SKFileWriter;
|
||||
|
||||
use crate::kmer_sort::chunk_size_from_ram;
|
||||
|
||||
use super::count::count_partition;
|
||||
use super::dereplicate::{dereplicate_partition, optimal_buckets};
|
||||
use super::{PARTITIONS_SUBDIR, SK_EXT};
|
||||
|
||||
pub struct KmerSpectrum {
|
||||
pub f0: u64,
|
||||
pub f1: u64,
|
||||
pub counts: BTreeMap<u32, u64>,
|
||||
}
|
||||
|
||||
pub struct KmerPartition {
|
||||
root_path: PathBuf,
|
||||
n_partitions: usize,
|
||||
partitions_mask: u64,
|
||||
kmer_size: usize,
|
||||
minimizer_size: usize,
|
||||
writers: Vec<Option<SKFileWriter>>,
|
||||
level: Level,
|
||||
closed: bool,
|
||||
}
|
||||
|
||||
impl KmerPartition {
|
||||
pub fn create<P: AsRef<Path>>(
|
||||
path: P,
|
||||
n_bits: usize,
|
||||
kmer_size: usize,
|
||||
minimizer_size: usize,
|
||||
force: bool,
|
||||
) -> SKResult<Self> {
|
||||
Self::create_with(path, n_bits, kmer_size, minimizer_size, Level::One, force)
|
||||
}
|
||||
|
||||
pub fn create_with<P: AsRef<Path>>(
|
||||
path: P,
|
||||
n_bits: usize,
|
||||
kmer_size: usize,
|
||||
minimizer_size: usize,
|
||||
level: Level,
|
||||
force: bool,
|
||||
) -> SKResult<Self> {
|
||||
let root_path = path.as_ref().to_owned();
|
||||
if root_path.exists() {
|
||||
if force {
|
||||
remove_dir_all(&root_path)?;
|
||||
} else {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::AlreadyExists,
|
||||
format!(
|
||||
"{}: partition directory already exists",
|
||||
root_path.display()
|
||||
),
|
||||
)
|
||||
.into());
|
||||
}
|
||||
}
|
||||
fs::create_dir_all(root_path.join(PARTITIONS_SUBDIR))?;
|
||||
let n_partitions = 1usize << n_bits;
|
||||
let writers = (0..n_partitions).map(|_| None).collect();
|
||||
let partition = Self {
|
||||
root_path,
|
||||
n_partitions,
|
||||
partitions_mask: (1u64 << n_bits) - 1,
|
||||
kmer_size,
|
||||
minimizer_size,
|
||||
writers,
|
||||
level,
|
||||
closed: false,
|
||||
};
|
||||
Ok(partition)
|
||||
}
|
||||
|
||||
pub fn open_with_config<P: AsRef<Path>>(
|
||||
path: P,
|
||||
kmer_size: usize,
|
||||
minimizer_size: usize,
|
||||
n_bits: usize,
|
||||
) -> SKResult<Self> {
|
||||
let root_path = path.as_ref().to_owned();
|
||||
if !root_path.exists() {
|
||||
return Err(io::Error::new(
|
||||
io::ErrorKind::NotFound,
|
||||
format!("{}: partition directory not found", root_path.display()),
|
||||
)
|
||||
.into());
|
||||
}
|
||||
let n_partitions = 1usize << n_bits;
|
||||
let writers = (0..n_partitions).map(|_| None).collect();
|
||||
Ok(Self {
|
||||
root_path,
|
||||
n_partitions,
|
||||
partitions_mask: (1u64 << n_bits) - 1,
|
||||
kmer_size,
|
||||
minimizer_size,
|
||||
writers,
|
||||
level: Level::One,
|
||||
closed: true,
|
||||
})
|
||||
}
|
||||
|
||||
/// Route and write one super-kmer to its partition file.
|
||||
pub fn write(&mut self, rsk: RoutableSuperKmer) -> SKResult<()> {
|
||||
self.check_not_closed()?;
|
||||
let partition = (rsk.minimizer().seq_hash() & self.partitions_mask) as usize;
|
||||
let sk = rsk.into_superkmer();
|
||||
self.ensure_writer(partition)?.write(&sk)
|
||||
}
|
||||
|
||||
/// Route and write a batch of super-kmers.
|
||||
pub fn write_batch(&mut self, rsks: Vec<RoutableSuperKmer>) -> SKResult<()> {
|
||||
self.check_not_closed()?;
|
||||
for rsk in rsks {
|
||||
let partition = (rsk.minimizer().seq_hash() & self.partitions_mask) as usize;
|
||||
let sk = rsk.into_superkmer();
|
||||
self.ensure_writer(partition)?.write(&sk)?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn flush(&mut self) -> SKResult<()> {
|
||||
self.check_not_closed()?;
|
||||
for writer in self.writers.iter_mut().flatten() {
|
||||
writer.flush()?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn close(&mut self) -> SKResult<()> {
|
||||
if self.closed {
|
||||
return Ok(());
|
||||
}
|
||||
self.closed = true;
|
||||
for writer in self.writers.iter_mut().flatten() {
|
||||
writer.close()?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn is_open(&self) -> bool {
|
||||
!self.closed
|
||||
}
|
||||
|
||||
pub fn path(&self) -> &Path {
|
||||
&self.root_path
|
||||
}
|
||||
|
||||
/// Path of partition `i` directory.
|
||||
pub fn part_dir(&self, i: usize) -> PathBuf {
|
||||
self.root_path.join(PARTITIONS_SUBDIR).join(format!("part_{i:05}"))
|
||||
}
|
||||
|
||||
pub fn kmer_size(&self) -> usize {
|
||||
self.kmer_size
|
||||
}
|
||||
|
||||
pub fn minimizer_size(&self) -> usize {
|
||||
self.minimizer_size
|
||||
}
|
||||
|
||||
pub fn n_partitions(&self) -> usize {
|
||||
self.n_partitions
|
||||
}
|
||||
|
||||
/// Deduplicate all `raw.{ext}` files in parallel, replacing each with a
|
||||
/// `dereplicated.{ext}` file where identical canonical sequences are merged
|
||||
/// and their counts summed.
|
||||
///
|
||||
/// Each partition file is processed in two phases to bound memory use:
|
||||
///
|
||||
/// 1. **Split** — the raw file is scattered into `2^temp_bits` temporary
|
||||
/// files routed by `hash(canonical_seq) & temp_mask`. Because duplicates
|
||||
/// always share the same hash, they always land in the same temp file.
|
||||
/// 2. **Merge** — each temp file is loaded fully into a `HashMap`, counts
|
||||
/// are accumulated in `u64` (no 24-bit overflow risk), and the result is
|
||||
/// appended to `dereplicated.{ext}`.
|
||||
///
|
||||
/// If a merged count exceeds the 24-bit header limit, the sequence is
|
||||
/// emitted as multiple records whose counts sum to the true total.
|
||||
///
|
||||
/// `temp_bits` controls the split fan-out (`2^temp_bits` temp files per
|
||||
/// partition). Higher values reduce per-temp-file memory at the cost of
|
||||
/// more temporary file descriptors — all managed by the global fd pool.
|
||||
pub fn dereplicate(&self) -> SKResult<()> {
|
||||
let level = self.level;
|
||||
let sys = System::new_all();
|
||||
// available_memory() can return 0 on macOS when the compressor page count exceeds
|
||||
// free+inactive+purgeable pages (sysinfo saturating_sub). Fall back to half of total.
|
||||
let available = match sys.available_memory() {
|
||||
0 => sys.total_memory() / 2,
|
||||
n => n,
|
||||
};
|
||||
let n_threads = rayon::current_num_threads().max(1) as u64;
|
||||
let available_per_thread = available / n_threads;
|
||||
|
||||
let pb = progress_bar("dereplication", self.n_partitions as u64, "partitions");
|
||||
|
||||
let results: Vec<SKResult<()>> = (0..self.n_partitions)
|
||||
.into_par_iter()
|
||||
.map(|i| {
|
||||
let dir = self.part_dir(i);
|
||||
if !dir.exists() {
|
||||
pb.inc(1);
|
||||
return Ok(());
|
||||
}
|
||||
let raw_path = dir.join(format!("raw.{SK_EXT}"));
|
||||
let t = Instant::now();
|
||||
let n_buckets = optimal_buckets(&raw_path, available_per_thread);
|
||||
let result = dereplicate_partition(&dir, level, n_buckets);
|
||||
pb.set_message(format!("last {:.0}ms", t.elapsed().as_millis()));
|
||||
pb.inc(1);
|
||||
result
|
||||
})
|
||||
.collect();
|
||||
|
||||
pb.finish_and_clear();
|
||||
for r in results {
|
||||
r?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// For each partition that has a `dereplicated.{ext}` file:
|
||||
/// 1. Enumerates all unique canonical kmers (two passes over the file).
|
||||
/// 2. Builds a provisional MPHF (FMPHGO) over those kmers.
|
||||
/// 3. Writes a flat binary count file (`counts1.bin`, one `u32` per slot,
|
||||
/// memory-mapped) accumulating kmer abundances from the superkmer counts.
|
||||
/// 4. Persists the MPHF to `mphf1.bin` for downstream use.
|
||||
///
|
||||
/// Returns the aggregated `KmerSpectrum`. Per-partition spectrum files are
|
||||
/// deleted after aggregation unless `keep_partial` is true.
|
||||
///
|
||||
/// Partitions are processed in parallel via Rayon (one task per thread).
|
||||
/// Peak memory per partition is ~80 MB, so n_threads partitions run simultaneously.
|
||||
pub fn count_kmer(&self, keep_partial: bool) -> SKResult<KmerSpectrum> {
|
||||
let sys = System::new_all();
|
||||
let available = match sys.available_memory() {
|
||||
0 => sys.total_memory() / 2,
|
||||
n => n,
|
||||
};
|
||||
let n_threads = rayon::current_num_threads().max(1) as u64;
|
||||
let chunk_kmers = chunk_size_from_ram(available / n_threads);
|
||||
|
||||
let pb = progress_bar("counting", self.n_partitions as u64, "partitions");
|
||||
|
||||
let results: Vec<SKResult<()>> = (0..self.n_partitions)
|
||||
.into_par_iter()
|
||||
.map(|i| {
|
||||
let dir = self.part_dir(i);
|
||||
let dedup_path = dir.join(format!("dereplicated.{SK_EXT}"));
|
||||
if !dedup_path.exists() {
|
||||
pb.inc(1);
|
||||
return Ok(());
|
||||
}
|
||||
let t = Instant::now();
|
||||
let result = count_partition(&dir, &dedup_path, chunk_kmers);
|
||||
pb.set_message(format!("last {:.0}ms", t.elapsed().as_millis()));
|
||||
pb.inc(1);
|
||||
result
|
||||
})
|
||||
.collect();
|
||||
|
||||
pb.finish_and_clear();
|
||||
for r in results {
|
||||
r?;
|
||||
}
|
||||
|
||||
// Aggregate per-partition spectra.
|
||||
let mut counts: BTreeMap<u32, u64> = BTreeMap::new();
|
||||
let mut f0: u64 = 0;
|
||||
let mut f1: u64 = 0;
|
||||
|
||||
for i in 0..self.n_partitions {
|
||||
let path = self.part_dir(i).join("kmer_spectrum_raw.json");
|
||||
if !path.exists() {
|
||||
continue;
|
||||
}
|
||||
let v: serde_json::Value =
|
||||
serde_json::from_str(&fs::read_to_string(&path)?).map_err(io::Error::other)?;
|
||||
f0 += v["f0"].as_u64().unwrap_or(0);
|
||||
f1 += v["f1"].as_u64().unwrap_or(0);
|
||||
if let Some(obj) = v["spectrum"].as_object() {
|
||||
for (c_str, freq) in obj {
|
||||
if let (Ok(c), Some(f)) = (c_str.parse::<u32>(), freq.as_u64()) {
|
||||
*counts.entry(c).or_insert(0) += f;
|
||||
}
|
||||
}
|
||||
}
|
||||
if !keep_partial {
|
||||
let _ = fs::remove_file(&path);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(KmerSpectrum { f0, f1, counts })
|
||||
}
|
||||
|
||||
// ── private ───────────────────────────────────────────────────────────────
|
||||
|
||||
fn check_not_closed(&self) -> SKResult<()> {
|
||||
if self.closed {
|
||||
Err(io::Error::new(io::ErrorKind::BrokenPipe, "write to closed KmerPartition").into())
|
||||
} else {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
fn ensure_writer(&mut self, partition: usize) -> SKResult<&mut SKFileWriter> {
|
||||
if self.writers[partition].is_none() {
|
||||
let dir = self.root_path.join(PARTITIONS_SUBDIR).join(format!("part_{:05}", partition));
|
||||
fs::create_dir_all(&dir)?;
|
||||
let file_path = dir.join(format!("raw.{SK_EXT}"));
|
||||
let writer = SKFileWriter::create_with(file_path, Format::Zstd, self.level)?;
|
||||
self.writers[partition] = Some(writer);
|
||||
}
|
||||
Ok(self.writers[partition].as_mut().unwrap())
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for KmerPartition {
|
||||
fn drop(&mut self) {
|
||||
let _ = self.close();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
//! K-mer partitioning: routing super-kmers into per-partition files,
|
||||
//! deduplicating them, and counting unique canonical k-mers.
|
||||
//!
|
||||
//! Submodules: [`kmer_partition`] (`KmerPartition`, `KmerSpectrum` and the
|
||||
//! routing/lifecycle API), [`dereplicate`] (two-phase split+merge
|
||||
//! deduplication), [`count`] (unique-kmer enumeration, MPHF, abundance
|
||||
//! counting).
|
||||
|
||||
mod count;
|
||||
mod dereplicate;
|
||||
mod kmer_partition;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
pub use kmer_partition::{KmerPartition, KmerSpectrum};
|
||||
|
||||
const SK_EXT: &str = "skmer.zst";
|
||||
pub const PARTITIONS_SUBDIR: &str = "partitions";
|
||||
@@ -0,0 +1,113 @@
|
||||
use std::collections::HashMap;
|
||||
use std::fs;
|
||||
|
||||
use obikrope::Rope;
|
||||
use obikseq::SuperKmer;
|
||||
use obiskbuilder::build_superkmers;
|
||||
|
||||
use super::count::count_partition;
|
||||
use super::{KmerPartition, PARTITIONS_SUBDIR};
|
||||
|
||||
const K: usize = 11;
|
||||
const M: usize = 5;
|
||||
|
||||
fn setup() {
|
||||
obikseq::params::set_k(K);
|
||||
obikseq::params::set_m(M);
|
||||
}
|
||||
|
||||
/// Direct canonical k-mer counts from ASCII sequences — ground truth.
|
||||
fn direct_counts(seqs: &[&[u8]]) -> (u64, u64) {
|
||||
let mut counts: HashMap<Vec<u8>, u64> = HashMap::new();
|
||||
for seq in seqs {
|
||||
for i in 0..seq.len().saturating_sub(K - 1) {
|
||||
let km = SuperKmer::from_ascii(&seq[i..i + K]).to_ascii();
|
||||
*counts.entry(km).or_insert(0) += 1;
|
||||
}
|
||||
}
|
||||
let f0 = counts.len() as u64;
|
||||
let f1: u64 = counts.values().sum();
|
||||
(f0, f1)
|
||||
}
|
||||
|
||||
/// Run the full pipeline on a list of sequences and return (f0, f1) from
|
||||
/// the `kmer_spectrum_raw.json` produced by `count_partition`.
|
||||
fn pipeline_counts(seqs: &[&[u8]]) -> (u64, u64) {
|
||||
setup();
|
||||
|
||||
let mut rope_data: Vec<u8> = Vec::new();
|
||||
for seq in seqs {
|
||||
rope_data.extend_from_slice(seq);
|
||||
rope_data.push(0x00);
|
||||
}
|
||||
let mut rope = Rope::new(None);
|
||||
rope.push(rope_data);
|
||||
|
||||
let superkmers: Vec<_> = build_superkmers(rope, K, 1, 0.0);
|
||||
|
||||
let dir = tempfile::tempdir().unwrap();
|
||||
let mut kp = KmerPartition::create(dir.path(), 0, K, M, true).unwrap();
|
||||
kp.write_batch(superkmers).unwrap();
|
||||
kp.close().unwrap();
|
||||
kp.dereplicate().unwrap();
|
||||
|
||||
let part_dir = dir.path().join(PARTITIONS_SUBDIR).join("part_00000");
|
||||
let dedup_path = part_dir.join("dereplicated.skmer.zst");
|
||||
if !dedup_path.exists() {
|
||||
return (0, 0);
|
||||
}
|
||||
count_partition(&part_dir, &dedup_path, 1 << 20).unwrap();
|
||||
|
||||
let spec: serde_json::Value = serde_json::from_reader(
|
||||
fs::File::open(part_dir.join("kmer_spectrum_raw.json")).unwrap(),
|
||||
).unwrap();
|
||||
let f0 = spec["f0"].as_u64().unwrap_or(0);
|
||||
let f1 = spec["f1"].as_u64().unwrap_or(0);
|
||||
(f0, f1)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn single_sequence_f0_f1_match() {
|
||||
let seqs: &[&[u8]] = &[b"ACGTACGTACGTACGTACGT"];
|
||||
let (ef0, ef1) = direct_counts(seqs);
|
||||
let (gf0, gf1) = pipeline_counts(seqs);
|
||||
assert_eq!(gf0, ef0, "f0 wrong: expected {ef0}, got {gf0}");
|
||||
assert_eq!(gf1, ef1, "f1 wrong: expected {ef1}, got {gf1}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn two_sequences_f0_f1_match() {
|
||||
let seqs: &[&[u8]] = &[
|
||||
b"ACGTACGTACGTACGTACGT",
|
||||
b"TGCATGCATGCATGCATGCA",
|
||||
];
|
||||
let (ef0, ef1) = direct_counts(seqs);
|
||||
let (gf0, gf1) = pipeline_counts(seqs);
|
||||
assert_eq!(gf0, ef0, "f0 wrong: expected {ef0}, got {gf0}");
|
||||
assert_eq!(gf1, ef1, "f1 wrong: expected {ef1}, got {gf1}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn repeated_sequence_f1_doubles() {
|
||||
let seq = b"ACGTACGTACGTACGTACGT";
|
||||
let seqs: &[&[u8]] = &[seq, seq];
|
||||
let (ef0, ef1) = direct_counts(seqs);
|
||||
let (gf0, gf1) = pipeline_counts(seqs);
|
||||
assert_eq!(gf0, ef0, "f0 wrong: expected {ef0}, got {gf0}");
|
||||
assert_eq!(gf1, ef1, "f1 wrong: expected {ef1}, got {gf1}");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn many_sequences_f0_f1_match() {
|
||||
// 20 distinct sequences of length 40 — forces multiple super-kmers and
|
||||
// multiple minimizer boundaries per sequence.
|
||||
let bases = b"ACGT";
|
||||
let seqs: Vec<Vec<u8>> = (0..20u32)
|
||||
.map(|i| (0..40).map(|j| bases[((i * 7 + j * 3) % 4) as usize]).collect())
|
||||
.collect();
|
||||
let seq_refs: Vec<&[u8]> = seqs.iter().map(|v| v.as_slice()).collect();
|
||||
let (ef0, ef1) = direct_counts(&seq_refs);
|
||||
let (gf0, gf1) = pipeline_counts(&seq_refs);
|
||||
assert_eq!(gf0, ef0, "f0 wrong: expected {ef0}, got {gf0}");
|
||||
assert_eq!(gf1, ef1, "f1 wrong: expected {ef1}, got {gf1}");
|
||||
}
|
||||
@@ -1,828 +0,0 @@
|
||||
//! Cursors for sequential and random access over a [`Rope`].
|
||||
//!
|
||||
//! # Design
|
||||
//!
|
||||
//! A cursor borrows a `&'a Rope` and keeps a small block cache so that
|
||||
//! consecutive accesses within the same block cost O(1). The first access to a
|
||||
//! new block costs O(log n) (binary search in [`Rope::lookup`]); subsequent
|
||||
//! accesses within that block are free.
|
||||
//!
|
||||
//! All mutable state (current position, cache) is stored in [`Cell`] fields,
|
||||
//! so every cursor method takes `&self` rather than `&mut self`. This means:
|
||||
//!
|
||||
//! - Two cursors can coexist on the same rope without lifetime conflicts.
|
||||
//! - The `iter()` method returns a lightweight wrapper that holds `&Cursor`,
|
||||
//! allowing `cursor.tell()` or `cursor.seek()` to be called **inside a `for`
|
||||
//! loop** over the same cursor.
|
||||
//!
|
||||
//! # Cursors
|
||||
//!
|
||||
//! | Type | Direction | First `read_next` | `seek(Relative, +n)` |
|
||||
//! |------|-----------|-------------------|----------------------|
|
||||
//! | [`ForwardCursor`] | start → end | index 0 | advances (+n) |
|
||||
//! | [`BackwardCursor`] | end → start | index `len-1` | retreats (+n) |
|
||||
//!
|
||||
//! # Example
|
||||
//!
|
||||
//! ```
|
||||
//! use obikrope::{Rope, RopeCursor};
|
||||
//!
|
||||
//! let mut rope = Rope::new(None);
|
||||
//! rope.push(b"ACGT".to_vec());
|
||||
//!
|
||||
//! let cursor = rope.fw_cursor();
|
||||
//! for byte in cursor.iter() {
|
||||
//! // cursor.tell() is valid here — iter() holds &cursor, not &mut cursor
|
||||
//! let _ = cursor.tell();
|
||||
//! }
|
||||
//! ```
|
||||
|
||||
use std::cell::Cell;
|
||||
|
||||
use crate::{Rope, RopeError};
|
||||
|
||||
/// Controls how the `pos` argument of [`RopeCursor::seek`] is interpreted.
|
||||
#[derive(Clone, Copy)]
|
||||
pub enum SeekMode {
|
||||
/// `pos` is an absolute byte index from the start of the rope.
|
||||
Absolute,
|
||||
/// `pos` is relative to the current position.
|
||||
/// Positive = forward for [`ForwardCursor`], backward for [`BackwardCursor`].
|
||||
Relative,
|
||||
/// `pos` is counted back from the end: target = `len - pos`.
|
||||
RelativeToEnd,
|
||||
/// `pos` is a rope index relative to the start of the rope.
|
||||
Rope,
|
||||
}
|
||||
|
||||
// ── shared state ──────────────────────────────────────────────────────────────
|
||||
|
||||
/// Per-cursor cache of the last accessed block, the current position, and the
|
||||
/// base offset that defines the cursor's local coordinate system.
|
||||
///
|
||||
/// All fields are [`Cell`]-wrapped so they can be mutated through a shared
|
||||
/// reference, enabling `&self` methods on cursors.
|
||||
#[derive(Clone)]
|
||||
pub struct CursorState<'a> {
|
||||
block_idx: Cell<usize>,
|
||||
block_start: Cell<usize>,
|
||||
block_end: Cell<usize>,
|
||||
block: Cell<&'a [Cell<u8>]>,
|
||||
initialized: Cell<bool>,
|
||||
current: Cell<Option<usize>>,
|
||||
/// Absolute rope index that maps to local position 0.
|
||||
/// All user-facing coordinates are relative to this value.
|
||||
offset: Cell<usize>,
|
||||
}
|
||||
|
||||
impl<'a> CursorState<'a> {
|
||||
fn new() -> Self {
|
||||
Self::with_offset(0)
|
||||
}
|
||||
|
||||
fn with_offset(offset: usize) -> Self {
|
||||
Self {
|
||||
block_idx: Cell::new(0),
|
||||
block_start: Cell::new(0),
|
||||
block_end: Cell::new(0),
|
||||
block: Cell::new(&[]),
|
||||
initialized: Cell::new(false),
|
||||
current: Cell::new(None),
|
||||
offset: Cell::new(offset),
|
||||
}
|
||||
}
|
||||
|
||||
fn get(&self, rope: &'a Rope, i: usize) -> Option<u8> {
|
||||
if !self.initialized.get() || i < self.block_start.get() || i >= self.block_end.get() {
|
||||
let (bi, bs, be) = rope.lookup(i)?;
|
||||
self.block_idx.set(bi);
|
||||
self.block_start.set(bs);
|
||||
self.block_end.set(be);
|
||||
self.block.set(rope.get_block(bi)?);
|
||||
self.initialized.set(true);
|
||||
}
|
||||
Some(self.block.get()[i - self.block_start.get()].get())
|
||||
}
|
||||
|
||||
fn set(&self, rope: &'a Rope, i: usize, value: u8) -> Result<(), RopeError> {
|
||||
if !self.initialized.get() || i < self.block_start.get() || i >= self.block_end.get() {
|
||||
let (bi, bs, be) = rope.lookup(i).ok_or(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} > {}",
|
||||
i,
|
||||
rope.len()
|
||||
)))?;
|
||||
self.block_idx.set(bi);
|
||||
self.block_start.set(bs);
|
||||
self.block_end.set(be);
|
||||
self.block
|
||||
.set(rope.get_block(bi).ok_or(RopeError::BlockNotFound(format!(
|
||||
"Cannot find block for index {}",
|
||||
i
|
||||
)))?);
|
||||
self.initialized.set(true);
|
||||
}
|
||||
self.block.get()[i - self.block_start.get()].set(value);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
// ── trait ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Common interface for all rope cursors.
|
||||
///
|
||||
/// # Required methods
|
||||
///
|
||||
/// Implementors must provide [`rope`](RopeCursor::rope),
|
||||
/// [`state`](RopeCursor::state), [`read_next`](RopeCursor::read_next) and
|
||||
/// [`seek`](RopeCursor::seek). Everything else has a default implementation.
|
||||
///
|
||||
/// The direction of `read_next` and the sign convention for
|
||||
/// [`SeekMode::Relative`] differ between [`ForwardCursor`] and
|
||||
/// [`BackwardCursor`]; all other methods are identical.
|
||||
pub trait RopeCursor<'a> {
|
||||
/// The rope this cursor is bound to.
|
||||
fn rope(&self) -> &'a Rope;
|
||||
/// Internal cache state — implementation detail exposed for default methods.
|
||||
fn state(&self) -> &CursorState<'a>;
|
||||
|
||||
/// Read the next byte in cursor direction and advance the position.
|
||||
/// Returns `Err` at the exhausted end.
|
||||
fn read_next(&self) -> Result<u8, RopeError>;
|
||||
|
||||
/// Move the cursor to a new position.
|
||||
///
|
||||
/// `pos` is interpreted according to `mode`:
|
||||
/// - [`Absolute`](SeekMode::Absolute): local coordinate (`pos + offset` in the rope).
|
||||
/// - [`Rope`](SeekMode::Rope): raw rope index, ignores the offset. Pass a value
|
||||
/// from [`rope_tell`](RopeCursor::rope_tell) to restore a saved position.
|
||||
/// - [`Relative`](SeekMode::Relative): delta from the current position.
|
||||
/// For [`ForwardCursor`], positive advances toward the end;
|
||||
/// for [`BackwardCursor`], positive retreats toward the start.
|
||||
/// - [`RelativeToEnd`](SeekMode::RelativeToEnd): `rope.len() - pos`.
|
||||
///
|
||||
/// Returns the new position as a **rope index** (same value as
|
||||
/// [`rope_tell`](RopeCursor::rope_tell) would return immediately after).
|
||||
fn seek(&self, pos: isize, mode: SeekMode) -> Result<usize, RopeError>;
|
||||
|
||||
// ── default methods ───────────────────────────────────────────────────────
|
||||
|
||||
/// Read the byte at **local** index `i` (relative to the cursor's offset)
|
||||
/// without moving the position.
|
||||
fn get(&self, i: usize) -> Option<u8> {
|
||||
self.state().get(self.rope(), i + self.state().offset.get())
|
||||
}
|
||||
|
||||
/// Write `value` at **local** index `i` without moving the position.
|
||||
fn set(&self, i: usize, value: u8) -> Result<(), RopeError> {
|
||||
self.state()
|
||||
.set(self.rope(), i + self.state().offset.get(), value)
|
||||
}
|
||||
|
||||
/// Current position relative to the cursor's offset, or `None` if the
|
||||
/// cursor has not moved yet.
|
||||
fn tell(&self) -> Option<usize> {
|
||||
let abs = self.state().current.get()?;
|
||||
Some(abs.saturating_sub(self.state().offset.get()))
|
||||
}
|
||||
|
||||
/// Current position as an absolute rope index.
|
||||
///
|
||||
/// Unlike [`tell`](RopeCursor::tell), this method **always** returns a
|
||||
/// value: if the cursor has not moved yet, it returns the cursor's offset
|
||||
/// (the rope index of local position 0).
|
||||
///
|
||||
/// Use the returned value with [`SeekMode::Rope`] to restore a position,
|
||||
/// or as a truncation point after a write pass.
|
||||
fn rope_tell(&self) -> usize {
|
||||
self.state()
|
||||
.current
|
||||
.get()
|
||||
.unwrap_or(self.state().offset.get())
|
||||
}
|
||||
|
||||
/// Number of bytes visible through this cursor (`rope.len() - offset`).
|
||||
fn len(&self) -> usize {
|
||||
self.rope().len().saturating_sub(self.state().offset.get())
|
||||
}
|
||||
|
||||
/// Reset the cursor to its initial state (positioned before the first
|
||||
/// byte of its local view). Equivalent to `seek(0, Absolute)` on a
|
||||
/// fresh cursor, but works even when `current` is `None`.
|
||||
fn reset(&self) {
|
||||
self.state().current.set(None);
|
||||
}
|
||||
|
||||
/// Read the byte at the current position without advancing.
|
||||
fn peek(&self) -> Option<u8> {
|
||||
self.state().get(self.rope(), self.state().current.get()?)
|
||||
}
|
||||
|
||||
/// Write `value` at the current position without advancing.
|
||||
fn poke(&self, value: u8) -> Result<(), RopeError> {
|
||||
let pos = self.state().current.get().ok_or(RopeError::CurrentNotSet)?;
|
||||
self.state().set(self.rope(), pos, value)
|
||||
}
|
||||
|
||||
/// Move backward by `go_back_of` steps (toward lower indices for
|
||||
/// [`ForwardCursor`], toward higher indices for [`BackwardCursor`]).
|
||||
fn rewind(&self, go_back_of: usize) -> Result<(), RopeError> {
|
||||
self.seek(-(go_back_of as isize), SeekMode::Relative)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Move forward by `ahead` steps (opposite of [`rewind`](RopeCursor::rewind)).
|
||||
fn forward(&self, ahead: usize) -> Result<(), RopeError> {
|
||||
self.seek(ahead as isize, SeekMode::Relative)?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
// ── ForwardCursor ─────────────────────────────────────────────────────────────
|
||||
|
||||
/// A cursor that reads from the start toward the end of the rope.
|
||||
///
|
||||
/// - `read_next`: first call reads index 0, then 1, 2, …
|
||||
/// - `seek(Relative, +n)`: advances by n.
|
||||
/// - `rewind(n)`: steps back by n.
|
||||
///
|
||||
/// Extra methods not in the trait: [`read_ahead`](ForwardCursor::read_ahead),
|
||||
/// [`write`](ForwardCursor::write), [`iter`](ForwardCursor::iter).
|
||||
#[derive(Clone)]
|
||||
pub struct ForwardCursor<'a> {
|
||||
rope: &'a Rope,
|
||||
state: CursorState<'a>,
|
||||
}
|
||||
|
||||
impl<'a> ForwardCursor<'a> {
|
||||
/// Create a new forward cursor positioned before the first byte.
|
||||
pub fn new(rope: &'a Rope) -> Self {
|
||||
Self {
|
||||
rope,
|
||||
state: CursorState::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Read the byte at `current + ahead` without moving the position.
|
||||
pub fn read_ahead(&self, ahead: usize) -> Result<u8, RopeError> {
|
||||
let pos = self.state.current.get().ok_or(RopeError::CurrentNotSet)?;
|
||||
self.state
|
||||
.get(self.rope, pos + ahead)
|
||||
.ok_or(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} + {} > {}",
|
||||
pos,
|
||||
ahead,
|
||||
self.rope.len()
|
||||
)))
|
||||
}
|
||||
|
||||
/// Write `value` at the current position and advance by one.
|
||||
///
|
||||
/// If the cursor has not moved yet, writes at the first byte of its local
|
||||
/// view (absolute index = offset).
|
||||
pub fn write(&self, value: u8) -> Result<(), RopeError> {
|
||||
let pos = self.state.current.get().unwrap_or(self.state.offset.get());
|
||||
self.state.set(self.rope, pos, value)?;
|
||||
self.state.current.set(Some(pos + 1));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Return a shared-borrow iterator that yields bytes forward.
|
||||
///
|
||||
/// Because the iterator holds `&self` rather than `&mut self`, methods
|
||||
/// such as [`tell`](RopeCursor::tell) and [`seek`](RopeCursor::seek) can
|
||||
/// be called on the cursor inside the loop body.
|
||||
pub fn iter(&self) -> ForwardIter<'a, '_> {
|
||||
ForwardIter { cursor: self }
|
||||
}
|
||||
|
||||
/// Create a new [`ForwardCursor`] whose local position 0 starts at the
|
||||
/// current absolute position of `self`.
|
||||
///
|
||||
/// The new cursor shares the same underlying [`Rope`] (with the same
|
||||
/// [`Cell`]-based interior mutability) but has an independent position and
|
||||
/// an `offset` equal to `self.absolute_tell()`. If `self` has not moved
|
||||
/// yet, the new cursor starts at the same offset as `self`.
|
||||
pub fn cursor(&self) -> ForwardCursor<'a> {
|
||||
let new_offset = self.rope_tell();
|
||||
ForwardCursor {
|
||||
rope: self.rope,
|
||||
state: CursorState::with_offset(new_offset),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> RopeCursor<'a> for ForwardCursor<'a> {
|
||||
fn rope(&self) -> &'a Rope {
|
||||
self.rope
|
||||
}
|
||||
fn state(&self) -> &CursorState<'a> {
|
||||
&self.state
|
||||
}
|
||||
|
||||
fn read_next(&self) -> Result<u8, RopeError> {
|
||||
let next_pos = match self.state.current.get() {
|
||||
Some(i) => i + 1,
|
||||
None => self.state.offset.get(),
|
||||
};
|
||||
let value = self
|
||||
.state
|
||||
.get(self.rope, next_pos)
|
||||
.ok_or(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} > {}",
|
||||
next_pos,
|
||||
self.rope.len()
|
||||
)))?;
|
||||
self.state.current.set(Some(next_pos));
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
fn seek(&self, pos: isize, mode: SeekMode) -> Result<usize, RopeError> {
|
||||
let offset = self.state.offset.get() as isize;
|
||||
let abs_pos = match mode {
|
||||
SeekMode::Absolute => pos + offset,
|
||||
SeekMode::Relative => {
|
||||
self.state.current.get().ok_or(RopeError::CurrentNotSet)? as isize + pos
|
||||
}
|
||||
SeekMode::RelativeToEnd => self.rope.len() as isize - pos,
|
||||
SeekMode::Rope => pos,
|
||||
};
|
||||
if abs_pos < 0 {
|
||||
return Err(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} < 0",
|
||||
abs_pos
|
||||
)));
|
||||
}
|
||||
self.state.current.set(Some(abs_pos as usize));
|
||||
Ok(abs_pos as usize)
|
||||
}
|
||||
}
|
||||
|
||||
impl Iterator for ForwardCursor<'_> {
|
||||
type Item = u8;
|
||||
fn next(&mut self) -> Option<Self::Item> {
|
||||
self.read_next().ok()
|
||||
}
|
||||
}
|
||||
|
||||
/// Shared-borrow iterator returned by [`ForwardCursor::iter`].
|
||||
pub struct ForwardIter<'a, 'b> {
|
||||
cursor: &'b ForwardCursor<'a>,
|
||||
}
|
||||
|
||||
impl Iterator for ForwardIter<'_, '_> {
|
||||
type Item = u8;
|
||||
fn next(&mut self) -> Option<u8> {
|
||||
self.cursor.read_next().ok()
|
||||
}
|
||||
}
|
||||
|
||||
// ── BackwardCursor ────────────────────────────────────────────────────────────
|
||||
|
||||
/// A cursor that reads from the end toward the start of the rope.
|
||||
///
|
||||
/// - `read_next`: first call reads index `len-1`, then `len-2`, …
|
||||
/// - `seek(Relative, +n)`: retreats by n (subtracts n from the index).
|
||||
/// - `rewind(n)`: advances toward the end by n.
|
||||
///
|
||||
/// Extra methods not in the trait: [`read_behind`](BackwardCursor::read_behind),
|
||||
/// [`iter`](BackwardCursor::iter).
|
||||
#[derive(Clone)]
|
||||
pub struct BackwardCursor<'a> {
|
||||
rope: &'a Rope,
|
||||
state: CursorState<'a>,
|
||||
}
|
||||
|
||||
impl<'a> BackwardCursor<'a> {
|
||||
/// Create a new backward cursor positioned past the last byte.
|
||||
pub fn new(rope: &'a Rope) -> Self {
|
||||
Self {
|
||||
rope,
|
||||
state: CursorState::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Read the byte at `current + behind` (toward higher indices) without moving.
|
||||
pub fn read_behind(&self, behind: usize) -> Result<u8, RopeError> {
|
||||
let pos = self.state.current.get().ok_or(RopeError::CurrentNotSet)?;
|
||||
let target = pos
|
||||
.checked_add(behind)
|
||||
.filter(|&t| t < self.rope.len())
|
||||
.ok_or(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} + {} > {}",
|
||||
pos,
|
||||
behind,
|
||||
self.rope.len()
|
||||
)))?;
|
||||
self.state
|
||||
.get(self.rope, target)
|
||||
.ok_or(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} + {} > {}",
|
||||
pos,
|
||||
behind,
|
||||
self.rope.len()
|
||||
)))
|
||||
}
|
||||
|
||||
/// Return a shared-borrow iterator that yields bytes backward.
|
||||
///
|
||||
/// Because the iterator holds `&self` rather than `&mut self`, methods
|
||||
/// such as [`tell`](RopeCursor::tell) and [`seek`](RopeCursor::seek) can
|
||||
/// be called on the cursor inside the loop body.
|
||||
pub fn iter(&self) -> BackwardIter<'a, '_> {
|
||||
BackwardIter { cursor: self }
|
||||
}
|
||||
|
||||
/// Create a new [`BackwardCursor`] that stops at the current absolute
|
||||
/// position of `self` (used as the lower bound / offset of the new cursor).
|
||||
///
|
||||
/// The new cursor scans from `rope.len() - 1` down to the current absolute
|
||||
/// position of `self`. If `self` has not moved yet, the new cursor has the
|
||||
/// same offset as `self` (no restriction).
|
||||
pub fn cursor(&self) -> BackwardCursor<'a> {
|
||||
let new_offset = self.rope_tell();
|
||||
BackwardCursor {
|
||||
rope: self.rope,
|
||||
state: CursorState::with_offset(new_offset),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> RopeCursor<'a> for BackwardCursor<'a> {
|
||||
fn rope(&self) -> &'a Rope {
|
||||
self.rope
|
||||
}
|
||||
fn state(&self) -> &CursorState<'a> {
|
||||
&self.state
|
||||
}
|
||||
|
||||
fn read_next(&self) -> Result<u8, RopeError> {
|
||||
let offset = self.state.offset.get();
|
||||
let next_pos = match self.state.current.get() {
|
||||
None => self
|
||||
.rope
|
||||
.len()
|
||||
.checked_sub(1)
|
||||
.ok_or(RopeError::OutOfBounds(
|
||||
"BackwardCursor: rope is empty".to_string(),
|
||||
))?,
|
||||
Some(i) if i <= offset => {
|
||||
return Err(RopeError::OutOfBounds(
|
||||
"BackwardCursor: already at beginning".to_string(),
|
||||
));
|
||||
}
|
||||
Some(i) => i - 1,
|
||||
};
|
||||
let value = self
|
||||
.state
|
||||
.get(self.rope, next_pos)
|
||||
.ok_or(RopeError::OutOfBounds(format!(
|
||||
"BackwardCursor: index out of bounds at i={}",
|
||||
next_pos
|
||||
)))?;
|
||||
self.state.current.set(Some(next_pos));
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
fn seek(&self, pos: isize, mode: SeekMode) -> Result<usize, RopeError> {
|
||||
let offset = self.state.offset.get() as isize;
|
||||
let abs_pos = match mode {
|
||||
SeekMode::Absolute => pos + offset,
|
||||
SeekMode::Relative => {
|
||||
self.state.current.get().ok_or(RopeError::CurrentNotSet)? as isize - pos
|
||||
}
|
||||
SeekMode::RelativeToEnd => self.rope.len() as isize - pos,
|
||||
SeekMode::Rope => pos,
|
||||
};
|
||||
if abs_pos < 0 {
|
||||
return Err(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} < 0",
|
||||
abs_pos
|
||||
)));
|
||||
}
|
||||
self.state.current.set(Some(abs_pos as usize));
|
||||
Ok(abs_pos as usize)
|
||||
}
|
||||
}
|
||||
|
||||
impl Iterator for BackwardCursor<'_> {
|
||||
type Item = u8;
|
||||
fn next(&mut self) -> Option<Self::Item> {
|
||||
self.read_next().ok()
|
||||
}
|
||||
}
|
||||
|
||||
/// Shared-borrow iterator returned by [`BackwardCursor::iter`].
|
||||
pub struct BackwardIter<'a, 'b> {
|
||||
cursor: &'b BackwardCursor<'a>,
|
||||
}
|
||||
|
||||
impl Iterator for BackwardIter<'_, '_> {
|
||||
type Item = u8;
|
||||
fn next(&mut self) -> Option<u8> {
|
||||
self.cursor.read_next().ok()
|
||||
}
|
||||
}
|
||||
|
||||
// ── tests ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::Rope;
|
||||
|
||||
fn rope(data: &[u8]) -> Rope {
|
||||
let mut r = Rope::new(None);
|
||||
r.push(data.to_vec());
|
||||
r
|
||||
}
|
||||
|
||||
fn rope2(a: &[u8], b: &[u8]) -> Rope {
|
||||
let mut r = Rope::new(None);
|
||||
r.push(a.to_vec());
|
||||
r.push(b.to_vec());
|
||||
r
|
||||
}
|
||||
|
||||
// ── ForwardCursor ─────────────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn forward_reads_all_bytes() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
let out: Vec<u8> = c.collect();
|
||||
assert_eq!(out, b"ACGT");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_tell_tracks_position() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
assert_eq!(c.tell(), None);
|
||||
c.read_next().unwrap();
|
||||
assert_eq!(c.tell(), Some(0));
|
||||
c.read_next().unwrap();
|
||||
assert_eq!(c.tell(), Some(1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_iter_with_tell_inside_loop() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
let mut positions = Vec::new();
|
||||
for _ in c.iter() {
|
||||
positions.push(c.tell());
|
||||
}
|
||||
assert_eq!(positions, vec![Some(0), Some(1), Some(2), Some(3)]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_read_ahead() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // at 0 = 'A'
|
||||
assert_eq!(c.read_ahead(1).unwrap(), b'C');
|
||||
assert_eq!(c.read_ahead(2).unwrap(), b'G');
|
||||
assert_eq!(c.tell(), Some(0)); // position unchanged
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_write_and_read_back() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
c.write(b'X').unwrap();
|
||||
c.write(b'Y').unwrap();
|
||||
let c2 = r.fw_cursor();
|
||||
assert_eq!(c2.read_next().unwrap(), b'X');
|
||||
assert_eq!(c2.read_next().unwrap(), b'Y');
|
||||
assert_eq!(c2.read_next().unwrap(), b'G');
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_rewind_and_reread() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // A → current = Some(0)
|
||||
c.read_next().unwrap(); // C → current = Some(1)
|
||||
c.read_next().unwrap(); // G → current = Some(2)
|
||||
c.rewind(1).unwrap(); // current = Some(1) → next read = index 2
|
||||
assert_eq!(c.read_next().unwrap(), b'G');
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_seek_absolute() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
c.seek(2, SeekMode::Absolute).unwrap();
|
||||
assert_eq!(c.read_next().unwrap(), b'T');
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_seek_relative_to_end() {
|
||||
let r = rope(b"ACGT");
|
||||
// seek(1, RelativeToEnd): current = len-1 = 3; peek() reads index 3 = T.
|
||||
let c = r.fw_cursor();
|
||||
c.seek(1, SeekMode::RelativeToEnd).unwrap();
|
||||
assert_eq!(c.peek().unwrap(), b'T');
|
||||
// seek(2, RelativeToEnd): current = len-2 = 2; read_next reads index 3 = T.
|
||||
let c2 = r.fw_cursor();
|
||||
c2.seek(2, SeekMode::RelativeToEnd).unwrap();
|
||||
assert_eq!(c2.read_next().unwrap(), b'T');
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_get_random_access() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
assert_eq!(c.get(0), Some(b'A'));
|
||||
assert_eq!(c.get(3), Some(b'T'));
|
||||
assert_eq!(c.get(4), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_crosses_block_boundary() {
|
||||
let r = rope2(b"AC", b"GT");
|
||||
let c = r.fw_cursor();
|
||||
let out: Vec<u8> = c.collect();
|
||||
assert_eq!(out, b"ACGT");
|
||||
}
|
||||
|
||||
// ── BackwardCursor ────────────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn backward_reads_all_bytes_in_reverse() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.bw_cursor();
|
||||
let out: Vec<u8> = c.collect();
|
||||
assert_eq!(out, b"TGCA");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_tell_tracks_position() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.bw_cursor();
|
||||
assert_eq!(c.tell(), None);
|
||||
c.read_next().unwrap(); // reads index 3
|
||||
assert_eq!(c.tell(), Some(3));
|
||||
c.read_next().unwrap(); // reads index 2
|
||||
assert_eq!(c.tell(), Some(2));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_iter_with_tell_and_seek_inside_loop() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.bw_cursor();
|
||||
let mut restart: usize = 0;
|
||||
for byte in c.iter() {
|
||||
if byte == b'G' {
|
||||
restart = c.tell().unwrap();
|
||||
}
|
||||
if byte == b'A' {
|
||||
// seek back to G and break
|
||||
c.seek(restart as isize, SeekMode::Absolute).ok();
|
||||
break;
|
||||
}
|
||||
}
|
||||
assert_eq!(c.tell(), Some(restart));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_rewind_moves_toward_end() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.bw_cursor();
|
||||
c.read_next().unwrap(); // index 3 = T
|
||||
c.read_next().unwrap(); // index 2 = G
|
||||
c.rewind(1).unwrap(); // back to index 3
|
||||
assert_eq!(c.tell(), Some(3));
|
||||
assert_eq!(c.read_next().unwrap(), b'G'); // reads index 2
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_crosses_block_boundary() {
|
||||
let r = rope2(b"AC", b"GT");
|
||||
let c = r.bw_cursor();
|
||||
let out: Vec<u8> = c.collect();
|
||||
assert_eq!(out, b"TGCA");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_empty_rope_returns_error() {
|
||||
let r = Rope::new(None);
|
||||
let c = r.bw_cursor();
|
||||
assert!(c.read_next().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_empty_rope_returns_error() {
|
||||
let r = Rope::new(None);
|
||||
let c = r.fw_cursor();
|
||||
assert!(c.read_next().is_err());
|
||||
}
|
||||
|
||||
// ── offset / sub-cursor ───────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn forward_cursor_reads_from_offset() {
|
||||
// cursor() at current=Some(2) → new cursor reads from index 2
|
||||
let r = rope(b"ABCDE");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // A → current=Some(0)
|
||||
c.read_next().unwrap(); // B → current=Some(1)
|
||||
c.read_next().unwrap(); // C → current=Some(2)
|
||||
let sub = c.cursor(); // offset=2 (absolute_tell=2)
|
||||
assert_eq!(sub.read_next().unwrap(), b'C'); // reads index 2
|
||||
assert_eq!(sub.tell(), Some(0)); // relative: 2-2=0
|
||||
assert_eq!(sub.rope_tell(), 2);
|
||||
assert_eq!(sub.read_next().unwrap(), b'D');
|
||||
assert_eq!(sub.tell(), Some(1)); // relative: 3-2=1
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_cursor_get_uses_relative_index() {
|
||||
let r = rope(b"ABCDE");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // A → current=Some(0), absolute_tell=0
|
||||
let _sub = c.cursor(); // offset=0 — created to show cursor() compiles; not used further
|
||||
// From sub2 with offset=2: get(0)=C, get(2)=E
|
||||
let c2 = r.fw_cursor();
|
||||
c2.read_next().unwrap(); // at 0
|
||||
c2.read_next().unwrap(); // at 1
|
||||
c2.read_next().unwrap(); // at 2, absolute=2
|
||||
let sub2 = c2.cursor(); // offset=2
|
||||
assert_eq!(sub2.get(0), Some(b'C')); // local 0 = absolute 2
|
||||
assert_eq!(sub2.get(2), Some(b'E')); // local 2 = absolute 4
|
||||
assert_eq!(sub2.get(3), None); // local 3 = absolute 5, OOB
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_cursor_len_reflects_offset() {
|
||||
let r = rope(b"ABCDE"); // len=5
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap();
|
||||
c.read_next().unwrap();
|
||||
c.read_next().unwrap(); // absolute_tell=2
|
||||
let sub = c.cursor(); // offset=2
|
||||
assert_eq!(sub.len(), 3); // 5 - 2
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_reset_goes_back_to_start() {
|
||||
let r = rope(b"ABCDE");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // A
|
||||
c.read_next().unwrap(); // B
|
||||
c.reset();
|
||||
assert_eq!(c.tell(), None);
|
||||
assert_eq!(c.read_next().unwrap(), b'A'); // starts over
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_sub_cursor_write_and_reset() {
|
||||
// Write two bytes, discard them via reset(), write again.
|
||||
let r = rope(b"XXXXX");
|
||||
let c = r.fw_cursor();
|
||||
c.write(b'A').unwrap(); // absolute 0 → current=Some(1)
|
||||
c.write(b'B').unwrap(); // absolute 1 → current=Some(2)
|
||||
let seg = c.cursor(); // absolute_tell=2, offset=2
|
||||
seg.write(b'C').unwrap(); // absolute 2 → current=Some(3), tell=3-2=1
|
||||
seg.write(b'D').unwrap(); // absolute 3 → current=Some(4), tell=4-2=2
|
||||
assert_eq!(seg.tell(), Some(2)); // 2 bytes written into this segment
|
||||
seg.reset();
|
||||
assert_eq!(seg.tell(), None);
|
||||
seg.write(b'E').unwrap(); // absolute 2 again
|
||||
let all: Vec<u8> = r.fw_cursor().collect();
|
||||
assert_eq!(&all[..3], b"ABE");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_cursor_stops_at_offset() {
|
||||
// BackwardCursor.cursor() creates a cursor with offset = absolute_tell.
|
||||
// offset = local position 0 (inclusive lower bound).
|
||||
// The cursor reads rope.len()-1 downto offset, then stops.
|
||||
let r = rope(b"ABCDE"); // 0=A 1=B 2=C 3=D 4=E
|
||||
let bw = r.bw_cursor();
|
||||
bw.read_next().unwrap(); // E=4, current=Some(4)
|
||||
bw.read_next().unwrap(); // D=3, current=Some(3), absolute_tell=3
|
||||
// sub: offset=3, reads from 4 down to 3 (inclusive), then stops.
|
||||
let sub = bw.cursor();
|
||||
assert_eq!(sub.read_next().unwrap(), b'E'); // index 4, tell=4-3=1
|
||||
assert_eq!(sub.read_next().unwrap(), b'D'); // index 3, tell=3-3=0 (local 0)
|
||||
assert!(sub.read_next().is_err()); // would go to 2 < offset=3
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_absolute_tell_unchanged_by_offset() {
|
||||
let r = rope(b"ABCDE");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // absolute=0
|
||||
let sub = c.cursor(); // offset=0
|
||||
sub.read_next().unwrap(); // reads index 0, absolute_tell=0
|
||||
sub.read_next().unwrap(); // reads index 1, absolute_tell=1
|
||||
assert_eq!(sub.tell(), Some(1));
|
||||
assert_eq!(sub.rope_tell(), 1);
|
||||
// sub2 with offset=1
|
||||
let sub2 = sub.cursor(); // offset=1
|
||||
sub2.read_next().unwrap(); // reads index 1, absolute=1
|
||||
assert_eq!(sub2.tell(), Some(0)); // relative: 1-1=0
|
||||
assert_eq!(sub2.rope_tell(), 1);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
use crate::{Rope, RopeError};
|
||||
|
||||
use super::state::CursorState;
|
||||
use super::traits::{RopeCursor, SeekMode};
|
||||
|
||||
/// A cursor that reads from the end toward the start of the rope.
|
||||
///
|
||||
/// - `read_next`: first call reads index `len-1`, then `len-2`, …
|
||||
/// - `seek(Relative, +n)`: retreats by n (subtracts n from the index).
|
||||
/// - `rewind(n)`: advances toward the end by n.
|
||||
///
|
||||
/// Extra methods not in the trait: [`read_behind`](BackwardCursor::read_behind),
|
||||
/// [`iter`](BackwardCursor::iter).
|
||||
#[derive(Clone)]
|
||||
pub struct BackwardCursor<'a> {
|
||||
rope: &'a Rope,
|
||||
state: CursorState<'a>,
|
||||
}
|
||||
|
||||
impl<'a> BackwardCursor<'a> {
|
||||
/// Create a new backward cursor positioned past the last byte.
|
||||
pub fn new(rope: &'a Rope) -> Self {
|
||||
Self {
|
||||
rope,
|
||||
state: CursorState::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Read the byte at `current + behind` (toward higher indices) without moving.
|
||||
pub fn read_behind(&self, behind: usize) -> Result<u8, RopeError> {
|
||||
let pos = self.state.current.get().ok_or(RopeError::CurrentNotSet)?;
|
||||
let target = pos
|
||||
.checked_add(behind)
|
||||
.filter(|&t| t < self.rope.len())
|
||||
.ok_or(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} + {} > {}",
|
||||
pos,
|
||||
behind,
|
||||
self.rope.len()
|
||||
)))?;
|
||||
self.state
|
||||
.get(self.rope, target)
|
||||
.ok_or(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} + {} > {}",
|
||||
pos,
|
||||
behind,
|
||||
self.rope.len()
|
||||
)))
|
||||
}
|
||||
|
||||
/// Return a shared-borrow iterator that yields bytes backward.
|
||||
///
|
||||
/// Because the iterator holds `&self` rather than `&mut self`, methods
|
||||
/// such as [`tell`](RopeCursor::tell) and [`seek`](RopeCursor::seek) can
|
||||
/// be called on the cursor inside the loop body.
|
||||
pub fn iter(&self) -> BackwardIter<'a, '_> {
|
||||
BackwardIter { cursor: self }
|
||||
}
|
||||
|
||||
/// Create a new [`BackwardCursor`] that stops at the current absolute
|
||||
/// position of `self` (used as the lower bound / offset of the new cursor).
|
||||
///
|
||||
/// The new cursor scans from `rope.len() - 1` down to the current absolute
|
||||
/// position of `self`. If `self` has not moved yet, the new cursor has the
|
||||
/// same offset as `self` (no restriction).
|
||||
pub fn cursor(&self) -> BackwardCursor<'a> {
|
||||
let new_offset = self.rope_tell();
|
||||
BackwardCursor {
|
||||
rope: self.rope,
|
||||
state: CursorState::with_offset(new_offset),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> RopeCursor<'a> for BackwardCursor<'a> {
|
||||
fn rope(&self) -> &'a Rope {
|
||||
self.rope
|
||||
}
|
||||
fn state(&self) -> &CursorState<'a> {
|
||||
&self.state
|
||||
}
|
||||
|
||||
fn read_next(&self) -> Result<u8, RopeError> {
|
||||
let offset = self.state.offset.get();
|
||||
let next_pos = match self.state.current.get() {
|
||||
None => self
|
||||
.rope
|
||||
.len()
|
||||
.checked_sub(1)
|
||||
.ok_or(RopeError::OutOfBounds(
|
||||
"BackwardCursor: rope is empty".to_string(),
|
||||
))?,
|
||||
Some(i) if i <= offset => {
|
||||
return Err(RopeError::OutOfBounds(
|
||||
"BackwardCursor: already at beginning".to_string(),
|
||||
));
|
||||
}
|
||||
Some(i) => i - 1,
|
||||
};
|
||||
let value = self
|
||||
.state
|
||||
.get(self.rope, next_pos)
|
||||
.ok_or(RopeError::OutOfBounds(format!(
|
||||
"BackwardCursor: index out of bounds at i={}",
|
||||
next_pos
|
||||
)))?;
|
||||
self.state.current.set(Some(next_pos));
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
fn seek(&self, pos: isize, mode: SeekMode) -> Result<usize, RopeError> {
|
||||
let offset = self.state.offset.get() as isize;
|
||||
let abs_pos = match mode {
|
||||
SeekMode::Absolute => pos + offset,
|
||||
SeekMode::Relative => {
|
||||
self.state.current.get().ok_or(RopeError::CurrentNotSet)? as isize - pos
|
||||
}
|
||||
SeekMode::RelativeToEnd => self.rope.len() as isize - pos,
|
||||
SeekMode::Rope => pos,
|
||||
};
|
||||
if abs_pos < 0 {
|
||||
return Err(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} < 0",
|
||||
abs_pos
|
||||
)));
|
||||
}
|
||||
self.state.current.set(Some(abs_pos as usize));
|
||||
Ok(abs_pos as usize)
|
||||
}
|
||||
}
|
||||
|
||||
impl Iterator for BackwardCursor<'_> {
|
||||
type Item = u8;
|
||||
fn next(&mut self) -> Option<Self::Item> {
|
||||
self.read_next().ok()
|
||||
}
|
||||
}
|
||||
|
||||
/// Shared-borrow iterator returned by [`BackwardCursor::iter`].
|
||||
pub struct BackwardIter<'a, 'b> {
|
||||
cursor: &'b BackwardCursor<'a>,
|
||||
}
|
||||
|
||||
impl Iterator for BackwardIter<'_, '_> {
|
||||
type Item = u8;
|
||||
fn next(&mut self) -> Option<u8> {
|
||||
self.cursor.read_next().ok()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,142 @@
|
||||
use crate::{Rope, RopeError};
|
||||
|
||||
use super::state::CursorState;
|
||||
use super::traits::{RopeCursor, SeekMode};
|
||||
|
||||
/// A cursor that reads from the start toward the end of the rope.
|
||||
///
|
||||
/// - `read_next`: first call reads index 0, then 1, 2, …
|
||||
/// - `seek(Relative, +n)`: advances by n.
|
||||
/// - `rewind(n)`: steps back by n.
|
||||
///
|
||||
/// Extra methods not in the trait: [`read_ahead`](ForwardCursor::read_ahead),
|
||||
/// [`write`](ForwardCursor::write), [`iter`](ForwardCursor::iter).
|
||||
#[derive(Clone)]
|
||||
pub struct ForwardCursor<'a> {
|
||||
rope: &'a Rope,
|
||||
state: CursorState<'a>,
|
||||
}
|
||||
|
||||
impl<'a> ForwardCursor<'a> {
|
||||
/// Create a new forward cursor positioned before the first byte.
|
||||
pub fn new(rope: &'a Rope) -> Self {
|
||||
Self {
|
||||
rope,
|
||||
state: CursorState::new(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Read the byte at `current + ahead` without moving the position.
|
||||
pub fn read_ahead(&self, ahead: usize) -> Result<u8, RopeError> {
|
||||
let pos = self.state.current.get().ok_or(RopeError::CurrentNotSet)?;
|
||||
self.state
|
||||
.get(self.rope, pos + ahead)
|
||||
.ok_or(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} + {} > {}",
|
||||
pos,
|
||||
ahead,
|
||||
self.rope.len()
|
||||
)))
|
||||
}
|
||||
|
||||
/// Write `value` at the current position and advance by one.
|
||||
///
|
||||
/// If the cursor has not moved yet, writes at the first byte of its local
|
||||
/// view (absolute index = offset).
|
||||
pub fn write(&self, value: u8) -> Result<(), RopeError> {
|
||||
let pos = self.state.current.get().unwrap_or(self.state.offset.get());
|
||||
self.state.set(self.rope, pos, value)?;
|
||||
self.state.current.set(Some(pos + 1));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Return a shared-borrow iterator that yields bytes forward.
|
||||
///
|
||||
/// Because the iterator holds `&self` rather than `&mut self`, methods
|
||||
/// such as [`tell`](RopeCursor::tell) and [`seek`](RopeCursor::seek) can
|
||||
/// be called on the cursor inside the loop body.
|
||||
pub fn iter(&self) -> ForwardIter<'a, '_> {
|
||||
ForwardIter { cursor: self }
|
||||
}
|
||||
|
||||
/// Create a new [`ForwardCursor`] whose local position 0 starts at the
|
||||
/// current absolute position of `self`.
|
||||
///
|
||||
/// The new cursor shares the same underlying [`Rope`] (with the same
|
||||
/// [`Cell`](std::cell::Cell)-based interior mutability) but has an
|
||||
/// independent position and an `offset` equal to `self.absolute_tell()`.
|
||||
/// If `self` has not moved yet, the new cursor starts at the same offset
|
||||
/// as `self`.
|
||||
pub fn cursor(&self) -> ForwardCursor<'a> {
|
||||
let new_offset = self.rope_tell();
|
||||
ForwardCursor {
|
||||
rope: self.rope,
|
||||
state: CursorState::with_offset(new_offset),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<'a> RopeCursor<'a> for ForwardCursor<'a> {
|
||||
fn rope(&self) -> &'a Rope {
|
||||
self.rope
|
||||
}
|
||||
fn state(&self) -> &CursorState<'a> {
|
||||
&self.state
|
||||
}
|
||||
|
||||
fn read_next(&self) -> Result<u8, RopeError> {
|
||||
let next_pos = match self.state.current.get() {
|
||||
Some(i) => i + 1,
|
||||
None => self.state.offset.get(),
|
||||
};
|
||||
let value = self
|
||||
.state
|
||||
.get(self.rope, next_pos)
|
||||
.ok_or(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} > {}",
|
||||
next_pos,
|
||||
self.rope.len()
|
||||
)))?;
|
||||
self.state.current.set(Some(next_pos));
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
fn seek(&self, pos: isize, mode: SeekMode) -> Result<usize, RopeError> {
|
||||
let offset = self.state.offset.get() as isize;
|
||||
let abs_pos = match mode {
|
||||
SeekMode::Absolute => pos + offset,
|
||||
SeekMode::Relative => {
|
||||
self.state.current.get().ok_or(RopeError::CurrentNotSet)? as isize + pos
|
||||
}
|
||||
SeekMode::RelativeToEnd => self.rope.len() as isize - pos,
|
||||
SeekMode::Rope => pos,
|
||||
};
|
||||
if abs_pos < 0 {
|
||||
return Err(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} < 0",
|
||||
abs_pos
|
||||
)));
|
||||
}
|
||||
self.state.current.set(Some(abs_pos as usize));
|
||||
Ok(abs_pos as usize)
|
||||
}
|
||||
}
|
||||
|
||||
impl Iterator for ForwardCursor<'_> {
|
||||
type Item = u8;
|
||||
fn next(&mut self) -> Option<Self::Item> {
|
||||
self.read_next().ok()
|
||||
}
|
||||
}
|
||||
|
||||
/// Shared-borrow iterator returned by [`ForwardCursor::iter`].
|
||||
pub struct ForwardIter<'a, 'b> {
|
||||
cursor: &'b ForwardCursor<'a>,
|
||||
}
|
||||
|
||||
impl Iterator for ForwardIter<'_, '_> {
|
||||
type Item = u8;
|
||||
fn next(&mut self) -> Option<u8> {
|
||||
self.cursor.read_next().ok()
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,53 @@
|
||||
//! Cursors for sequential and random access over a [`Rope`].
|
||||
//!
|
||||
//! # Design
|
||||
//!
|
||||
//! A cursor borrows a `&'a Rope` and keeps a small block cache so that
|
||||
//! consecutive accesses within the same block cost O(1). The first access to a
|
||||
//! new block costs O(log n) (binary search in [`Rope::lookup`]); subsequent
|
||||
//! accesses within that block are free.
|
||||
//!
|
||||
//! All mutable state (current position, cache) is stored in [`Cell`](std::cell::Cell)
|
||||
//! fields, so every cursor method takes `&self` rather than `&mut self`. This means:
|
||||
//!
|
||||
//! - Two cursors can coexist on the same rope without lifetime conflicts.
|
||||
//! - The `iter()` method returns a lightweight wrapper that holds `&Cursor`,
|
||||
//! allowing `cursor.tell()` or `cursor.seek()` to be called **inside a `for`
|
||||
//! loop** over the same cursor.
|
||||
//!
|
||||
//! # Cursors
|
||||
//!
|
||||
//! | Type | Direction | First `read_next` | `seek(Relative, +n)` |
|
||||
//! |------|-----------|-------------------|----------------------|
|
||||
//! | [`ForwardCursor`] | start → end | index 0 | advances (+n) |
|
||||
//! | [`BackwardCursor`] | end → start | index `len-1` | retreats (+n) |
|
||||
//!
|
||||
//! # Example
|
||||
//!
|
||||
//! ```
|
||||
//! use obikrope::{Rope, RopeCursor};
|
||||
//!
|
||||
//! let mut rope = Rope::new(None);
|
||||
//! rope.push(b"ACGT".to_vec());
|
||||
//!
|
||||
//! let cursor = rope.fw_cursor();
|
||||
//! for byte in cursor.iter() {
|
||||
//! // cursor.tell() is valid here — iter() holds &cursor, not &mut cursor
|
||||
//! let _ = cursor.tell();
|
||||
//! }
|
||||
//! ```
|
||||
//!
|
||||
//! Submodules: [`state`] (shared block-cache state), [`traits`] (`SeekMode`,
|
||||
//! `RopeCursor`), [`forward`]/[`backward`] (the two cursor implementations).
|
||||
|
||||
mod backward;
|
||||
mod forward;
|
||||
mod state;
|
||||
mod traits;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
pub use backward::BackwardCursor;
|
||||
pub use forward::ForwardCursor;
|
||||
pub use traits::{RopeCursor, SeekMode};
|
||||
@@ -0,0 +1,72 @@
|
||||
use std::cell::Cell;
|
||||
|
||||
use crate::{Rope, RopeError};
|
||||
|
||||
/// Per-cursor cache of the last accessed block, the current position, and the
|
||||
/// base offset that defines the cursor's local coordinate system.
|
||||
///
|
||||
/// All fields are [`Cell`]-wrapped so they can be mutated through a shared
|
||||
/// reference, enabling `&self` methods on cursors.
|
||||
#[derive(Clone)]
|
||||
pub struct CursorState<'a> {
|
||||
block_idx: Cell<usize>,
|
||||
block_start: Cell<usize>,
|
||||
block_end: Cell<usize>,
|
||||
block: Cell<&'a [Cell<u8>]>,
|
||||
initialized: Cell<bool>,
|
||||
pub(super) current: Cell<Option<usize>>,
|
||||
/// Absolute rope index that maps to local position 0.
|
||||
/// All user-facing coordinates are relative to this value.
|
||||
pub(super) offset: Cell<usize>,
|
||||
}
|
||||
|
||||
impl<'a> CursorState<'a> {
|
||||
pub(super) fn new() -> Self {
|
||||
Self::with_offset(0)
|
||||
}
|
||||
|
||||
pub(super) fn with_offset(offset: usize) -> Self {
|
||||
Self {
|
||||
block_idx: Cell::new(0),
|
||||
block_start: Cell::new(0),
|
||||
block_end: Cell::new(0),
|
||||
block: Cell::new(&[]),
|
||||
initialized: Cell::new(false),
|
||||
current: Cell::new(None),
|
||||
offset: Cell::new(offset),
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn get(&self, rope: &'a Rope, i: usize) -> Option<u8> {
|
||||
if !self.initialized.get() || i < self.block_start.get() || i >= self.block_end.get() {
|
||||
let (bi, bs, be) = rope.lookup(i)?;
|
||||
self.block_idx.set(bi);
|
||||
self.block_start.set(bs);
|
||||
self.block_end.set(be);
|
||||
self.block.set(rope.get_block(bi)?);
|
||||
self.initialized.set(true);
|
||||
}
|
||||
Some(self.block.get()[i - self.block_start.get()].get())
|
||||
}
|
||||
|
||||
pub(super) fn set(&self, rope: &'a Rope, i: usize, value: u8) -> Result<(), RopeError> {
|
||||
if !self.initialized.get() || i < self.block_start.get() || i >= self.block_end.get() {
|
||||
let (bi, bs, be) = rope.lookup(i).ok_or(RopeError::OutOfBounds(format!(
|
||||
"index out of bounds: i={} > {}",
|
||||
i,
|
||||
rope.len()
|
||||
)))?;
|
||||
self.block_idx.set(bi);
|
||||
self.block_start.set(bs);
|
||||
self.block_end.set(be);
|
||||
self.block
|
||||
.set(rope.get_block(bi).ok_or(RopeError::BlockNotFound(format!(
|
||||
"Cannot find block for index {}",
|
||||
i
|
||||
)))?);
|
||||
self.initialized.set(true);
|
||||
}
|
||||
self.block.get()[i - self.block_start.get()].set(value);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,298 @@
|
||||
use super::*;
|
||||
use crate::Rope;
|
||||
|
||||
fn rope(data: &[u8]) -> Rope {
|
||||
let mut r = Rope::new(None);
|
||||
r.push(data.to_vec());
|
||||
r
|
||||
}
|
||||
|
||||
fn rope2(a: &[u8], b: &[u8]) -> Rope {
|
||||
let mut r = Rope::new(None);
|
||||
r.push(a.to_vec());
|
||||
r.push(b.to_vec());
|
||||
r
|
||||
}
|
||||
|
||||
// ── ForwardCursor ─────────────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn forward_reads_all_bytes() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
let out: Vec<u8> = c.collect();
|
||||
assert_eq!(out, b"ACGT");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_tell_tracks_position() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
assert_eq!(c.tell(), None);
|
||||
c.read_next().unwrap();
|
||||
assert_eq!(c.tell(), Some(0));
|
||||
c.read_next().unwrap();
|
||||
assert_eq!(c.tell(), Some(1));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_iter_with_tell_inside_loop() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
let mut positions = Vec::new();
|
||||
for _ in c.iter() {
|
||||
positions.push(c.tell());
|
||||
}
|
||||
assert_eq!(positions, vec![Some(0), Some(1), Some(2), Some(3)]);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_read_ahead() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // at 0 = 'A'
|
||||
assert_eq!(c.read_ahead(1).unwrap(), b'C');
|
||||
assert_eq!(c.read_ahead(2).unwrap(), b'G');
|
||||
assert_eq!(c.tell(), Some(0)); // position unchanged
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_write_and_read_back() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
c.write(b'X').unwrap();
|
||||
c.write(b'Y').unwrap();
|
||||
let c2 = r.fw_cursor();
|
||||
assert_eq!(c2.read_next().unwrap(), b'X');
|
||||
assert_eq!(c2.read_next().unwrap(), b'Y');
|
||||
assert_eq!(c2.read_next().unwrap(), b'G');
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_rewind_and_reread() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // A → current = Some(0)
|
||||
c.read_next().unwrap(); // C → current = Some(1)
|
||||
c.read_next().unwrap(); // G → current = Some(2)
|
||||
c.rewind(1).unwrap(); // current = Some(1) → next read = index 2
|
||||
assert_eq!(c.read_next().unwrap(), b'G');
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_seek_absolute() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
c.seek(2, SeekMode::Absolute).unwrap();
|
||||
assert_eq!(c.read_next().unwrap(), b'T');
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_seek_relative_to_end() {
|
||||
let r = rope(b"ACGT");
|
||||
// seek(1, RelativeToEnd): current = len-1 = 3; peek() reads index 3 = T.
|
||||
let c = r.fw_cursor();
|
||||
c.seek(1, SeekMode::RelativeToEnd).unwrap();
|
||||
assert_eq!(c.peek().unwrap(), b'T');
|
||||
// seek(2, RelativeToEnd): current = len-2 = 2; read_next reads index 3 = T.
|
||||
let c2 = r.fw_cursor();
|
||||
c2.seek(2, SeekMode::RelativeToEnd).unwrap();
|
||||
assert_eq!(c2.read_next().unwrap(), b'T');
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_get_random_access() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.fw_cursor();
|
||||
assert_eq!(c.get(0), Some(b'A'));
|
||||
assert_eq!(c.get(3), Some(b'T'));
|
||||
assert_eq!(c.get(4), None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_crosses_block_boundary() {
|
||||
let r = rope2(b"AC", b"GT");
|
||||
let c = r.fw_cursor();
|
||||
let out: Vec<u8> = c.collect();
|
||||
assert_eq!(out, b"ACGT");
|
||||
}
|
||||
|
||||
// ── BackwardCursor ────────────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn backward_reads_all_bytes_in_reverse() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.bw_cursor();
|
||||
let out: Vec<u8> = c.collect();
|
||||
assert_eq!(out, b"TGCA");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_tell_tracks_position() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.bw_cursor();
|
||||
assert_eq!(c.tell(), None);
|
||||
c.read_next().unwrap(); // reads index 3
|
||||
assert_eq!(c.tell(), Some(3));
|
||||
c.read_next().unwrap(); // reads index 2
|
||||
assert_eq!(c.tell(), Some(2));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_iter_with_tell_and_seek_inside_loop() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.bw_cursor();
|
||||
let mut restart: usize = 0;
|
||||
for byte in c.iter() {
|
||||
if byte == b'G' {
|
||||
restart = c.tell().unwrap();
|
||||
}
|
||||
if byte == b'A' {
|
||||
// seek back to G and break
|
||||
c.seek(restart as isize, SeekMode::Absolute).ok();
|
||||
break;
|
||||
}
|
||||
}
|
||||
assert_eq!(c.tell(), Some(restart));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_rewind_moves_toward_end() {
|
||||
let r = rope(b"ACGT");
|
||||
let c = r.bw_cursor();
|
||||
c.read_next().unwrap(); // index 3 = T
|
||||
c.read_next().unwrap(); // index 2 = G
|
||||
c.rewind(1).unwrap(); // back to index 3
|
||||
assert_eq!(c.tell(), Some(3));
|
||||
assert_eq!(c.read_next().unwrap(), b'G'); // reads index 2
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_crosses_block_boundary() {
|
||||
let r = rope2(b"AC", b"GT");
|
||||
let c = r.bw_cursor();
|
||||
let out: Vec<u8> = c.collect();
|
||||
assert_eq!(out, b"TGCA");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_empty_rope_returns_error() {
|
||||
let r = Rope::new(None);
|
||||
let c = r.bw_cursor();
|
||||
assert!(c.read_next().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_empty_rope_returns_error() {
|
||||
let r = Rope::new(None);
|
||||
let c = r.fw_cursor();
|
||||
assert!(c.read_next().is_err());
|
||||
}
|
||||
|
||||
// ── offset / sub-cursor ───────────────────────────────────────────────────
|
||||
|
||||
#[test]
|
||||
fn forward_cursor_reads_from_offset() {
|
||||
// cursor() at current=Some(2) → new cursor reads from index 2
|
||||
let r = rope(b"ABCDE");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // A → current=Some(0)
|
||||
c.read_next().unwrap(); // B → current=Some(1)
|
||||
c.read_next().unwrap(); // C → current=Some(2)
|
||||
let sub = c.cursor(); // offset=2 (absolute_tell=2)
|
||||
assert_eq!(sub.read_next().unwrap(), b'C'); // reads index 2
|
||||
assert_eq!(sub.tell(), Some(0)); // relative: 2-2=0
|
||||
assert_eq!(sub.rope_tell(), 2);
|
||||
assert_eq!(sub.read_next().unwrap(), b'D');
|
||||
assert_eq!(sub.tell(), Some(1)); // relative: 3-2=1
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_cursor_get_uses_relative_index() {
|
||||
let r = rope(b"ABCDE");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // A → current=Some(0), absolute_tell=0
|
||||
let _sub = c.cursor(); // offset=0 — created to show cursor() compiles; not used further
|
||||
// From sub2 with offset=2: get(0)=C, get(2)=E
|
||||
let c2 = r.fw_cursor();
|
||||
c2.read_next().unwrap(); // at 0
|
||||
c2.read_next().unwrap(); // at 1
|
||||
c2.read_next().unwrap(); // at 2, absolute=2
|
||||
let sub2 = c2.cursor(); // offset=2
|
||||
assert_eq!(sub2.get(0), Some(b'C')); // local 0 = absolute 2
|
||||
assert_eq!(sub2.get(2), Some(b'E')); // local 2 = absolute 4
|
||||
assert_eq!(sub2.get(3), None); // local 3 = absolute 5, OOB
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_cursor_len_reflects_offset() {
|
||||
let r = rope(b"ABCDE"); // len=5
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap();
|
||||
c.read_next().unwrap();
|
||||
c.read_next().unwrap(); // absolute_tell=2
|
||||
let sub = c.cursor(); // offset=2
|
||||
assert_eq!(sub.len(), 3); // 5 - 2
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_reset_goes_back_to_start() {
|
||||
let r = rope(b"ABCDE");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // A
|
||||
c.read_next().unwrap(); // B
|
||||
c.reset();
|
||||
assert_eq!(c.tell(), None);
|
||||
assert_eq!(c.read_next().unwrap(), b'A'); // starts over
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_sub_cursor_write_and_reset() {
|
||||
// Write two bytes, discard them via reset(), write again.
|
||||
let r = rope(b"XXXXX");
|
||||
let c = r.fw_cursor();
|
||||
c.write(b'A').unwrap(); // absolute 0 → current=Some(1)
|
||||
c.write(b'B').unwrap(); // absolute 1 → current=Some(2)
|
||||
let seg = c.cursor(); // absolute_tell=2, offset=2
|
||||
seg.write(b'C').unwrap(); // absolute 2 → current=Some(3), tell=3-2=1
|
||||
seg.write(b'D').unwrap(); // absolute 3 → current=Some(4), tell=4-2=2
|
||||
assert_eq!(seg.tell(), Some(2)); // 2 bytes written into this segment
|
||||
seg.reset();
|
||||
assert_eq!(seg.tell(), None);
|
||||
seg.write(b'E').unwrap(); // absolute 2 again
|
||||
let all: Vec<u8> = r.fw_cursor().collect();
|
||||
assert_eq!(&all[..3], b"ABE");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn backward_cursor_stops_at_offset() {
|
||||
// BackwardCursor.cursor() creates a cursor with offset = absolute_tell.
|
||||
// offset = local position 0 (inclusive lower bound).
|
||||
// The cursor reads rope.len()-1 downto offset, then stops.
|
||||
let r = rope(b"ABCDE"); // 0=A 1=B 2=C 3=D 4=E
|
||||
let bw = r.bw_cursor();
|
||||
bw.read_next().unwrap(); // E=4, current=Some(4)
|
||||
bw.read_next().unwrap(); // D=3, current=Some(3), absolute_tell=3
|
||||
// sub: offset=3, reads from 4 down to 3 (inclusive), then stops.
|
||||
let sub = bw.cursor();
|
||||
assert_eq!(sub.read_next().unwrap(), b'E'); // index 4, tell=4-3=1
|
||||
assert_eq!(sub.read_next().unwrap(), b'D'); // index 3, tell=3-3=0 (local 0)
|
||||
assert!(sub.read_next().is_err()); // would go to 2 < offset=3
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn forward_absolute_tell_unchanged_by_offset() {
|
||||
let r = rope(b"ABCDE");
|
||||
let c = r.fw_cursor();
|
||||
c.read_next().unwrap(); // absolute=0
|
||||
let sub = c.cursor(); // offset=0
|
||||
sub.read_next().unwrap(); // reads index 0, absolute_tell=0
|
||||
sub.read_next().unwrap(); // reads index 1, absolute_tell=1
|
||||
assert_eq!(sub.tell(), Some(1));
|
||||
assert_eq!(sub.rope_tell(), 1);
|
||||
// sub2 with offset=1
|
||||
let sub2 = sub.cursor(); // offset=1
|
||||
sub2.read_next().unwrap(); // reads index 1, absolute=1
|
||||
assert_eq!(sub2.tell(), Some(0)); // relative: 1-1=0
|
||||
assert_eq!(sub2.rope_tell(), 1);
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
use crate::{Rope, RopeError};
|
||||
|
||||
use super::state::CursorState;
|
||||
|
||||
/// Controls how the `pos` argument of [`RopeCursor::seek`] is interpreted.
|
||||
#[derive(Clone, Copy)]
|
||||
pub enum SeekMode {
|
||||
/// `pos` is an absolute byte index from the start of the rope.
|
||||
Absolute,
|
||||
/// `pos` is relative to the current position.
|
||||
/// Positive = forward for [`ForwardCursor`](super::ForwardCursor), backward for [`BackwardCursor`](super::BackwardCursor).
|
||||
Relative,
|
||||
/// `pos` is counted back from the end: target = `len - pos`.
|
||||
RelativeToEnd,
|
||||
/// `pos` is a rope index relative to the start of the rope.
|
||||
Rope,
|
||||
}
|
||||
|
||||
/// Common interface for all rope cursors.
|
||||
///
|
||||
/// # Required methods
|
||||
///
|
||||
/// Implementors must provide [`rope`](RopeCursor::rope),
|
||||
/// [`state`](RopeCursor::state), [`read_next`](RopeCursor::read_next) and
|
||||
/// [`seek`](RopeCursor::seek). Everything else has a default implementation.
|
||||
///
|
||||
/// The direction of `read_next` and the sign convention for
|
||||
/// [`SeekMode::Relative`] differ between [`ForwardCursor`](super::ForwardCursor)
|
||||
/// and [`BackwardCursor`](super::BackwardCursor); all other methods are identical.
|
||||
pub trait RopeCursor<'a> {
|
||||
/// The rope this cursor is bound to.
|
||||
fn rope(&self) -> &'a Rope;
|
||||
/// Internal cache state — implementation detail exposed for default methods.
|
||||
fn state(&self) -> &CursorState<'a>;
|
||||
|
||||
/// Read the next byte in cursor direction and advance the position.
|
||||
/// Returns `Err` at the exhausted end.
|
||||
fn read_next(&self) -> Result<u8, RopeError>;
|
||||
|
||||
/// Move the cursor to a new position.
|
||||
///
|
||||
/// `pos` is interpreted according to `mode`:
|
||||
/// - [`Absolute`](SeekMode::Absolute): local coordinate (`pos + offset` in the rope).
|
||||
/// - [`Rope`](SeekMode::Rope): raw rope index, ignores the offset. Pass a value
|
||||
/// from [`rope_tell`](RopeCursor::rope_tell) to restore a saved position.
|
||||
/// - [`Relative`](SeekMode::Relative): delta from the current position.
|
||||
/// For [`ForwardCursor`](super::ForwardCursor), positive advances toward the end;
|
||||
/// for [`BackwardCursor`](super::BackwardCursor), positive retreats toward the start.
|
||||
/// - [`RelativeToEnd`](SeekMode::RelativeToEnd): `rope.len() - pos`.
|
||||
///
|
||||
/// Returns the new position as a **rope index** (same value as
|
||||
/// [`rope_tell`](RopeCursor::rope_tell) would return immediately after).
|
||||
fn seek(&self, pos: isize, mode: SeekMode) -> Result<usize, RopeError>;
|
||||
|
||||
// ── default methods ───────────────────────────────────────────────────────
|
||||
|
||||
/// Read the byte at **local** index `i` (relative to the cursor's offset)
|
||||
/// without moving the position.
|
||||
fn get(&self, i: usize) -> Option<u8> {
|
||||
self.state().get(self.rope(), i + self.state().offset.get())
|
||||
}
|
||||
|
||||
/// Write `value` at **local** index `i` without moving the position.
|
||||
fn set(&self, i: usize, value: u8) -> Result<(), RopeError> {
|
||||
self.state()
|
||||
.set(self.rope(), i + self.state().offset.get(), value)
|
||||
}
|
||||
|
||||
/// Current position relative to the cursor's offset, or `None` if the
|
||||
/// cursor has not moved yet.
|
||||
fn tell(&self) -> Option<usize> {
|
||||
let abs = self.state().current.get()?;
|
||||
Some(abs.saturating_sub(self.state().offset.get()))
|
||||
}
|
||||
|
||||
/// Current position as an absolute rope index.
|
||||
///
|
||||
/// Unlike [`tell`](RopeCursor::tell), this method **always** returns a
|
||||
/// value: if the cursor has not moved yet, it returns the cursor's offset
|
||||
/// (the rope index of local position 0).
|
||||
///
|
||||
/// Use the returned value with [`SeekMode::Rope`] to restore a position,
|
||||
/// or as a truncation point after a write pass.
|
||||
fn rope_tell(&self) -> usize {
|
||||
self.state()
|
||||
.current
|
||||
.get()
|
||||
.unwrap_or(self.state().offset.get())
|
||||
}
|
||||
|
||||
/// Number of bytes visible through this cursor (`rope.len() - offset`).
|
||||
fn len(&self) -> usize {
|
||||
self.rope().len().saturating_sub(self.state().offset.get())
|
||||
}
|
||||
|
||||
/// Reset the cursor to its initial state (positioned before the first
|
||||
/// byte of its local view). Equivalent to `seek(0, Absolute)` on a
|
||||
/// fresh cursor, but works even when `current` is `None`.
|
||||
fn reset(&self) {
|
||||
self.state().current.set(None);
|
||||
}
|
||||
|
||||
/// Read the byte at the current position without advancing.
|
||||
fn peek(&self) -> Option<u8> {
|
||||
self.state().get(self.rope(), self.state().current.get()?)
|
||||
}
|
||||
|
||||
/// Write `value` at the current position without advancing.
|
||||
fn poke(&self, value: u8) -> Result<(), RopeError> {
|
||||
let pos = self.state().current.get().ok_or(RopeError::CurrentNotSet)?;
|
||||
self.state().set(self.rope(), pos, value)
|
||||
}
|
||||
|
||||
/// Move backward by `go_back_of` steps (toward lower indices for
|
||||
/// [`ForwardCursor`](super::ForwardCursor), toward higher indices for
|
||||
/// [`BackwardCursor`](super::BackwardCursor)).
|
||||
fn rewind(&self, go_back_of: usize) -> Result<(), RopeError> {
|
||||
self.seek(-(go_back_of as isize), SeekMode::Relative)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Move forward by `ahead` steps (opposite of [`rewind`](RopeCursor::rewind)).
|
||||
fn forward(&self, ahead: usize) -> Result<(), RopeError> {
|
||||
self.seek(ahead as isize, SeekMode::Relative)?;
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
@@ -1,878 +0,0 @@
|
||||
use crossbeam_channel::{Receiver, Select, Sender, bounded};
|
||||
use std::error::Error;
|
||||
use std::fmt;
|
||||
use std::marker::PhantomData;
|
||||
use std::sync::Arc;
|
||||
use std::thread;
|
||||
|
||||
/// Error type for pipeline operations.
|
||||
#[derive(Debug)]
|
||||
pub enum PipelineError {
|
||||
/// A stage received a `PipelineData` variant it did not expect.
|
||||
TypeMismatch,
|
||||
/// The step kind is not compatible with the data type.
|
||||
StepKindMismatch(&'static str),
|
||||
/// The source has no more data to produce.
|
||||
EndOfStream,
|
||||
/// An error occurred inside a stage (e.g., I/O, parsing, custom logic).
|
||||
StepError(Box<dyn Error + Send + Sync>),
|
||||
}
|
||||
|
||||
impl fmt::Display for PipelineError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
PipelineError::TypeMismatch => write!(f, "data type mismatch in pipeline stage"),
|
||||
PipelineError::StepKindMismatch(s) => write!(f, "step kind mismatch: {}", s),
|
||||
PipelineError::EndOfStream => write!(f, "end of input stream"),
|
||||
PipelineError::StepError(e) => write!(f, "stage error: {}", e),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Error for PipelineError {
|
||||
fn source(&self) -> Option<&(dyn Error + 'static)> {
|
||||
match self {
|
||||
PipelineError::StepError(e) => Some(e.as_ref()),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── Function types ────────────────────────────────────────────────────────────
|
||||
|
||||
/// Fonction source : appelée répétitivement, retourne le prochain item ou EndOfStream.
|
||||
/// `FnMut` car elle maintient un état interne (position dans l'itérateur).
|
||||
pub type SourceFn<D> = Box<dyn FnMut() -> Result<D, PipelineError> + Send>;
|
||||
|
||||
/// Fonction sink : consomme un item final, peut échouer (erreur d'I/O, etc.).
|
||||
pub type SinkFn<D> = Box<dyn Fn(D) -> Result<(), PipelineError> + Send>;
|
||||
|
||||
/// Fonction de transformation partagée entre workers via Arc.
|
||||
pub type SharedFn<D> = Arc<dyn Fn(D) -> Result<D, PipelineError> + Send + Sync>;
|
||||
|
||||
/// Fonction de transformation 1→N (flat map) partagée entre workers via Arc.
|
||||
///
|
||||
/// La fonction reçoit l'item d'entrée, un canal `push` pour envoyer chaque item
|
||||
/// produit, et un canal `delta` pour signaler au scheduler combien d'items
|
||||
/// supplémentaires sont entrés dans le pipeline (N-1 si N items produits).
|
||||
/// Elle doit appeler `delta.send(N - 1)` **après** avoir poussé tous les items.
|
||||
pub type SharedFlatFn<D> =
|
||||
Arc<dyn Fn(D, &Sender<Result<D, PipelineError>>, &Sender<isize>) + Send + Sync>;
|
||||
|
||||
// ── Stage enum ────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Une étape du pipeline : transform classique (1→1) ou flat transform (1→N).
|
||||
pub enum Stage<D> {
|
||||
Transform(SharedFn<D>),
|
||||
Flat(SharedFlatFn<D>),
|
||||
}
|
||||
|
||||
impl<D> Clone for Stage<D> {
|
||||
fn clone(&self) -> Self {
|
||||
match self {
|
||||
Stage::Transform(f) => Stage::Transform(Arc::clone(f)),
|
||||
Stage::Flat(f) => Stage::Flat(Arc::clone(f)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── Worker task ───────────────────────────────────────────────────────────────
|
||||
|
||||
enum WorkerTask<D> {
|
||||
Transform(D, usize),
|
||||
Flat(D, usize),
|
||||
}
|
||||
|
||||
// ── Thread runners ────────────────────────────────────────────────────────────
|
||||
|
||||
fn source_runner<DATA>(
|
||||
mut source: SourceFn<DATA>,
|
||||
capacity: usize,
|
||||
) -> (
|
||||
Receiver<Result<DATA, PipelineError>>,
|
||||
thread::JoinHandle<()>,
|
||||
)
|
||||
where
|
||||
DATA: Send + Sync + 'static,
|
||||
{
|
||||
let (tx, rx) = bounded(capacity);
|
||||
let handle = thread::spawn(move || {
|
||||
loop {
|
||||
match source() {
|
||||
Ok(data) => {
|
||||
if tx.send(Ok(data)).is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Err(PipelineError::EndOfStream) => break,
|
||||
Err(e) => {
|
||||
eprintln!("Source error: {:?}", e);
|
||||
let _ = tx.send(Err(e));
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
(rx, handle)
|
||||
}
|
||||
|
||||
/// Lance un thread worker du pool.
|
||||
///
|
||||
/// Gère deux types de tâches :
|
||||
/// - `Transform` : applique `f(data)` et envoie le résultat dans `result_tx`.
|
||||
/// - `Flat` : appelle `f(data, &push_tx, &delta_tx)` ; la fonction elle-même
|
||||
/// pousse ses items dans `push_tx` et envoie `N-1` dans `delta_tx`.
|
||||
fn transform_runner<DATA>(
|
||||
task_rx: Receiver<WorkerTask<DATA>>,
|
||||
stages: Vec<Stage<DATA>>,
|
||||
stage_txs: Vec<Sender<Result<DATA, PipelineError>>>,
|
||||
flat_delta_tx: Sender<isize>,
|
||||
) -> thread::JoinHandle<()>
|
||||
where
|
||||
DATA: Send + Sync + 'static,
|
||||
{
|
||||
thread::spawn(move || {
|
||||
while let Ok(task) = task_rx.recv() {
|
||||
match task {
|
||||
WorkerTask::Transform(data, idx) => {
|
||||
if let Stage::Transform(f) = &stages[idx] {
|
||||
let _ = stage_txs[idx].send(f(data));
|
||||
}
|
||||
}
|
||||
WorkerTask::Flat(data, idx) => {
|
||||
if let Stage::Flat(f) = &stages[idx] {
|
||||
f(data, &stage_txs[idx], &flat_delta_tx);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/// Lance le thread sink.
|
||||
fn sink_runner<DATA>(
|
||||
sink: SinkFn<DATA>,
|
||||
capacity: usize,
|
||||
) -> (
|
||||
Sender<DATA>,
|
||||
Receiver<PipelineError>,
|
||||
thread::JoinHandle<()>,
|
||||
)
|
||||
where
|
||||
DATA: Send + Sync + 'static,
|
||||
{
|
||||
let (data_tx, data_rx) = bounded(capacity);
|
||||
let (err_tx, err_rx) = bounded(capacity);
|
||||
let handle = thread::spawn(move || {
|
||||
for data in data_rx {
|
||||
if let Err(e) = sink(data) {
|
||||
let _ = err_tx.send(e);
|
||||
break;
|
||||
}
|
||||
}
|
||||
});
|
||||
(data_tx, err_rx, handle)
|
||||
}
|
||||
|
||||
// ── Pipeline ──────────────────────────────────────────────────────────────────
|
||||
|
||||
pub struct Pipeline<DATA> {
|
||||
source: SourceFn<DATA>,
|
||||
stages: Vec<Stage<DATA>>,
|
||||
sink: SinkFn<DATA>,
|
||||
}
|
||||
|
||||
impl<DATA> Pipeline<DATA> {
|
||||
pub fn new(
|
||||
source: SourceFn<DATA>,
|
||||
stages: Vec<Stage<DATA>>,
|
||||
sink: SinkFn<DATA>,
|
||||
) -> Self {
|
||||
Self { source, stages, sink }
|
||||
}
|
||||
}
|
||||
|
||||
// ── WorkerPool ────────────────────────────────────────────────────────────────
|
||||
|
||||
pub struct WorkerPool<DATA> {
|
||||
pipeline: Pipeline<DATA>,
|
||||
handles: Vec<std::thread::JoinHandle<()>>,
|
||||
n_workers: usize,
|
||||
capacity: usize,
|
||||
}
|
||||
|
||||
impl<DATA> WorkerPool<DATA>
|
||||
where
|
||||
DATA: Send + Sync + 'static,
|
||||
{
|
||||
pub fn new(pipeline: Pipeline<DATA>, n_workers: usize, capacity: usize) -> Self {
|
||||
Self {
|
||||
pipeline,
|
||||
handles: Vec::new(),
|
||||
n_workers,
|
||||
capacity,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn run(mut self) {
|
||||
let n = self.pipeline.stages.len();
|
||||
|
||||
// ── Canaux inter-stages ────────────────────────────────────────────
|
||||
// stage_txs[i] / stage_rxs[i] : sortie du stage i
|
||||
let mut stage_txs: Vec<Sender<Result<DATA, PipelineError>>> = Vec::new();
|
||||
let mut stage_rxs: Vec<Receiver<Result<DATA, PipelineError>>> = Vec::new();
|
||||
for _ in 0..n {
|
||||
let (tx, rx) = bounded(self.capacity);
|
||||
stage_txs.push(tx);
|
||||
stage_rxs.push(rx);
|
||||
}
|
||||
|
||||
// ── Source thread ──────────────────────────────────────────────────
|
||||
let (source_rx, src_handle) = source_runner(self.pipeline.source, self.capacity);
|
||||
self.handles.push(src_handle);
|
||||
|
||||
let stages = self.pipeline.stages;
|
||||
|
||||
// ── Canal delta pour les flat stages ───────────────────────────────
|
||||
// Chaque flat worker envoie `N-1` ici après avoir poussé N items.
|
||||
// Le scheduler ajuste `in_flight` en conséquence.
|
||||
let (flat_delta_tx, flat_delta_rx) = bounded::<isize>(self.capacity);
|
||||
|
||||
// ── Worker pool ────────────────────────────────────────────────────
|
||||
let (worker_tx, worker_rx): (Sender<WorkerTask<DATA>>, Receiver<WorkerTask<DATA>>) =
|
||||
bounded(self.capacity);
|
||||
|
||||
for _ in 0..self.n_workers {
|
||||
self.handles.push(transform_runner(
|
||||
worker_rx.clone(),
|
||||
stages.iter().map(Stage::clone).collect(),
|
||||
stage_txs.clone(),
|
||||
flat_delta_tx.clone(),
|
||||
));
|
||||
}
|
||||
// Le scheduler ne tient plus flat_delta_tx : les workers le détiennent.
|
||||
// On le drop ici pour que le canal se ferme quand les workers terminent.
|
||||
drop(flat_delta_tx);
|
||||
|
||||
// ── Sink thread ────────────────────────────────────────────────────
|
||||
let (sink_tx, sink_err_rx, sink_handle) = sink_runner(self.pipeline.sink, self.capacity);
|
||||
self.handles.push(sink_handle);
|
||||
|
||||
// ── Boucle principale ──────────────────────────────────────────────
|
||||
//
|
||||
// `in_flight` (isize) = nb d'items qui doivent encore atteindre le sink.
|
||||
// Peut temporairement être négatif si un flat worker a poussé ses items
|
||||
// avant que le scheduler ait reçu le delta correspondant.
|
||||
//
|
||||
// `flat_workers_active` = nb de flat workers en cours d'exécution.
|
||||
// Empêche la terminaison prématurée quand in_flight vaut 0 mais qu'un
|
||||
// flat worker n'a pas encore envoyé son delta.
|
||||
//
|
||||
// Priorités du Select biaisé (index le plus bas = priorité la plus haute) :
|
||||
// 0 → sink_err_rx (arrêt immédiat sur erreur sink)
|
||||
// 1 → flat_delta_rx (mettre à jour in_flight avant de dispatcher)
|
||||
// 2..=n+1 → stage_rxs[n-1..0] (vider le pipeline en priorité)
|
||||
// n+2 → source_rx (dernier recours : nouvelles données)
|
||||
//
|
||||
// Quand k = 0 : erreur du sink
|
||||
// Quand k = 1 : delta d'un flat worker
|
||||
// Quand 2 ≤ k ≤ n+1 : résultat du stage n+1-k
|
||||
// Quand k = n+2 : item source
|
||||
//
|
||||
// Terminaison : source tarie ET in_flight == 0 ET aucun flat worker actif.
|
||||
{
|
||||
let mut source_done = false;
|
||||
let mut in_flight: isize = 0;
|
||||
let mut flat_workers_active: usize = 0;
|
||||
|
||||
loop {
|
||||
if source_done && in_flight == 0 && flat_workers_active == 0 {
|
||||
break;
|
||||
}
|
||||
|
||||
let mut sel = Select::new_biased();
|
||||
sel.recv(&sink_err_rx); // index 0
|
||||
sel.recv(&flat_delta_rx); // index 1
|
||||
for rx in stage_rxs.iter().rev() {
|
||||
sel.recv(rx); // indices 2..=n+1
|
||||
}
|
||||
let src_idx = if !source_done {
|
||||
Some(sel.recv(&source_rx)) // index n+2
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let oper = sel.select();
|
||||
let k = oper.index();
|
||||
|
||||
if k == 0 {
|
||||
// ── Erreur du sink ────────────────────────────────────
|
||||
match oper.recv(&sink_err_rx) {
|
||||
Ok(e) => { eprintln!("Sink error: {:?}", e); break; }
|
||||
Err(_) => break,
|
||||
}
|
||||
} else if k == 1 {
|
||||
// ── Delta d'un flat worker ────────────────────────────
|
||||
// delta = N - 1 (N items poussés, 1 item consommé)
|
||||
match oper.recv(&flat_delta_rx) {
|
||||
Ok(delta) => {
|
||||
in_flight += delta;
|
||||
flat_workers_active -= 1;
|
||||
}
|
||||
Err(_) => {}
|
||||
}
|
||||
} else if src_idx == Some(k) {
|
||||
// ── Nouvel item depuis la source ──────────────────────
|
||||
match oper.recv(&source_rx) {
|
||||
Ok(Ok(data)) => {
|
||||
if n == 0 {
|
||||
let _ = sink_tx.send(data);
|
||||
} else {
|
||||
in_flight += 1;
|
||||
dispatch(
|
||||
data, 0,
|
||||
&stages, &worker_tx,
|
||||
&mut flat_workers_active,
|
||||
);
|
||||
}
|
||||
}
|
||||
Ok(Err(e)) => eprintln!("Source error: {:?}", e),
|
||||
Err(_) => source_done = true,
|
||||
}
|
||||
} else {
|
||||
// ── Résultat d'un stage intermédiaire ─────────────────
|
||||
// k ∈ [2, n+1] → stage = n+1 - k
|
||||
let stage = n + 1 - k;
|
||||
match oper.recv(&stage_rxs[stage]) {
|
||||
Ok(Ok(data)) => {
|
||||
if stage == n - 1 {
|
||||
in_flight -= 1;
|
||||
let _ = sink_tx.send(data);
|
||||
} else {
|
||||
dispatch(
|
||||
data, stage + 1,
|
||||
&stages, &worker_tx,
|
||||
&mut flat_workers_active,
|
||||
);
|
||||
}
|
||||
}
|
||||
Ok(Err(e)) => eprintln!("Stage {} error: {:?}", stage, e),
|
||||
Err(_) => break,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
drop(worker_tx);
|
||||
drop(sink_tx);
|
||||
|
||||
for h in self.handles {
|
||||
let _ = h.join();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── Pipe ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Typed, composable iterator transformer.
|
||||
///
|
||||
/// A `Pipe<D, In, Out>` is a pure description of pipeline stages — no threads,
|
||||
/// no channels, no scheduler. Call `.apply(iter, n_workers, capacity)` to start
|
||||
/// execution and get back a `PipeIter<Out>`.
|
||||
///
|
||||
/// Compose two pipes with `.then()`: the resulting `Pipe` holds the concatenated
|
||||
/// stage list, so a single scheduler is created when `.apply()` is eventually called.
|
||||
pub struct Pipe<D, In, Out> {
|
||||
stages: Vec<Stage<D>>,
|
||||
wrap: Arc<dyn Fn(In) -> D + Send + Sync>,
|
||||
unwrap: Arc<dyn Fn(D) -> Out + Send + Sync>,
|
||||
_phantom: PhantomData<(In, Out)>,
|
||||
}
|
||||
|
||||
impl<D, In, Out> Pipe<D, In, Out> {
|
||||
/// Build a `Pipe` from stages and wrap/unwrap converters.
|
||||
/// Prefer the `make_pipe!` macro.
|
||||
pub fn new(
|
||||
stages: Vec<Stage<D>>,
|
||||
wrap: Arc<dyn Fn(In) -> D + Send + Sync>,
|
||||
unwrap: Arc<dyn Fn(D) -> Out + Send + Sync>,
|
||||
) -> Self {
|
||||
Self { stages, wrap, unwrap, _phantom: PhantomData }
|
||||
}
|
||||
|
||||
/// Concatenate stages from two pipes into one.
|
||||
///
|
||||
/// Requires `Out` of `self` == `In` of `other`. The single scheduler
|
||||
/// created at `.apply()` time sees the full combined stage list.
|
||||
pub fn then<Next>(self, other: Pipe<D, Out, Next>) -> Pipe<D, In, Next> {
|
||||
Pipe {
|
||||
stages: self.stages.into_iter().chain(other.stages).collect(),
|
||||
wrap: self.wrap,
|
||||
unwrap: other.unwrap,
|
||||
_phantom: PhantomData,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<D, In, Out> Pipe<D, In, Out>
|
||||
where
|
||||
D: Send + Sync + 'static,
|
||||
In: Send + 'static,
|
||||
Out: Send + 'static,
|
||||
{
|
||||
/// Run the pipeline in a background thread; returns an iterator over the output.
|
||||
pub fn apply(
|
||||
self,
|
||||
input: impl Iterator<Item = In> + Send + 'static,
|
||||
n_workers: usize,
|
||||
capacity: usize,
|
||||
) -> PipeIter<Out> {
|
||||
let wrap = Arc::clone(&self.wrap);
|
||||
let unwrap = Arc::clone(&self.unwrap);
|
||||
|
||||
let mut iter = input;
|
||||
let source: SourceFn<D> = Box::new(move || match iter.next() {
|
||||
Some(x) => Ok(wrap(x)),
|
||||
None => Err(PipelineError::EndOfStream),
|
||||
});
|
||||
|
||||
let (out_tx, out_rx) = bounded::<Out>(capacity);
|
||||
let sink: SinkFn<D> = Box::new(move |data: D| {
|
||||
out_tx.send(unwrap(data)).map_err(|_| {
|
||||
PipelineError::StepError(Box::new(std::io::Error::new(
|
||||
std::io::ErrorKind::BrokenPipe,
|
||||
"output channel closed",
|
||||
)))
|
||||
})
|
||||
});
|
||||
|
||||
let pipeline = Pipeline::new(source, self.stages, sink);
|
||||
let handle = thread::spawn(move || {
|
||||
WorkerPool::new(pipeline, n_workers, capacity).run();
|
||||
});
|
||||
|
||||
PipeIter { rx: out_rx, handle: Some(handle) }
|
||||
}
|
||||
}
|
||||
|
||||
// ── PipeIter ──────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Iterator over the output of `Pipe::apply()`.
|
||||
pub struct PipeIter<Out> {
|
||||
rx: Receiver<Out>,
|
||||
handle: Option<thread::JoinHandle<()>>,
|
||||
}
|
||||
|
||||
impl<Out> Iterator for PipeIter<Out> {
|
||||
type Item = Out;
|
||||
|
||||
fn next(&mut self) -> Option<Out> {
|
||||
self.rx.recv().ok()
|
||||
}
|
||||
}
|
||||
|
||||
impl<Out> Drop for PipeIter<Out> {
|
||||
fn drop(&mut self) {
|
||||
// Drain buffered items so the scheduler can unblock if the channel is full.
|
||||
while self.rx.try_recv().is_ok() {}
|
||||
if let Some(h) = self.handle.take() {
|
||||
let _ = h.join();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Envoie `data` au stage `stage_idx`.
|
||||
/// Pour un `Transform`, empile une `WorkerTask::Transform`.
|
||||
/// Pour un `Flat`, incrémente `flat_workers_active` et empile une `WorkerTask::Flat`.
|
||||
#[inline]
|
||||
fn dispatch<DATA>(
|
||||
data: DATA,
|
||||
stage_idx: usize,
|
||||
stages: &[Stage<DATA>],
|
||||
worker_tx: &Sender<WorkerTask<DATA>>,
|
||||
flat_workers_active: &mut usize,
|
||||
) {
|
||||
match &stages[stage_idx] {
|
||||
Stage::Transform(_) => {
|
||||
let _ = worker_tx.send(WorkerTask::Transform(data, stage_idx));
|
||||
}
|
||||
Stage::Flat(_) => {
|
||||
*flat_workers_active += 1;
|
||||
let _ = worker_tx.send(WorkerTask::Flat(data, stage_idx));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── Macros ────────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Creates a `SourceFn` from an iterator of plain values.
|
||||
#[macro_export]
|
||||
macro_rules! make_source {
|
||||
($enum:ident, $iterator:expr, $output:ident) => {{
|
||||
let mut iter = $iterator.into_iter();
|
||||
Box::new(
|
||||
move || -> ::std::result::Result<$enum, $crate::PipelineError> {
|
||||
match iter.next() {
|
||||
Some(x) => Ok($enum::$output(x)),
|
||||
None => Err($crate::PipelineError::EndOfStream),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<dyn FnMut() -> ::std::result::Result<$enum, $crate::PipelineError> + Send>
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `SourceFn` from an iterator of `Result<T, E>`.
|
||||
#[macro_export]
|
||||
macro_rules! make_source_fallible {
|
||||
($enum:ident, $iterator:expr, $output:ident) => {{
|
||||
let mut iter = $iterator.into_iter();
|
||||
Box::new(
|
||||
move || -> ::std::result::Result<$enum, $crate::PipelineError> {
|
||||
match iter.next() {
|
||||
Some(Ok(x)) => Ok($enum::$output(x)),
|
||||
Some(Err(e)) => Err($crate::PipelineError::StepError(Box::new(e))),
|
||||
None => Err($crate::PipelineError::EndOfStream),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<dyn FnMut() -> ::std::result::Result<$enum, $crate::PipelineError> + Send>
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `Stage::Transform` from a pure (non-fallible) function `Fn(T) -> U`.
|
||||
#[macro_export]
|
||||
macro_rules! make_transform {
|
||||
($enum:ident, $func:tt, $input:ident, $output:ident) => {{
|
||||
let __f = $func;
|
||||
$crate::Stage::Transform(
|
||||
::std::sync::Arc::from(Box::new(
|
||||
move |data: $enum| -> ::std::result::Result<$enum, $crate::PipelineError> {
|
||||
match data {
|
||||
$enum::$input(x) => Ok($enum::$output(__f(x))),
|
||||
_ => Err($crate::PipelineError::TypeMismatch),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<
|
||||
dyn Fn($enum) -> ::std::result::Result<$enum, $crate::PipelineError>
|
||||
+ Send
|
||||
+ Sync,
|
||||
>)
|
||||
)
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `Stage::Transform` from a fallible function `Fn(T) -> Result<U, E>`.
|
||||
#[macro_export]
|
||||
macro_rules! make_transform_fallible {
|
||||
($enum:ident, $func:tt, $input:ident, $output:ident) => {{
|
||||
let __f = $func;
|
||||
$crate::Stage::Transform(
|
||||
::std::sync::Arc::from(Box::new(
|
||||
move |data: $enum| -> ::std::result::Result<$enum, $crate::PipelineError> {
|
||||
match data {
|
||||
$enum::$input(inner) => {
|
||||
let result = __f(inner)
|
||||
.map_err(|e| $crate::PipelineError::StepError(Box::new(e)))?;
|
||||
Ok($enum::$output(result))
|
||||
}
|
||||
_ => Err($crate::PipelineError::TypeMismatch),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<
|
||||
dyn Fn($enum) -> ::std::result::Result<$enum, $crate::PipelineError>
|
||||
+ Send
|
||||
+ Sync,
|
||||
>)
|
||||
)
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `Stage::Flat` from a function `Fn(T) -> impl IntoIterator<Item = U>`.
|
||||
///
|
||||
/// Pour chaque item produit par l'itérateur, il est poussé individuellement dans
|
||||
/// le canal de sortie, permettant au scheduler de dispatcher les items en parallèle
|
||||
/// dès qu'un worker est disponible.
|
||||
#[macro_export]
|
||||
macro_rules! make_flat_transform {
|
||||
($enum:ident, $func:tt, $input:ident, $output:ident) => {{
|
||||
let __f = $func;
|
||||
$crate::Stage::Flat(
|
||||
::std::sync::Arc::new(
|
||||
move |data: $enum,
|
||||
push: &$crate::PipelineSender<
|
||||
::std::result::Result<$enum, $crate::PipelineError>,
|
||||
>,
|
||||
delta: &$crate::PipelineSender<isize>| {
|
||||
match data {
|
||||
$enum::$input(inner) => {
|
||||
let mut count: isize = 0;
|
||||
for item in __f(inner) {
|
||||
push.send(Ok($enum::$output(item))).ok();
|
||||
count += 1;
|
||||
}
|
||||
delta.send(count - 1).ok();
|
||||
}
|
||||
_ => {
|
||||
push.send(Err($crate::PipelineError::TypeMismatch)).ok();
|
||||
delta.send(0).ok();
|
||||
}
|
||||
}
|
||||
},
|
||||
) as $crate::SharedFlatFn<$enum>
|
||||
)
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `Stage::Flat` from a fallible function
|
||||
/// `Fn(T) -> Result<impl IntoIterator<Item = U>, E>`.
|
||||
///
|
||||
/// Si la fonction retourne `Err`, une erreur est poussée dans le canal et aucun
|
||||
/// item normal n'est produit.
|
||||
#[macro_export]
|
||||
macro_rules! make_flat_transform_fallible {
|
||||
($enum:ident, $func:tt, $input:ident, $output:ident) => {{
|
||||
let __f = $func;
|
||||
$crate::Stage::Flat(
|
||||
::std::sync::Arc::new(
|
||||
move |data: $enum,
|
||||
push: &$crate::PipelineSender<
|
||||
::std::result::Result<$enum, $crate::PipelineError>,
|
||||
>,
|
||||
delta: &$crate::PipelineSender<isize>| {
|
||||
match data {
|
||||
$enum::$input(inner) => match __f(inner) {
|
||||
Ok(iter) => {
|
||||
let mut count: isize = 0;
|
||||
for item in iter {
|
||||
push.send(Ok($enum::$output(item))).ok();
|
||||
count += 1;
|
||||
}
|
||||
delta.send(count - 1).ok();
|
||||
}
|
||||
Err(e) => {
|
||||
push.send(Err($crate::PipelineError::StepError(Box::new(e))))
|
||||
.ok();
|
||||
delta.send(0).ok();
|
||||
}
|
||||
},
|
||||
_ => {
|
||||
push.send(Err($crate::PipelineError::TypeMismatch)).ok();
|
||||
delta.send(0).ok();
|
||||
}
|
||||
}
|
||||
},
|
||||
) as $crate::SharedFlatFn<$enum>
|
||||
)
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `SinkFn` from a function that consumes a concrete value and returns `()`.
|
||||
#[macro_export]
|
||||
macro_rules! make_sink {
|
||||
($enum:ident, $func:tt, $input:ident) => {{
|
||||
let __f = $func;
|
||||
Box::new(
|
||||
move |data: $enum| -> ::std::result::Result<(), $crate::PipelineError> {
|
||||
match data {
|
||||
$enum::$input(x) => {
|
||||
__f(x);
|
||||
Ok(())
|
||||
}
|
||||
_ => Err($crate::PipelineError::TypeMismatch),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<dyn Fn($enum) -> ::std::result::Result<(), $crate::PipelineError> + Send>
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `SinkFn` from a fallible function that returns `Result<(), E>`.
|
||||
#[macro_export]
|
||||
macro_rules! make_sink_fallible {
|
||||
($enum:ident, $func:tt, $input:ident) => {{
|
||||
let __f = $func;
|
||||
Box::new(
|
||||
move |data: $enum| -> ::std::result::Result<(), $crate::PipelineError> {
|
||||
match data {
|
||||
$enum::$input(inner) => {
|
||||
__f(inner).map_err(|e| $crate::PipelineError::StepError(Box::new(e)))
|
||||
}
|
||||
_ => Err($crate::PipelineError::TypeMismatch),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<dyn Fn($enum) -> ::std::result::Result<(), $crate::PipelineError> + Send>
|
||||
}};
|
||||
}
|
||||
|
||||
/// Construit un `Pipeline` à partir d'une source, d'une liste de stages et d'un sink.
|
||||
///
|
||||
/// Syntaxe :
|
||||
/// ```ignore
|
||||
/// make_pipeline! {
|
||||
/// MyData,
|
||||
/// source my_iter => Variant, // source non-fallible
|
||||
/// source? my_iter => Variant, // source fallible (Result<T, E>)
|
||||
/// | func: In => Out, // transform 1→1 non-fallible
|
||||
/// |? func: In => Out, // transform 1→1 fallible
|
||||
/// || func: In => Out, // flat transform 1→N non-fallible
|
||||
/// ||? func: In => Out, // flat transform 1→N fallible
|
||||
/// sink my_func @ Variant, // sink non-fallible
|
||||
/// sink? my_func @ Variant, // sink fallible
|
||||
/// }
|
||||
/// ```
|
||||
#[macro_export]
|
||||
macro_rules! make_pipeline {
|
||||
// ── Points d'entrée ──────────────────────────────────────────────────
|
||||
|
||||
($enum:ident, source $src:expr => $src_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum,
|
||||
{ $crate::make_source!($enum, $src, $src_out) },
|
||||
[],
|
||||
$($rest)*)
|
||||
};
|
||||
($enum:ident, source? $src:expr => $src_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum,
|
||||
{ $crate::make_source_fallible!($enum, $src, $src_out) },
|
||||
[],
|
||||
$($rest)*)
|
||||
};
|
||||
|
||||
// ── Accumulation des stages ──────────────────────────────────────────
|
||||
|
||||
// transform 1→1 non-fallible
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
| $tf:tt : $t_in:ident => $t_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum, $source,
|
||||
[$($acc)* $crate::make_transform!($enum, $tf, $t_in, $t_out),],
|
||||
$($rest)*)
|
||||
};
|
||||
|
||||
// transform 1→1 fallible
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
|? $tf:tt : $t_in:ident => $t_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum, $source,
|
||||
[$($acc)* $crate::make_transform_fallible!($enum, $tf, $t_in, $t_out),],
|
||||
$($rest)*)
|
||||
};
|
||||
|
||||
// flat transform 1→N non-fallible
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
|| $tf:tt : $t_in:ident => $t_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum, $source,
|
||||
[$($acc)* $crate::make_flat_transform!($enum, $tf, $t_in, $t_out),],
|
||||
$($rest)*)
|
||||
};
|
||||
|
||||
// flat transform 1→N fallible
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
||? $tf:tt : $t_in:ident => $t_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum, $source,
|
||||
[$($acc)* $crate::make_flat_transform_fallible!($enum, $tf, $t_in, $t_out),],
|
||||
$($rest)*)
|
||||
};
|
||||
|
||||
// ── Terminaison : sink ───────────────────────────────────────────────
|
||||
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
sink $sink_fn:tt @ $sink_in:ident $(,)?) => {
|
||||
$crate::Pipeline::new(
|
||||
$source,
|
||||
vec![$($acc)*],
|
||||
$crate::make_sink!($enum, $sink_fn, $sink_in),
|
||||
)
|
||||
};
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
sink? $sink_fn:tt @ $sink_in:ident $(,)?) => {
|
||||
$crate::Pipeline::new(
|
||||
$source,
|
||||
vec![$($acc)*],
|
||||
$crate::make_sink_fallible!($enum, $sink_fn, $sink_in),
|
||||
)
|
||||
};
|
||||
}
|
||||
|
||||
/// Builds a typed `Pipe<D, In, Out>` — sourceless and sinkless.
|
||||
///
|
||||
/// Syntax:
|
||||
/// ```ignore
|
||||
/// make_pipe! {
|
||||
/// MyData : InType => OutType,
|
||||
/// | func : InVariant => OutVariant, // transform 1→1
|
||||
/// |? func : InVariant => OutVariant, // transform 1→1 fallible
|
||||
/// || func : InVariant => OutVariant, // flat transform 1→N
|
||||
/// ||? func : InVariant => OutVariant, // flat transform 1→N fallible
|
||||
/// }
|
||||
/// ```
|
||||
#[macro_export]
|
||||
macro_rules! make_pipe {
|
||||
// ── Entry: first stage | ─────────────────────────────────────────────
|
||||
($enum:ident : $in_ty:ty => $out_ty:ty,
|
||||
| $tf:tt : $fi:ident => $fo:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$crate::make_transform!($enum, $tf, $fi, $fo),], $fo, $($rest)*)
|
||||
};
|
||||
// ── Entry: first stage |? ────────────────────────────────────────────
|
||||
($enum:ident : $in_ty:ty => $out_ty:ty,
|
||||
|? $tf:tt : $fi:ident => $fo:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$crate::make_transform_fallible!($enum, $tf, $fi, $fo),], $fo, $($rest)*)
|
||||
};
|
||||
// ── Entry: first stage || ────────────────────────────────────────────
|
||||
($enum:ident : $in_ty:ty => $out_ty:ty,
|
||||
|| $tf:tt : $fi:ident => $fo:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$crate::make_flat_transform!($enum, $tf, $fi, $fo),], $fo, $($rest)*)
|
||||
};
|
||||
// ── Entry: first stage ||? ───────────────────────────────────────────
|
||||
($enum:ident : $in_ty:ty => $out_ty:ty,
|
||||
||? $tf:tt : $fi:ident => $fo:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$crate::make_flat_transform_fallible!($enum, $tf, $fi, $fo),], $fo, $($rest)*)
|
||||
};
|
||||
|
||||
// ── Accumulation: | ──────────────────────────────────────────────────
|
||||
(@build $enum:ident : $in_ty:ty => $out_ty:ty, $fi:ident,
|
||||
[$($acc:tt)*], $lo:ident,
|
||||
| $tf:tt : $ti:ident => $to:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$($acc)* $crate::make_transform!($enum, $tf, $ti, $to),], $to, $($rest)*)
|
||||
};
|
||||
// ── Accumulation: |? ─────────────────────────────────────────────────
|
||||
(@build $enum:ident : $in_ty:ty => $out_ty:ty, $fi:ident,
|
||||
[$($acc:tt)*], $lo:ident,
|
||||
|? $tf:tt : $ti:ident => $to:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$($acc)* $crate::make_transform_fallible!($enum, $tf, $ti, $to),], $to, $($rest)*)
|
||||
};
|
||||
// ── Accumulation: || ─────────────────────────────────────────────────
|
||||
(@build $enum:ident : $in_ty:ty => $out_ty:ty, $fi:ident,
|
||||
[$($acc:tt)*], $lo:ident,
|
||||
|| $tf:tt : $ti:ident => $to:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$($acc)* $crate::make_flat_transform!($enum, $tf, $ti, $to),], $to, $($rest)*)
|
||||
};
|
||||
// ── Accumulation: ||? ────────────────────────────────────────────────
|
||||
(@build $enum:ident : $in_ty:ty => $out_ty:ty, $fi:ident,
|
||||
[$($acc:tt)*], $lo:ident,
|
||||
||? $tf:tt : $ti:ident => $to:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$($acc)* $crate::make_flat_transform_fallible!($enum, $tf, $ti, $to),], $to, $($rest)*)
|
||||
};
|
||||
|
||||
// ── Termination ───────────────────────────────────────────────────────
|
||||
(@build $enum:ident : $in_ty:ty => $out_ty:ty, $fi:ident,
|
||||
[$($acc:tt)*], $lo:ident $(,)?) => {
|
||||
$crate::Pipe::new(
|
||||
vec![$($acc)*],
|
||||
::std::sync::Arc::new(|x: $in_ty| $enum::$fi(x)),
|
||||
::std::sync::Arc::new(|d: $enum| -> $out_ty {
|
||||
if let $enum::$lo(x) = d { x }
|
||||
else { ::std::unreachable!("unexpected pipeline data variant in make_pipe!") }
|
||||
}),
|
||||
)
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,35 @@
|
||||
use std::error::Error;
|
||||
use std::fmt;
|
||||
|
||||
/// Error type for pipeline operations.
|
||||
#[derive(Debug)]
|
||||
pub enum PipelineError {
|
||||
/// A stage received a `PipelineData` variant it did not expect.
|
||||
TypeMismatch,
|
||||
/// The step kind is not compatible with the data type.
|
||||
StepKindMismatch(&'static str),
|
||||
/// The source has no more data to produce.
|
||||
EndOfStream,
|
||||
/// An error occurred inside a stage (e.g., I/O, parsing, custom logic).
|
||||
StepError(Box<dyn Error + Send + Sync>),
|
||||
}
|
||||
|
||||
impl fmt::Display for PipelineError {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
match self {
|
||||
PipelineError::TypeMismatch => write!(f, "data type mismatch in pipeline stage"),
|
||||
PipelineError::StepKindMismatch(s) => write!(f, "step kind mismatch: {}", s),
|
||||
PipelineError::EndOfStream => write!(f, "end of input stream"),
|
||||
PipelineError::StepError(e) => write!(f, "stage error: {}", e),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Error for PipelineError {
|
||||
fn source(&self) -> Option<&(dyn Error + 'static)> {
|
||||
match self {
|
||||
PipelineError::StepError(e) => Some(e.as_ref()),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,373 @@
|
||||
// ── Macros ────────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Creates a `SourceFn` from an iterator of plain values.
|
||||
#[macro_export]
|
||||
macro_rules! make_source {
|
||||
($enum:ident, $iterator:expr, $output:ident) => {{
|
||||
let mut iter = $iterator.into_iter();
|
||||
Box::new(
|
||||
move || -> ::std::result::Result<$enum, $crate::PipelineError> {
|
||||
match iter.next() {
|
||||
Some(x) => Ok($enum::$output(x)),
|
||||
None => Err($crate::PipelineError::EndOfStream),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<dyn FnMut() -> ::std::result::Result<$enum, $crate::PipelineError> + Send>
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `SourceFn` from an iterator of `Result<T, E>`.
|
||||
#[macro_export]
|
||||
macro_rules! make_source_fallible {
|
||||
($enum:ident, $iterator:expr, $output:ident) => {{
|
||||
let mut iter = $iterator.into_iter();
|
||||
Box::new(
|
||||
move || -> ::std::result::Result<$enum, $crate::PipelineError> {
|
||||
match iter.next() {
|
||||
Some(Ok(x)) => Ok($enum::$output(x)),
|
||||
Some(Err(e)) => Err($crate::PipelineError::StepError(Box::new(e))),
|
||||
None => Err($crate::PipelineError::EndOfStream),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<dyn FnMut() -> ::std::result::Result<$enum, $crate::PipelineError> + Send>
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `Stage::Transform` from a pure (non-fallible) function `Fn(T) -> U`.
|
||||
#[macro_export]
|
||||
macro_rules! make_transform {
|
||||
($enum:ident, $func:tt, $input:ident, $output:ident) => {{
|
||||
let __f = $func;
|
||||
$crate::Stage::Transform(
|
||||
::std::sync::Arc::from(Box::new(
|
||||
move |data: $enum| -> ::std::result::Result<$enum, $crate::PipelineError> {
|
||||
match data {
|
||||
$enum::$input(x) => Ok($enum::$output(__f(x))),
|
||||
_ => Err($crate::PipelineError::TypeMismatch),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<
|
||||
dyn Fn($enum) -> ::std::result::Result<$enum, $crate::PipelineError>
|
||||
+ Send
|
||||
+ Sync,
|
||||
>)
|
||||
)
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `Stage::Transform` from a fallible function `Fn(T) -> Result<U, E>`.
|
||||
#[macro_export]
|
||||
macro_rules! make_transform_fallible {
|
||||
($enum:ident, $func:tt, $input:ident, $output:ident) => {{
|
||||
let __f = $func;
|
||||
$crate::Stage::Transform(
|
||||
::std::sync::Arc::from(Box::new(
|
||||
move |data: $enum| -> ::std::result::Result<$enum, $crate::PipelineError> {
|
||||
match data {
|
||||
$enum::$input(inner) => {
|
||||
let result = __f(inner)
|
||||
.map_err(|e| $crate::PipelineError::StepError(Box::new(e)))?;
|
||||
Ok($enum::$output(result))
|
||||
}
|
||||
_ => Err($crate::PipelineError::TypeMismatch),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<
|
||||
dyn Fn($enum) -> ::std::result::Result<$enum, $crate::PipelineError>
|
||||
+ Send
|
||||
+ Sync,
|
||||
>)
|
||||
)
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `Stage::Flat` from a function `Fn(T) -> impl IntoIterator<Item = U>`.
|
||||
///
|
||||
/// Pour chaque item produit par l'itérateur, il est poussé individuellement dans
|
||||
/// le canal de sortie, permettant au scheduler de dispatcher les items en parallèle
|
||||
/// dès qu'un worker est disponible.
|
||||
#[macro_export]
|
||||
macro_rules! make_flat_transform {
|
||||
($enum:ident, $func:tt, $input:ident, $output:ident) => {{
|
||||
let __f = $func;
|
||||
$crate::Stage::Flat(
|
||||
::std::sync::Arc::new(
|
||||
move |data: $enum,
|
||||
push: &$crate::PipelineSender<
|
||||
::std::result::Result<$enum, $crate::PipelineError>,
|
||||
>,
|
||||
delta: &$crate::PipelineSender<isize>| {
|
||||
match data {
|
||||
$enum::$input(inner) => {
|
||||
let mut count: isize = 0;
|
||||
for item in __f(inner) {
|
||||
push.send(Ok($enum::$output(item))).ok();
|
||||
count += 1;
|
||||
}
|
||||
delta.send(count - 1).ok();
|
||||
}
|
||||
_ => {
|
||||
push.send(Err($crate::PipelineError::TypeMismatch)).ok();
|
||||
delta.send(0).ok();
|
||||
}
|
||||
}
|
||||
},
|
||||
) as $crate::SharedFlatFn<$enum>
|
||||
)
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `Stage::Flat` from a fallible function
|
||||
/// `Fn(T) -> Result<impl IntoIterator<Item = U>, E>`.
|
||||
///
|
||||
/// Si la fonction retourne `Err`, une erreur est poussée dans le canal et aucun
|
||||
/// item normal n'est produit.
|
||||
#[macro_export]
|
||||
macro_rules! make_flat_transform_fallible {
|
||||
($enum:ident, $func:tt, $input:ident, $output:ident) => {{
|
||||
let __f = $func;
|
||||
$crate::Stage::Flat(
|
||||
::std::sync::Arc::new(
|
||||
move |data: $enum,
|
||||
push: &$crate::PipelineSender<
|
||||
::std::result::Result<$enum, $crate::PipelineError>,
|
||||
>,
|
||||
delta: &$crate::PipelineSender<isize>| {
|
||||
match data {
|
||||
$enum::$input(inner) => match __f(inner) {
|
||||
Ok(iter) => {
|
||||
let mut count: isize = 0;
|
||||
for item in iter {
|
||||
push.send(Ok($enum::$output(item))).ok();
|
||||
count += 1;
|
||||
}
|
||||
delta.send(count - 1).ok();
|
||||
}
|
||||
Err(e) => {
|
||||
push.send(Err($crate::PipelineError::StepError(Box::new(e))))
|
||||
.ok();
|
||||
delta.send(0).ok();
|
||||
}
|
||||
},
|
||||
_ => {
|
||||
push.send(Err($crate::PipelineError::TypeMismatch)).ok();
|
||||
delta.send(0).ok();
|
||||
}
|
||||
}
|
||||
},
|
||||
) as $crate::SharedFlatFn<$enum>
|
||||
)
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `SinkFn` from a function that consumes a concrete value and returns `()`.
|
||||
#[macro_export]
|
||||
macro_rules! make_sink {
|
||||
($enum:ident, $func:tt, $input:ident) => {{
|
||||
let __f = $func;
|
||||
Box::new(
|
||||
move |data: $enum| -> ::std::result::Result<(), $crate::PipelineError> {
|
||||
match data {
|
||||
$enum::$input(x) => {
|
||||
__f(x);
|
||||
Ok(())
|
||||
}
|
||||
_ => Err($crate::PipelineError::TypeMismatch),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<dyn Fn($enum) -> ::std::result::Result<(), $crate::PipelineError> + Send>
|
||||
}};
|
||||
}
|
||||
|
||||
/// Creates a `SinkFn` from a fallible function that returns `Result<(), E>`.
|
||||
#[macro_export]
|
||||
macro_rules! make_sink_fallible {
|
||||
($enum:ident, $func:tt, $input:ident) => {{
|
||||
let __f = $func;
|
||||
Box::new(
|
||||
move |data: $enum| -> ::std::result::Result<(), $crate::PipelineError> {
|
||||
match data {
|
||||
$enum::$input(inner) => {
|
||||
__f(inner).map_err(|e| $crate::PipelineError::StepError(Box::new(e)))
|
||||
}
|
||||
_ => Err($crate::PipelineError::TypeMismatch),
|
||||
}
|
||||
},
|
||||
)
|
||||
as Box<dyn Fn($enum) -> ::std::result::Result<(), $crate::PipelineError> + Send>
|
||||
}};
|
||||
}
|
||||
|
||||
/// Construit un `Pipeline` à partir d'une source, d'une liste de stages et d'un sink.
|
||||
///
|
||||
/// Syntaxe :
|
||||
/// ```ignore
|
||||
/// make_pipeline! {
|
||||
/// MyData,
|
||||
/// source my_iter => Variant, // source non-fallible
|
||||
/// source? my_iter => Variant, // source fallible (Result<T, E>)
|
||||
/// | func: In => Out, // transform 1→1 non-fallible
|
||||
/// |? func: In => Out, // transform 1→1 fallible
|
||||
/// || func: In => Out, // flat transform 1→N non-fallible
|
||||
/// ||? func: In => Out, // flat transform 1→N fallible
|
||||
/// sink my_func @ Variant, // sink non-fallible
|
||||
/// sink? my_func @ Variant, // sink fallible
|
||||
/// }
|
||||
/// ```
|
||||
#[macro_export]
|
||||
macro_rules! make_pipeline {
|
||||
// ── Points d'entrée ──────────────────────────────────────────────────
|
||||
|
||||
($enum:ident, source $src:expr => $src_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum,
|
||||
{ $crate::make_source!($enum, $src, $src_out) },
|
||||
[],
|
||||
$($rest)*)
|
||||
};
|
||||
($enum:ident, source? $src:expr => $src_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum,
|
||||
{ $crate::make_source_fallible!($enum, $src, $src_out) },
|
||||
[],
|
||||
$($rest)*)
|
||||
};
|
||||
|
||||
// ── Accumulation des stages ──────────────────────────────────────────
|
||||
|
||||
// transform 1→1 non-fallible
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
| $tf:tt : $t_in:ident => $t_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum, $source,
|
||||
[$($acc)* $crate::make_transform!($enum, $tf, $t_in, $t_out),],
|
||||
$($rest)*)
|
||||
};
|
||||
|
||||
// transform 1→1 fallible
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
|? $tf:tt : $t_in:ident => $t_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum, $source,
|
||||
[$($acc)* $crate::make_transform_fallible!($enum, $tf, $t_in, $t_out),],
|
||||
$($rest)*)
|
||||
};
|
||||
|
||||
// flat transform 1→N non-fallible
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
|| $tf:tt : $t_in:ident => $t_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum, $source,
|
||||
[$($acc)* $crate::make_flat_transform!($enum, $tf, $t_in, $t_out),],
|
||||
$($rest)*)
|
||||
};
|
||||
|
||||
// flat transform 1→N fallible
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
||? $tf:tt : $t_in:ident => $t_out:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipeline!(@build $enum, $source,
|
||||
[$($acc)* $crate::make_flat_transform_fallible!($enum, $tf, $t_in, $t_out),],
|
||||
$($rest)*)
|
||||
};
|
||||
|
||||
// ── Terminaison : sink ───────────────────────────────────────────────
|
||||
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
sink $sink_fn:tt @ $sink_in:ident $(,)?) => {
|
||||
$crate::Pipeline::new(
|
||||
$source,
|
||||
vec![$($acc)*],
|
||||
$crate::make_sink!($enum, $sink_fn, $sink_in),
|
||||
)
|
||||
};
|
||||
(@build $enum:ident, $source:tt, [$($acc:tt)*],
|
||||
sink? $sink_fn:tt @ $sink_in:ident $(,)?) => {
|
||||
$crate::Pipeline::new(
|
||||
$source,
|
||||
vec![$($acc)*],
|
||||
$crate::make_sink_fallible!($enum, $sink_fn, $sink_in),
|
||||
)
|
||||
};
|
||||
}
|
||||
|
||||
/// Builds a typed `Pipe<D, In, Out>` — sourceless and sinkless.
|
||||
///
|
||||
/// Syntax:
|
||||
/// ```ignore
|
||||
/// make_pipe! {
|
||||
/// MyData : InType => OutType,
|
||||
/// | func : InVariant => OutVariant, // transform 1→1
|
||||
/// |? func : InVariant => OutVariant, // transform 1→1 fallible
|
||||
/// || func : InVariant => OutVariant, // flat transform 1→N
|
||||
/// ||? func : InVariant => OutVariant, // flat transform 1→N fallible
|
||||
/// }
|
||||
/// ```
|
||||
#[macro_export]
|
||||
macro_rules! make_pipe {
|
||||
// ── Entry: first stage | ─────────────────────────────────────────────
|
||||
($enum:ident : $in_ty:ty => $out_ty:ty,
|
||||
| $tf:tt : $fi:ident => $fo:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$crate::make_transform!($enum, $tf, $fi, $fo),], $fo, $($rest)*)
|
||||
};
|
||||
// ── Entry: first stage |? ────────────────────────────────────────────
|
||||
($enum:ident : $in_ty:ty => $out_ty:ty,
|
||||
|? $tf:tt : $fi:ident => $fo:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$crate::make_transform_fallible!($enum, $tf, $fi, $fo),], $fo, $($rest)*)
|
||||
};
|
||||
// ── Entry: first stage || ────────────────────────────────────────────
|
||||
($enum:ident : $in_ty:ty => $out_ty:ty,
|
||||
|| $tf:tt : $fi:ident => $fo:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$crate::make_flat_transform!($enum, $tf, $fi, $fo),], $fo, $($rest)*)
|
||||
};
|
||||
// ── Entry: first stage ||? ───────────────────────────────────────────
|
||||
($enum:ident : $in_ty:ty => $out_ty:ty,
|
||||
||? $tf:tt : $fi:ident => $fo:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$crate::make_flat_transform_fallible!($enum, $tf, $fi, $fo),], $fo, $($rest)*)
|
||||
};
|
||||
|
||||
// ── Accumulation: | ──────────────────────────────────────────────────
|
||||
(@build $enum:ident : $in_ty:ty => $out_ty:ty, $fi:ident,
|
||||
[$($acc:tt)*], $lo:ident,
|
||||
| $tf:tt : $ti:ident => $to:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$($acc)* $crate::make_transform!($enum, $tf, $ti, $to),], $to, $($rest)*)
|
||||
};
|
||||
// ── Accumulation: |? ─────────────────────────────────────────────────
|
||||
(@build $enum:ident : $in_ty:ty => $out_ty:ty, $fi:ident,
|
||||
[$($acc:tt)*], $lo:ident,
|
||||
|? $tf:tt : $ti:ident => $to:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$($acc)* $crate::make_transform_fallible!($enum, $tf, $ti, $to),], $to, $($rest)*)
|
||||
};
|
||||
// ── Accumulation: || ─────────────────────────────────────────────────
|
||||
(@build $enum:ident : $in_ty:ty => $out_ty:ty, $fi:ident,
|
||||
[$($acc:tt)*], $lo:ident,
|
||||
|| $tf:tt : $ti:ident => $to:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$($acc)* $crate::make_flat_transform!($enum, $tf, $ti, $to),], $to, $($rest)*)
|
||||
};
|
||||
// ── Accumulation: ||? ────────────────────────────────────────────────
|
||||
(@build $enum:ident : $in_ty:ty => $out_ty:ty, $fi:ident,
|
||||
[$($acc:tt)*], $lo:ident,
|
||||
||? $tf:tt : $ti:ident => $to:ident, $($rest:tt)*) => {
|
||||
$crate::make_pipe!(@build $enum : $in_ty => $out_ty, $fi,
|
||||
[$($acc)* $crate::make_flat_transform_fallible!($enum, $tf, $ti, $to),], $to, $($rest)*)
|
||||
};
|
||||
|
||||
// ── Termination ───────────────────────────────────────────────────────
|
||||
(@build $enum:ident : $in_ty:ty => $out_ty:ty, $fi:ident,
|
||||
[$($acc:tt)*], $lo:ident $(,)?) => {
|
||||
$crate::Pipe::new(
|
||||
vec![$($acc)*],
|
||||
::std::sync::Arc::new(|x: $in_ty| $enum::$fi(x)),
|
||||
::std::sync::Arc::new(|d: $enum| -> $out_ty {
|
||||
if let $enum::$lo(x) = d { x }
|
||||
else { ::std::unreachable!("unexpected pipeline data variant in make_pipe!") }
|
||||
}),
|
||||
)
|
||||
};
|
||||
}
|
||||
@@ -0,0 +1,19 @@
|
||||
//! Scheduler: a channel/thread-based pipeline runtime with typed, composable
|
||||
//! stages (`Pipe`), plus the macro-based `PipelineData`-enum runtime it grew
|
||||
//! out of (`Pipeline`/`WorkerPool`).
|
||||
//!
|
||||
//! Submodules: [`error`] (`PipelineError`), [`types`] (function types, `Stage`),
|
||||
//! [`runner`] (thread bodies), [`pool`] (`Pipeline`/`WorkerPool` scheduler
|
||||
//! loop), [`pipe`] (`Pipe`/`PipeIter`), [`macros`] (`make_pipe!` and friends).
|
||||
|
||||
mod error;
|
||||
mod macros;
|
||||
mod pipe;
|
||||
mod pool;
|
||||
mod runner;
|
||||
mod types;
|
||||
|
||||
pub use error::PipelineError;
|
||||
pub use pipe::{Pipe, PipeIter};
|
||||
pub use pool::{Pipeline, WorkerPool};
|
||||
pub use types::{SharedFlatFn, SharedFn, SinkFn, SourceFn, Stage};
|
||||
@@ -0,0 +1,117 @@
|
||||
use crossbeam_channel::{Receiver, bounded};
|
||||
use std::marker::PhantomData;
|
||||
use std::sync::Arc;
|
||||
use std::thread;
|
||||
|
||||
use super::error::PipelineError;
|
||||
use super::pool::{Pipeline, WorkerPool};
|
||||
use super::types::{SinkFn, SourceFn, Stage};
|
||||
|
||||
// ── Pipe ──────────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Typed, composable iterator transformer.
|
||||
///
|
||||
/// A `Pipe<D, In, Out>` is a pure description of pipeline stages — no threads,
|
||||
/// no channels, no scheduler. Call `.apply(iter, n_workers, capacity)` to start
|
||||
/// execution and get back a `PipeIter<Out>`.
|
||||
///
|
||||
/// Compose two pipes with `.then()`: the resulting `Pipe` holds the concatenated
|
||||
/// stage list, so a single scheduler is created when `.apply()` is eventually called.
|
||||
pub struct Pipe<D, In, Out> {
|
||||
stages: Vec<Stage<D>>,
|
||||
wrap: Arc<dyn Fn(In) -> D + Send + Sync>,
|
||||
unwrap: Arc<dyn Fn(D) -> Out + Send + Sync>,
|
||||
_phantom: PhantomData<(In, Out)>,
|
||||
}
|
||||
|
||||
impl<D, In, Out> Pipe<D, In, Out> {
|
||||
/// Build a `Pipe` from stages and wrap/unwrap converters.
|
||||
/// Prefer the `make_pipe!` macro.
|
||||
pub fn new(
|
||||
stages: Vec<Stage<D>>,
|
||||
wrap: Arc<dyn Fn(In) -> D + Send + Sync>,
|
||||
unwrap: Arc<dyn Fn(D) -> Out + Send + Sync>,
|
||||
) -> Self {
|
||||
Self { stages, wrap, unwrap, _phantom: PhantomData }
|
||||
}
|
||||
|
||||
/// Concatenate stages from two pipes into one.
|
||||
///
|
||||
/// Requires `Out` of `self` == `In` of `other`. The single scheduler
|
||||
/// created at `.apply()` time sees the full combined stage list.
|
||||
pub fn then<Next>(self, other: Pipe<D, Out, Next>) -> Pipe<D, In, Next> {
|
||||
Pipe {
|
||||
stages: self.stages.into_iter().chain(other.stages).collect(),
|
||||
wrap: self.wrap,
|
||||
unwrap: other.unwrap,
|
||||
_phantom: PhantomData,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<D, In, Out> Pipe<D, In, Out>
|
||||
where
|
||||
D: Send + Sync + 'static,
|
||||
In: Send + 'static,
|
||||
Out: Send + 'static,
|
||||
{
|
||||
/// Run the pipeline in a background thread; returns an iterator over the output.
|
||||
pub fn apply(
|
||||
self,
|
||||
input: impl Iterator<Item = In> + Send + 'static,
|
||||
n_workers: usize,
|
||||
capacity: usize,
|
||||
) -> PipeIter<Out> {
|
||||
let wrap = Arc::clone(&self.wrap);
|
||||
let unwrap = Arc::clone(&self.unwrap);
|
||||
|
||||
let mut iter = input;
|
||||
let source: SourceFn<D> = Box::new(move || match iter.next() {
|
||||
Some(x) => Ok(wrap(x)),
|
||||
None => Err(PipelineError::EndOfStream),
|
||||
});
|
||||
|
||||
let (out_tx, out_rx) = bounded::<Out>(capacity);
|
||||
let sink: SinkFn<D> = Box::new(move |data: D| {
|
||||
out_tx.send(unwrap(data)).map_err(|_| {
|
||||
PipelineError::StepError(Box::new(std::io::Error::new(
|
||||
std::io::ErrorKind::BrokenPipe,
|
||||
"output channel closed",
|
||||
)))
|
||||
})
|
||||
});
|
||||
|
||||
let pipeline = Pipeline::new(source, self.stages, sink);
|
||||
let handle = thread::spawn(move || {
|
||||
WorkerPool::new(pipeline, n_workers, capacity).run();
|
||||
});
|
||||
|
||||
PipeIter { rx: out_rx, handle: Some(handle) }
|
||||
}
|
||||
}
|
||||
|
||||
// ── PipeIter ──────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Iterator over the output of `Pipe::apply()`.
|
||||
pub struct PipeIter<Out> {
|
||||
rx: Receiver<Out>,
|
||||
handle: Option<thread::JoinHandle<()>>,
|
||||
}
|
||||
|
||||
impl<Out> Iterator for PipeIter<Out> {
|
||||
type Item = Out;
|
||||
|
||||
fn next(&mut self) -> Option<Out> {
|
||||
self.rx.recv().ok()
|
||||
}
|
||||
}
|
||||
|
||||
impl<Out> Drop for PipeIter<Out> {
|
||||
fn drop(&mut self) {
|
||||
// Drain buffered items so the scheduler can unblock if the channel is full.
|
||||
while self.rx.try_recv().is_ok() {}
|
||||
if let Some(h) = self.handle.take() {
|
||||
let _ = h.join();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,203 @@
|
||||
use crossbeam_channel::{Receiver, Select, Sender, bounded};
|
||||
|
||||
use super::error::PipelineError;
|
||||
use super::runner::{dispatch, sink_runner, source_runner, transform_runner};
|
||||
use super::types::{SinkFn, SourceFn, Stage, WorkerTask};
|
||||
|
||||
// ── Pipeline ──────────────────────────────────────────────────────────────────
|
||||
|
||||
pub struct Pipeline<DATA> {
|
||||
source: SourceFn<DATA>,
|
||||
stages: Vec<Stage<DATA>>,
|
||||
sink: SinkFn<DATA>,
|
||||
}
|
||||
|
||||
impl<DATA> Pipeline<DATA> {
|
||||
pub fn new(
|
||||
source: SourceFn<DATA>,
|
||||
stages: Vec<Stage<DATA>>,
|
||||
sink: SinkFn<DATA>,
|
||||
) -> Self {
|
||||
Self { source, stages, sink }
|
||||
}
|
||||
}
|
||||
|
||||
// ── WorkerPool ────────────────────────────────────────────────────────────────
|
||||
|
||||
pub struct WorkerPool<DATA> {
|
||||
pipeline: Pipeline<DATA>,
|
||||
handles: Vec<std::thread::JoinHandle<()>>,
|
||||
n_workers: usize,
|
||||
capacity: usize,
|
||||
}
|
||||
|
||||
impl<DATA> WorkerPool<DATA>
|
||||
where
|
||||
DATA: Send + Sync + 'static,
|
||||
{
|
||||
pub fn new(pipeline: Pipeline<DATA>, n_workers: usize, capacity: usize) -> Self {
|
||||
Self {
|
||||
pipeline,
|
||||
handles: Vec::new(),
|
||||
n_workers,
|
||||
capacity,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn run(mut self) {
|
||||
let n = self.pipeline.stages.len();
|
||||
|
||||
// ── Canaux inter-stages ────────────────────────────────────────────
|
||||
// stage_txs[i] / stage_rxs[i] : sortie du stage i
|
||||
let mut stage_txs: Vec<Sender<Result<DATA, PipelineError>>> = Vec::new();
|
||||
let mut stage_rxs: Vec<Receiver<Result<DATA, PipelineError>>> = Vec::new();
|
||||
for _ in 0..n {
|
||||
let (tx, rx) = bounded(self.capacity);
|
||||
stage_txs.push(tx);
|
||||
stage_rxs.push(rx);
|
||||
}
|
||||
|
||||
// ── Source thread ──────────────────────────────────────────────────
|
||||
let (source_rx, src_handle) = source_runner(self.pipeline.source, self.capacity);
|
||||
self.handles.push(src_handle);
|
||||
|
||||
let stages = self.pipeline.stages;
|
||||
|
||||
// ── Canal delta pour les flat stages ───────────────────────────────
|
||||
// Chaque flat worker envoie `N-1` ici après avoir poussé N items.
|
||||
// Le scheduler ajuste `in_flight` en conséquence.
|
||||
let (flat_delta_tx, flat_delta_rx) = bounded::<isize>(self.capacity);
|
||||
|
||||
// ── Worker pool ────────────────────────────────────────────────────
|
||||
let (worker_tx, worker_rx): (Sender<WorkerTask<DATA>>, Receiver<WorkerTask<DATA>>) =
|
||||
bounded(self.capacity);
|
||||
|
||||
for _ in 0..self.n_workers {
|
||||
self.handles.push(transform_runner(
|
||||
worker_rx.clone(),
|
||||
stages.iter().map(Stage::clone).collect(),
|
||||
stage_txs.clone(),
|
||||
flat_delta_tx.clone(),
|
||||
));
|
||||
}
|
||||
// Le scheduler ne tient plus flat_delta_tx : les workers le détiennent.
|
||||
// On le drop ici pour que le canal se ferme quand les workers terminent.
|
||||
drop(flat_delta_tx);
|
||||
|
||||
// ── Sink thread ────────────────────────────────────────────────────
|
||||
let (sink_tx, sink_err_rx, sink_handle) = sink_runner(self.pipeline.sink, self.capacity);
|
||||
self.handles.push(sink_handle);
|
||||
|
||||
// ── Boucle principale ──────────────────────────────────────────────
|
||||
//
|
||||
// `in_flight` (isize) = nb d'items qui doivent encore atteindre le sink.
|
||||
// Peut temporairement être négatif si un flat worker a poussé ses items
|
||||
// avant que le scheduler ait reçu le delta correspondant.
|
||||
//
|
||||
// `flat_workers_active` = nb de flat workers en cours d'exécution.
|
||||
// Empêche la terminaison prématurée quand in_flight vaut 0 mais qu'un
|
||||
// flat worker n'a pas encore envoyé son delta.
|
||||
//
|
||||
// Priorités du Select biaisé (index le plus bas = priorité la plus haute) :
|
||||
// 0 → sink_err_rx (arrêt immédiat sur erreur sink)
|
||||
// 1 → flat_delta_rx (mettre à jour in_flight avant de dispatcher)
|
||||
// 2..=n+1 → stage_rxs[n-1..0] (vider le pipeline en priorité)
|
||||
// n+2 → source_rx (dernier recours : nouvelles données)
|
||||
//
|
||||
// Quand k = 0 : erreur du sink
|
||||
// Quand k = 1 : delta d'un flat worker
|
||||
// Quand 2 ≤ k ≤ n+1 : résultat du stage n+1-k
|
||||
// Quand k = n+2 : item source
|
||||
//
|
||||
// Terminaison : source tarie ET in_flight == 0 ET aucun flat worker actif.
|
||||
{
|
||||
let mut source_done = false;
|
||||
let mut in_flight: isize = 0;
|
||||
let mut flat_workers_active: usize = 0;
|
||||
|
||||
loop {
|
||||
if source_done && in_flight == 0 && flat_workers_active == 0 {
|
||||
break;
|
||||
}
|
||||
|
||||
let mut sel = Select::new_biased();
|
||||
sel.recv(&sink_err_rx); // index 0
|
||||
sel.recv(&flat_delta_rx); // index 1
|
||||
for rx in stage_rxs.iter().rev() {
|
||||
sel.recv(rx); // indices 2..=n+1
|
||||
}
|
||||
let src_idx = if !source_done {
|
||||
Some(sel.recv(&source_rx)) // index n+2
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
let oper = sel.select();
|
||||
let k = oper.index();
|
||||
|
||||
if k == 0 {
|
||||
// ── Erreur du sink ────────────────────────────────────
|
||||
match oper.recv(&sink_err_rx) {
|
||||
Ok(e) => { eprintln!("Sink error: {:?}", e); break; }
|
||||
Err(_) => break,
|
||||
}
|
||||
} else if k == 1 {
|
||||
// ── Delta d'un flat worker ────────────────────────────
|
||||
// delta = N - 1 (N items poussés, 1 item consommé)
|
||||
match oper.recv(&flat_delta_rx) {
|
||||
Ok(delta) => {
|
||||
in_flight += delta;
|
||||
flat_workers_active -= 1;
|
||||
}
|
||||
Err(_) => {}
|
||||
}
|
||||
} else if src_idx == Some(k) {
|
||||
// ── Nouvel item depuis la source ──────────────────────
|
||||
match oper.recv(&source_rx) {
|
||||
Ok(Ok(data)) => {
|
||||
if n == 0 {
|
||||
let _ = sink_tx.send(data);
|
||||
} else {
|
||||
in_flight += 1;
|
||||
dispatch(
|
||||
data, 0,
|
||||
&stages, &worker_tx,
|
||||
&mut flat_workers_active,
|
||||
);
|
||||
}
|
||||
}
|
||||
Ok(Err(e)) => eprintln!("Source error: {:?}", e),
|
||||
Err(_) => source_done = true,
|
||||
}
|
||||
} else {
|
||||
// ── Résultat d'un stage intermédiaire ─────────────────
|
||||
// k ∈ [2, n+1] → stage = n+1 - k
|
||||
let stage = n + 1 - k;
|
||||
match oper.recv(&stage_rxs[stage]) {
|
||||
Ok(Ok(data)) => {
|
||||
if stage == n - 1 {
|
||||
in_flight -= 1;
|
||||
let _ = sink_tx.send(data);
|
||||
} else {
|
||||
dispatch(
|
||||
data, stage + 1,
|
||||
&stages, &worker_tx,
|
||||
&mut flat_workers_active,
|
||||
);
|
||||
}
|
||||
}
|
||||
Ok(Err(e)) => eprintln!("Stage {} error: {:?}", stage, e),
|
||||
Err(_) => break,
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
drop(worker_tx);
|
||||
drop(sink_tx);
|
||||
|
||||
for h in self.handles {
|
||||
let _ = h.join();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,118 @@
|
||||
use crossbeam_channel::{Receiver, Sender, bounded};
|
||||
use std::thread;
|
||||
|
||||
use super::error::PipelineError;
|
||||
use super::types::{SinkFn, SourceFn, Stage, WorkerTask};
|
||||
|
||||
// ── Thread runners ────────────────────────────────────────────────────────────
|
||||
|
||||
pub(super) fn source_runner<DATA>(
|
||||
mut source: SourceFn<DATA>,
|
||||
capacity: usize,
|
||||
) -> (
|
||||
Receiver<Result<DATA, PipelineError>>,
|
||||
thread::JoinHandle<()>,
|
||||
)
|
||||
where
|
||||
DATA: Send + Sync + 'static,
|
||||
{
|
||||
let (tx, rx) = bounded(capacity);
|
||||
let handle = thread::spawn(move || {
|
||||
loop {
|
||||
match source() {
|
||||
Ok(data) => {
|
||||
if tx.send(Ok(data)).is_err() {
|
||||
break;
|
||||
}
|
||||
}
|
||||
Err(PipelineError::EndOfStream) => break,
|
||||
Err(e) => {
|
||||
eprintln!("Source error: {:?}", e);
|
||||
let _ = tx.send(Err(e));
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
});
|
||||
(rx, handle)
|
||||
}
|
||||
|
||||
/// Lance un thread worker du pool.
|
||||
///
|
||||
/// Gère deux types de tâches :
|
||||
/// - `Transform` : applique `f(data)` et envoie le résultat dans `result_tx`.
|
||||
/// - `Flat` : appelle `f(data, &push_tx, &delta_tx)` ; la fonction elle-même
|
||||
/// pousse ses items dans `push_tx` et envoie `N-1` dans `delta_tx`.
|
||||
pub(super) fn transform_runner<DATA>(
|
||||
task_rx: Receiver<WorkerTask<DATA>>,
|
||||
stages: Vec<Stage<DATA>>,
|
||||
stage_txs: Vec<Sender<Result<DATA, PipelineError>>>,
|
||||
flat_delta_tx: Sender<isize>,
|
||||
) -> thread::JoinHandle<()>
|
||||
where
|
||||
DATA: Send + Sync + 'static,
|
||||
{
|
||||
thread::spawn(move || {
|
||||
while let Ok(task) = task_rx.recv() {
|
||||
match task {
|
||||
WorkerTask::Transform(data, idx) => {
|
||||
if let Stage::Transform(f) = &stages[idx] {
|
||||
let _ = stage_txs[idx].send(f(data));
|
||||
}
|
||||
}
|
||||
WorkerTask::Flat(data, idx) => {
|
||||
if let Stage::Flat(f) = &stages[idx] {
|
||||
f(data, &stage_txs[idx], &flat_delta_tx);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/// Lance le thread sink.
|
||||
pub(super) fn sink_runner<DATA>(
|
||||
sink: SinkFn<DATA>,
|
||||
capacity: usize,
|
||||
) -> (
|
||||
Sender<DATA>,
|
||||
Receiver<PipelineError>,
|
||||
thread::JoinHandle<()>,
|
||||
)
|
||||
where
|
||||
DATA: Send + Sync + 'static,
|
||||
{
|
||||
let (data_tx, data_rx) = bounded(capacity);
|
||||
let (err_tx, err_rx) = bounded(capacity);
|
||||
let handle = thread::spawn(move || {
|
||||
for data in data_rx {
|
||||
if let Err(e) = sink(data) {
|
||||
let _ = err_tx.send(e);
|
||||
break;
|
||||
}
|
||||
}
|
||||
});
|
||||
(data_tx, err_rx, handle)
|
||||
}
|
||||
|
||||
/// Envoie `data` au stage `stage_idx`.
|
||||
/// Pour un `Transform`, empile une `WorkerTask::Transform`.
|
||||
/// Pour un `Flat`, incrémente `flat_workers_active` et empile une `WorkerTask::Flat`.
|
||||
#[inline]
|
||||
pub(super) fn dispatch<DATA>(
|
||||
data: DATA,
|
||||
stage_idx: usize,
|
||||
stages: &[Stage<DATA>],
|
||||
worker_tx: &Sender<WorkerTask<DATA>>,
|
||||
flat_workers_active: &mut usize,
|
||||
) {
|
||||
match &stages[stage_idx] {
|
||||
Stage::Transform(_) => {
|
||||
let _ = worker_tx.send(WorkerTask::Transform(data, stage_idx));
|
||||
}
|
||||
Stage::Flat(_) => {
|
||||
*flat_workers_active += 1;
|
||||
let _ = worker_tx.send(WorkerTask::Flat(data, stage_idx));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
use crossbeam_channel::Sender;
|
||||
use std::sync::Arc;
|
||||
|
||||
use super::error::PipelineError;
|
||||
|
||||
// ── Function types ────────────────────────────────────────────────────────────
|
||||
|
||||
/// Fonction source : appelée répétitivement, retourne le prochain item ou EndOfStream.
|
||||
/// `FnMut` car elle maintient un état interne (position dans l'itérateur).
|
||||
pub type SourceFn<D> = Box<dyn FnMut() -> Result<D, PipelineError> + Send>;
|
||||
|
||||
/// Fonction sink : consomme un item final, peut échouer (erreur d'I/O, etc.).
|
||||
pub type SinkFn<D> = Box<dyn Fn(D) -> Result<(), PipelineError> + Send>;
|
||||
|
||||
/// Fonction de transformation partagée entre workers via Arc.
|
||||
pub type SharedFn<D> = Arc<dyn Fn(D) -> Result<D, PipelineError> + Send + Sync>;
|
||||
|
||||
/// Fonction de transformation 1→N (flat map) partagée entre workers via Arc.
|
||||
///
|
||||
/// La fonction reçoit l'item d'entrée, un canal `push` pour envoyer chaque item
|
||||
/// produit, et un canal `delta` pour signaler au scheduler combien d'items
|
||||
/// supplémentaires sont entrés dans le pipeline (N-1 si N items produits).
|
||||
/// Elle doit appeler `delta.send(N - 1)` **après** avoir poussé tous les items.
|
||||
pub type SharedFlatFn<D> =
|
||||
Arc<dyn Fn(D, &Sender<Result<D, PipelineError>>, &Sender<isize>) + Send + Sync>;
|
||||
|
||||
// ── Stage enum ────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Une étape du pipeline : transform classique (1→1) ou flat transform (1→N).
|
||||
pub enum Stage<D> {
|
||||
Transform(SharedFn<D>),
|
||||
Flat(SharedFlatFn<D>),
|
||||
}
|
||||
|
||||
impl<D> Clone for Stage<D> {
|
||||
fn clone(&self) -> Self {
|
||||
match self {
|
||||
Stage::Transform(f) => Stage::Transform(Arc::clone(f)),
|
||||
Stage::Flat(f) => Stage::Flat(Arc::clone(f)),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ── Worker task ───────────────────────────────────────────────────────────────
|
||||
|
||||
pub(super) enum WorkerTask<D> {
|
||||
Transform(D, usize),
|
||||
Flat(D, usize),
|
||||
}
|
||||
@@ -1,733 +0,0 @@
|
||||
use std::io::{self, Read};
|
||||
use std::mem::ManuallyDrop;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use crate::mimetype::MimeTypeGuesser;
|
||||
use crate::xopen::open_raw;
|
||||
|
||||
pub const MAX_K: usize = 31;
|
||||
const PAGE_SIZE: usize = 65536;
|
||||
// overlap (MAX_K - 1) + page data (PAGE_SIZE) + 1 byte for the end-of-page terminating 0
|
||||
const BUF_SIZE: usize = MAX_K + PAGE_SIZE;
|
||||
|
||||
// ─── OverlapState ─────────────────────────────────────────────────────────────
|
||||
|
||||
pub(crate) struct OverlapState {
|
||||
data: [u8; MAX_K],
|
||||
len: usize,
|
||||
k: usize,
|
||||
}
|
||||
|
||||
impl OverlapState {
|
||||
pub(crate) fn new(k: usize) -> Self {
|
||||
assert!(k > 0 && k <= MAX_K);
|
||||
Self {
|
||||
data: [0u8; MAX_K],
|
||||
len: 0,
|
||||
k,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ─── NucParser trait ──────────────────────────────────────────────────────────
|
||||
|
||||
// Transforms a raw page into a compacted nucleotide stream in-place.
|
||||
//
|
||||
// Buffer layout on each call:
|
||||
// buf[0..overlap_len()] — overlap bytes copied by write_overlap()
|
||||
// buf[overlap_len()..overlap_len()+n] — raw bytes just read from the source
|
||||
//
|
||||
// Returns the number of output bytes in buf[0..returned].
|
||||
pub(crate) trait NucParser {
|
||||
// required: format-specific
|
||||
fn new(k: usize) -> Self
|
||||
where
|
||||
Self: Sized;
|
||||
fn overlap_state(&self) -> &OverlapState;
|
||||
fn overlap_state_mut(&mut self) -> &mut OverlapState;
|
||||
fn is_in_seq(&self) -> bool;
|
||||
fn parse_inplace(&mut self, buf: &mut [u8], n: usize) -> usize;
|
||||
|
||||
// provided: format-independent overlap management
|
||||
fn overlap_len(&self) -> usize {
|
||||
self.overlap_state().len
|
||||
}
|
||||
|
||||
fn write_overlap(&self, buf: &mut [u8]) {
|
||||
let ol = &self.overlap_state();
|
||||
buf[..ol.len].copy_from_slice(&ol.data[..ol.len]);
|
||||
}
|
||||
|
||||
// Called at end of parse_inplace: saves overlap state and returns adjusted j.
|
||||
// seq_start is the j-position where the last sequence started in this call's output.
|
||||
fn save_overlap(&mut self, buf: &mut [u8], j: usize, seq_start: usize) -> usize {
|
||||
if !self.is_in_seq() {
|
||||
self.overlap_state_mut().len = 0;
|
||||
return j;
|
||||
}
|
||||
let seq_len = j - seq_start;
|
||||
let k = self.overlap_state().k;
|
||||
if seq_len >= k {
|
||||
// Sequence long enough: save last k-1 nucleotides, terminate with 0.
|
||||
let ol = k - 1;
|
||||
self.overlap_state_mut().data[..ol].copy_from_slice(&buf[j - ol..j]);
|
||||
self.overlap_state_mut().len = ol;
|
||||
// SAFETY: j <= total - 1 < BUF_SIZE = buf.len()
|
||||
// (total = overlap_len + n <= (MAX_K-1) + PAGE_SIZE = BUF_SIZE - 1)
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j + 1
|
||||
} else if seq_len > 0 {
|
||||
// Short sequence (< k): save whole fragment, strip from output.
|
||||
self.overlap_state_mut().data[..seq_len].copy_from_slice(&buf[seq_start..j]);
|
||||
self.overlap_state_mut().len = seq_len;
|
||||
seq_start
|
||||
} else {
|
||||
self.overlap_state_mut().len = 0;
|
||||
j
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ─── FASTA parser ─────────────────────────────────────────────────────────────
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum FastaState {
|
||||
OutSeq,
|
||||
InTitle,
|
||||
InSeq,
|
||||
InAmbiguous,
|
||||
}
|
||||
|
||||
pub(crate) struct FastaParser {
|
||||
state: FastaState,
|
||||
overlap: OverlapState,
|
||||
}
|
||||
|
||||
impl NucParser for FastaParser {
|
||||
fn new(k: usize) -> Self {
|
||||
Self {
|
||||
state: FastaState::OutSeq,
|
||||
overlap: OverlapState::new(k),
|
||||
}
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state(&self) -> &OverlapState {
|
||||
&self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state_mut(&mut self) -> &mut OverlapState {
|
||||
&mut self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn is_in_seq(&self) -> bool {
|
||||
matches!(self.state, FastaState::InSeq)
|
||||
}
|
||||
|
||||
fn parse_inplace(&mut self, buf: &mut [u8], n: usize) -> usize {
|
||||
let total = self.overlap.len + n;
|
||||
let mut i = 0; // read index
|
||||
let mut j = 0; // write index (invariant: j <= i always)
|
||||
// j-position where the current sequence started in this call's output;
|
||||
// meaningful only when state is InSeq.
|
||||
let mut seq_start: usize = 0;
|
||||
|
||||
while i < total {
|
||||
// SAFETY: i < total <= BUF_SIZE = buf.len()
|
||||
let byte = unsafe { *buf.get_unchecked(i) };
|
||||
|
||||
match self.state {
|
||||
FastaState::OutSeq => {
|
||||
if byte == b'>' {
|
||||
self.state = FastaState::InTitle;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastaState::InTitle => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = FastaState::InSeq;
|
||||
seq_start = j;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastaState::InSeq => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
let nuc = byte & 0xDF; // to uppercase
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
i += 1;
|
||||
} else if byte == b'>' {
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = FastaState::InTitle;
|
||||
i += 1;
|
||||
} else {
|
||||
// first ambiguous base: end current sequence if non-empty
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = FastaState::InAmbiguous;
|
||||
i += 1;
|
||||
}
|
||||
}
|
||||
FastaState::InAmbiguous => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
if byte == b'>' {
|
||||
self.state = FastaState::InTitle;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
let nuc = byte & 0xDF;
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
seq_start = j;
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
self.state = FastaState::InSeq;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
self.save_overlap(buf, j, seq_start)
|
||||
}
|
||||
}
|
||||
|
||||
// ─── FASTQ parser ─────────────────────────────────────────────────────────────
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum FastqState {
|
||||
OutSeq,
|
||||
InTitle,
|
||||
InSeq,
|
||||
InAmbiguous,
|
||||
InQualTitle,
|
||||
InQual,
|
||||
}
|
||||
|
||||
pub(crate) struct FastqParser {
|
||||
state: FastqState,
|
||||
overlap: OverlapState,
|
||||
}
|
||||
|
||||
impl NucParser for FastqParser {
|
||||
fn new(k: usize) -> Self {
|
||||
Self {
|
||||
state: FastqState::OutSeq,
|
||||
overlap: OverlapState::new(k),
|
||||
}
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state(&self) -> &OverlapState {
|
||||
&self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state_mut(&mut self) -> &mut OverlapState {
|
||||
&mut self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn is_in_seq(&self) -> bool {
|
||||
matches!(self.state, FastqState::InSeq)
|
||||
}
|
||||
|
||||
fn parse_inplace(&mut self, buf: &mut [u8], n: usize) -> usize {
|
||||
let total = self.overlap.len + n;
|
||||
let mut i = 0;
|
||||
let mut j = 0;
|
||||
let mut seq_start: usize = 0;
|
||||
|
||||
while i < total {
|
||||
// SAFETY: i < total <= BUF_SIZE = buf.len()
|
||||
let byte = unsafe { *buf.get_unchecked(i) };
|
||||
|
||||
match self.state {
|
||||
FastqState::OutSeq => {
|
||||
if byte == b'@' {
|
||||
self.state = FastqState::InTitle;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastqState::InTitle => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = FastqState::InSeq;
|
||||
seq_start = j;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastqState::InSeq => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = FastqState::InQualTitle;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
let nuc = byte & 0xDF;
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
} else {
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = FastqState::InAmbiguous;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastqState::InAmbiguous => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = FastqState::InQualTitle;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
let nuc = byte & 0xDF;
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
seq_start = j;
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
self.state = FastqState::InSeq;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastqState::InQualTitle => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = FastqState::InQual;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastqState::InQual => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = FastqState::OutSeq;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
self.save_overlap(buf, j, seq_start)
|
||||
}
|
||||
}
|
||||
|
||||
// ─── GenBank parser ───────────────────────────────────────────────────────────
|
||||
|
||||
const ORIGIN_TAIL: &[u8] = b"RIGIN";
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum GenbankState {
|
||||
OutSeq,
|
||||
MatchOrigin,
|
||||
SkipOriginLine,
|
||||
InSeq,
|
||||
InSlash,
|
||||
InAmbiguous,
|
||||
}
|
||||
|
||||
pub(crate) struct GenbankParser {
|
||||
state: GenbankState,
|
||||
overlap: OverlapState,
|
||||
keyword_pos: usize,
|
||||
at_line_start: bool,
|
||||
}
|
||||
|
||||
impl NucParser for GenbankParser {
|
||||
fn new(k: usize) -> Self {
|
||||
Self {
|
||||
state: GenbankState::OutSeq,
|
||||
overlap: OverlapState::new(k),
|
||||
keyword_pos: 0,
|
||||
at_line_start: true,
|
||||
}
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state(&self) -> &OverlapState {
|
||||
&self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state_mut(&mut self) -> &mut OverlapState {
|
||||
&mut self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn is_in_seq(&self) -> bool {
|
||||
matches!(self.state, GenbankState::InSeq)
|
||||
}
|
||||
|
||||
fn parse_inplace(&mut self, buf: &mut [u8], n: usize) -> usize {
|
||||
let total = self.overlap.len + n;
|
||||
let mut i = 0;
|
||||
let mut j = 0;
|
||||
let mut seq_start: usize = 0;
|
||||
|
||||
while i < total {
|
||||
// SAFETY: i < total <= BUF_SIZE = buf.len()
|
||||
let byte = unsafe { *buf.get_unchecked(i) };
|
||||
|
||||
match self.state {
|
||||
GenbankState::OutSeq => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.at_line_start = true;
|
||||
} else if self.at_line_start && byte == b'O' {
|
||||
self.state = GenbankState::MatchOrigin;
|
||||
self.keyword_pos = 1;
|
||||
self.at_line_start = false;
|
||||
} else {
|
||||
self.at_line_start = false;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
GenbankState::MatchOrigin => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = GenbankState::OutSeq;
|
||||
self.at_line_start = true;
|
||||
} else if byte == ORIGIN_TAIL[self.keyword_pos - 1] {
|
||||
self.keyword_pos += 1;
|
||||
if self.keyword_pos == 6 {
|
||||
self.state = GenbankState::SkipOriginLine;
|
||||
}
|
||||
} else {
|
||||
self.state = GenbankState::OutSeq;
|
||||
self.at_line_start = false;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
GenbankState::SkipOriginLine => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = GenbankState::InSeq;
|
||||
seq_start = j;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
GenbankState::InSeq => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.at_line_start = true;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
if self.at_line_start && byte == b'/' {
|
||||
self.state = GenbankState::InSlash;
|
||||
self.at_line_start = false;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
self.at_line_start = false;
|
||||
let nuc = byte & 0xDF;
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
} else if byte.is_ascii_digit() || byte == b' ' {
|
||||
// position numbers and spacing between groups: skip
|
||||
} else {
|
||||
// ambiguous base: end current sequence if non-empty
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = GenbankState::InAmbiguous;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
GenbankState::InSlash => {
|
||||
if byte == b'/' {
|
||||
// confirmed "//": end of sequence record
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = GenbankState::OutSeq;
|
||||
self.at_line_start = false;
|
||||
} else if byte == b'\n' || byte == b'\r' {
|
||||
// single '/' line: back to sequence
|
||||
self.state = GenbankState::InSeq;
|
||||
self.at_line_start = true;
|
||||
} else {
|
||||
// false positive: single '/' mid-line, resume sequence
|
||||
self.state = GenbankState::InSeq;
|
||||
self.at_line_start = false;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
GenbankState::InAmbiguous => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.at_line_start = true;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
if self.at_line_start && byte == b'/' {
|
||||
self.state = GenbankState::InSlash;
|
||||
self.at_line_start = false;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
self.at_line_start = false;
|
||||
let nuc = byte & 0xDF;
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
seq_start = j;
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
self.state = GenbankState::InSeq;
|
||||
}
|
||||
// digits, spaces, other ambiguous codes: skip
|
||||
i += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
self.save_overlap(buf, j, seq_start)
|
||||
}
|
||||
}
|
||||
|
||||
// ─── NucPage ──────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Owned page of compacted nucleotides: uppercase A/C/G/T bytes separated by `0`
|
||||
/// at sequence boundaries. Automatically returns its buffer to the pool on drop.
|
||||
pub struct NucPage {
|
||||
data: ManuallyDrop<Vec<u8>>,
|
||||
len: usize,
|
||||
pool: Arc<Mutex<Vec<Vec<u8>>>>,
|
||||
}
|
||||
|
||||
impl std::ops::Deref for NucPage {
|
||||
type Target = [u8];
|
||||
fn deref(&self) -> &[u8] {
|
||||
&self.data[..self.len]
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for NucPage {
|
||||
fn drop(&mut self) {
|
||||
// SAFETY: data is never accessed after this point
|
||||
let buf = unsafe { ManuallyDrop::take(&mut self.data) };
|
||||
self.pool.lock().unwrap().push(buf);
|
||||
}
|
||||
}
|
||||
|
||||
// ─── NucPageCursor ────────────────────────────────────────────────────────────
|
||||
|
||||
/// A forward cursor over the normalised bytes of a [`NucPage`].
|
||||
///
|
||||
/// Provides the `next_byte` / `rewind` interface consumed by
|
||||
/// [`obiskbuilder::SuperKmerStreamIter`].
|
||||
pub struct NucPageCursor<'a> {
|
||||
data: &'a [u8],
|
||||
pos: usize,
|
||||
}
|
||||
|
||||
impl NucPageCursor<'_> {
|
||||
/// Returns the next byte in the page, or `None` at end.
|
||||
#[inline]
|
||||
pub fn next_byte(&mut self) -> Option<u8> {
|
||||
if self.pos < self.data.len() {
|
||||
let b = self.data[self.pos];
|
||||
self.pos += 1;
|
||||
Some(b)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
/// Steps the cursor back by `n` bytes.
|
||||
///
|
||||
/// The caller guarantees that the last `n` bytes were all `ACGT`
|
||||
/// (no `0x00` separators), so they are still in the page buffer.
|
||||
#[inline]
|
||||
pub fn rewind(&mut self, n: usize) {
|
||||
self.pos -= n;
|
||||
}
|
||||
|
||||
/// Total number of bytes in the underlying page.
|
||||
#[inline]
|
||||
pub fn len(&self) -> usize {
|
||||
self.data.len()
|
||||
}
|
||||
|
||||
/// Returns `true` if the page contains no bytes.
|
||||
#[inline]
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.data.is_empty()
|
||||
}
|
||||
}
|
||||
|
||||
impl NucPage {
|
||||
/// Creates a forward cursor positioned at the start of this page.
|
||||
pub fn cursor(&self) -> NucPageCursor<'_> {
|
||||
NucPageCursor { data: self, pos: 0 }
|
||||
}
|
||||
}
|
||||
|
||||
// ─── NucStream ────────────────────────────────────────────────────────────────
|
||||
|
||||
pub(crate) struct NucStream<R: Read, P: NucParser> {
|
||||
reader: R,
|
||||
parser: P,
|
||||
pool: Arc<Mutex<Vec<Vec<u8>>>>,
|
||||
eof: bool,
|
||||
}
|
||||
|
||||
impl<R: Read, P: NucParser> NucStream<R, P> {
|
||||
pub(crate) fn new(reader: R, k: usize) -> Self {
|
||||
Self {
|
||||
reader,
|
||||
parser: P::new(k),
|
||||
pool: Arc::new(Mutex::new(Vec::new())),
|
||||
eof: false,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn read_page(&mut self) -> Option<NucPage> {
|
||||
loop {
|
||||
if self.eof {
|
||||
return None;
|
||||
}
|
||||
// take a buffer from the pool, or allocate fresh if all are in-flight
|
||||
let mut buf = self
|
||||
.pool
|
||||
.lock()
|
||||
.unwrap()
|
||||
.pop()
|
||||
.unwrap_or_else(|| vec![0u8; BUF_SIZE]);
|
||||
|
||||
let ol = self.parser.overlap_len();
|
||||
self.parser.write_overlap(&mut buf[..ol]);
|
||||
let n = self.reader.read(&mut buf[ol..ol + PAGE_SIZE]).unwrap_or(0);
|
||||
if n == 0 {
|
||||
self.eof = true;
|
||||
if ol == 0 {
|
||||
self.pool.lock().unwrap().push(buf);
|
||||
return None;
|
||||
}
|
||||
}
|
||||
let out_len = self.parser.parse_inplace(&mut buf, n);
|
||||
if out_len > 0 {
|
||||
return Some(NucPage {
|
||||
data: ManuallyDrop::new(buf),
|
||||
len: out_len,
|
||||
pool: Arc::clone(&self.pool),
|
||||
});
|
||||
}
|
||||
// empty page (all headers/ambiguous): return buf to pool and loop
|
||||
self.pool.lock().unwrap().push(buf);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<R: Read, P: NucParser> Iterator for NucStream<R, P> {
|
||||
type Item = NucPage;
|
||||
fn next(&mut self) -> Option<NucPage> {
|
||||
self.read_page()
|
||||
}
|
||||
}
|
||||
|
||||
// ─── FastaNucStream ───────────────────────────────────────────────────────────
|
||||
|
||||
pub(crate) type FastaNucStream<R> = NucStream<R, FastaParser>;
|
||||
pub(crate) type FastqNucStream<R> = NucStream<R, FastqParser>;
|
||||
pub(crate) type GenbankNucStream<R> = NucStream<R, GenbankParser>;
|
||||
|
||||
// ─── AnyNucStream ─────────────────────────────────────────────────────────────
|
||||
|
||||
pub(crate) enum AnyNucStream<R: Read> {
|
||||
Fasta(FastaNucStream<R>),
|
||||
Fastq(FastqNucStream<R>),
|
||||
Genbank(GenbankNucStream<R>),
|
||||
}
|
||||
|
||||
impl<R: Read> Iterator for AnyNucStream<R> {
|
||||
type Item = NucPage;
|
||||
fn next(&mut self) -> Option<NucPage> {
|
||||
match self {
|
||||
AnyNucStream::Fasta(s) => s.next(),
|
||||
AnyNucStream::Fastq(s) => s.next(),
|
||||
AnyNucStream::Genbank(s) => s.next(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn dispatch<R: Read>(
|
||||
mut guesser: MimeTypeGuesser<R>,
|
||||
k: usize,
|
||||
) -> Option<AnyNucStream<MimeTypeGuesser<R>>> {
|
||||
match guesser.mime_type() {
|
||||
Some("text/fasta") => Some(AnyNucStream::Fasta(NucStream::new(guesser, k))),
|
||||
Some("text/fastq") => Some(AnyNucStream::Fastq(NucStream::new(guesser, k))),
|
||||
Some("text/gbff") => Some(AnyNucStream::Genbank(NucStream::new(guesser, k))),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Wraps an already-open reader in a nucleotide stream, detecting its format.
|
||||
/// Returns `None` if the format is not recognised.
|
||||
pub(crate) fn nuc_stream<R: Read>(reader: R, k: usize) -> Option<AnyNucStream<MimeTypeGuesser<R>>> {
|
||||
dispatch(MimeTypeGuesser::new(reader), k)
|
||||
}
|
||||
|
||||
/// Opens a nucleotide stream from any source (file path, URL, or `-` for stdin),
|
||||
/// with transparent decompression and automatic format detection.
|
||||
///
|
||||
/// # Errors
|
||||
/// Returns an `io::Error` if the source cannot be opened, decompression fails,
|
||||
/// or the format is not recognised.
|
||||
pub fn open_nuc_stream(
|
||||
source: &str,
|
||||
k: usize,
|
||||
) -> io::Result<Box<dyn Iterator<Item = NucPage> + Send>> {
|
||||
let reader = open_raw(source)?;
|
||||
nuc_stream(reader, k)
|
||||
.map(|s| Box::new(s) as Box<dyn Iterator<Item = NucPage> + Send>)
|
||||
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "unknown sequence format"))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "tests/nucstream.rs"]
|
||||
mod tests;
|
||||
@@ -0,0 +1,128 @@
|
||||
use super::overlap::{NucParser, OverlapState};
|
||||
|
||||
// ─── FASTA parser ─────────────────────────────────────────────────────────────
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum FastaState {
|
||||
OutSeq,
|
||||
InTitle,
|
||||
InSeq,
|
||||
InAmbiguous,
|
||||
}
|
||||
|
||||
pub(crate) struct FastaParser {
|
||||
state: FastaState,
|
||||
overlap: OverlapState,
|
||||
}
|
||||
|
||||
impl NucParser for FastaParser {
|
||||
fn new(k: usize) -> Self {
|
||||
Self {
|
||||
state: FastaState::OutSeq,
|
||||
overlap: OverlapState::new(k),
|
||||
}
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state(&self) -> &OverlapState {
|
||||
&self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state_mut(&mut self) -> &mut OverlapState {
|
||||
&mut self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn is_in_seq(&self) -> bool {
|
||||
matches!(self.state, FastaState::InSeq)
|
||||
}
|
||||
|
||||
fn parse_inplace(&mut self, buf: &mut [u8], n: usize) -> usize {
|
||||
let total = self.overlap.len + n;
|
||||
let mut i = 0; // read index
|
||||
let mut j = 0; // write index (invariant: j <= i always)
|
||||
// j-position where the current sequence started in this call's output;
|
||||
// meaningful only when state is InSeq.
|
||||
let mut seq_start: usize = 0;
|
||||
|
||||
while i < total {
|
||||
// SAFETY: i < total <= BUF_SIZE = buf.len()
|
||||
let byte = unsafe { *buf.get_unchecked(i) };
|
||||
|
||||
match self.state {
|
||||
FastaState::OutSeq => {
|
||||
if byte == b'>' {
|
||||
self.state = FastaState::InTitle;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastaState::InTitle => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = FastaState::InSeq;
|
||||
seq_start = j;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastaState::InSeq => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
let nuc = byte & 0xDF; // to uppercase
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
i += 1;
|
||||
} else if byte == b'>' {
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = FastaState::InTitle;
|
||||
i += 1;
|
||||
} else {
|
||||
// first ambiguous base: end current sequence if non-empty
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = FastaState::InAmbiguous;
|
||||
i += 1;
|
||||
}
|
||||
}
|
||||
FastaState::InAmbiguous => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
if byte == b'>' {
|
||||
self.state = FastaState::InTitle;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
let nuc = byte & 0xDF;
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
seq_start = j;
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
self.state = FastaState::InSeq;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
self.save_overlap(buf, j, seq_start)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,132 @@
|
||||
use super::overlap::{NucParser, OverlapState};
|
||||
|
||||
// ─── FASTQ parser ─────────────────────────────────────────────────────────────
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum FastqState {
|
||||
OutSeq,
|
||||
InTitle,
|
||||
InSeq,
|
||||
InAmbiguous,
|
||||
InQualTitle,
|
||||
InQual,
|
||||
}
|
||||
|
||||
pub(crate) struct FastqParser {
|
||||
state: FastqState,
|
||||
overlap: OverlapState,
|
||||
}
|
||||
|
||||
impl NucParser for FastqParser {
|
||||
fn new(k: usize) -> Self {
|
||||
Self {
|
||||
state: FastqState::OutSeq,
|
||||
overlap: OverlapState::new(k),
|
||||
}
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state(&self) -> &OverlapState {
|
||||
&self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state_mut(&mut self) -> &mut OverlapState {
|
||||
&mut self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn is_in_seq(&self) -> bool {
|
||||
matches!(self.state, FastqState::InSeq)
|
||||
}
|
||||
|
||||
fn parse_inplace(&mut self, buf: &mut [u8], n: usize) -> usize {
|
||||
let total = self.overlap.len + n;
|
||||
let mut i = 0;
|
||||
let mut j = 0;
|
||||
let mut seq_start: usize = 0;
|
||||
|
||||
while i < total {
|
||||
// SAFETY: i < total <= BUF_SIZE = buf.len()
|
||||
let byte = unsafe { *buf.get_unchecked(i) };
|
||||
|
||||
match self.state {
|
||||
FastqState::OutSeq => {
|
||||
if byte == b'@' {
|
||||
self.state = FastqState::InTitle;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastqState::InTitle => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = FastqState::InSeq;
|
||||
seq_start = j;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastqState::InSeq => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = FastqState::InQualTitle;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
let nuc = byte & 0xDF;
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
} else {
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = FastqState::InAmbiguous;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastqState::InAmbiguous => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = FastqState::InQualTitle;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
let nuc = byte & 0xDF;
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
seq_start = j;
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
self.state = FastqState::InSeq;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastqState::InQualTitle => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = FastqState::InQual;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
FastqState::InQual => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = FastqState::OutSeq;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
self.save_overlap(buf, j, seq_start)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,181 @@
|
||||
use super::overlap::{NucParser, OverlapState};
|
||||
|
||||
// ─── GenBank parser ───────────────────────────────────────────────────────────
|
||||
|
||||
const ORIGIN_TAIL: &[u8] = b"RIGIN";
|
||||
|
||||
#[derive(Clone, Copy)]
|
||||
enum GenbankState {
|
||||
OutSeq,
|
||||
MatchOrigin,
|
||||
SkipOriginLine,
|
||||
InSeq,
|
||||
InSlash,
|
||||
InAmbiguous,
|
||||
}
|
||||
|
||||
pub(crate) struct GenbankParser {
|
||||
state: GenbankState,
|
||||
overlap: OverlapState,
|
||||
keyword_pos: usize,
|
||||
at_line_start: bool,
|
||||
}
|
||||
|
||||
impl NucParser for GenbankParser {
|
||||
fn new(k: usize) -> Self {
|
||||
Self {
|
||||
state: GenbankState::OutSeq,
|
||||
overlap: OverlapState::new(k),
|
||||
keyword_pos: 0,
|
||||
at_line_start: true,
|
||||
}
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state(&self) -> &OverlapState {
|
||||
&self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn overlap_state_mut(&mut self) -> &mut OverlapState {
|
||||
&mut self.overlap
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn is_in_seq(&self) -> bool {
|
||||
matches!(self.state, GenbankState::InSeq)
|
||||
}
|
||||
|
||||
fn parse_inplace(&mut self, buf: &mut [u8], n: usize) -> usize {
|
||||
let total = self.overlap.len + n;
|
||||
let mut i = 0;
|
||||
let mut j = 0;
|
||||
let mut seq_start: usize = 0;
|
||||
|
||||
while i < total {
|
||||
// SAFETY: i < total <= BUF_SIZE = buf.len()
|
||||
let byte = unsafe { *buf.get_unchecked(i) };
|
||||
|
||||
match self.state {
|
||||
GenbankState::OutSeq => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.at_line_start = true;
|
||||
} else if self.at_line_start && byte == b'O' {
|
||||
self.state = GenbankState::MatchOrigin;
|
||||
self.keyword_pos = 1;
|
||||
self.at_line_start = false;
|
||||
} else {
|
||||
self.at_line_start = false;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
GenbankState::MatchOrigin => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = GenbankState::OutSeq;
|
||||
self.at_line_start = true;
|
||||
} else if byte == ORIGIN_TAIL[self.keyword_pos - 1] {
|
||||
self.keyword_pos += 1;
|
||||
if self.keyword_pos == 6 {
|
||||
self.state = GenbankState::SkipOriginLine;
|
||||
}
|
||||
} else {
|
||||
self.state = GenbankState::OutSeq;
|
||||
self.at_line_start = false;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
GenbankState::SkipOriginLine => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.state = GenbankState::InSeq;
|
||||
seq_start = j;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
GenbankState::InSeq => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.at_line_start = true;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
if self.at_line_start && byte == b'/' {
|
||||
self.state = GenbankState::InSlash;
|
||||
self.at_line_start = false;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
self.at_line_start = false;
|
||||
let nuc = byte & 0xDF;
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
} else if byte.is_ascii_digit() || byte == b' ' {
|
||||
// position numbers and spacing between groups: skip
|
||||
} else {
|
||||
// ambiguous base: end current sequence if non-empty
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = GenbankState::InAmbiguous;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
GenbankState::InSlash => {
|
||||
if byte == b'/' {
|
||||
// confirmed "//": end of sequence record
|
||||
if j > seq_start {
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j += 1;
|
||||
}
|
||||
self.state = GenbankState::OutSeq;
|
||||
self.at_line_start = false;
|
||||
} else if byte == b'\n' || byte == b'\r' {
|
||||
// single '/' line: back to sequence
|
||||
self.state = GenbankState::InSeq;
|
||||
self.at_line_start = true;
|
||||
} else {
|
||||
// false positive: single '/' mid-line, resume sequence
|
||||
self.state = GenbankState::InSeq;
|
||||
self.at_line_start = false;
|
||||
}
|
||||
i += 1;
|
||||
}
|
||||
GenbankState::InAmbiguous => {
|
||||
if byte == b'\n' || byte == b'\r' {
|
||||
self.at_line_start = true;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
if self.at_line_start && byte == b'/' {
|
||||
self.state = GenbankState::InSlash;
|
||||
self.at_line_start = false;
|
||||
i += 1;
|
||||
continue;
|
||||
}
|
||||
self.at_line_start = false;
|
||||
let nuc = byte & 0xDF;
|
||||
if nuc == b'A' || nuc == b'C' || nuc == b'G' || nuc == b'T' {
|
||||
seq_start = j;
|
||||
// SAFETY: j <= i < total <= BUF_SIZE = buf.len()
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = nuc;
|
||||
}
|
||||
j += 1;
|
||||
self.state = GenbankState::InSeq;
|
||||
}
|
||||
// digits, spaces, other ambiguous codes: skip
|
||||
i += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
self.save_overlap(buf, j, seq_start)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,40 @@
|
||||
//! Streaming, in-place normalisation of raw sequence bytes into compacted
|
||||
//! nucleotide pages: uppercase A/C/G/T separated by `0` at sequence
|
||||
//! boundaries, ready for k-mer extraction without re-scanning for case or
|
||||
//! ambiguity codes.
|
||||
//!
|
||||
//! Submodules: [`overlap`] (format-independent k-1 overlap bookkeeping
|
||||
//! shared by every parser), [`fasta`]/[`fastq`]/[`genbank`] (one
|
||||
//! format-specific in-place state machine each), [`page`] (the pooled
|
||||
//! output buffer, [`NucPage`]), [`stream`] (format dispatch and the public
|
||||
//! [`open_nuc_stream`] entry point).
|
||||
|
||||
mod fasta;
|
||||
mod fastq;
|
||||
mod genbank;
|
||||
mod overlap;
|
||||
mod page;
|
||||
mod stream;
|
||||
|
||||
pub use page::{NucPage, NucPageCursor};
|
||||
pub use stream::open_nuc_stream;
|
||||
|
||||
// Only used by `tests.rs` (`use super::*`) below — the crate itself always
|
||||
// reaches these through their defining submodule directly.
|
||||
#[cfg(test)]
|
||||
pub(crate) use fasta::FastaParser;
|
||||
#[cfg(test)]
|
||||
pub(crate) use fastq::FastqParser;
|
||||
#[cfg(test)]
|
||||
pub(crate) use genbank::GenbankParser;
|
||||
#[cfg(test)]
|
||||
pub(crate) use stream::NucStream;
|
||||
|
||||
pub(crate) const MAX_K: usize = 31;
|
||||
pub(crate) const PAGE_SIZE: usize = 65536;
|
||||
// overlap (MAX_K - 1) + page data (PAGE_SIZE) + 1 byte for the end-of-page terminating 0
|
||||
pub(crate) const BUF_SIZE: usize = MAX_K + PAGE_SIZE;
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "../tests/nucstream.rs"]
|
||||
mod tests;
|
||||
@@ -0,0 +1,81 @@
|
||||
use super::MAX_K;
|
||||
|
||||
// ─── OverlapState ─────────────────────────────────────────────────────────────
|
||||
|
||||
pub(crate) struct OverlapState {
|
||||
data: [u8; MAX_K],
|
||||
pub(super) len: usize,
|
||||
k: usize,
|
||||
}
|
||||
|
||||
impl OverlapState {
|
||||
pub(crate) fn new(k: usize) -> Self {
|
||||
assert!(k > 0 && k <= MAX_K);
|
||||
Self {
|
||||
data: [0u8; MAX_K],
|
||||
len: 0,
|
||||
k,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ─── NucParser trait ──────────────────────────────────────────────────────────
|
||||
|
||||
// Transforms a raw page into a compacted nucleotide stream in-place.
|
||||
//
|
||||
// Buffer layout on each call:
|
||||
// buf[0..overlap_len()] — overlap bytes copied by write_overlap()
|
||||
// buf[overlap_len()..overlap_len()+n] — raw bytes just read from the source
|
||||
//
|
||||
// Returns the number of output bytes in buf[0..returned].
|
||||
pub(crate) trait NucParser {
|
||||
// required: format-specific
|
||||
fn new(k: usize) -> Self
|
||||
where
|
||||
Self: Sized;
|
||||
fn overlap_state(&self) -> &OverlapState;
|
||||
fn overlap_state_mut(&mut self) -> &mut OverlapState;
|
||||
fn is_in_seq(&self) -> bool;
|
||||
fn parse_inplace(&mut self, buf: &mut [u8], n: usize) -> usize;
|
||||
|
||||
// provided: format-independent overlap management
|
||||
fn overlap_len(&self) -> usize {
|
||||
self.overlap_state().len
|
||||
}
|
||||
|
||||
fn write_overlap(&self, buf: &mut [u8]) {
|
||||
let ol = &self.overlap_state();
|
||||
buf[..ol.len].copy_from_slice(&ol.data[..ol.len]);
|
||||
}
|
||||
|
||||
// Called at end of parse_inplace: saves overlap state and returns adjusted j.
|
||||
// seq_start is the j-position where the last sequence started in this call's output.
|
||||
fn save_overlap(&mut self, buf: &mut [u8], j: usize, seq_start: usize) -> usize {
|
||||
if !self.is_in_seq() {
|
||||
self.overlap_state_mut().len = 0;
|
||||
return j;
|
||||
}
|
||||
let seq_len = j - seq_start;
|
||||
let k = self.overlap_state().k;
|
||||
if seq_len >= k {
|
||||
// Sequence long enough: save last k-1 nucleotides, terminate with 0.
|
||||
let ol = k - 1;
|
||||
self.overlap_state_mut().data[..ol].copy_from_slice(&buf[j - ol..j]);
|
||||
self.overlap_state_mut().len = ol;
|
||||
// SAFETY: j <= total - 1 < BUF_SIZE = buf.len()
|
||||
// (total = overlap_len + n <= (MAX_K-1) + PAGE_SIZE = BUF_SIZE - 1)
|
||||
unsafe {
|
||||
*buf.get_unchecked_mut(j) = 0;
|
||||
}
|
||||
j + 1
|
||||
} else if seq_len > 0 {
|
||||
// Short sequence (< k): save whole fragment, strip from output.
|
||||
self.overlap_state_mut().data[..seq_len].copy_from_slice(&buf[seq_start..j]);
|
||||
self.overlap_state_mut().len = seq_len;
|
||||
seq_start
|
||||
} else {
|
||||
self.overlap_state_mut().len = 0;
|
||||
j
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,80 @@
|
||||
use std::mem::ManuallyDrop;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
// ─── NucPage ──────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Owned page of compacted nucleotides: uppercase A/C/G/T bytes separated by `0`
|
||||
/// at sequence boundaries. Automatically returns its buffer to the pool on drop.
|
||||
pub struct NucPage {
|
||||
pub(super) data: ManuallyDrop<Vec<u8>>,
|
||||
pub(super) len: usize,
|
||||
pub(super) pool: Arc<Mutex<Vec<Vec<u8>>>>,
|
||||
}
|
||||
|
||||
impl std::ops::Deref for NucPage {
|
||||
type Target = [u8];
|
||||
fn deref(&self) -> &[u8] {
|
||||
&self.data[..self.len]
|
||||
}
|
||||
}
|
||||
|
||||
impl Drop for NucPage {
|
||||
fn drop(&mut self) {
|
||||
// SAFETY: data is never accessed after this point
|
||||
let buf = unsafe { ManuallyDrop::take(&mut self.data) };
|
||||
self.pool.lock().unwrap().push(buf);
|
||||
}
|
||||
}
|
||||
|
||||
// ─── NucPageCursor ────────────────────────────────────────────────────────────
|
||||
|
||||
/// A forward cursor over the normalised bytes of a [`NucPage`].
|
||||
///
|
||||
/// Provides the `next_byte` / `rewind` interface consumed by
|
||||
/// [`obiskbuilder::SuperKmerStreamIter`].
|
||||
pub struct NucPageCursor<'a> {
|
||||
data: &'a [u8],
|
||||
pos: usize,
|
||||
}
|
||||
|
||||
impl NucPageCursor<'_> {
|
||||
/// Returns the next byte in the page, or `None` at end.
|
||||
#[inline]
|
||||
pub fn next_byte(&mut self) -> Option<u8> {
|
||||
if self.pos < self.data.len() {
|
||||
let b = self.data[self.pos];
|
||||
self.pos += 1;
|
||||
Some(b)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
/// Steps the cursor back by `n` bytes.
|
||||
///
|
||||
/// The caller guarantees that the last `n` bytes were all `ACGT`
|
||||
/// (no `0x00` separators), so they are still in the page buffer.
|
||||
#[inline]
|
||||
pub fn rewind(&mut self, n: usize) {
|
||||
self.pos -= n;
|
||||
}
|
||||
|
||||
/// Total number of bytes in the underlying page.
|
||||
#[inline]
|
||||
pub fn len(&self) -> usize {
|
||||
self.data.len()
|
||||
}
|
||||
|
||||
/// Returns `true` if the page contains no bytes.
|
||||
#[inline]
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.data.is_empty()
|
||||
}
|
||||
}
|
||||
|
||||
impl NucPage {
|
||||
/// Creates a forward cursor positioned at the start of this page.
|
||||
pub fn cursor(&self) -> NucPageCursor<'_> {
|
||||
NucPageCursor { data: self, pos: 0 }
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,135 @@
|
||||
use std::io::{self, Read};
|
||||
use std::mem::ManuallyDrop;
|
||||
use std::sync::{Arc, Mutex};
|
||||
|
||||
use crate::mimetype::MimeTypeGuesser;
|
||||
use crate::xopen::open_raw;
|
||||
|
||||
use super::fasta::FastaParser;
|
||||
use super::fastq::FastqParser;
|
||||
use super::genbank::GenbankParser;
|
||||
use super::overlap::NucParser;
|
||||
use super::page::NucPage;
|
||||
use super::{BUF_SIZE, PAGE_SIZE};
|
||||
|
||||
// ─── NucStream ────────────────────────────────────────────────────────────────
|
||||
|
||||
pub(crate) struct NucStream<R: Read, P: NucParser> {
|
||||
reader: R,
|
||||
parser: P,
|
||||
pool: Arc<Mutex<Vec<Vec<u8>>>>,
|
||||
eof: bool,
|
||||
}
|
||||
|
||||
impl<R: Read, P: NucParser> NucStream<R, P> {
|
||||
pub(crate) fn new(reader: R, k: usize) -> Self {
|
||||
Self {
|
||||
reader,
|
||||
parser: P::new(k),
|
||||
pool: Arc::new(Mutex::new(Vec::new())),
|
||||
eof: false,
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn read_page(&mut self) -> Option<NucPage> {
|
||||
loop {
|
||||
if self.eof {
|
||||
return None;
|
||||
}
|
||||
// take a buffer from the pool, or allocate fresh if all are in-flight
|
||||
let mut buf = self
|
||||
.pool
|
||||
.lock()
|
||||
.unwrap()
|
||||
.pop()
|
||||
.unwrap_or_else(|| vec![0u8; BUF_SIZE]);
|
||||
|
||||
let ol = self.parser.overlap_len();
|
||||
self.parser.write_overlap(&mut buf[..ol]);
|
||||
let n = self.reader.read(&mut buf[ol..ol + PAGE_SIZE]).unwrap_or(0);
|
||||
if n == 0 {
|
||||
self.eof = true;
|
||||
if ol == 0 {
|
||||
self.pool.lock().unwrap().push(buf);
|
||||
return None;
|
||||
}
|
||||
}
|
||||
let out_len = self.parser.parse_inplace(&mut buf, n);
|
||||
if out_len > 0 {
|
||||
return Some(NucPage {
|
||||
data: ManuallyDrop::new(buf),
|
||||
len: out_len,
|
||||
pool: Arc::clone(&self.pool),
|
||||
});
|
||||
}
|
||||
// empty page (all headers/ambiguous): return buf to pool and loop
|
||||
self.pool.lock().unwrap().push(buf);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl<R: Read, P: NucParser> Iterator for NucStream<R, P> {
|
||||
type Item = NucPage;
|
||||
fn next(&mut self) -> Option<NucPage> {
|
||||
self.read_page()
|
||||
}
|
||||
}
|
||||
|
||||
// ─── FastaNucStream ───────────────────────────────────────────────────────────
|
||||
|
||||
pub(crate) type FastaNucStream<R> = NucStream<R, FastaParser>;
|
||||
pub(crate) type FastqNucStream<R> = NucStream<R, FastqParser>;
|
||||
pub(crate) type GenbankNucStream<R> = NucStream<R, GenbankParser>;
|
||||
|
||||
// ─── AnyNucStream ─────────────────────────────────────────────────────────────
|
||||
|
||||
pub(crate) enum AnyNucStream<R: Read> {
|
||||
Fasta(FastaNucStream<R>),
|
||||
Fastq(FastqNucStream<R>),
|
||||
Genbank(GenbankNucStream<R>),
|
||||
}
|
||||
|
||||
impl<R: Read> Iterator for AnyNucStream<R> {
|
||||
type Item = NucPage;
|
||||
fn next(&mut self) -> Option<NucPage> {
|
||||
match self {
|
||||
AnyNucStream::Fasta(s) => s.next(),
|
||||
AnyNucStream::Fastq(s) => s.next(),
|
||||
AnyNucStream::Genbank(s) => s.next(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn dispatch<R: Read>(
|
||||
mut guesser: MimeTypeGuesser<R>,
|
||||
k: usize,
|
||||
) -> Option<AnyNucStream<MimeTypeGuesser<R>>> {
|
||||
match guesser.mime_type() {
|
||||
Some("text/fasta") => Some(AnyNucStream::Fasta(NucStream::new(guesser, k))),
|
||||
Some("text/fastq") => Some(AnyNucStream::Fastq(NucStream::new(guesser, k))),
|
||||
Some("text/gbff") => Some(AnyNucStream::Genbank(NucStream::new(guesser, k))),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Wraps an already-open reader in a nucleotide stream, detecting its format.
|
||||
/// Returns `None` if the format is not recognised.
|
||||
pub(crate) fn nuc_stream<R: Read>(reader: R, k: usize) -> Option<AnyNucStream<MimeTypeGuesser<R>>> {
|
||||
dispatch(MimeTypeGuesser::new(reader), k)
|
||||
}
|
||||
|
||||
/// Opens a nucleotide stream from any source (file path, URL, or `-` for stdin),
|
||||
/// with transparent decompression and automatic format detection.
|
||||
///
|
||||
/// # Errors
|
||||
/// Returns an `io::Error` if the source cannot be opened, decompression fails,
|
||||
/// or the format is not recognised.
|
||||
pub fn open_nuc_stream(
|
||||
source: &str,
|
||||
k: usize,
|
||||
) -> io::Result<Box<dyn Iterator<Item = NucPage> + Send>> {
|
||||
let reader = open_raw(source)?;
|
||||
nuc_stream(reader, k)
|
||||
.map(|s| Box::new(s) as Box<dyn Iterator<Item = NucPage> + Send>)
|
||||
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "unknown sequence format"))
|
||||
}
|
||||
@@ -1,550 +0,0 @@
|
||||
use std::fs::File;
|
||||
use std::io::{BufWriter, Write as _};
|
||||
use std::path::{Path, PathBuf};
|
||||
use std::sync::Arc;
|
||||
|
||||
use memmap2::Mmap;
|
||||
use obikseq::{CanonicalKmer, Kmer, Unitig};
|
||||
|
||||
pub use obikseq::MAX_KMERS_PER_CHUNK;
|
||||
|
||||
use crate::error::{SKError, SKResult};
|
||||
|
||||
// ── Block index parameters ────────────────────────────────────────────────────
|
||||
//
|
||||
// BLOCK_SIZE = 1 << block_bits chunks share one offset entry in the index.
|
||||
// block_bits=0 → one entry per chunk (exact offsets, no scan).
|
||||
// block_bits=6 → one entry per 64 chunks (default; O(64) scan per lookup).
|
||||
//
|
||||
// block_bits is stored in the index file so the reader derives all parameters
|
||||
// at runtime — no compile-time constant constrains the format.
|
||||
|
||||
const MAGIC: [u8; 4] = *b"UIX3";
|
||||
|
||||
/// Default block granularity used by [`UnitigFileWriter::create`].
|
||||
pub const DEFAULT_BLOCK_BITS: u8 = 0;
|
||||
|
||||
fn idx_path(path: &Path) -> PathBuf {
|
||||
crate::append_path_suffix(path, ".idx")
|
||||
}
|
||||
|
||||
// ── Writer ────────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Writes a sequence of [`Unitig`] to an uncompressed binary file and builds
|
||||
/// a block-sampled offset index at close time.
|
||||
///
|
||||
/// One offset is stored every `1 << block_bits` chunks; random access to chunk
|
||||
/// `i` costs at most `(1 << block_bits) − 1` sequential chunk scans after the
|
||||
/// block lookup.
|
||||
///
|
||||
/// Unitigs with more than [`MAX_KMERS_PER_CHUNK`] k-mers are transparently split
|
||||
/// into overlapping chunks (k−1 nucleotide overlap) so no k-mer is lost.
|
||||
pub struct UnitigFileWriter {
|
||||
file: BufWriter<File>,
|
||||
block_offsets: Vec<u32>,
|
||||
chunk_count: usize,
|
||||
next_offset: u32,
|
||||
n_kmers: usize,
|
||||
k: usize,
|
||||
block_bits: u8,
|
||||
mask: usize, // (1 << block_bits) - 1
|
||||
}
|
||||
|
||||
impl UnitigFileWriter {
|
||||
/// Create a writer with the default block size (`DEFAULT_BLOCK_BITS = 6`).
|
||||
pub fn create(path: &Path) -> SKResult<Self> {
|
||||
Self::create_with_block_bits(path, DEFAULT_BLOCK_BITS)
|
||||
}
|
||||
|
||||
/// Create a writer with a custom block size.
|
||||
///
|
||||
/// `block_bits` must be in 0..=31. `block_bits=0` stores one offset per
|
||||
/// chunk (exact, no scan); larger values trade index size for scan length.
|
||||
pub fn create_with_block_bits(path: &Path, block_bits: u8) -> SKResult<Self> {
|
||||
assert!(block_bits <= 31, "block_bits must be ≤ 31");
|
||||
let file = File::create(path).map_err(SKError::Io)?;
|
||||
Ok(Self {
|
||||
file: BufWriter::new(file),
|
||||
block_offsets: Vec::new(),
|
||||
chunk_count: 0,
|
||||
next_offset: 0,
|
||||
n_kmers: 0,
|
||||
k: obikseq::params::k(),
|
||||
block_bits,
|
||||
mask: (1usize << block_bits) - 1,
|
||||
})
|
||||
}
|
||||
|
||||
/// Write a unitig, splitting into overlapping chunks if it exceeds
|
||||
/// [`MAX_KMERS_PER_CHUNK`].
|
||||
pub fn write(&mut self, unitig: &Unitig) -> SKResult<()> {
|
||||
let seql = unitig.seql();
|
||||
let k = self.k;
|
||||
|
||||
if seql < k {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let n_kmers = seql - k + 1;
|
||||
if n_kmers <= MAX_KMERS_PER_CHUNK {
|
||||
return self.write_chunk(unitig);
|
||||
}
|
||||
|
||||
let chunk_nucl = MAX_KMERS_PER_CHUNK + k - 1;
|
||||
let stride = MAX_KMERS_PER_CHUNK;
|
||||
let mut start = 0;
|
||||
while start < seql {
|
||||
let end = (start + chunk_nucl).min(seql);
|
||||
self.write_chunk(&unitig.sub(start, end))?;
|
||||
if end == seql { break; }
|
||||
start += stride;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn write_chunk(&mut self, unitig: &Unitig) -> SKResult<()> {
|
||||
let seql = unitig.seql();
|
||||
let byte_len = (seql + 3) / 4;
|
||||
|
||||
debug_assert!(seql - self.k <= u8::MAX as usize, "chunk exceeds MAX_KMERS_PER_CHUNK");
|
||||
|
||||
if self.chunk_count & self.mask == 0 {
|
||||
self.block_offsets.push(self.next_offset);
|
||||
}
|
||||
|
||||
self.n_kmers += seql - self.k + 1;
|
||||
self.chunk_count += 1;
|
||||
|
||||
unitig.write_to_binary(&mut self.file).map_err(SKError::Io)?;
|
||||
|
||||
self.next_offset += 1 + byte_len as u32;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Flush and close the binary sequence file.
|
||||
///
|
||||
/// The companion `.idx` file is **not** written here; call
|
||||
/// [`build_unitig_idx`] separately when exact evidence is needed.
|
||||
pub fn close(mut self) -> SKResult<()> {
|
||||
self.file.flush().map_err(SKError::Io)?;
|
||||
drop(self.file);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn len(&self) -> usize { self.chunk_count }
|
||||
pub fn is_empty(&self) -> bool { self.chunk_count == 0 }
|
||||
pub fn block_bits(&self) -> u8 { self.block_bits }
|
||||
}
|
||||
|
||||
fn write_idx(
|
||||
path: &Path,
|
||||
n_unitigs: u32,
|
||||
n_kmers: u64,
|
||||
block_bits: u8,
|
||||
block_offsets: &[u32],
|
||||
) -> SKResult<()> {
|
||||
let mut w = BufWriter::new(File::create(path).map_err(SKError::Io)?);
|
||||
w.write_all(&MAGIC).map_err(SKError::Io)?;
|
||||
w.write_all(&(block_bits as u32).to_le_bytes()).map_err(SKError::Io)?;
|
||||
w.write_all(&n_unitigs.to_le_bytes()).map_err(SKError::Io)?;
|
||||
w.write_all(&n_kmers.to_le_bytes()).map_err(SKError::Io)?;
|
||||
for &off in block_offsets {
|
||||
w.write_all(&off.to_le_bytes()).map_err(SKError::Io)?;
|
||||
}
|
||||
w.flush().map_err(SKError::Io)
|
||||
}
|
||||
|
||||
/// Scan an existing `unitigs.bin` file and write its companion `.idx`.
|
||||
///
|
||||
/// Called by the exact-evidence construction route after the sequence file is
|
||||
/// closed. `block_bits` controls index granularity (1 << block_bits chunks per
|
||||
/// offset entry); use [`DEFAULT_BLOCK_BITS`] for the default.
|
||||
pub fn build_unitig_idx(unitigs_path: &Path, block_bits: u8) -> SKResult<()> {
|
||||
assert!(block_bits <= 31, "block_bits must be ≤ 31");
|
||||
|
||||
let file = File::open(unitigs_path).map_err(SKError::Io)?;
|
||||
let mmap = unsafe { Mmap::map(&file).map_err(SKError::Io)? };
|
||||
|
||||
let k = obikseq::params::k();
|
||||
let block_size = 1usize << block_bits;
|
||||
let mask = block_size - 1;
|
||||
|
||||
let mut block_offsets: Vec<u32> = Vec::new();
|
||||
let mut offset = 0usize;
|
||||
let mut chunk_count = 0usize;
|
||||
let mut n_kmers = 0usize;
|
||||
|
||||
while offset < mmap.len() {
|
||||
if chunk_count & mask == 0 {
|
||||
block_offsets.push(offset as u32);
|
||||
}
|
||||
let seql_minus_k = mmap[offset] as usize;
|
||||
let byte_len = (seql_minus_k + k + 3) / 4;
|
||||
n_kmers += seql_minus_k + 1;
|
||||
offset += 1 + byte_len;
|
||||
chunk_count += 1;
|
||||
}
|
||||
|
||||
block_offsets.push(offset as u32); // sentinel
|
||||
|
||||
write_idx(
|
||||
&idx_path(unitigs_path),
|
||||
chunk_count as u32,
|
||||
n_kmers as u64,
|
||||
block_bits,
|
||||
&block_offsets,
|
||||
)
|
||||
}
|
||||
|
||||
// ── Reader ────────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Memory-mapped view of a unitig file, with optional direct-access index.
|
||||
///
|
||||
/// Three constructors select the operating mode:
|
||||
/// - [`open`](Self::open) — smart default: direct access if `.idx` exists, sequential otherwise.
|
||||
/// - [`open_sequential`](Self::open_sequential) — always sequential, ignores `.idx`.
|
||||
/// - [`open_direct_access`](Self::open_direct_access) — requires `.idx`, errors if absent.
|
||||
///
|
||||
/// All positional methods (`chunk_start`, `verify_canonical_kmer`, …) work in
|
||||
/// both modes. Without `.idx` they fall back to an O(i) sequential scan —
|
||||
/// correct but slower.
|
||||
pub struct UnitigFileReader {
|
||||
mmap: Mmap,
|
||||
block_offsets: Vec<u32>,
|
||||
n_unitigs: usize,
|
||||
n_kmers: usize,
|
||||
k: usize,
|
||||
block_bits: u8,
|
||||
mask: usize, // (1 << block_bits) - 1
|
||||
}
|
||||
|
||||
impl UnitigFileReader {
|
||||
/// Smart default: opens with direct access if `.idx` is present, sequential otherwise.
|
||||
pub fn open(path: &Path) -> SKResult<Self> {
|
||||
if idx_path(path).exists() {
|
||||
Self::open_direct_access(path)
|
||||
} else {
|
||||
Self::open_sequential(path)
|
||||
}
|
||||
}
|
||||
|
||||
/// Always sequential — never reads `.idx` even if present.
|
||||
///
|
||||
/// Scans the binary file once to count chunks and k-mers.
|
||||
/// Positional access (`chunk_start`, `verify_canonical_kmer`) falls back to
|
||||
/// O(i) sequential scan.
|
||||
pub fn open_sequential(path: &Path) -> SKResult<Self> {
|
||||
let file = File::open(path).map_err(SKError::Io)?;
|
||||
let mmap = unsafe { Mmap::map(&file).map_err(SKError::Io)? };
|
||||
let k = obikseq::params::k();
|
||||
|
||||
let mut offset = 0usize;
|
||||
let mut n_unitigs = 0usize;
|
||||
let mut n_kmers = 0usize;
|
||||
while offset < mmap.len() {
|
||||
let seql_minus_k = mmap[offset] as usize;
|
||||
n_kmers += seql_minus_k + 1;
|
||||
offset += 1 + (seql_minus_k + k + 3) / 4;
|
||||
n_unitigs += 1;
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
mmap,
|
||||
block_offsets: Vec::new(),
|
||||
n_unitigs,
|
||||
n_kmers,
|
||||
k,
|
||||
block_bits: DEFAULT_BLOCK_BITS,
|
||||
mask: (1usize << DEFAULT_BLOCK_BITS) - 1,
|
||||
})
|
||||
}
|
||||
|
||||
/// Requires `.idx` — errors if the companion index file is absent.
|
||||
///
|
||||
/// Enables O(1 << block_bits) positional access to any chunk.
|
||||
/// Use only when direct access is architecturally required (query-time
|
||||
/// verification on an exact-evidence layer).
|
||||
pub fn open_direct_access(path: &Path) -> SKResult<Self> {
|
||||
let file = File::open(path).map_err(SKError::Io)?;
|
||||
let mmap = unsafe { Mmap::map(&file).map_err(SKError::Io)? };
|
||||
let (n_unitigs, n_kmers, block_bits, block_offsets) = read_idx(&idx_path(path))?;
|
||||
let k = obikseq::params::k();
|
||||
Ok(Self {
|
||||
mmap,
|
||||
block_offsets,
|
||||
n_unitigs,
|
||||
n_kmers,
|
||||
k,
|
||||
block_bits,
|
||||
mask: (1usize << block_bits) - 1,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn len(&self) -> usize { self.n_unitigs }
|
||||
pub fn is_empty(&self) -> bool { self.n_unitigs == 0 }
|
||||
pub fn n_kmers(&self) -> usize { self.n_kmers }
|
||||
pub fn block_bits(&self) -> u8 { self.block_bits }
|
||||
pub fn has_direct_access(&self) -> bool { !self.block_offsets.is_empty() }
|
||||
|
||||
/// Byte offset of record `i` in the mmap.
|
||||
///
|
||||
/// Fast path (O(1 << block_bits)) when `.idx` is loaded; degraded O(i)
|
||||
/// sequential scan otherwise.
|
||||
#[inline]
|
||||
fn chunk_start(&self, i: usize) -> usize {
|
||||
if !self.block_offsets.is_empty() {
|
||||
if self.block_bits == 0 {
|
||||
return self.block_offsets[i] as usize;
|
||||
}
|
||||
let block = i >> self.block_bits;
|
||||
let rem = i & self.mask;
|
||||
let mut offset = self.block_offsets[block] as usize;
|
||||
for _ in 0..rem {
|
||||
let seql_minus_k = self.mmap[offset] as usize;
|
||||
offset += 1 + (seql_minus_k + self.k + 3) / 4;
|
||||
}
|
||||
offset
|
||||
} else {
|
||||
let mut offset = 0usize;
|
||||
for _ in 0..i {
|
||||
let seql_minus_k = self.mmap[offset] as usize;
|
||||
offset += 1 + (seql_minus_k + self.k + 3) / 4;
|
||||
}
|
||||
offset
|
||||
}
|
||||
}
|
||||
|
||||
/// Nucleotide length of chunk `i`.
|
||||
#[inline]
|
||||
pub fn seql(&self, i: usize) -> usize {
|
||||
self.mmap[self.chunk_start(i)] as usize + self.k
|
||||
}
|
||||
|
||||
/// Reconstruct chunk `i` as a [`Unitig`].
|
||||
pub fn unitig(&self, i: usize) -> Unitig {
|
||||
let offset = self.chunk_start(i);
|
||||
let seql = self.mmap[offset] as usize + self.k;
|
||||
let byte_len = (seql + 3) / 4;
|
||||
let bytes = self.mmap[offset + 1..offset + 1 + byte_len].to_vec().into_boxed_slice();
|
||||
Unitig::new((seql % 4) as u8, bytes)
|
||||
}
|
||||
|
||||
/// Raw left-aligned u64 of the k-mer at position `j` within chunk `i`.
|
||||
#[inline]
|
||||
pub fn raw_kmer(&self, i: usize, j: usize) -> u64 {
|
||||
let offset = self.chunk_start(i);
|
||||
extract_kmer_raw(&self.mmap[offset + 1..], j, self.k)
|
||||
}
|
||||
|
||||
/// `true` iff the k-mer at position `j` of chunk `i` matches `query`.
|
||||
///
|
||||
/// Works in both modes; O(i) scan when `.idx` is absent.
|
||||
#[inline]
|
||||
pub fn verify_canonical_kmer(&self, i: usize, j: usize, query: CanonicalKmer) -> bool {
|
||||
canonical_raw(self.raw_kmer(i, j), self.k) == query.raw()
|
||||
}
|
||||
|
||||
// ── Sequential iterators (O(n) running-offset cursor) ─────────────────────
|
||||
|
||||
fn iter_chunks_sequential(&self) -> impl Iterator<Item = (usize, Unitig)> + '_ {
|
||||
let k = self.k;
|
||||
let mmap = &*self.mmap;
|
||||
let n = self.n_unitigs;
|
||||
let mut offset = 0usize;
|
||||
(0..n).map(move |chunk_id| {
|
||||
let seql = mmap[offset] as usize + k;
|
||||
let byte_len = (seql + 3) / 4;
|
||||
let bytes = mmap[offset + 1..offset + 1 + byte_len].to_vec().into_boxed_slice();
|
||||
offset += 1 + byte_len;
|
||||
(chunk_id, Unitig::new((seql % 4) as u8, bytes))
|
||||
})
|
||||
}
|
||||
|
||||
/// Iterate all unitigs sequentially. Works without `.idx` (sequential open).
|
||||
pub fn iter_unitigs(&self) -> impl Iterator<Item = (usize, Unitig)> + '_ {
|
||||
self.iter_chunks_sequential()
|
||||
}
|
||||
|
||||
pub fn iter_kmers(&self) -> impl Iterator<Item = Kmer> + '_ {
|
||||
self.iter_chunks_sequential()
|
||||
.flat_map(|(_, u)| u.into_kmers())
|
||||
}
|
||||
|
||||
pub fn iter_indexed_canonical_kmers(
|
||||
&self,
|
||||
) -> impl Iterator<Item = (CanonicalKmer, usize, usize)> + '_ {
|
||||
self.iter_chunks_sequential()
|
||||
.flat_map(|(chunk_id, u)| {
|
||||
u.into_canonical_kmers()
|
||||
.enumerate()
|
||||
.map(move |(rank, kmer)| (kmer, chunk_id, rank))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn read_idx(path: &Path) -> SKResult<(usize, usize, u8, Vec<u32>)> {
|
||||
let data = std::fs::read(path).map_err(SKError::Io)?;
|
||||
let mut pos = 0;
|
||||
|
||||
let magic_bytes = data.get(pos..pos + 4)
|
||||
.ok_or(SKError::Truncated { context: "unitig index: magic" })?;
|
||||
if magic_bytes != &MAGIC {
|
||||
return Err(SKError::BadMagic {
|
||||
expected: "UIX3",
|
||||
got: magic_bytes.try_into().unwrap(),
|
||||
});
|
||||
}
|
||||
pos += 4;
|
||||
|
||||
let bb_bytes = data.get(pos..pos + 4)
|
||||
.ok_or(SKError::Truncated { context: "unitig index: block_bits" })?;
|
||||
let block_bits_u32 = u32::from_le_bytes(bb_bytes.try_into().unwrap());
|
||||
if block_bits_u32 > 31 {
|
||||
return Err(SKError::InvalidData {
|
||||
context: "unitig index",
|
||||
detail: format!("block_bits out of range: {block_bits_u32}"),
|
||||
});
|
||||
}
|
||||
let block_bits = block_bits_u32 as u8;
|
||||
pos += 4;
|
||||
|
||||
let n_bytes = data.get(pos..pos + 4)
|
||||
.ok_or(SKError::Truncated { context: "unitig index: n_unitigs" })?;
|
||||
let n_unitigs = u32::from_le_bytes(n_bytes.try_into().unwrap()) as usize;
|
||||
pos += 4;
|
||||
|
||||
let nk_bytes = data.get(pos..pos + 8)
|
||||
.ok_or(SKError::Truncated { context: "unitig index: n_kmers" })?;
|
||||
let n_kmers = u64::from_le_bytes(nk_bytes.try_into().unwrap()) as usize;
|
||||
pos += 8;
|
||||
|
||||
let block_size = 1usize << block_bits;
|
||||
let n_blocks = (n_unitigs + block_size - 1) >> block_bits;
|
||||
let n_offsets = n_blocks + 1;
|
||||
let mut block_offsets = Vec::with_capacity(n_offsets);
|
||||
for _ in 0..n_offsets {
|
||||
let off_bytes = data.get(pos..pos + 4)
|
||||
.ok_or(SKError::Truncated { context: "unitig index: block_offsets" })?;
|
||||
block_offsets.push(u32::from_le_bytes(off_bytes.try_into().unwrap()));
|
||||
pos += 4;
|
||||
}
|
||||
|
||||
Ok((n_unitigs, n_kmers, block_bits, block_offsets))
|
||||
}
|
||||
|
||||
// ── Kmer utilities ────────────────────────────────────────────────────────────
|
||||
|
||||
#[inline]
|
||||
fn revcomp_raw(raw: u64, k: usize) -> u64 {
|
||||
let x = !raw;
|
||||
let x = x.swap_bytes();
|
||||
let x = ((x >> 4) & 0x0F0F0F0F0F0F0F0F) | ((x & 0x0F0F0F0F0F0F0F0F) << 4);
|
||||
let x = ((x >> 2) & 0x3333333333333333) | ((x & 0x3333333333333333) << 2);
|
||||
x << (64 - 2 * k)
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn canonical_raw(raw: u64, k: usize) -> u64 {
|
||||
raw.min(revcomp_raw(raw, k))
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn extract_kmer_raw(bytes: &[u8], j: usize, k: usize) -> u64 {
|
||||
let bit_start = j * 2;
|
||||
let byte_start = bit_start / 8;
|
||||
let bit_offset = bit_start % 8;
|
||||
let bytes_needed = (bit_offset + 2 * k + 7) / 8;
|
||||
|
||||
let mut acc = 0u128;
|
||||
for idx in 0..bytes_needed {
|
||||
acc = (acc << 8) | bytes.get(byte_start + idx).copied().unwrap_or(0) as u128;
|
||||
}
|
||||
|
||||
let shift = bytes_needed * 8 - bit_offset - 2 * k;
|
||||
let mask = !0u64 >> (64 - 2 * k);
|
||||
let raw = (acc >> shift) as u64 & mask;
|
||||
raw << (64 - 2 * k)
|
||||
}
|
||||
|
||||
// ── CanonicalKmerRawIter ──────────────────────────────────────────────────────
|
||||
|
||||
// ── CanonicalKmerIter ─────────────────────────────────────────────────────────
|
||||
|
||||
/// Sequential iterator over [`CanonicalKmer`] from a `unitigs.bin` file.
|
||||
///
|
||||
/// Holds an `Arc<Mmap>` so that `Clone` is O(1): both copies share the same
|
||||
/// memory-mapped pages. Cloning resets the cursor to position 0 — this lets
|
||||
/// ptr_hash's `new_from_par_iter` (which requires a `Clone`-able parallel
|
||||
/// iterator via `par_bridge()`) make multiple passes without ever creating
|
||||
/// a `.idx` file.
|
||||
pub struct CanonicalKmerIter {
|
||||
mmap: Arc<Mmap>,
|
||||
k: usize,
|
||||
chunk_pos: usize, // byte offset of the current chunk header
|
||||
data_pos: usize, // byte offset of the current chunk's sequence bytes
|
||||
n_kmers: usize, // kmers in current chunk
|
||||
kmer_idx: usize, // next kmer index to yield within the current chunk
|
||||
}
|
||||
|
||||
impl CanonicalKmerIter {
|
||||
pub fn new(path: &Path) -> SKResult<Self> {
|
||||
let file = File::open(path).map_err(SKError::Io)?;
|
||||
let mmap = Arc::new(unsafe { Mmap::map(&file).map_err(SKError::Io)? });
|
||||
let k = obikseq::params::k();
|
||||
let mut s = Self { mmap, k, chunk_pos: 0, data_pos: 0, n_kmers: 0, kmer_idx: 0 };
|
||||
s.load_chunk();
|
||||
Ok(s)
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn load_chunk(&mut self) {
|
||||
if self.chunk_pos < self.mmap.len() {
|
||||
let seql_minus_k = self.mmap[self.chunk_pos] as usize;
|
||||
self.n_kmers = seql_minus_k + 1;
|
||||
self.data_pos = self.chunk_pos + 1;
|
||||
self.kmer_idx = 0;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Clone for CanonicalKmerIter {
|
||||
fn clone(&self) -> Self {
|
||||
let mut c = Self {
|
||||
mmap: Arc::clone(&self.mmap),
|
||||
k: self.k,
|
||||
chunk_pos: 0,
|
||||
data_pos: 0,
|
||||
n_kmers: 0,
|
||||
kmer_idx: 0,
|
||||
};
|
||||
c.load_chunk();
|
||||
c
|
||||
}
|
||||
}
|
||||
|
||||
impl Iterator for CanonicalKmerIter {
|
||||
type Item = CanonicalKmer;
|
||||
|
||||
#[inline]
|
||||
fn next(&mut self) -> Option<CanonicalKmer> {
|
||||
loop {
|
||||
if self.chunk_pos >= self.mmap.len() {
|
||||
return None;
|
||||
}
|
||||
if self.kmer_idx < self.n_kmers {
|
||||
let raw = extract_kmer_raw(&self.mmap[self.data_pos..], self.kmer_idx, self.k);
|
||||
let canon = canonical_raw(raw, self.k);
|
||||
self.kmer_idx += 1;
|
||||
return Some(CanonicalKmer::from_raw_unchecked(canon));
|
||||
}
|
||||
let seql_minus_k = self.mmap[self.chunk_pos] as usize;
|
||||
let byte_len = (seql_minus_k + self.k + 3) / 4;
|
||||
self.chunk_pos += 1 + byte_len;
|
||||
self.load_chunk();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "tests/unitig_index.rs"]
|
||||
mod tests;
|
||||
@@ -0,0 +1,86 @@
|
||||
use std::path::Path;
|
||||
use std::sync::Arc;
|
||||
|
||||
use memmap2::Mmap;
|
||||
use obikseq::CanonicalKmer;
|
||||
|
||||
use crate::error::{SKError, SKResult};
|
||||
|
||||
use super::kmer_raw::{canonical_raw, extract_kmer_raw};
|
||||
|
||||
// ── CanonicalKmerIter ─────────────────────────────────────────────────────────
|
||||
|
||||
/// Sequential iterator over [`CanonicalKmer`] from a `unitigs.bin` file.
|
||||
///
|
||||
/// Holds an `Arc<Mmap>` so that `Clone` is O(1): both copies share the same
|
||||
/// memory-mapped pages. Cloning resets the cursor to position 0 — this lets
|
||||
/// ptr_hash's `new_from_par_iter` (which requires a `Clone`-able parallel
|
||||
/// iterator via `par_bridge()`) make multiple passes without ever creating
|
||||
/// a `.idx` file.
|
||||
pub struct CanonicalKmerIter {
|
||||
mmap: Arc<Mmap>,
|
||||
k: usize,
|
||||
chunk_pos: usize, // byte offset of the current chunk header
|
||||
data_pos: usize, // byte offset of the current chunk's sequence bytes
|
||||
n_kmers: usize, // kmers in current chunk
|
||||
kmer_idx: usize, // next kmer index to yield within the current chunk
|
||||
}
|
||||
|
||||
impl CanonicalKmerIter {
|
||||
pub fn new(path: &Path) -> SKResult<Self> {
|
||||
let file = std::fs::File::open(path).map_err(SKError::Io)?;
|
||||
let mmap = Arc::new(unsafe { Mmap::map(&file).map_err(SKError::Io)? });
|
||||
let k = obikseq::params::k();
|
||||
let mut s = Self { mmap, k, chunk_pos: 0, data_pos: 0, n_kmers: 0, kmer_idx: 0 };
|
||||
s.load_chunk();
|
||||
Ok(s)
|
||||
}
|
||||
|
||||
#[inline]
|
||||
fn load_chunk(&mut self) {
|
||||
if self.chunk_pos < self.mmap.len() {
|
||||
let seql_minus_k = self.mmap[self.chunk_pos] as usize;
|
||||
self.n_kmers = seql_minus_k + 1;
|
||||
self.data_pos = self.chunk_pos + 1;
|
||||
self.kmer_idx = 0;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl Clone for CanonicalKmerIter {
|
||||
fn clone(&self) -> Self {
|
||||
let mut c = Self {
|
||||
mmap: Arc::clone(&self.mmap),
|
||||
k: self.k,
|
||||
chunk_pos: 0,
|
||||
data_pos: 0,
|
||||
n_kmers: 0,
|
||||
kmer_idx: 0,
|
||||
};
|
||||
c.load_chunk();
|
||||
c
|
||||
}
|
||||
}
|
||||
|
||||
impl Iterator for CanonicalKmerIter {
|
||||
type Item = CanonicalKmer;
|
||||
|
||||
#[inline]
|
||||
fn next(&mut self) -> Option<CanonicalKmer> {
|
||||
loop {
|
||||
if self.chunk_pos >= self.mmap.len() {
|
||||
return None;
|
||||
}
|
||||
if self.kmer_idx < self.n_kmers {
|
||||
let raw = extract_kmer_raw(&self.mmap[self.data_pos..], self.kmer_idx, self.k);
|
||||
let canon = canonical_raw(raw, self.k);
|
||||
self.kmer_idx += 1;
|
||||
return Some(CanonicalKmer::from_raw_unchecked(canon));
|
||||
}
|
||||
let seql_minus_k = self.mmap[self.chunk_pos] as usize;
|
||||
let byte_len = (seql_minus_k + self.k + 3) / 4;
|
||||
self.chunk_pos += 1 + byte_len;
|
||||
self.load_chunk();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,33 @@
|
||||
// ── Kmer utilities ────────────────────────────────────────────────────────────
|
||||
|
||||
#[inline]
|
||||
pub(super) fn revcomp_raw(raw: u64, k: usize) -> u64 {
|
||||
let x = !raw;
|
||||
let x = x.swap_bytes();
|
||||
let x = ((x >> 4) & 0x0F0F0F0F0F0F0F0F) | ((x & 0x0F0F0F0F0F0F0F0F) << 4);
|
||||
let x = ((x >> 2) & 0x3333333333333333) | ((x & 0x3333333333333333) << 2);
|
||||
x << (64 - 2 * k)
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub(super) fn canonical_raw(raw: u64, k: usize) -> u64 {
|
||||
raw.min(revcomp_raw(raw, k))
|
||||
}
|
||||
|
||||
#[inline]
|
||||
pub(super) fn extract_kmer_raw(bytes: &[u8], j: usize, k: usize) -> u64 {
|
||||
let bit_start = j * 2;
|
||||
let byte_start = bit_start / 8;
|
||||
let bit_offset = bit_start % 8;
|
||||
let bytes_needed = (bit_offset + 2 * k + 7) / 8;
|
||||
|
||||
let mut acc = 0u128;
|
||||
for idx in 0..bytes_needed {
|
||||
acc = (acc << 8) | bytes.get(byte_start + idx).copied().unwrap_or(0) as u128;
|
||||
}
|
||||
|
||||
let shift = bytes_needed * 8 - bit_offset - 2 * k;
|
||||
let mask = !0u64 >> (64 - 2 * k);
|
||||
let raw = (acc >> shift) as u64 & mask;
|
||||
raw << (64 - 2 * k)
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
//! Binary unitig storage: an append-only sequence file (`unitigs.bin`) plus
|
||||
//! an optional block-sampled `.idx` for O(1 << block_bits) random access.
|
||||
//!
|
||||
//! Submodules: [`writer`] ([`UnitigFileWriter`], chunk splitting, `.idx`
|
||||
//! construction), [`reader`] ([`UnitigFileReader`], sequential or
|
||||
//! direct-access), [`kmer_raw`] (packed 2-bit k-mer extraction/canonicalisation,
|
||||
//! shared by both the reader and [`CanonicalKmerIter`]), [`kmer_iter`]
|
||||
//! (cheaply-cloneable sequential k-mer iterator for `ptr_hash`).
|
||||
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
mod kmer_iter;
|
||||
mod kmer_raw;
|
||||
mod reader;
|
||||
mod writer;
|
||||
|
||||
pub use kmer_iter::CanonicalKmerIter;
|
||||
pub use obikseq::MAX_KMERS_PER_CHUNK;
|
||||
pub use reader::UnitigFileReader;
|
||||
pub use writer::{UnitigFileWriter, build_unitig_idx};
|
||||
|
||||
// Only used by `tests/unitig_index.rs` (`use super::*`) below.
|
||||
#[cfg(test)]
|
||||
use obikseq::CanonicalKmer;
|
||||
#[cfg(test)]
|
||||
use kmer_raw::{canonical_raw, extract_kmer_raw, revcomp_raw};
|
||||
|
||||
// ── Block index parameters ────────────────────────────────────────────────────
|
||||
//
|
||||
// BLOCK_SIZE = 1 << block_bits chunks share one offset entry in the index.
|
||||
// block_bits=0 → one entry per chunk (exact offsets, no scan).
|
||||
// block_bits=6 → one entry per 64 chunks (default; O(64) scan per lookup).
|
||||
//
|
||||
// block_bits is stored in the index file so the reader derives all parameters
|
||||
// at runtime — no compile-time constant constrains the format.
|
||||
|
||||
const MAGIC: [u8; 4] = *b"UIX3";
|
||||
|
||||
/// Default block granularity used by [`UnitigFileWriter::create`].
|
||||
pub const DEFAULT_BLOCK_BITS: u8 = 0;
|
||||
|
||||
fn idx_path(path: &Path) -> PathBuf {
|
||||
crate::append_path_suffix(path, ".idx")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
#[path = "../tests/unitig_index.rs"]
|
||||
mod tests;
|
||||
@@ -0,0 +1,246 @@
|
||||
use std::fs::File;
|
||||
use std::path::Path;
|
||||
|
||||
use memmap2::Mmap;
|
||||
use obikseq::{CanonicalKmer, Kmer, Unitig};
|
||||
|
||||
use crate::error::{SKError, SKResult};
|
||||
|
||||
use super::kmer_raw::{canonical_raw, extract_kmer_raw};
|
||||
use super::{DEFAULT_BLOCK_BITS, MAGIC, idx_path};
|
||||
|
||||
// ── Reader ────────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Memory-mapped view of a unitig file, with optional direct-access index.
|
||||
///
|
||||
/// Three constructors select the operating mode:
|
||||
/// - [`open`](Self::open) — smart default: direct access if `.idx` exists, sequential otherwise.
|
||||
/// - [`open_sequential`](Self::open_sequential) — always sequential, ignores `.idx`.
|
||||
/// - [`open_direct_access`](Self::open_direct_access) — requires `.idx`, errors if absent.
|
||||
///
|
||||
/// All positional methods (`chunk_start`, `verify_canonical_kmer`, …) work in
|
||||
/// both modes. Without `.idx` they fall back to an O(i) sequential scan —
|
||||
/// correct but slower.
|
||||
pub struct UnitigFileReader {
|
||||
mmap: Mmap,
|
||||
block_offsets: Vec<u32>,
|
||||
n_unitigs: usize,
|
||||
n_kmers: usize,
|
||||
k: usize,
|
||||
block_bits: u8,
|
||||
mask: usize, // (1 << block_bits) - 1
|
||||
}
|
||||
|
||||
impl UnitigFileReader {
|
||||
/// Smart default: opens with direct access if `.idx` is present, sequential otherwise.
|
||||
pub fn open(path: &Path) -> SKResult<Self> {
|
||||
if idx_path(path).exists() {
|
||||
Self::open_direct_access(path)
|
||||
} else {
|
||||
Self::open_sequential(path)
|
||||
}
|
||||
}
|
||||
|
||||
/// Always sequential — never reads `.idx` even if present.
|
||||
///
|
||||
/// Scans the binary file once to count chunks and k-mers.
|
||||
/// Positional access (`chunk_start`, `verify_canonical_kmer`) falls back to
|
||||
/// O(i) sequential scan.
|
||||
pub fn open_sequential(path: &Path) -> SKResult<Self> {
|
||||
let file = File::open(path).map_err(SKError::Io)?;
|
||||
let mmap = unsafe { Mmap::map(&file).map_err(SKError::Io)? };
|
||||
let k = obikseq::params::k();
|
||||
|
||||
let mut offset = 0usize;
|
||||
let mut n_unitigs = 0usize;
|
||||
let mut n_kmers = 0usize;
|
||||
while offset < mmap.len() {
|
||||
let seql_minus_k = mmap[offset] as usize;
|
||||
n_kmers += seql_minus_k + 1;
|
||||
offset += 1 + (seql_minus_k + k + 3) / 4;
|
||||
n_unitigs += 1;
|
||||
}
|
||||
|
||||
Ok(Self {
|
||||
mmap,
|
||||
block_offsets: Vec::new(),
|
||||
n_unitigs,
|
||||
n_kmers,
|
||||
k,
|
||||
block_bits: DEFAULT_BLOCK_BITS,
|
||||
mask: (1usize << DEFAULT_BLOCK_BITS) - 1,
|
||||
})
|
||||
}
|
||||
|
||||
/// Requires `.idx` — errors if the companion index file is absent.
|
||||
///
|
||||
/// Enables O(1 << block_bits) positional access to any chunk.
|
||||
/// Use only when direct access is architecturally required (query-time
|
||||
/// verification on an exact-evidence layer).
|
||||
pub fn open_direct_access(path: &Path) -> SKResult<Self> {
|
||||
let file = File::open(path).map_err(SKError::Io)?;
|
||||
let mmap = unsafe { Mmap::map(&file).map_err(SKError::Io)? };
|
||||
let (n_unitigs, n_kmers, block_bits, block_offsets) = read_idx(&idx_path(path))?;
|
||||
let k = obikseq::params::k();
|
||||
Ok(Self {
|
||||
mmap,
|
||||
block_offsets,
|
||||
n_unitigs,
|
||||
n_kmers,
|
||||
k,
|
||||
block_bits,
|
||||
mask: (1usize << block_bits) - 1,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn len(&self) -> usize { self.n_unitigs }
|
||||
pub fn is_empty(&self) -> bool { self.n_unitigs == 0 }
|
||||
pub fn n_kmers(&self) -> usize { self.n_kmers }
|
||||
pub fn block_bits(&self) -> u8 { self.block_bits }
|
||||
pub fn has_direct_access(&self) -> bool { !self.block_offsets.is_empty() }
|
||||
|
||||
/// Byte offset of record `i` in the mmap.
|
||||
///
|
||||
/// Fast path (O(1 << block_bits)) when `.idx` is loaded; degraded O(i)
|
||||
/// sequential scan otherwise.
|
||||
#[inline]
|
||||
fn chunk_start(&self, i: usize) -> usize {
|
||||
if !self.block_offsets.is_empty() {
|
||||
if self.block_bits == 0 {
|
||||
return self.block_offsets[i] as usize;
|
||||
}
|
||||
let block = i >> self.block_bits;
|
||||
let rem = i & self.mask;
|
||||
let mut offset = self.block_offsets[block] as usize;
|
||||
for _ in 0..rem {
|
||||
let seql_minus_k = self.mmap[offset] as usize;
|
||||
offset += 1 + (seql_minus_k + self.k + 3) / 4;
|
||||
}
|
||||
offset
|
||||
} else {
|
||||
let mut offset = 0usize;
|
||||
for _ in 0..i {
|
||||
let seql_minus_k = self.mmap[offset] as usize;
|
||||
offset += 1 + (seql_minus_k + self.k + 3) / 4;
|
||||
}
|
||||
offset
|
||||
}
|
||||
}
|
||||
|
||||
/// Nucleotide length of chunk `i`.
|
||||
#[inline]
|
||||
pub fn seql(&self, i: usize) -> usize {
|
||||
self.mmap[self.chunk_start(i)] as usize + self.k
|
||||
}
|
||||
|
||||
/// Reconstruct chunk `i` as a [`Unitig`].
|
||||
pub fn unitig(&self, i: usize) -> Unitig {
|
||||
let offset = self.chunk_start(i);
|
||||
let seql = self.mmap[offset] as usize + self.k;
|
||||
let byte_len = (seql + 3) / 4;
|
||||
let bytes = self.mmap[offset + 1..offset + 1 + byte_len].to_vec().into_boxed_slice();
|
||||
Unitig::new((seql % 4) as u8, bytes)
|
||||
}
|
||||
|
||||
/// Raw left-aligned u64 of the k-mer at position `j` within chunk `i`.
|
||||
#[inline]
|
||||
pub fn raw_kmer(&self, i: usize, j: usize) -> u64 {
|
||||
let offset = self.chunk_start(i);
|
||||
extract_kmer_raw(&self.mmap[offset + 1..], j, self.k)
|
||||
}
|
||||
|
||||
/// `true` iff the k-mer at position `j` of chunk `i` matches `query`.
|
||||
///
|
||||
/// Works in both modes; O(i) scan when `.idx` is absent.
|
||||
#[inline]
|
||||
pub fn verify_canonical_kmer(&self, i: usize, j: usize, query: CanonicalKmer) -> bool {
|
||||
canonical_raw(self.raw_kmer(i, j), self.k) == query.raw()
|
||||
}
|
||||
|
||||
// ── Sequential iterators (O(n) running-offset cursor) ─────────────────────
|
||||
|
||||
pub(crate) fn iter_chunks_sequential(&self) -> impl Iterator<Item = (usize, Unitig)> + '_ {
|
||||
let k = self.k;
|
||||
let mmap = &*self.mmap;
|
||||
let n = self.n_unitigs;
|
||||
let mut offset = 0usize;
|
||||
(0..n).map(move |chunk_id| {
|
||||
let seql = mmap[offset] as usize + k;
|
||||
let byte_len = (seql + 3) / 4;
|
||||
let bytes = mmap[offset + 1..offset + 1 + byte_len].to_vec().into_boxed_slice();
|
||||
offset += 1 + byte_len;
|
||||
(chunk_id, Unitig::new((seql % 4) as u8, bytes))
|
||||
})
|
||||
}
|
||||
|
||||
/// Iterate all unitigs sequentially. Works without `.idx` (sequential open).
|
||||
pub fn iter_unitigs(&self) -> impl Iterator<Item = (usize, Unitig)> + '_ {
|
||||
self.iter_chunks_sequential()
|
||||
}
|
||||
|
||||
pub fn iter_kmers(&self) -> impl Iterator<Item = Kmer> + '_ {
|
||||
self.iter_chunks_sequential()
|
||||
.flat_map(|(_, u)| u.into_kmers())
|
||||
}
|
||||
|
||||
pub fn iter_indexed_canonical_kmers(
|
||||
&self,
|
||||
) -> impl Iterator<Item = (CanonicalKmer, usize, usize)> + '_ {
|
||||
self.iter_chunks_sequential()
|
||||
.flat_map(|(chunk_id, u)| {
|
||||
u.into_canonical_kmers()
|
||||
.enumerate()
|
||||
.map(move |(rank, kmer)| (kmer, chunk_id, rank))
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn read_idx(path: &Path) -> SKResult<(usize, usize, u8, Vec<u32>)> {
|
||||
let data = std::fs::read(path).map_err(SKError::Io)?;
|
||||
let mut pos = 0;
|
||||
|
||||
let magic_bytes = data.get(pos..pos + 4)
|
||||
.ok_or(SKError::Truncated { context: "unitig index: magic" })?;
|
||||
if magic_bytes != &MAGIC {
|
||||
return Err(SKError::BadMagic {
|
||||
expected: "UIX3",
|
||||
got: magic_bytes.try_into().unwrap(),
|
||||
});
|
||||
}
|
||||
pos += 4;
|
||||
|
||||
let bb_bytes = data.get(pos..pos + 4)
|
||||
.ok_or(SKError::Truncated { context: "unitig index: block_bits" })?;
|
||||
let block_bits_u32 = u32::from_le_bytes(bb_bytes.try_into().unwrap());
|
||||
if block_bits_u32 > 31 {
|
||||
return Err(SKError::InvalidData {
|
||||
context: "unitig index",
|
||||
detail: format!("block_bits out of range: {block_bits_u32}"),
|
||||
});
|
||||
}
|
||||
let block_bits = block_bits_u32 as u8;
|
||||
pos += 4;
|
||||
|
||||
let n_bytes = data.get(pos..pos + 4)
|
||||
.ok_or(SKError::Truncated { context: "unitig index: n_unitigs" })?;
|
||||
let n_unitigs = u32::from_le_bytes(n_bytes.try_into().unwrap()) as usize;
|
||||
pos += 4;
|
||||
|
||||
let nk_bytes = data.get(pos..pos + 8)
|
||||
.ok_or(SKError::Truncated { context: "unitig index: n_kmers" })?;
|
||||
let n_kmers = u64::from_le_bytes(nk_bytes.try_into().unwrap()) as usize;
|
||||
pos += 8;
|
||||
|
||||
let block_size = 1usize << block_bits;
|
||||
let n_blocks = (n_unitigs + block_size - 1) >> block_bits;
|
||||
let n_offsets = n_blocks + 1;
|
||||
let mut block_offsets = Vec::with_capacity(n_offsets);
|
||||
for _ in 0..n_offsets {
|
||||
let off_bytes = data.get(pos..pos + 4)
|
||||
.ok_or(SKError::Truncated { context: "unitig index: block_offsets" })?;
|
||||
block_offsets.push(u32::from_le_bytes(off_bytes.try_into().unwrap()));
|
||||
pos += 4;
|
||||
}
|
||||
|
||||
Ok((n_unitigs, n_kmers, block_bits, block_offsets))
|
||||
}
|
||||
@@ -0,0 +1,178 @@
|
||||
use std::fs::File;
|
||||
use std::io::{BufWriter, Write as _};
|
||||
use std::path::Path;
|
||||
|
||||
use memmap2::Mmap;
|
||||
use obikseq::{MAX_KMERS_PER_CHUNK, Unitig};
|
||||
|
||||
use crate::error::{SKError, SKResult};
|
||||
|
||||
use super::{DEFAULT_BLOCK_BITS, MAGIC, idx_path};
|
||||
|
||||
// ── Writer ────────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Writes a sequence of [`Unitig`] to an uncompressed binary file and builds
|
||||
/// a block-sampled offset index at close time.
|
||||
///
|
||||
/// One offset is stored every `1 << block_bits` chunks; random access to chunk
|
||||
/// `i` costs at most `(1 << block_bits) − 1` sequential chunk scans after the
|
||||
/// block lookup.
|
||||
///
|
||||
/// Unitigs with more than [`MAX_KMERS_PER_CHUNK`] k-mers are transparently split
|
||||
/// into overlapping chunks (k−1 nucleotide overlap) so no k-mer is lost.
|
||||
pub struct UnitigFileWriter {
|
||||
file: BufWriter<File>,
|
||||
block_offsets: Vec<u32>,
|
||||
chunk_count: usize,
|
||||
next_offset: u32,
|
||||
n_kmers: usize,
|
||||
k: usize,
|
||||
block_bits: u8,
|
||||
mask: usize, // (1 << block_bits) - 1
|
||||
}
|
||||
|
||||
impl UnitigFileWriter {
|
||||
/// Create a writer with the default block size (`DEFAULT_BLOCK_BITS = 6`).
|
||||
pub fn create(path: &Path) -> SKResult<Self> {
|
||||
Self::create_with_block_bits(path, DEFAULT_BLOCK_BITS)
|
||||
}
|
||||
|
||||
/// Create a writer with a custom block size.
|
||||
///
|
||||
/// `block_bits` must be in 0..=31. `block_bits=0` stores one offset per
|
||||
/// chunk (exact, no scan); larger values trade index size for scan length.
|
||||
pub fn create_with_block_bits(path: &Path, block_bits: u8) -> SKResult<Self> {
|
||||
assert!(block_bits <= 31, "block_bits must be ≤ 31");
|
||||
let file = File::create(path).map_err(SKError::Io)?;
|
||||
Ok(Self {
|
||||
file: BufWriter::new(file),
|
||||
block_offsets: Vec::new(),
|
||||
chunk_count: 0,
|
||||
next_offset: 0,
|
||||
n_kmers: 0,
|
||||
k: obikseq::params::k(),
|
||||
block_bits,
|
||||
mask: (1usize << block_bits) - 1,
|
||||
})
|
||||
}
|
||||
|
||||
/// Write a unitig, splitting into overlapping chunks if it exceeds
|
||||
/// [`MAX_KMERS_PER_CHUNK`].
|
||||
pub fn write(&mut self, unitig: &Unitig) -> SKResult<()> {
|
||||
let seql = unitig.seql();
|
||||
let k = self.k;
|
||||
|
||||
if seql < k {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let n_kmers = seql - k + 1;
|
||||
if n_kmers <= MAX_KMERS_PER_CHUNK {
|
||||
return self.write_chunk(unitig);
|
||||
}
|
||||
|
||||
let chunk_nucl = MAX_KMERS_PER_CHUNK + k - 1;
|
||||
let stride = MAX_KMERS_PER_CHUNK;
|
||||
let mut start = 0;
|
||||
while start < seql {
|
||||
let end = (start + chunk_nucl).min(seql);
|
||||
self.write_chunk(&unitig.sub(start, end))?;
|
||||
if end == seql { break; }
|
||||
start += stride;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn write_chunk(&mut self, unitig: &Unitig) -> SKResult<()> {
|
||||
let seql = unitig.seql();
|
||||
let byte_len = (seql + 3) / 4;
|
||||
|
||||
debug_assert!(seql - self.k <= u8::MAX as usize, "chunk exceeds MAX_KMERS_PER_CHUNK");
|
||||
|
||||
if self.chunk_count & self.mask == 0 {
|
||||
self.block_offsets.push(self.next_offset);
|
||||
}
|
||||
|
||||
self.n_kmers += seql - self.k + 1;
|
||||
self.chunk_count += 1;
|
||||
|
||||
unitig.write_to_binary(&mut self.file).map_err(SKError::Io)?;
|
||||
|
||||
self.next_offset += 1 + byte_len as u32;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Flush and close the binary sequence file.
|
||||
///
|
||||
/// The companion `.idx` file is **not** written here; call
|
||||
/// [`build_unitig_idx`] separately when exact evidence is needed.
|
||||
pub fn close(mut self) -> SKResult<()> {
|
||||
self.file.flush().map_err(SKError::Io)?;
|
||||
drop(self.file);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub fn len(&self) -> usize { self.chunk_count }
|
||||
pub fn is_empty(&self) -> bool { self.chunk_count == 0 }
|
||||
pub fn block_bits(&self) -> u8 { self.block_bits }
|
||||
}
|
||||
|
||||
fn write_idx(
|
||||
path: &Path,
|
||||
n_unitigs: u32,
|
||||
n_kmers: u64,
|
||||
block_bits: u8,
|
||||
block_offsets: &[u32],
|
||||
) -> SKResult<()> {
|
||||
let mut w = BufWriter::new(File::create(path).map_err(SKError::Io)?);
|
||||
w.write_all(&MAGIC).map_err(SKError::Io)?;
|
||||
w.write_all(&(block_bits as u32).to_le_bytes()).map_err(SKError::Io)?;
|
||||
w.write_all(&n_unitigs.to_le_bytes()).map_err(SKError::Io)?;
|
||||
w.write_all(&n_kmers.to_le_bytes()).map_err(SKError::Io)?;
|
||||
for &off in block_offsets {
|
||||
w.write_all(&off.to_le_bytes()).map_err(SKError::Io)?;
|
||||
}
|
||||
w.flush().map_err(SKError::Io)
|
||||
}
|
||||
|
||||
/// Scan an existing `unitigs.bin` file and write its companion `.idx`.
|
||||
///
|
||||
/// Called by the exact-evidence construction route after the sequence file is
|
||||
/// closed. `block_bits` controls index granularity (1 << block_bits chunks per
|
||||
/// offset entry); use [`DEFAULT_BLOCK_BITS`] for the default.
|
||||
pub fn build_unitig_idx(unitigs_path: &Path, block_bits: u8) -> SKResult<()> {
|
||||
assert!(block_bits <= 31, "block_bits must be ≤ 31");
|
||||
|
||||
let file = File::open(unitigs_path).map_err(SKError::Io)?;
|
||||
let mmap = unsafe { Mmap::map(&file).map_err(SKError::Io)? };
|
||||
|
||||
let k = obikseq::params::k();
|
||||
let block_size = 1usize << block_bits;
|
||||
let mask = block_size - 1;
|
||||
|
||||
let mut block_offsets: Vec<u32> = Vec::new();
|
||||
let mut offset = 0usize;
|
||||
let mut chunk_count = 0usize;
|
||||
let mut n_kmers = 0usize;
|
||||
|
||||
while offset < mmap.len() {
|
||||
if chunk_count & mask == 0 {
|
||||
block_offsets.push(offset as u32);
|
||||
}
|
||||
let seql_minus_k = mmap[offset] as usize;
|
||||
let byte_len = (seql_minus_k + k + 3) / 4;
|
||||
n_kmers += seql_minus_k + 1;
|
||||
offset += 1 + byte_len;
|
||||
chunk_count += 1;
|
||||
}
|
||||
|
||||
block_offsets.push(offset as u32); // sentinel
|
||||
|
||||
write_idx(
|
||||
&idx_path(unitigs_path),
|
||||
chunk_count as u32,
|
||||
n_kmers as u64,
|
||||
block_bits,
|
||||
&block_offsets,
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,65 @@
|
||||
use std::sync::{Condvar, Mutex};
|
||||
|
||||
struct BudgetInner {
|
||||
remaining: u64,
|
||||
active: usize,
|
||||
peak_active: usize,
|
||||
}
|
||||
|
||||
/// Counting semaphore that limits total concurrent estimated memory usage.
|
||||
///
|
||||
/// Each worker acquires a cost (bytes) before starting and releases it on
|
||||
/// completion. Non-deadlock guarantee: when no worker is active the next
|
||||
/// acquire always succeeds regardless of cost vs. remaining budget.
|
||||
pub struct MemoryBudget {
|
||||
total: u64,
|
||||
inner: Mutex<BudgetInner>,
|
||||
condvar: Condvar,
|
||||
}
|
||||
|
||||
impl MemoryBudget {
|
||||
pub fn new(total: u64) -> Self {
|
||||
Self {
|
||||
total,
|
||||
inner: Mutex::new(BudgetInner {
|
||||
remaining: total,
|
||||
active: 0,
|
||||
peak_active: 0,
|
||||
}),
|
||||
condvar: Condvar::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn acquire(&self, cost: u64) {
|
||||
let mut g = self.inner.lock().unwrap();
|
||||
loop {
|
||||
if g.active == 0 || g.remaining >= cost {
|
||||
g.remaining = g.remaining.saturating_sub(cost);
|
||||
g.active += 1;
|
||||
g.peak_active = g.peak_active.max(g.active);
|
||||
return;
|
||||
}
|
||||
g = self.condvar.wait(g).unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
pub fn release(&self, cost: u64) {
|
||||
let mut g = self.inner.lock().unwrap();
|
||||
g.remaining = (g.remaining + cost).min(self.total);
|
||||
g.active -= 1;
|
||||
self.condvar.notify_all();
|
||||
}
|
||||
|
||||
pub fn total(&self) -> u64 {
|
||||
self.total
|
||||
}
|
||||
pub fn active(&self) -> usize {
|
||||
self.inner.lock().unwrap().active
|
||||
}
|
||||
pub fn remaining(&self) -> u64 {
|
||||
self.inner.lock().unwrap().remaining
|
||||
}
|
||||
pub fn peak_active(&self) -> usize {
|
||||
self.inner.lock().unwrap().peak_active
|
||||
}
|
||||
}
|
||||
+14
-867
@@ -1,867 +1,14 @@
|
||||
use std::fmt;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::{Condvar, Mutex};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use indicatif::{ProgressBar, ProgressStyle};
|
||||
use tracing::{debug, info, warn};
|
||||
|
||||
const BRAILLE: &[&str] = &["⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"];
|
||||
|
||||
// ── DirLock ──────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Exclusive advisory lock on an index directory, held for the duration of
|
||||
/// any command that writes into an already-existing index (building the
|
||||
/// sibling annex, merging into a destination, filtering/selecting in place,
|
||||
/// ...). Two such commands racing on the same directory can otherwise
|
||||
/// corrupt each other's writes with no error from either side.
|
||||
///
|
||||
/// Only the directory actually being *written to* needs a lock — a command
|
||||
/// like `merge` that reads several source indexes to build one destination
|
||||
/// only needs to lock the destination.
|
||||
///
|
||||
/// Uses the OS's advisory file lock (`flock` on Unix, `LockFileEx` on
|
||||
/// Windows) via `fs4`, not a hand-rolled PID file: the OS releases it
|
||||
/// automatically on process exit, including a crash — no stale-lock cleanup
|
||||
/// logic needed.
|
||||
pub struct DirLock {
|
||||
_file: std::fs::File,
|
||||
}
|
||||
|
||||
impl DirLock {
|
||||
/// Block until the exclusive lock on `dir` is acquired (creating `dir`
|
||||
/// and the lock file within it if needed). Logs once if the wait is
|
||||
/// non-trivial, so a blocked command doesn't look silently hung.
|
||||
pub fn acquire(dir: &std::path::Path) -> std::io::Result<Self> {
|
||||
use fs4::fs_std::FileExt;
|
||||
|
||||
std::fs::create_dir_all(dir)?;
|
||||
let lock_path = dir.join(".obikmer.lock");
|
||||
let file = std::fs::OpenOptions::new()
|
||||
.create(true)
|
||||
.truncate(false)
|
||||
.write(true)
|
||||
.open(&lock_path)?;
|
||||
|
||||
if file.try_lock_exclusive().is_err() {
|
||||
info!(dir = %dir.display(), "waiting for another obikmer process to release this index");
|
||||
file.lock_exclusive()?;
|
||||
}
|
||||
Ok(Self { _file: file })
|
||||
}
|
||||
}
|
||||
|
||||
// ── TracedBar ──────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Wrapper around `ProgressBar` that emits `tracing` events when stderr is not
|
||||
/// a TTY (e.g. HPC job logs): every 10% for bounded bars, every ~10 s for
|
||||
/// spinners (throttled on `set_message`).
|
||||
pub struct TracedBar {
|
||||
pb: ProgressBar,
|
||||
label: String,
|
||||
unit: String,
|
||||
total: u64, // 0 for spinners
|
||||
start: Instant, // creation time, for spinner throttling
|
||||
last_pct: AtomicU64, // last emitted 10%-bucket (1..=10), 0 = none yet
|
||||
last_log_ms: AtomicU64, // ms since `start` at last spinner log
|
||||
}
|
||||
|
||||
impl TracedBar {
|
||||
pub fn inc(&self, delta: u64) {
|
||||
self.pb.inc(delta);
|
||||
if self.pb.is_hidden() && self.total > 0 {
|
||||
let pos = self.pb.position();
|
||||
let pct10 = (pos * 10) / self.total; // 0..=10
|
||||
let last = self.last_pct.load(Ordering::Relaxed);
|
||||
if pct10 > last
|
||||
&& self
|
||||
.last_pct
|
||||
.compare_exchange(last, pct10, Ordering::Relaxed, Ordering::Relaxed)
|
||||
.is_ok()
|
||||
{
|
||||
info!(
|
||||
stage = %self.label,
|
||||
progress = format_args!("{}%", pct10 * 10),
|
||||
"{}/{} {}",
|
||||
pos, self.total, self.unit
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn set_message(&self, msg: impl Into<String>) {
|
||||
let msg = msg.into();
|
||||
if self.pb.is_hidden() {
|
||||
if self.total > 0 {
|
||||
debug!(stage = %self.label, "{msg}");
|
||||
} else {
|
||||
// spinner: throttle to ~10 s
|
||||
let now_ms = self.start.elapsed().as_millis() as u64;
|
||||
let last = self.last_log_ms.load(Ordering::Relaxed);
|
||||
if now_ms >= last + 10_000
|
||||
&& self
|
||||
.last_log_ms
|
||||
.compare_exchange(last, now_ms, Ordering::Relaxed, Ordering::Relaxed)
|
||||
.is_ok()
|
||||
{
|
||||
info!(stage = %self.label, "{msg}");
|
||||
}
|
||||
}
|
||||
}
|
||||
self.pb.set_message(msg);
|
||||
}
|
||||
|
||||
pub fn finish_and_clear(&self) {
|
||||
self.pb.finish_and_clear();
|
||||
}
|
||||
}
|
||||
|
||||
/// Spinner with the standard project look: `⠋ label — msg 0s`.
|
||||
/// Caller updates the message with `pb.set_message(...)`.
|
||||
pub fn spinner(label: &str) -> TracedBar {
|
||||
let pb = ProgressBar::new_spinner();
|
||||
pb.set_style(
|
||||
ProgressStyle::with_template(&format!("{{spinner}} {label} — {{msg}} {{elapsed}}"))
|
||||
.unwrap()
|
||||
.tick_strings(BRAILLE),
|
||||
);
|
||||
pb.enable_steady_tick(Duration::from_millis(100));
|
||||
TracedBar {
|
||||
pb,
|
||||
label: label.to_string(),
|
||||
unit: String::new(),
|
||||
total: 0,
|
||||
start: Instant::now(),
|
||||
last_pct: AtomicU64::new(0),
|
||||
last_log_ms: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
/// Progress bar with the standard project look:
|
||||
/// `⠋ label — [████░░░░] pos/len unit elapsed`.
|
||||
pub fn progress_bar(label: &str, n: u64, unit: &str) -> TracedBar {
|
||||
let pb = ProgressBar::new(n);
|
||||
pb.set_style(
|
||||
ProgressStyle::with_template(&format!(
|
||||
"{{spinner}} {label} — {{bar:40.cyan/blue}} {{pos}}/{{len}} {unit} {{elapsed}}"
|
||||
))
|
||||
.unwrap()
|
||||
.tick_strings(BRAILLE),
|
||||
);
|
||||
pb.enable_steady_tick(Duration::from_millis(100));
|
||||
TracedBar {
|
||||
pb,
|
||||
label: label.to_string(),
|
||||
unit: unit.to_string(),
|
||||
total: n,
|
||||
start: Instant::now(),
|
||||
last_pct: AtomicU64::new(0),
|
||||
last_log_ms: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
use libc::{RUSAGE_SELF, getrusage, rusage, timeval};
|
||||
use sysinfo::System;
|
||||
|
||||
// ── Memory query ──────────────────────────────────────────────────────────────
|
||||
|
||||
/// Returns the number of bytes available for allocation in the current process context.
|
||||
///
|
||||
/// On Linux, cgroup memory limits (SLURM, containers) are checked first: the
|
||||
/// process may be constrained to far less than the host's available RAM.
|
||||
/// Returns `min(cgroup_available, host_available)` when a finite limit is found.
|
||||
///
|
||||
/// On macOS, `available_memory()` can return 0 when the memory compressor
|
||||
/// inflates the page count; in that case we fall back to half of total memory.
|
||||
/// Returns the process peak RSS (high-water mark since process start).
|
||||
/// Monotonically increasing — use delta before/after a phase to measure its RAM cost.
|
||||
pub fn peak_rss_bytes() -> u64 {
|
||||
rss_to_bytes(&get_rusage())
|
||||
}
|
||||
|
||||
pub fn available_memory_bytes() -> u64 {
|
||||
let sys = System::new_all();
|
||||
let host_avail = match sys.available_memory() {
|
||||
0 => sys.total_memory() / 2,
|
||||
n => n,
|
||||
};
|
||||
#[cfg(target_os = "linux")]
|
||||
if let Some(cg) = cgroup_v2_available().or_else(cgroup_v1_available) {
|
||||
return cg.min(host_avail);
|
||||
}
|
||||
host_avail
|
||||
}
|
||||
|
||||
/// cgroup v2 (unified hierarchy): reads memory.max and memory.current for the
|
||||
/// current process's cgroup. Returns None if unlimited or on any parse error.
|
||||
#[cfg(target_os = "linux")]
|
||||
fn cgroup_v2_available() -> Option<u64> {
|
||||
let cgroup = std::fs::read_to_string("/proc/self/cgroup").ok()?;
|
||||
let rel = cgroup
|
||||
.lines()
|
||||
.find(|l| l.starts_with("0::"))?
|
||||
.strip_prefix("0::")?
|
||||
.trim();
|
||||
let base = format!("/sys/fs/cgroup{rel}");
|
||||
// "max" means no limit → parse::<u64>() fails → None
|
||||
let limit: u64 = std::fs::read_to_string(format!("{base}/memory.max"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
let used: u64 = std::fs::read_to_string(format!("{base}/memory.current"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
Some(limit.saturating_sub(used))
|
||||
}
|
||||
|
||||
/// cgroup v1 (memory subsystem): reads memory.limit_in_bytes and
|
||||
/// memory.usage_in_bytes. Returns None if unlimited or on any parse error.
|
||||
#[cfg(target_os = "linux")]
|
||||
fn cgroup_v1_available() -> Option<u64> {
|
||||
let cgroup = std::fs::read_to_string("/proc/self/cgroup").ok()?;
|
||||
let path = cgroup
|
||||
.lines()
|
||||
.find(|l| l.contains(":memory:"))?
|
||||
.split(':')
|
||||
.nth(2)?;
|
||||
let base = format!("/sys/fs/cgroup/memory{path}");
|
||||
let limit: u64 = std::fs::read_to_string(format!("{base}/memory.limit_in_bytes"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
// Kernel uses 2^63 (rounded to page) as "no limit" sentinel
|
||||
if limit > (1u64 << 62) {
|
||||
return None;
|
||||
}
|
||||
let used: u64 = std::fs::read_to_string(format!("{base}/memory.usage_in_bytes"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
Some(limit.saturating_sub(used))
|
||||
}
|
||||
|
||||
// ── CPU parallelism query ────────────────────────────────────────────────────
|
||||
|
||||
/// Returns the number of cores this process can actually use concurrently.
|
||||
///
|
||||
/// `std::thread::available_parallelism()` reads CPU affinity
|
||||
/// (`sched_getaffinity`), not the container's CPU quota — a Docker/cgroup
|
||||
/// container commonly reports the *host's* full core count this way while
|
||||
/// actually being throttled (via `cpu.max`/`cpu.cfs_quota_us`) to a fraction
|
||||
/// of a core. Sizing a thread/worker pool off the unthrottled count causes
|
||||
/// severe oversubscription: dozens of threads contending for a sliver of
|
||||
/// real CPU time, which can look indistinguishable from a hang for minutes
|
||||
/// or hours (observed in CI). On Linux, this reads the cgroup CPU quota
|
||||
/// first and returns `min(cgroup_quota, host_parallelism)` when a finite
|
||||
/// quota is found; falls back to `available_parallelism()` otherwise (same
|
||||
/// convention as [`available_memory_bytes`]).
|
||||
pub fn effective_parallelism() -> usize {
|
||||
let host = std::thread::available_parallelism().map(|n| n.get()).unwrap_or(1);
|
||||
#[cfg(target_os = "linux")]
|
||||
{
|
||||
if let Some(quota) = cgroup_v2_cpu_quota() {
|
||||
let effective = quota.clamp(1, host);
|
||||
tracing::debug!(host, quota, effective, source = "cgroup v2", "effective_parallelism");
|
||||
return effective;
|
||||
}
|
||||
if let Some(quota) = cgroup_v1_cpu_quota() {
|
||||
let effective = quota.clamp(1, host);
|
||||
tracing::debug!(host, quota, effective, source = "cgroup v1", "effective_parallelism");
|
||||
return effective;
|
||||
}
|
||||
}
|
||||
tracing::debug!(host, effective = host, source = "available_parallelism (no cgroup quota found)", "effective_parallelism");
|
||||
host
|
||||
}
|
||||
|
||||
/// cgroup v2 (unified hierarchy): reads `cpu.max` ("<quota> <period>", or
|
||||
/// "max <period>" when unlimited) for the current process's cgroup, rounded
|
||||
/// up to whole cores. Returns `None` if unlimited or on any parse error.
|
||||
#[cfg(target_os = "linux")]
|
||||
fn cgroup_v2_cpu_quota() -> Option<usize> {
|
||||
let cgroup = std::fs::read_to_string("/proc/self/cgroup").ok()?;
|
||||
let rel = cgroup
|
||||
.lines()
|
||||
.find(|l| l.starts_with("0::"))?
|
||||
.strip_prefix("0::")?
|
||||
.trim();
|
||||
let base = format!("/sys/fs/cgroup{rel}");
|
||||
let raw = std::fs::read_to_string(format!("{base}/cpu.max")).ok()?;
|
||||
let mut parts = raw.split_whitespace();
|
||||
let quota_str = parts.next()?;
|
||||
let period: f64 = parts.next()?.parse().ok()?;
|
||||
if quota_str == "max" {
|
||||
return None; // unlimited
|
||||
}
|
||||
let quota: f64 = quota_str.parse().ok()?;
|
||||
Some((quota / period).ceil().max(1.0) as usize)
|
||||
}
|
||||
|
||||
/// cgroup v1 (cpu subsystem): reads `cpu.cfs_quota_us`/`cpu.cfs_period_us`,
|
||||
/// rounded up to whole cores. Returns `None` if unlimited (quota <= 0) or on
|
||||
/// any parse error.
|
||||
#[cfg(target_os = "linux")]
|
||||
fn cgroup_v1_cpu_quota() -> Option<usize> {
|
||||
let cgroup = std::fs::read_to_string("/proc/self/cgroup").ok()?;
|
||||
let path = cgroup
|
||||
.lines()
|
||||
.find(|l| l.contains(":cpu:") || l.contains(":cpu,cpuacct:"))?
|
||||
.split(':')
|
||||
.nth(2)?;
|
||||
let base = format!("/sys/fs/cgroup/cpu{path}");
|
||||
let quota: i64 = std::fs::read_to_string(format!("{base}/cpu.cfs_quota_us"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
if quota <= 0 {
|
||||
return None; // unlimited
|
||||
}
|
||||
let period: i64 = std::fs::read_to_string(format!("{base}/cpu.cfs_period_us"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
if period <= 0 {
|
||||
return None;
|
||||
}
|
||||
Some(((quota as f64) / (period as f64)).ceil().max(1.0) as usize)
|
||||
}
|
||||
|
||||
// ── raw helpers ───────────────────────────────────────────────────────────────
|
||||
|
||||
fn get_rusage() -> rusage {
|
||||
let mut ru = unsafe { std::mem::zeroed::<rusage>() };
|
||||
unsafe { getrusage(RUSAGE_SELF, &mut ru) };
|
||||
ru
|
||||
}
|
||||
|
||||
fn tv_to_secs(tv: timeval) -> f64 {
|
||||
tv.tv_sec as f64 + tv.tv_usec as f64 * 1e-6
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
fn rss_to_bytes(ru: &rusage) -> u64 {
|
||||
ru.ru_maxrss as u64
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "macos"))]
|
||||
fn rss_to_bytes(ru: &rusage) -> u64 {
|
||||
ru.ru_maxrss as u64 * 1024
|
||||
}
|
||||
|
||||
// Monotonically increasing counters — negative delta would be a kernel bug.
|
||||
fn delta(end: i64, start: i64) -> u64 {
|
||||
(end - start).max(0) as u64
|
||||
}
|
||||
|
||||
// ── CpuSample ─────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Snapshot of process-wide CPU time + wall clock at a point in time.
|
||||
/// Use [`cpu_efficiency`](Self::cpu_efficiency) to measure the fraction of
|
||||
/// available cores used since the snapshot was taken.
|
||||
pub struct CpuSample {
|
||||
wall: Instant,
|
||||
user_secs: f64,
|
||||
sys_secs: f64,
|
||||
previous: f64,
|
||||
}
|
||||
|
||||
impl CpuSample {
|
||||
pub fn now() -> Self {
|
||||
let ru = get_rusage();
|
||||
Self {
|
||||
wall: Instant::now(),
|
||||
user_secs: tv_to_secs(ru.ru_utime),
|
||||
sys_secs: tv_to_secs(ru.ru_stime),
|
||||
previous: 0.0,
|
||||
}
|
||||
}
|
||||
|
||||
/// (user_delta + sys_delta) / (wall_delta × n_cores) since this snapshot.
|
||||
/// Returns 0.0 if less than 100 ms have elapsed (too noisy).
|
||||
pub fn cpu_efficiency(&self, n_cores: usize) -> f64 {
|
||||
let ru = get_rusage();
|
||||
let wall = self.wall.elapsed().as_secs_f64();
|
||||
if wall < 0.1 {
|
||||
return 0.0;
|
||||
}
|
||||
let cpu =
|
||||
(tv_to_secs(ru.ru_utime) - self.user_secs) + (tv_to_secs(ru.ru_stime) - self.sys_secs);
|
||||
cpu / (wall * n_cores as f64)
|
||||
}
|
||||
|
||||
pub fn do_i_activate(&mut self, threshold: f64) -> bool {
|
||||
let delta_wall = self.wall.elapsed().as_secs_f64();
|
||||
if delta_wall < 0.1 {
|
||||
// Window too short to be meaningful — leave state untouched so it
|
||||
// keeps accumulating until a real sample can be taken.
|
||||
return false;
|
||||
}
|
||||
|
||||
let n = CpuSample::now();
|
||||
let delta_ru = (n.user_secs - self.user_secs) + (n.sys_secs - self.sys_secs);
|
||||
|
||||
let efficiency = delta_ru / delta_wall;
|
||||
let activate = 0f64.max(efficiency - self.previous) >= threshold;
|
||||
|
||||
debug!(
|
||||
"Do I activate : {} -> {} = {} Activate: {}",
|
||||
self.previous,
|
||||
efficiency,
|
||||
0f64.max(efficiency - self.previous),
|
||||
activate
|
||||
);
|
||||
self.previous = efficiency;
|
||||
self.user_secs = n.user_secs;
|
||||
self.sys_secs = n.sys_secs;
|
||||
self.wall = n.wall;
|
||||
|
||||
activate
|
||||
}
|
||||
}
|
||||
|
||||
// ── IoSample ──────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Snapshot of process-wide block I/O (bytes read + written) + wall clock.
|
||||
///
|
||||
/// Same activation protocol as [`CpuSample`], but the growth check in
|
||||
/// [`do_i_activate`](Self::do_i_activate) is *relative* rather than absolute:
|
||||
/// raw I/O throughput has no portable scale across storage devices, unlike a
|
||||
/// core count.
|
||||
pub struct IoSample {
|
||||
wall: Instant,
|
||||
bytes: u64,
|
||||
previous_rate: f64,
|
||||
}
|
||||
|
||||
impl IoSample {
|
||||
pub fn now() -> Self {
|
||||
Self {
|
||||
wall: Instant::now(),
|
||||
bytes: Self::read_bytes(),
|
||||
previous_rate: 0.0,
|
||||
}
|
||||
}
|
||||
|
||||
/// Bytes actually submitted to the block layer (read + write), summed
|
||||
/// process-wide. Returns 0 if unavailable — degrades gracefully to a
|
||||
/// signal that never triggers activation (CPU-only heuristic).
|
||||
#[cfg(target_os = "linux")]
|
||||
fn read_bytes() -> u64 {
|
||||
let Ok(io) = std::fs::read_to_string("/proc/self/io") else {
|
||||
return 0;
|
||||
};
|
||||
io.lines()
|
||||
.filter_map(|l| {
|
||||
l.strip_prefix("read_bytes: ")
|
||||
.or_else(|| l.strip_prefix("write_bytes: "))
|
||||
})
|
||||
.filter_map(|v| v.trim().parse::<u64>().ok())
|
||||
.sum()
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
fn read_bytes() -> u64 {
|
||||
use libc::{RUSAGE_INFO_V4, getpid, proc_pid_rusage, rusage_info_v4};
|
||||
let mut info: rusage_info_v4 = unsafe { std::mem::zeroed() };
|
||||
let ret =
|
||||
unsafe { proc_pid_rusage(getpid(), RUSAGE_INFO_V4, &mut info as *mut _ as *mut _) };
|
||||
if ret != 0 {
|
||||
return 0;
|
||||
}
|
||||
info.ri_diskio_bytesread + info.ri_diskio_byteswritten
|
||||
}
|
||||
|
||||
#[cfg(not(any(target_os = "linux", target_os = "macos")))]
|
||||
fn read_bytes() -> u64 {
|
||||
0
|
||||
}
|
||||
|
||||
/// Same protocol as [`CpuSample::do_i_activate`] (0.1 s minimum window,
|
||||
/// state untouched on early return), but growth is measured relative to
|
||||
/// the previous rate. `threshold` is a fraction, e.g. `0.2` for a 20 %
|
||||
/// increase in throughput since the last real sample.
|
||||
pub fn do_i_activate(&mut self, threshold: f64) -> bool {
|
||||
let elapsed = self.wall.elapsed().as_secs_f64();
|
||||
if elapsed < 0.1 {
|
||||
return false;
|
||||
}
|
||||
|
||||
let n = Self::read_bytes();
|
||||
let rate = n.saturating_sub(self.bytes) as f64 / elapsed;
|
||||
let activate = if self.previous_rate == 0.0 {
|
||||
rate > 0.0 // bootstrap: any measured throughput is signal enough
|
||||
} else {
|
||||
(rate - self.previous_rate) / self.previous_rate >= threshold
|
||||
};
|
||||
|
||||
debug!(
|
||||
"Do I activate (I/O) : {} -> {} Activate: {}",
|
||||
self.previous_rate, rate, activate
|
||||
);
|
||||
self.previous_rate = rate;
|
||||
self.bytes = n;
|
||||
self.wall = Instant::now();
|
||||
|
||||
activate
|
||||
}
|
||||
}
|
||||
|
||||
// ── public API ────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Snapshot taken at the start of a pipeline stage.
|
||||
#[must_use = "call .stop() to record the stage"]
|
||||
pub struct Stage {
|
||||
label: String,
|
||||
wall: Instant,
|
||||
ru: rusage,
|
||||
}
|
||||
|
||||
impl Stage {
|
||||
pub fn start(label: impl Into<String>) -> Self {
|
||||
let label = label.into();
|
||||
info!(stage = %label, "started");
|
||||
Self {
|
||||
label,
|
||||
wall: Instant::now(),
|
||||
ru: get_rusage(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn stop(self) -> StageStats {
|
||||
let wall_secs = self.wall.elapsed().as_secs_f64();
|
||||
let end = get_rusage();
|
||||
let stats = StageStats {
|
||||
label: self.label,
|
||||
wall_secs,
|
||||
user_secs: tv_to_secs(end.ru_utime) - tv_to_secs(self.ru.ru_utime),
|
||||
sys_secs: tv_to_secs(end.ru_stime) - tv_to_secs(self.ru.ru_stime),
|
||||
max_rss_bytes: rss_to_bytes(&end),
|
||||
minor_faults: delta(end.ru_minflt as i64, self.ru.ru_minflt as i64),
|
||||
major_faults: delta(end.ru_majflt as i64, self.ru.ru_majflt as i64),
|
||||
vol_ctx: delta(end.ru_nvcsw as i64, self.ru.ru_nvcsw as i64),
|
||||
invol_ctx: delta(end.ru_nivcsw as i64, self.ru.ru_nivcsw as i64),
|
||||
in_blocks: delta(end.ru_inblock as i64, self.ru.ru_inblock as i64),
|
||||
out_blocks: delta(end.ru_oublock as i64, self.ru.ru_oublock as i64),
|
||||
swaps: delta(end.ru_nswap as i64, self.ru.ru_nswap as i64),
|
||||
};
|
||||
info!(
|
||||
stage = %stats.label,
|
||||
wall_secs = format_args!("{:.3}", stats.wall_secs),
|
||||
rss = %fmt_bytes(stats.max_rss_bytes),
|
||||
swaps = stats.swaps,
|
||||
"done"
|
||||
);
|
||||
if stats.swaps > 0 {
|
||||
warn!(
|
||||
stage = %stats.label,
|
||||
swaps = stats.swaps,
|
||||
"working set exceeds available RAM"
|
||||
);
|
||||
}
|
||||
stats
|
||||
}
|
||||
}
|
||||
|
||||
/// Per-stage efficiency metrics collected from `getrusage(RUSAGE_SELF)` deltas.
|
||||
pub struct StageStats {
|
||||
pub label: String,
|
||||
pub wall_secs: f64,
|
||||
pub user_secs: f64,
|
||||
pub sys_secs: f64,
|
||||
/// Peak RSS at end of stage (bytes). ru_maxrss is a process-lifetime maximum,
|
||||
/// so this reflects the high-water mark up to and including this stage.
|
||||
pub max_rss_bytes: u64,
|
||||
pub minor_faults: u64,
|
||||
pub major_faults: u64,
|
||||
pub vol_ctx: u64, // voluntary context switches
|
||||
pub invol_ctx: u64, // involuntary context switches
|
||||
pub in_blocks: u64, // filesystem block reads (after page cache)
|
||||
pub out_blocks: u64, // filesystem block writes
|
||||
pub swaps: u64,
|
||||
}
|
||||
|
||||
impl StageStats {
|
||||
/// (user + sys) / wall — effective thread count utilisation.
|
||||
pub fn parallelism(&self) -> f64 {
|
||||
if self.wall_secs > 1e-9 {
|
||||
(self.user_secs + self.sys_secs) / self.wall_secs
|
||||
} else {
|
||||
0.0
|
||||
}
|
||||
}
|
||||
|
||||
/// parallelism / n_cores — fraction of available CPU power used (0..1+).
|
||||
pub fn efficiency(&self, n_cores: usize) -> f64 {
|
||||
self.parallelism() / n_cores as f64
|
||||
}
|
||||
}
|
||||
|
||||
/// Accumulates stage stats and prints a human-readable summary table.
|
||||
#[derive(Default)]
|
||||
pub struct Reporter {
|
||||
stages: Vec<StageStats>,
|
||||
}
|
||||
|
||||
impl Reporter {
|
||||
pub fn new() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
pub fn push(&mut self, stats: StageStats) {
|
||||
self.stages.push(stats);
|
||||
}
|
||||
pub fn stages(&self) -> &[StageStats] {
|
||||
&self.stages
|
||||
}
|
||||
/// Print the summary to stderr.
|
||||
pub fn print(&self) {
|
||||
eprint!("{self}");
|
||||
}
|
||||
}
|
||||
|
||||
// ── diagnosis ─────────────────────────────────────────────────────────────────
|
||||
|
||||
struct Diagnosis {
|
||||
tag: &'static str,
|
||||
detail: Option<String>,
|
||||
}
|
||||
|
||||
// Thresholds are intentionally conservative to avoid false positives.
|
||||
fn diagnose(s: &StageStats, n_cores: usize) -> Diagnosis {
|
||||
let eff = s.efficiency(n_cores);
|
||||
let cpu_pct = eff * 100.0;
|
||||
let io_ops = s.in_blocks + s.out_blocks;
|
||||
|
||||
// swaps > 0 is the only reliable cross-platform indicator of true RAM exhaustion.
|
||||
// ru_majflt is intentionally excluded: on macOS it counts all file-backed mmap
|
||||
// page-ins (even from page cache), making it useless as a memory-pressure signal
|
||||
// for mmap-heavy code. On Linux it is more meaningful, but swaps covers the
|
||||
// severe case on both platforms.
|
||||
if s.swaps > 0 {
|
||||
return Diagnosis {
|
||||
tag: "swapping",
|
||||
detail: Some(format!(
|
||||
"swapped {} time(s) — working set exceeds available RAM",
|
||||
s.swaps,
|
||||
)),
|
||||
};
|
||||
}
|
||||
if eff < 0.3 && io_ops > 100 {
|
||||
return Diagnosis {
|
||||
tag: "disk I/O",
|
||||
detail: Some(format!(
|
||||
"{} block reads + {} writes — CPU at {:.0}%, stage is I/O-bound",
|
||||
s.in_blocks, s.out_blocks, cpu_pct,
|
||||
)),
|
||||
};
|
||||
}
|
||||
if eff < 0.3 && s.vol_ctx > 200 {
|
||||
return Diagnosis {
|
||||
tag: "contention",
|
||||
detail: Some(format!(
|
||||
"{} voluntary context switches — CPU at {:.0}%, possible lock contention or I/O wait",
|
||||
s.vol_ctx, cpu_pct,
|
||||
)),
|
||||
};
|
||||
}
|
||||
Diagnosis {
|
||||
tag: "—",
|
||||
detail: None,
|
||||
}
|
||||
}
|
||||
|
||||
// ── display helpers ───────────────────────────────────────────────────────────
|
||||
|
||||
fn fmt_secs(s: f64) -> String {
|
||||
if s >= 100.0 {
|
||||
format!("{:.0}s", s)
|
||||
} else if s >= 10.0 {
|
||||
format!("{:.1}s", s)
|
||||
} else if s >= 1.0 {
|
||||
format!("{:.2}s", s)
|
||||
} else {
|
||||
format!("{:.0}ms", s * 1000.0)
|
||||
}
|
||||
}
|
||||
|
||||
fn fmt_bytes(b: u64) -> String {
|
||||
if b >= 1 << 30 {
|
||||
format!("{:.1} GB", b as f64 / (1u64 << 30) as f64)
|
||||
} else if b >= 1 << 20 {
|
||||
format!("{:.0} MB", b as f64 / (1u64 << 20) as f64)
|
||||
} else {
|
||||
format!("{:.0} KB", b as f64 / 1024.0)
|
||||
}
|
||||
}
|
||||
|
||||
fn fmt_efficiency(par: f64, n_cores: usize) -> String {
|
||||
format!(
|
||||
"{:.1}×/{} ({:.0}%)",
|
||||
par,
|
||||
n_cores,
|
||||
par / n_cores as f64 * 100.0
|
||||
)
|
||||
}
|
||||
|
||||
// ── Display ───────────────────────────────────────────────────────────────────
|
||||
|
||||
// ── MemoryBudget ──────────────────────────────────────────────────────────────
|
||||
|
||||
struct BudgetInner {
|
||||
remaining: u64,
|
||||
active: usize,
|
||||
peak_active: usize,
|
||||
}
|
||||
|
||||
/// Counting semaphore that limits total concurrent estimated memory usage.
|
||||
///
|
||||
/// Each worker acquires a cost (bytes) before starting and releases it on
|
||||
/// completion. Non-deadlock guarantee: when no worker is active the next
|
||||
/// acquire always succeeds regardless of cost vs. remaining budget.
|
||||
pub struct MemoryBudget {
|
||||
total: u64,
|
||||
inner: Mutex<BudgetInner>,
|
||||
condvar: Condvar,
|
||||
}
|
||||
|
||||
impl MemoryBudget {
|
||||
pub fn new(total: u64) -> Self {
|
||||
Self {
|
||||
total,
|
||||
inner: Mutex::new(BudgetInner {
|
||||
remaining: total,
|
||||
active: 0,
|
||||
peak_active: 0,
|
||||
}),
|
||||
condvar: Condvar::new(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn acquire(&self, cost: u64) {
|
||||
let mut g = self.inner.lock().unwrap();
|
||||
loop {
|
||||
if g.active == 0 || g.remaining >= cost {
|
||||
g.remaining = g.remaining.saturating_sub(cost);
|
||||
g.active += 1;
|
||||
g.peak_active = g.peak_active.max(g.active);
|
||||
return;
|
||||
}
|
||||
g = self.condvar.wait(g).unwrap();
|
||||
}
|
||||
}
|
||||
|
||||
pub fn release(&self, cost: u64) {
|
||||
let mut g = self.inner.lock().unwrap();
|
||||
g.remaining = (g.remaining + cost).min(self.total);
|
||||
g.active -= 1;
|
||||
self.condvar.notify_all();
|
||||
}
|
||||
|
||||
pub fn total(&self) -> u64 {
|
||||
self.total
|
||||
}
|
||||
pub fn active(&self) -> usize {
|
||||
self.inner.lock().unwrap().active
|
||||
}
|
||||
pub fn remaining(&self) -> u64 {
|
||||
self.inner.lock().unwrap().remaining
|
||||
}
|
||||
pub fn peak_active(&self) -> usize {
|
||||
self.inner.lock().unwrap().peak_active
|
||||
}
|
||||
}
|
||||
|
||||
// ── Display ───────────────────────────────────────────────────────────────────
|
||||
|
||||
impl fmt::Display for Reporter {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
if self.stages.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let n_cores = effective_parallelism();
|
||||
|
||||
// column widths
|
||||
let nw = self
|
||||
.stages
|
||||
.iter()
|
||||
.map(|s| s.label.len())
|
||||
.max()
|
||||
.unwrap_or(5)
|
||||
.max(5);
|
||||
// efficiency col: worst-case width for this run's n_cores value
|
||||
let ew = format!("{:.1}×/{} (100%)", 99.9f64, n_cores).len();
|
||||
|
||||
let sep_w = nw + 2 + 7 + 2 + ew + 2 + 8 + 2 + 12;
|
||||
let sep = "─".repeat(sep_w);
|
||||
|
||||
// header
|
||||
writeln!(
|
||||
f,
|
||||
"{:<nw$} {:>7} {:>ew$} {:>8} status",
|
||||
"stage", "wall", "efficiency", "peak RSS"
|
||||
)?;
|
||||
writeln!(f, "{sep}")?;
|
||||
|
||||
// compute all diagnoses up front (needed for both table and footnotes)
|
||||
let diagnoses: Vec<Diagnosis> = self.stages.iter().map(|s| diagnose(s, n_cores)).collect();
|
||||
|
||||
// per-stage rows
|
||||
for (s, d) in self.stages.iter().zip(diagnoses.iter()) {
|
||||
writeln!(
|
||||
f,
|
||||
"{:<nw$} {:>7} {:>ew$} {:>8} {}",
|
||||
s.label,
|
||||
fmt_secs(s.wall_secs),
|
||||
fmt_efficiency(s.parallelism(), n_cores),
|
||||
fmt_bytes(s.max_rss_bytes),
|
||||
d.tag,
|
||||
)?;
|
||||
}
|
||||
|
||||
// totals
|
||||
let tw = self.stages.iter().map(|s| s.wall_secs).sum::<f64>();
|
||||
let tu = self.stages.iter().map(|s| s.user_secs).sum::<f64>();
|
||||
let ts = self.stages.iter().map(|s| s.sys_secs).sum::<f64>();
|
||||
let trss = self
|
||||
.stages
|
||||
.iter()
|
||||
.map(|s| s.max_rss_bytes)
|
||||
.max()
|
||||
.unwrap_or(0);
|
||||
let tpar = if tw > 1e-9 { (tu + ts) / tw } else { 0.0 };
|
||||
|
||||
writeln!(f, "{sep}")?;
|
||||
writeln!(
|
||||
f,
|
||||
"{:<nw$} {:>7} {:>ew$} {:>8}",
|
||||
"TOTAL",
|
||||
fmt_secs(tw),
|
||||
fmt_efficiency(tpar, n_cores),
|
||||
fmt_bytes(trss),
|
||||
)?;
|
||||
|
||||
// bottleneck footnotes (only if at least one anomaly detected)
|
||||
let bottlenecks: Vec<(&str, &str)> = self
|
||||
.stages
|
||||
.iter()
|
||||
.zip(diagnoses.iter())
|
||||
.filter_map(|(s, d)| d.detail.as_deref().map(|det| (s.label.as_str(), det)))
|
||||
.collect();
|
||||
|
||||
if !bottlenecks.is_empty() {
|
||||
writeln!(f, "\nBottlenecks:")?;
|
||||
for (label, detail) in &bottlenecks {
|
||||
writeln!(f, " {label} — {detail}")?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
//! Cross-cutting system utilities: directory locking, progress/logging,
|
||||
//! resource (CPU/memory) introspection, and per-stage profiling.
|
||||
|
||||
mod budget;
|
||||
mod lock;
|
||||
mod progress;
|
||||
mod resources;
|
||||
mod stage;
|
||||
|
||||
pub use budget::MemoryBudget;
|
||||
pub use lock::DirLock;
|
||||
pub use progress::{TracedBar, progress_bar, spinner};
|
||||
pub use resources::{CpuSample, IoSample, available_memory_bytes, effective_parallelism, peak_rss_bytes};
|
||||
pub use stage::{Reporter, Stage, StageStats};
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
use tracing::info;
|
||||
|
||||
/// Exclusive advisory lock on an index directory, held for the duration of
|
||||
/// any command that writes into an already-existing index (building the
|
||||
/// sibling annex, merging into a destination, filtering/selecting in place,
|
||||
/// ...). Two such commands racing on the same directory can otherwise
|
||||
/// corrupt each other's writes with no error from either side.
|
||||
///
|
||||
/// Only the directory actually being *written to* needs a lock — a command
|
||||
/// like `merge` that reads several source indexes to build one destination
|
||||
/// only needs to lock the destination.
|
||||
///
|
||||
/// Uses the OS's advisory file lock (`flock` on Unix, `LockFileEx` on
|
||||
/// Windows) via `fs4`, not a hand-rolled PID file: the OS releases it
|
||||
/// automatically on process exit, including a crash — no stale-lock cleanup
|
||||
/// logic needed.
|
||||
pub struct DirLock {
|
||||
_file: std::fs::File,
|
||||
}
|
||||
|
||||
impl DirLock {
|
||||
/// Block until the exclusive lock on `dir` is acquired (creating `dir`
|
||||
/// and the lock file within it if needed). Logs once if the wait is
|
||||
/// non-trivial, so a blocked command doesn't look silently hung.
|
||||
pub fn acquire(dir: &std::path::Path) -> std::io::Result<Self> {
|
||||
use fs4::fs_std::FileExt;
|
||||
|
||||
std::fs::create_dir_all(dir)?;
|
||||
let lock_path = dir.join(".obikmer.lock");
|
||||
let file = std::fs::OpenOptions::new()
|
||||
.create(true)
|
||||
.truncate(false)
|
||||
.write(true)
|
||||
.open(&lock_path)?;
|
||||
|
||||
if file.try_lock_exclusive().is_err() {
|
||||
info!(dir = %dir.display(), "waiting for another obikmer process to release this index");
|
||||
file.lock_exclusive()?;
|
||||
}
|
||||
Ok(Self { _file: file })
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::time::{Duration, Instant};
|
||||
|
||||
use indicatif::{ProgressBar, ProgressStyle};
|
||||
use tracing::{debug, info};
|
||||
|
||||
const BRAILLE: &[&str] = &["⠋", "⠙", "⠹", "⠸", "⠼", "⠴", "⠦", "⠧", "⠇", "⠏"];
|
||||
|
||||
/// Wrapper around `ProgressBar` that emits `tracing` events when stderr is not
|
||||
/// a TTY (e.g. HPC job logs): every 10% for bounded bars, every ~10 s for
|
||||
/// spinners (throttled on `set_message`).
|
||||
pub struct TracedBar {
|
||||
pb: ProgressBar,
|
||||
label: String,
|
||||
unit: String,
|
||||
total: u64, // 0 for spinners
|
||||
start: Instant, // creation time, for spinner throttling
|
||||
last_pct: AtomicU64, // last emitted 10%-bucket (1..=10), 0 = none yet
|
||||
last_log_ms: AtomicU64, // ms since `start` at last spinner log
|
||||
}
|
||||
|
||||
impl TracedBar {
|
||||
pub fn inc(&self, delta: u64) {
|
||||
self.pb.inc(delta);
|
||||
if self.pb.is_hidden() && self.total > 0 {
|
||||
let pos = self.pb.position();
|
||||
let pct10 = (pos * 10) / self.total; // 0..=10
|
||||
let last = self.last_pct.load(Ordering::Relaxed);
|
||||
if pct10 > last
|
||||
&& self
|
||||
.last_pct
|
||||
.compare_exchange(last, pct10, Ordering::Relaxed, Ordering::Relaxed)
|
||||
.is_ok()
|
||||
{
|
||||
info!(
|
||||
stage = %self.label,
|
||||
progress = format_args!("{}%", pct10 * 10),
|
||||
"{}/{} {}",
|
||||
pos, self.total, self.unit
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn set_message(&self, msg: impl Into<String>) {
|
||||
let msg = msg.into();
|
||||
if self.pb.is_hidden() {
|
||||
if self.total > 0 {
|
||||
debug!(stage = %self.label, "{msg}");
|
||||
} else {
|
||||
// spinner: throttle to ~10 s
|
||||
let now_ms = self.start.elapsed().as_millis() as u64;
|
||||
let last = self.last_log_ms.load(Ordering::Relaxed);
|
||||
if now_ms >= last + 10_000
|
||||
&& self
|
||||
.last_log_ms
|
||||
.compare_exchange(last, now_ms, Ordering::Relaxed, Ordering::Relaxed)
|
||||
.is_ok()
|
||||
{
|
||||
info!(stage = %self.label, "{msg}");
|
||||
}
|
||||
}
|
||||
}
|
||||
self.pb.set_message(msg);
|
||||
}
|
||||
|
||||
pub fn finish_and_clear(&self) {
|
||||
self.pb.finish_and_clear();
|
||||
}
|
||||
}
|
||||
|
||||
/// Spinner with the standard project look: `⠋ label — msg 0s`.
|
||||
/// Caller updates the message with `pb.set_message(...)`.
|
||||
pub fn spinner(label: &str) -> TracedBar {
|
||||
let pb = ProgressBar::new_spinner();
|
||||
pb.set_style(
|
||||
ProgressStyle::with_template(&format!("{{spinner}} {label} — {{msg}} {{elapsed}}"))
|
||||
.unwrap()
|
||||
.tick_strings(BRAILLE),
|
||||
);
|
||||
pb.enable_steady_tick(Duration::from_millis(100));
|
||||
TracedBar {
|
||||
pb,
|
||||
label: label.to_string(),
|
||||
unit: String::new(),
|
||||
total: 0,
|
||||
start: Instant::now(),
|
||||
last_pct: AtomicU64::new(0),
|
||||
last_log_ms: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
|
||||
/// Progress bar with the standard project look:
|
||||
/// `⠋ label — [████░░░░] pos/len unit elapsed`.
|
||||
pub fn progress_bar(label: &str, n: u64, unit: &str) -> TracedBar {
|
||||
let pb = ProgressBar::new(n);
|
||||
pb.set_style(
|
||||
ProgressStyle::with_template(&format!(
|
||||
"{{spinner}} {label} — {{bar:40.cyan/blue}} {{pos}}/{{len}} {unit} {{elapsed}}"
|
||||
))
|
||||
.unwrap()
|
||||
.tick_strings(BRAILLE),
|
||||
);
|
||||
pb.enable_steady_tick(Duration::from_millis(100));
|
||||
TracedBar {
|
||||
pb,
|
||||
label: label.to_string(),
|
||||
unit: unit.to_string(),
|
||||
total: n,
|
||||
start: Instant::now(),
|
||||
last_pct: AtomicU64::new(0),
|
||||
last_log_ms: AtomicU64::new(0),
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,355 @@
|
||||
use std::time::Instant;
|
||||
|
||||
use libc::{RUSAGE_SELF, getrusage, rusage, timeval};
|
||||
use sysinfo::System;
|
||||
use tracing::debug;
|
||||
|
||||
// ── Memory query ──────────────────────────────────────────────────────────────
|
||||
|
||||
/// Returns the number of bytes available for allocation in the current process context.
|
||||
///
|
||||
/// On Linux, cgroup memory limits (SLURM, containers) are checked first: the
|
||||
/// process may be constrained to far less than the host's available RAM.
|
||||
/// Returns `min(cgroup_available, host_available)` when a finite limit is found.
|
||||
///
|
||||
/// On macOS, `available_memory()` can return 0 when the memory compressor
|
||||
/// inflates the page count; in that case we fall back to half of total memory.
|
||||
/// Returns the process peak RSS (high-water mark since process start).
|
||||
/// Monotonically increasing — use delta before/after a phase to measure its RAM cost.
|
||||
pub fn peak_rss_bytes() -> u64 {
|
||||
rss_to_bytes(&get_rusage())
|
||||
}
|
||||
|
||||
pub fn available_memory_bytes() -> u64 {
|
||||
let sys = System::new_all();
|
||||
let host_avail = match sys.available_memory() {
|
||||
0 => sys.total_memory() / 2,
|
||||
n => n,
|
||||
};
|
||||
#[cfg(target_os = "linux")]
|
||||
if let Some(cg) = cgroup_v2_available().or_else(cgroup_v1_available) {
|
||||
return cg.min(host_avail);
|
||||
}
|
||||
host_avail
|
||||
}
|
||||
|
||||
/// cgroup v2 (unified hierarchy): reads memory.max and memory.current for the
|
||||
/// current process's cgroup. Returns None if unlimited or on any parse error.
|
||||
#[cfg(target_os = "linux")]
|
||||
fn cgroup_v2_available() -> Option<u64> {
|
||||
let cgroup = std::fs::read_to_string("/proc/self/cgroup").ok()?;
|
||||
let rel = cgroup
|
||||
.lines()
|
||||
.find(|l| l.starts_with("0::"))?
|
||||
.strip_prefix("0::")?
|
||||
.trim();
|
||||
let base = format!("/sys/fs/cgroup{rel}");
|
||||
// "max" means no limit → parse::<u64>() fails → None
|
||||
let limit: u64 = std::fs::read_to_string(format!("{base}/memory.max"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
let used: u64 = std::fs::read_to_string(format!("{base}/memory.current"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
Some(limit.saturating_sub(used))
|
||||
}
|
||||
|
||||
/// cgroup v1 (memory subsystem): reads memory.limit_in_bytes and
|
||||
/// memory.usage_in_bytes. Returns None if unlimited or on any parse error.
|
||||
#[cfg(target_os = "linux")]
|
||||
fn cgroup_v1_available() -> Option<u64> {
|
||||
let cgroup = std::fs::read_to_string("/proc/self/cgroup").ok()?;
|
||||
let path = cgroup
|
||||
.lines()
|
||||
.find(|l| l.contains(":memory:"))?
|
||||
.split(':')
|
||||
.nth(2)?;
|
||||
let base = format!("/sys/fs/cgroup/memory{path}");
|
||||
let limit: u64 = std::fs::read_to_string(format!("{base}/memory.limit_in_bytes"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
// Kernel uses 2^63 (rounded to page) as "no limit" sentinel
|
||||
if limit > (1u64 << 62) {
|
||||
return None;
|
||||
}
|
||||
let used: u64 = std::fs::read_to_string(format!("{base}/memory.usage_in_bytes"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
Some(limit.saturating_sub(used))
|
||||
}
|
||||
|
||||
// ── CPU parallelism query ────────────────────────────────────────────────────
|
||||
|
||||
/// Returns the number of cores this process can actually use concurrently.
|
||||
///
|
||||
/// `std::thread::available_parallelism()` reads CPU affinity
|
||||
/// (`sched_getaffinity`), not the container's CPU quota — a Docker/cgroup
|
||||
/// container commonly reports the *host's* full core count this way while
|
||||
/// actually being throttled (via `cpu.max`/`cpu.cfs_quota_us`) to a fraction
|
||||
/// of a core. Sizing a thread/worker pool off the unthrottled count causes
|
||||
/// severe oversubscription: dozens of threads contending for a sliver of
|
||||
/// real CPU time, which can look indistinguishable from a hang for minutes
|
||||
/// or hours (observed in CI). On Linux, this reads the cgroup CPU quota
|
||||
/// first and returns `min(cgroup_quota, host_parallelism)` when a finite
|
||||
/// quota is found; falls back to `available_parallelism()` otherwise (same
|
||||
/// convention as [`available_memory_bytes`]).
|
||||
pub fn effective_parallelism() -> usize {
|
||||
let host = std::thread::available_parallelism().map(|n| n.get()).unwrap_or(1);
|
||||
#[cfg(target_os = "linux")]
|
||||
{
|
||||
if let Some(quota) = cgroup_v2_cpu_quota() {
|
||||
let effective = quota.clamp(1, host);
|
||||
tracing::debug!(host, quota, effective, source = "cgroup v2", "effective_parallelism");
|
||||
return effective;
|
||||
}
|
||||
if let Some(quota) = cgroup_v1_cpu_quota() {
|
||||
let effective = quota.clamp(1, host);
|
||||
tracing::debug!(host, quota, effective, source = "cgroup v1", "effective_parallelism");
|
||||
return effective;
|
||||
}
|
||||
}
|
||||
tracing::debug!(host, effective = host, source = "available_parallelism (no cgroup quota found)", "effective_parallelism");
|
||||
host
|
||||
}
|
||||
|
||||
/// cgroup v2 (unified hierarchy): reads `cpu.max` ("<quota> <period>", or
|
||||
/// "max <period>" when unlimited) for the current process's cgroup, rounded
|
||||
/// up to whole cores. Returns `None` if unlimited or on any parse error.
|
||||
#[cfg(target_os = "linux")]
|
||||
fn cgroup_v2_cpu_quota() -> Option<usize> {
|
||||
let cgroup = std::fs::read_to_string("/proc/self/cgroup").ok()?;
|
||||
let rel = cgroup
|
||||
.lines()
|
||||
.find(|l| l.starts_with("0::"))?
|
||||
.strip_prefix("0::")?
|
||||
.trim();
|
||||
let base = format!("/sys/fs/cgroup{rel}");
|
||||
let raw = std::fs::read_to_string(format!("{base}/cpu.max")).ok()?;
|
||||
let mut parts = raw.split_whitespace();
|
||||
let quota_str = parts.next()?;
|
||||
let period: f64 = parts.next()?.parse().ok()?;
|
||||
if quota_str == "max" {
|
||||
return None; // unlimited
|
||||
}
|
||||
let quota: f64 = quota_str.parse().ok()?;
|
||||
Some((quota / period).ceil().max(1.0) as usize)
|
||||
}
|
||||
|
||||
/// cgroup v1 (cpu subsystem): reads `cpu.cfs_quota_us`/`cpu.cfs_period_us`,
|
||||
/// rounded up to whole cores. Returns `None` if unlimited (quota <= 0) or on
|
||||
/// any parse error.
|
||||
#[cfg(target_os = "linux")]
|
||||
fn cgroup_v1_cpu_quota() -> Option<usize> {
|
||||
let cgroup = std::fs::read_to_string("/proc/self/cgroup").ok()?;
|
||||
let path = cgroup
|
||||
.lines()
|
||||
.find(|l| l.contains(":cpu:") || l.contains(":cpu,cpuacct:"))?
|
||||
.split(':')
|
||||
.nth(2)?;
|
||||
let base = format!("/sys/fs/cgroup/cpu{path}");
|
||||
let quota: i64 = std::fs::read_to_string(format!("{base}/cpu.cfs_quota_us"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
if quota <= 0 {
|
||||
return None; // unlimited
|
||||
}
|
||||
let period: i64 = std::fs::read_to_string(format!("{base}/cpu.cfs_period_us"))
|
||||
.ok()?
|
||||
.trim()
|
||||
.parse()
|
||||
.ok()?;
|
||||
if period <= 0 {
|
||||
return None;
|
||||
}
|
||||
Some(((quota as f64) / (period as f64)).ceil().max(1.0) as usize)
|
||||
}
|
||||
|
||||
// ── raw helpers ───────────────────────────────────────────────────────────────
|
||||
|
||||
pub(crate) fn get_rusage() -> rusage {
|
||||
let mut ru = unsafe { std::mem::zeroed::<rusage>() };
|
||||
unsafe { getrusage(RUSAGE_SELF, &mut ru) };
|
||||
ru
|
||||
}
|
||||
|
||||
pub(crate) fn tv_to_secs(tv: timeval) -> f64 {
|
||||
tv.tv_sec as f64 + tv.tv_usec as f64 * 1e-6
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
pub(crate) fn rss_to_bytes(ru: &rusage) -> u64 {
|
||||
ru.ru_maxrss as u64
|
||||
}
|
||||
|
||||
#[cfg(not(target_os = "macos"))]
|
||||
pub(crate) fn rss_to_bytes(ru: &rusage) -> u64 {
|
||||
ru.ru_maxrss as u64 * 1024
|
||||
}
|
||||
|
||||
// Monotonically increasing counters — negative delta would be a kernel bug.
|
||||
pub(crate) fn delta(end: i64, start: i64) -> u64 {
|
||||
(end - start).max(0) as u64
|
||||
}
|
||||
|
||||
// ── CpuSample ─────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Snapshot of process-wide CPU time + wall clock at a point in time.
|
||||
/// Use [`cpu_efficiency`](Self::cpu_efficiency) to measure the fraction of
|
||||
/// available cores used since the snapshot was taken.
|
||||
pub struct CpuSample {
|
||||
wall: Instant,
|
||||
user_secs: f64,
|
||||
sys_secs: f64,
|
||||
previous: f64,
|
||||
}
|
||||
|
||||
impl CpuSample {
|
||||
pub fn now() -> Self {
|
||||
let ru = get_rusage();
|
||||
Self {
|
||||
wall: Instant::now(),
|
||||
user_secs: tv_to_secs(ru.ru_utime),
|
||||
sys_secs: tv_to_secs(ru.ru_stime),
|
||||
previous: 0.0,
|
||||
}
|
||||
}
|
||||
|
||||
/// (user_delta + sys_delta) / (wall_delta × n_cores) since this snapshot.
|
||||
/// Returns 0.0 if less than 100 ms have elapsed (too noisy).
|
||||
pub fn cpu_efficiency(&self, n_cores: usize) -> f64 {
|
||||
let ru = get_rusage();
|
||||
let wall = self.wall.elapsed().as_secs_f64();
|
||||
if wall < 0.1 {
|
||||
return 0.0;
|
||||
}
|
||||
let cpu =
|
||||
(tv_to_secs(ru.ru_utime) - self.user_secs) + (tv_to_secs(ru.ru_stime) - self.sys_secs);
|
||||
cpu / (wall * n_cores as f64)
|
||||
}
|
||||
|
||||
pub fn do_i_activate(&mut self, threshold: f64) -> bool {
|
||||
let delta_wall = self.wall.elapsed().as_secs_f64();
|
||||
if delta_wall < 0.1 {
|
||||
// Window too short to be meaningful — leave state untouched so it
|
||||
// keeps accumulating until a real sample can be taken.
|
||||
return false;
|
||||
}
|
||||
|
||||
let n = CpuSample::now();
|
||||
let delta_ru = (n.user_secs - self.user_secs) + (n.sys_secs - self.sys_secs);
|
||||
|
||||
let efficiency = delta_ru / delta_wall;
|
||||
let activate = 0f64.max(efficiency - self.previous) >= threshold;
|
||||
|
||||
debug!(
|
||||
"Do I activate : {} -> {} = {} Activate: {}",
|
||||
self.previous,
|
||||
efficiency,
|
||||
0f64.max(efficiency - self.previous),
|
||||
activate
|
||||
);
|
||||
self.previous = efficiency;
|
||||
self.user_secs = n.user_secs;
|
||||
self.sys_secs = n.sys_secs;
|
||||
self.wall = n.wall;
|
||||
|
||||
activate
|
||||
}
|
||||
}
|
||||
|
||||
// ── IoSample ──────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Snapshot of process-wide block I/O (bytes read + written) + wall clock.
|
||||
///
|
||||
/// Same activation protocol as [`CpuSample`], but the growth check in
|
||||
/// [`do_i_activate`](Self::do_i_activate) is *relative* rather than absolute:
|
||||
/// raw I/O throughput has no portable scale across storage devices, unlike a
|
||||
/// core count.
|
||||
pub struct IoSample {
|
||||
wall: Instant,
|
||||
bytes: u64,
|
||||
previous_rate: f64,
|
||||
}
|
||||
|
||||
impl IoSample {
|
||||
pub fn now() -> Self {
|
||||
Self {
|
||||
wall: Instant::now(),
|
||||
bytes: Self::read_bytes(),
|
||||
previous_rate: 0.0,
|
||||
}
|
||||
}
|
||||
|
||||
/// Bytes actually submitted to the block layer (read + write), summed
|
||||
/// process-wide. Returns 0 if unavailable — degrades gracefully to a
|
||||
/// signal that never triggers activation (CPU-only heuristic).
|
||||
#[cfg(target_os = "linux")]
|
||||
fn read_bytes() -> u64 {
|
||||
let Ok(io) = std::fs::read_to_string("/proc/self/io") else {
|
||||
return 0;
|
||||
};
|
||||
io.lines()
|
||||
.filter_map(|l| {
|
||||
l.strip_prefix("read_bytes: ")
|
||||
.or_else(|| l.strip_prefix("write_bytes: "))
|
||||
})
|
||||
.filter_map(|v| v.trim().parse::<u64>().ok())
|
||||
.sum()
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
fn read_bytes() -> u64 {
|
||||
use libc::{RUSAGE_INFO_V4, getpid, proc_pid_rusage, rusage_info_v4};
|
||||
let mut info: rusage_info_v4 = unsafe { std::mem::zeroed() };
|
||||
let ret =
|
||||
unsafe { proc_pid_rusage(getpid(), RUSAGE_INFO_V4, &mut info as *mut _ as *mut _) };
|
||||
if ret != 0 {
|
||||
return 0;
|
||||
}
|
||||
info.ri_diskio_bytesread + info.ri_diskio_byteswritten
|
||||
}
|
||||
|
||||
#[cfg(not(any(target_os = "linux", target_os = "macos")))]
|
||||
fn read_bytes() -> u64 {
|
||||
0
|
||||
}
|
||||
|
||||
/// Same protocol as [`CpuSample::do_i_activate`] (0.1 s minimum window,
|
||||
/// state untouched on early return), but growth is measured relative to
|
||||
/// the previous rate. `threshold` is a fraction, e.g. `0.2` for a 20 %
|
||||
/// increase in throughput since the last real sample.
|
||||
pub fn do_i_activate(&mut self, threshold: f64) -> bool {
|
||||
let elapsed = self.wall.elapsed().as_secs_f64();
|
||||
if elapsed < 0.1 {
|
||||
return false;
|
||||
}
|
||||
|
||||
let n = Self::read_bytes();
|
||||
let rate = n.saturating_sub(self.bytes) as f64 / elapsed;
|
||||
let activate = if self.previous_rate == 0.0 {
|
||||
rate > 0.0 // bootstrap: any measured throughput is signal enough
|
||||
} else {
|
||||
(rate - self.previous_rate) / self.previous_rate >= threshold
|
||||
};
|
||||
|
||||
debug!(
|
||||
"Do I activate (I/O) : {} -> {} Activate: {}",
|
||||
self.previous_rate, rate, activate
|
||||
);
|
||||
self.previous_rate = rate;
|
||||
self.bytes = n;
|
||||
self.wall = Instant::now();
|
||||
|
||||
activate
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,292 @@
|
||||
use std::fmt;
|
||||
use std::time::Instant;
|
||||
|
||||
use libc::rusage;
|
||||
use tracing::{info, warn};
|
||||
|
||||
use crate::resources::{delta, effective_parallelism, get_rusage, rss_to_bytes, tv_to_secs};
|
||||
|
||||
// ── public API ────────────────────────────────────────────────────────────────
|
||||
|
||||
/// Snapshot taken at the start of a pipeline stage.
|
||||
#[must_use = "call .stop() to record the stage"]
|
||||
pub struct Stage {
|
||||
label: String,
|
||||
wall: Instant,
|
||||
ru: rusage,
|
||||
}
|
||||
|
||||
impl Stage {
|
||||
pub fn start(label: impl Into<String>) -> Self {
|
||||
let label = label.into();
|
||||
info!(stage = %label, "started");
|
||||
Self {
|
||||
label,
|
||||
wall: Instant::now(),
|
||||
ru: get_rusage(),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn stop(self) -> StageStats {
|
||||
let wall_secs = self.wall.elapsed().as_secs_f64();
|
||||
let end = get_rusage();
|
||||
let stats = StageStats {
|
||||
label: self.label,
|
||||
wall_secs,
|
||||
user_secs: tv_to_secs(end.ru_utime) - tv_to_secs(self.ru.ru_utime),
|
||||
sys_secs: tv_to_secs(end.ru_stime) - tv_to_secs(self.ru.ru_stime),
|
||||
max_rss_bytes: rss_to_bytes(&end),
|
||||
minor_faults: delta(end.ru_minflt as i64, self.ru.ru_minflt as i64),
|
||||
major_faults: delta(end.ru_majflt as i64, self.ru.ru_majflt as i64),
|
||||
vol_ctx: delta(end.ru_nvcsw as i64, self.ru.ru_nvcsw as i64),
|
||||
invol_ctx: delta(end.ru_nivcsw as i64, self.ru.ru_nivcsw as i64),
|
||||
in_blocks: delta(end.ru_inblock as i64, self.ru.ru_inblock as i64),
|
||||
out_blocks: delta(end.ru_oublock as i64, self.ru.ru_oublock as i64),
|
||||
swaps: delta(end.ru_nswap as i64, self.ru.ru_nswap as i64),
|
||||
};
|
||||
info!(
|
||||
stage = %stats.label,
|
||||
wall_secs = format_args!("{:.3}", stats.wall_secs),
|
||||
rss = %fmt_bytes(stats.max_rss_bytes),
|
||||
swaps = stats.swaps,
|
||||
"done"
|
||||
);
|
||||
if stats.swaps > 0 {
|
||||
warn!(
|
||||
stage = %stats.label,
|
||||
swaps = stats.swaps,
|
||||
"working set exceeds available RAM"
|
||||
);
|
||||
}
|
||||
stats
|
||||
}
|
||||
}
|
||||
|
||||
/// Per-stage efficiency metrics collected from `getrusage(RUSAGE_SELF)` deltas.
|
||||
pub struct StageStats {
|
||||
pub label: String,
|
||||
pub wall_secs: f64,
|
||||
pub user_secs: f64,
|
||||
pub sys_secs: f64,
|
||||
/// Peak RSS at end of stage (bytes). ru_maxrss is a process-lifetime maximum,
|
||||
/// so this reflects the high-water mark up to and including this stage.
|
||||
pub max_rss_bytes: u64,
|
||||
pub minor_faults: u64,
|
||||
pub major_faults: u64,
|
||||
pub vol_ctx: u64, // voluntary context switches
|
||||
pub invol_ctx: u64, // involuntary context switches
|
||||
pub in_blocks: u64, // filesystem block reads (after page cache)
|
||||
pub out_blocks: u64, // filesystem block writes
|
||||
pub swaps: u64,
|
||||
}
|
||||
|
||||
impl StageStats {
|
||||
/// (user + sys) / wall — effective thread count utilisation.
|
||||
pub fn parallelism(&self) -> f64 {
|
||||
if self.wall_secs > 1e-9 {
|
||||
(self.user_secs + self.sys_secs) / self.wall_secs
|
||||
} else {
|
||||
0.0
|
||||
}
|
||||
}
|
||||
|
||||
/// parallelism / n_cores — fraction of available CPU power used (0..1+).
|
||||
pub fn efficiency(&self, n_cores: usize) -> f64 {
|
||||
self.parallelism() / n_cores as f64
|
||||
}
|
||||
}
|
||||
|
||||
/// Accumulates stage stats and prints a human-readable summary table.
|
||||
#[derive(Default)]
|
||||
pub struct Reporter {
|
||||
stages: Vec<StageStats>,
|
||||
}
|
||||
|
||||
impl Reporter {
|
||||
pub fn new() -> Self {
|
||||
Self::default()
|
||||
}
|
||||
pub fn push(&mut self, stats: StageStats) {
|
||||
self.stages.push(stats);
|
||||
}
|
||||
pub fn stages(&self) -> &[StageStats] {
|
||||
&self.stages
|
||||
}
|
||||
/// Print the summary to stderr.
|
||||
pub fn print(&self) {
|
||||
eprint!("{self}");
|
||||
}
|
||||
}
|
||||
|
||||
// ── diagnosis ─────────────────────────────────────────────────────────────────
|
||||
|
||||
struct Diagnosis {
|
||||
tag: &'static str,
|
||||
detail: Option<String>,
|
||||
}
|
||||
|
||||
// Thresholds are intentionally conservative to avoid false positives.
|
||||
fn diagnose(s: &StageStats, n_cores: usize) -> Diagnosis {
|
||||
let eff = s.efficiency(n_cores);
|
||||
let cpu_pct = eff * 100.0;
|
||||
let io_ops = s.in_blocks + s.out_blocks;
|
||||
|
||||
// swaps > 0 is the only reliable cross-platform indicator of true RAM exhaustion.
|
||||
// ru_majflt is intentionally excluded: on macOS it counts all file-backed mmap
|
||||
// page-ins (even from page cache), making it useless as a memory-pressure signal
|
||||
// for mmap-heavy code. On Linux it is more meaningful, but swaps covers the
|
||||
// severe case on both platforms.
|
||||
if s.swaps > 0 {
|
||||
return Diagnosis {
|
||||
tag: "swapping",
|
||||
detail: Some(format!(
|
||||
"swapped {} time(s) — working set exceeds available RAM",
|
||||
s.swaps,
|
||||
)),
|
||||
};
|
||||
}
|
||||
if eff < 0.3 && io_ops > 100 {
|
||||
return Diagnosis {
|
||||
tag: "disk I/O",
|
||||
detail: Some(format!(
|
||||
"{} block reads + {} writes — CPU at {:.0}%, stage is I/O-bound",
|
||||
s.in_blocks, s.out_blocks, cpu_pct,
|
||||
)),
|
||||
};
|
||||
}
|
||||
if eff < 0.3 && s.vol_ctx > 200 {
|
||||
return Diagnosis {
|
||||
tag: "contention",
|
||||
detail: Some(format!(
|
||||
"{} voluntary context switches — CPU at {:.0}%, possible lock contention or I/O wait",
|
||||
s.vol_ctx, cpu_pct,
|
||||
)),
|
||||
};
|
||||
}
|
||||
Diagnosis {
|
||||
tag: "—",
|
||||
detail: None,
|
||||
}
|
||||
}
|
||||
|
||||
// ── display helpers ───────────────────────────────────────────────────────────
|
||||
|
||||
fn fmt_secs(s: f64) -> String {
|
||||
if s >= 100.0 {
|
||||
format!("{:.0}s", s)
|
||||
} else if s >= 10.0 {
|
||||
format!("{:.1}s", s)
|
||||
} else if s >= 1.0 {
|
||||
format!("{:.2}s", s)
|
||||
} else {
|
||||
format!("{:.0}ms", s * 1000.0)
|
||||
}
|
||||
}
|
||||
|
||||
fn fmt_bytes(b: u64) -> String {
|
||||
if b >= 1 << 30 {
|
||||
format!("{:.1} GB", b as f64 / (1u64 << 30) as f64)
|
||||
} else if b >= 1 << 20 {
|
||||
format!("{:.0} MB", b as f64 / (1u64 << 20) as f64)
|
||||
} else {
|
||||
format!("{:.0} KB", b as f64 / 1024.0)
|
||||
}
|
||||
}
|
||||
|
||||
fn fmt_efficiency(par: f64, n_cores: usize) -> String {
|
||||
format!(
|
||||
"{:.1}×/{} ({:.0}%)",
|
||||
par,
|
||||
n_cores,
|
||||
par / n_cores as f64 * 100.0
|
||||
)
|
||||
}
|
||||
|
||||
// ── Display ───────────────────────────────────────────────────────────────────
|
||||
|
||||
impl fmt::Display for Reporter {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
if self.stages.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let n_cores = effective_parallelism();
|
||||
|
||||
// column widths
|
||||
let nw = self
|
||||
.stages
|
||||
.iter()
|
||||
.map(|s| s.label.len())
|
||||
.max()
|
||||
.unwrap_or(5)
|
||||
.max(5);
|
||||
// efficiency col: worst-case width for this run's n_cores value
|
||||
let ew = format!("{:.1}×/{} (100%)", 99.9f64, n_cores).len();
|
||||
|
||||
let sep_w = nw + 2 + 7 + 2 + ew + 2 + 8 + 2 + 12;
|
||||
let sep = "─".repeat(sep_w);
|
||||
|
||||
// header
|
||||
writeln!(
|
||||
f,
|
||||
"{:<nw$} {:>7} {:>ew$} {:>8} status",
|
||||
"stage", "wall", "efficiency", "peak RSS"
|
||||
)?;
|
||||
writeln!(f, "{sep}")?;
|
||||
|
||||
// compute all diagnoses up front (needed for both table and footnotes)
|
||||
let diagnoses: Vec<Diagnosis> = self.stages.iter().map(|s| diagnose(s, n_cores)).collect();
|
||||
|
||||
// per-stage rows
|
||||
for (s, d) in self.stages.iter().zip(diagnoses.iter()) {
|
||||
writeln!(
|
||||
f,
|
||||
"{:<nw$} {:>7} {:>ew$} {:>8} {}",
|
||||
s.label,
|
||||
fmt_secs(s.wall_secs),
|
||||
fmt_efficiency(s.parallelism(), n_cores),
|
||||
fmt_bytes(s.max_rss_bytes),
|
||||
d.tag,
|
||||
)?;
|
||||
}
|
||||
|
||||
// totals
|
||||
let tw = self.stages.iter().map(|s| s.wall_secs).sum::<f64>();
|
||||
let tu = self.stages.iter().map(|s| s.user_secs).sum::<f64>();
|
||||
let ts = self.stages.iter().map(|s| s.sys_secs).sum::<f64>();
|
||||
let trss = self
|
||||
.stages
|
||||
.iter()
|
||||
.map(|s| s.max_rss_bytes)
|
||||
.max()
|
||||
.unwrap_or(0);
|
||||
let tpar = if tw > 1e-9 { (tu + ts) / tw } else { 0.0 };
|
||||
|
||||
writeln!(f, "{sep}")?;
|
||||
writeln!(
|
||||
f,
|
||||
"{:<nw$} {:>7} {:>ew$} {:>8}",
|
||||
"TOTAL",
|
||||
fmt_secs(tw),
|
||||
fmt_efficiency(tpar, n_cores),
|
||||
fmt_bytes(trss),
|
||||
)?;
|
||||
|
||||
// bottleneck footnotes (only if at least one anomaly detected)
|
||||
let bottlenecks: Vec<(&str, &str)> = self
|
||||
.stages
|
||||
.iter()
|
||||
.zip(diagnoses.iter())
|
||||
.filter_map(|(s, d)| d.detail.as_deref().map(|det| (s.label.as_str(), det)))
|
||||
.collect();
|
||||
|
||||
if !bottlenecks.is_empty() {
|
||||
writeln!(f, "\nBottlenecks:")?;
|
||||
for (label, detail) in &bottlenecks {
|
||||
writeln!(f, " {label} — {detail}")?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user