use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tenshift_core::sample::{Sample, Tensor};
use tenshift_core::sources::MemorySource;
use tenshift_core::{ErrorPolicy, Pipeline};
#[test]
fn test_empty_pipeline() {
let source = MemorySource::new("empty", Vec::<Sample>::new());
let mut iter = Pipeline::from_source(source).workers(2).start().unwrap();
let count = iter.by_ref().count();
assert_eq!(count, 0, "Empty pipeline should yield zero batches");
let stats = iter.stats();
assert_eq!(stats.items_yielded, 0);
assert_eq!(stats.errors_skipped, 0);
}
#[test]
fn test_single_item_pipeline() {
let samples = vec![Sample::new().with("x", Tensor::i64(&[1], vec![1]))];
let source = MemorySource::new("single", samples);
let mut iter = Pipeline::from_source(source).workers(2).start().unwrap();
let batch = iter.next();
assert!(batch.is_some(), "Single item should yield one batch");
assert_eq!(batch.unwrap().len(), 1);
let stats = iter.stats();
assert_eq!(stats.items_yielded, 1);
}
#[test]
fn test_two_items_no_batch() {
let samples = vec![
Sample::new().with("x", Tensor::i64(&[1], vec![1])),
Sample::new().with("x", Tensor::i64(&[2], vec![1])),
];
let source = MemorySource::new("two", samples);
let mut iter = Pipeline::from_source(source).workers(2).start().unwrap();
let count = iter.by_ref().count();
assert_eq!(
count, 2,
"Two items should yield two batches without batching"
);
}
#[test]
fn test_100k_items() {
let samples: Vec<Sample> = (0..100_000)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("100k", samples);
let start = Instant::now();
let mut iter = Pipeline::from_source(source).workers(4).start().unwrap();
let total: usize = iter.by_ref().map(|batch| batch.len()).sum();
let elapsed = start.elapsed();
assert_eq!(total, 100_000, "All 100K items should be processed");
let throughput = 100_000.0 / elapsed.as_secs_f64();
eprintln!("100K items throughput: {:.0} items/sec", throughput);
}
#[test]
fn test_zero_workers_auto_correct() {
let samples = vec![Sample::new().with("x", Tensor::i64(&[1], vec![1]))];
let source = MemorySource::new("zero_workers", samples);
let pipeline = Pipeline::from_source(source).workers(0);
let mut iter = pipeline.start().unwrap();
let batch = iter.next();
assert!(batch.is_some(), "Pipeline should work even with workers(0)");
}
#[test]
fn test_default_scaling() {
let samples = vec![Sample::new().with("x", Tensor::i64(&[1], vec![1]))];
let source = MemorySource::new("scaling", samples);
let pipeline = Pipeline::from_source(source);
let mut iter = pipeline.start().unwrap();
assert!(iter.next().is_some());
}
#[test]
fn test_backpressure_slow_consumer() {
let samples: Vec<Sample> = (0..1000)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("backpressure", samples);
let mut iter = Pipeline::from_source(source)
.workers(4)
.prefetch(2) .start()
.unwrap();
let mut count = 0;
for _ in &mut iter {
count += 1;
std::thread::sleep(Duration::from_millis(1));
}
assert!(count > 0, "Slow consumer should still receive all batches");
}
#[test]
fn test_shutdown_mid_processing() {
let samples: Vec<Sample> = (0..10_000)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("shutdown", samples);
let mut iter = Pipeline::from_source(source).workers(4).start().unwrap();
let _ = iter.next();
let _ = iter.next();
iter.stop();
drop(iter);
}
#[test]
fn test_rapid_create_destroy() {
let samples: Vec<Sample> = (0..100)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
for i in 0..100 {
let source = MemorySource::new(format!("rapid_{}", i), samples.clone());
let iter = Pipeline::from_source(source).workers(2).start().unwrap();
drop(iter);
}
}
#[test]
fn test_throughput_baseline_1m_items() {
let samples: Vec<Sample> = (0..1_000_000)
.map(|i| Sample::new().with("x", Tensor::i64(&[i % 1000], vec![1])))
.collect();
let source = MemorySource::new("1m_baseline", samples);
let start = Instant::now();
let mut iter = Pipeline::from_source(source).workers(4).start().unwrap();
let total: usize = iter.by_ref().map(|batch| batch.len()).sum();
let elapsed = start.elapsed();
assert_eq!(total, 1_000_000);
let items_per_sec = 1_000_000.0 / elapsed.as_secs_f64();
eprintln!("1M items throughput: {:.0} items/sec", items_per_sec);
assert!(
items_per_sec > 10_000.0,
"Throughput should exceed 10K items/sec, got {:.0}",
items_per_sec
);
}
#[test]
fn test_config_zero_batch_auto_correct() {
let samples = vec![Sample::new().with("x", Tensor::i64(&[1], vec![1]))];
let source = MemorySource::new("batch_zero", samples);
let mut iter = Pipeline::from_source(source).batch(0).start().unwrap();
let batch = iter.next();
assert!(batch.is_some());
}
#[test]
fn test_config_zero_shuffle_auto_correct() {
let samples = vec![Sample::new().with("x", Tensor::i64(&[1], vec![1]))];
let source = MemorySource::new("shuffle_zero", samples);
let mut iter = Pipeline::from_source(source).shuffle(0).start().unwrap();
let batch = iter.next();
assert!(batch.is_some());
}
#[test]
fn test_config_zero_chunk_auto_correct() {
let samples = vec![Sample::new().with("x", Tensor::i64(&[1], vec![1]))];
let source = MemorySource::new("chunk_zero", samples);
let mut iter = Pipeline::from_source(source).chunk_size(0).start().unwrap();
let batch = iter.next();
assert!(batch.is_some());
}
#[test]
fn test_config_zero_prefetch_auto_correct() {
let samples = vec![Sample::new().with("x", Tensor::i64(&[1], vec![1]))];
let source = MemorySource::new("prefetch_zero", samples);
let mut iter = Pipeline::from_source(source).prefetch(0).start().unwrap();
let batch = iter.next();
assert!(batch.is_some());
}
#[test]
fn test_error_policy_skip() {
let samples: Vec<Sample> = (0..100)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("error_skip", samples);
let counter = Arc::new(AtomicUsize::new(0));
let counter_clone = Arc::clone(&counter);
let mut iter = Pipeline::from_source(source)
.workers(2)
.on_error(ErrorPolicy::Skip)
.map(move |mut s| {
let count = counter_clone.fetch_add(1, Ordering::SeqCst);
if count == 50 {
return Err(tenshift_core::error::Error::TransformFailed {
index: count as u64,
reason: "intentional test error".to_string(),
});
}
s.insert("processed", Tensor::i64(&[1], vec![1]));
Ok(s)
})
.start()
.unwrap();
let total: usize = iter.by_ref().map(|batch| batch.len()).sum();
assert_eq!(total, 99, "Skip policy should continue after error");
assert_eq!(iter.stats().errors_skipped, 1);
}
#[test]
fn test_flat_map_expansion() {
let samples: Vec<Sample> = (0..10)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("flat_map", samples);
let mut iter = Pipeline::from_source(source)
.workers(2)
.flat_map(|s| {
let mut results = Vec::new();
for j in 0..3 {
let mut new_s = s.clone();
new_s.insert("expanded", Tensor::i64(&[j], vec![1]));
results.push(new_s);
}
Ok(results)
})
.start()
.unwrap();
let total: usize = iter.by_ref().map(|batch| batch.len()).sum();
assert_eq!(total, 30, "Each of 10 samples should expand to 3");
}
#[test]
fn test_filter_removes_samples() {
let samples: Vec<Sample> = (0..100)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("filter", samples);
let mut iter = Pipeline::from_source(source)
.workers(2)
.filter(|s| {
if let Some(t) = s.get("x") {
if let Ok(vals) = t.try_as_i64() {
return vals.first().map(|&v| v % 2 == 0).unwrap_or(false);
}
}
false
})
.start()
.unwrap();
let total: usize = iter.by_ref().map(|batch| batch.len()).sum();
assert_eq!(total, 50, "Filter should keep exactly half the samples");
}
#[test]
fn test_multi_epoch() {
let samples: Vec<Sample> = (0..10)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("epochs", samples);
let mut iter = Pipeline::from_source(source)
.workers(2)
.epochs(3)
.start()
.unwrap();
let total: usize = iter.by_ref().map(|batch| batch.len()).sum();
assert_eq!(total, 30, "3 epochs of 10 samples = 30 total");
}
#[test]
fn test_batch_collation_shape() {
let samples: Vec<Sample> = (0..32)
.map(|i| {
Sample::new()
.with("image", Tensor::f32(&vec![i as f32; 784], vec![784]))
.with("label", Tensor::i64(&[i % 10], vec![1]))
})
.collect();
let source = MemorySource::new("batch_shape", samples);
let mut iter = Pipeline::from_source(source)
.workers(2)
.batch(16)
.start()
.unwrap();
let batch = iter.next().unwrap();
assert_eq!(
batch.len(),
1,
"With collation, each batch is one collated sample"
);
let collated = &batch[0];
let image = collated.get("image").expect("image field should exist");
let label = collated.get("label").expect("label field should exist");
assert_eq!(
image.shape(),
&[16, 784],
"Image should be [batch, features]"
);
assert_eq!(label.shape(), &[16, 1], "Label should be [batch, 1]");
}
#[test]
fn test_drop_last() {
let samples: Vec<Sample> = (0..35)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("drop_last", samples);
let mut iter = Pipeline::from_source(source)
.workers(2)
.batch(16)
.drop_last(true)
.start()
.unwrap();
let total: usize = iter.by_ref().map(|batch| batch.len()).sum();
assert_eq!(total, 2, "Should have 2 batches with drop_last");
}
#[test]
fn test_empty_batch_after_filter() {
let samples: Vec<Sample> = (0..10)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("empty_batch", samples);
let mut iter = Pipeline::from_source(source)
.workers(2)
.filter(|_| false) .batch(5)
.start()
.unwrap();
let count = iter.by_ref().count();
assert_eq!(count, 0, "All filtered out should yield nothing");
}
#[test]
fn test_stats_accuracy() {
let samples: Vec<Sample> = (0..100)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("stats", samples);
let mut iter = Pipeline::from_source(source)
.workers(2)
.batch(10)
.start()
.unwrap();
let count = iter.by_ref().count();
assert_eq!(count, 10, "Should have 10 batches of 10 samples");
let stats = iter.stats();
assert_eq!(stats.items_yielded, 10);
assert_eq!(stats.errors_skipped, 0);
assert!(stats.elapsed > Duration::ZERO);
assert!(stats.throughput > 0.0);
}
#[test]
fn test_map_sample_integrity() {
let samples: Vec<Sample> = (0..50)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("map_integrity", samples);
let mut iter = Pipeline::from_source(source)
.workers(4)
.map(|mut s| {
if let Some(t) = s.get("x") {
if let Ok(vals) = t.try_as_i64() {
let new_val = vals[0] * 2;
s.insert("x", Tensor::i64(&[new_val], vec![1]));
}
}
Ok(s)
})
.start()
.unwrap();
let mut seen_values = std::collections::HashSet::new();
for batch in &mut iter {
for sample in batch {
if let Some(t) = sample.get("x") {
if let Ok(vals) = t.try_as_i64() {
let val = vals[0];
assert!(val % 2 == 0, "All values should be even (doubled)");
assert!(val < 100, "Doubled values should be < 100");
seen_values.insert(val);
}
}
}
}
assert_eq!(
seen_values.len(),
50,
"Should have 50 unique doubled values"
);
}
#[test]
fn test_multiple_stages() {
let samples: Vec<Sample> = (0..100)
.map(|i| Sample::new().with("x", Tensor::i64(&[i], vec![1])))
.collect();
let source = MemorySource::new("multi_stage", samples);
let mut iter = Pipeline::from_source(source)
.workers(2)
.map(Ok)
.filter(|_| true)
.map(Ok)
.batch(20)
.start()
.unwrap();
let total: usize = iter.by_ref().map(|batch| batch.len()).sum();
assert_eq!(total, 5, "100 samples / 20 batch = 5 batches");
}
#[test]
fn test_large_samples_bounded_memory() {
let large_data = vec![0u8; 1024 * 1024];
let samples: Vec<Sample> = (0..10)
.map(|i| {
Sample::new()
.with("data", Tensor::u8(large_data.clone(), vec![1024 * 1024]))
.with("idx", Tensor::i64(&[i], vec![1]))
})
.collect();
let source = MemorySource::new("large_samples", samples);
let mut iter = Pipeline::from_source(source)
.workers(2)
.prefetch(2) .start()
.unwrap();
let total: usize = iter.by_ref().map(|batch| batch.len()).sum();
assert_eq!(total, 10, "Should process all 10 large samples");
}