use std::fmt;
use std::sync::Arc;
use crate::logging::{DualObserver, ExtractionObserver};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum Parallelism {
#[default]
Sequential,
Parallel(Option<usize>),
}
#[derive(Debug, Clone, Copy)]
pub enum ExtractionLevel {
Index(usize),
TargetMpp(f64),
}
#[derive(Clone, Default)]
pub struct ExtractionOptions {
pub parallelism: Parallelism,
pub min_tissue_fraction: Option<f32>,
pub observer: Option<Arc<dyn ExtractionObserver>>,
pub normalize_stain: bool,
pub extraction_level: Option<ExtractionLevel>,
}
impl fmt::Debug for ExtractionOptions {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("ExtractionOptions")
.field("parallelism", &self.parallelism)
.field("min_tissue_fraction", &self.min_tissue_fraction)
.field("observer", &self.observer.is_some())
.field("normalize_stain", &self.normalize_stain)
.field("extraction_level", &self.extraction_level)
.finish()
}
}
impl ExtractionOptions {
pub fn sequential() -> Self {
Self {
parallelism: Parallelism::Sequential,
min_tissue_fraction: None,
observer: None,
normalize_stain: false,
extraction_level: None,
}
}
pub fn parallel() -> Self {
Self {
parallelism: Parallelism::Parallel(None),
min_tissue_fraction: None,
observer: None,
normalize_stain: false,
extraction_level: None,
}
}
pub fn parallel_with_threads(threads: usize) -> Self {
Self {
parallelism: Parallelism::Parallel(Some(threads)),
min_tissue_fraction: None,
observer: None,
normalize_stain: false,
extraction_level: None,
}
}
pub fn with_min_tissue_fraction(mut self, min_fraction: f32) -> Self {
self.min_tissue_fraction = Some(min_fraction);
self
}
pub fn with_observer(mut self, observer: impl ExtractionObserver + 'static) -> Self {
self.add_observer(Arc::new(observer));
self
}
pub fn with_shared_observer(mut self, observer: Arc<dyn ExtractionObserver>) -> Self {
self.add_observer(observer);
self
}
pub fn with_logging(self) -> Self {
self.with_observer(crate::logging::SimpleLogging)
}
pub fn with_progress_bar(self) -> Self {
self.with_observer(crate::logging::ProgressBar::new())
}
pub fn with_stain_normalization(mut self) -> Self {
self.normalize_stain = true;
self
}
pub fn with_level(mut self, level_idx: usize) -> Self {
self.extraction_level = Some(ExtractionLevel::Index(level_idx));
self
}
pub fn with_target_mpp(mut self, target_mpp: f64) -> Self {
self.extraction_level = Some(ExtractionLevel::TargetMpp(target_mpp));
self
}
fn add_observer(&mut self, observer: Arc<dyn ExtractionObserver>) {
self.observer = Some(match self.observer.take() {
Some(existing) => Arc::new(DualObserver(existing, observer)),
None => observer,
})
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicUsize, Ordering};
#[derive(Default)]
struct CountingObserver(AtomicUsize);
impl ExtractionObserver for CountingObserver {
fn on_extraction_start(&self, _level_idx: usize, _total_tiles: usize) {
self.0.fetch_add(1, Ordering::Relaxed);
}
}
#[test]
fn sequential_and_parallel_constructors_set_expected_defaults() {
assert_eq!(
ExtractionOptions::sequential().parallelism,
Parallelism::Sequential
);
assert_eq!(
ExtractionOptions::parallel().parallelism,
Parallelism::Parallel(None)
);
assert_eq!(
ExtractionOptions::parallel_with_threads(4).parallelism,
Parallelism::Parallel(Some(4))
);
}
#[test]
fn with_level_and_with_target_mpp_overwrite_each_other() {
let options = ExtractionOptions::sequential()
.with_level(2)
.with_target_mpp(0.5);
assert!(matches!(
options.extraction_level,
Some(ExtractionLevel::TargetMpp(mpp)) if mpp == 0.5
));
}
#[test]
fn chained_observers_all_receive_events() {
let first = Arc::new(CountingObserver::default());
let second = Arc::new(CountingObserver::default());
let options = ExtractionOptions::sequential()
.with_shared_observer(first.clone())
.with_shared_observer(second.clone());
options
.observer
.as_ref()
.expect("observer should be set")
.on_extraction_start(0, 10);
assert_eq!(first.0.load(Ordering::Relaxed), 1);
assert_eq!(second.0.load(Ordering::Relaxed), 1);
}
}