#![cfg_attr(test, allow(clippy::unwrap_used, reason = "test scope"))]
use moirai_utils::CACHE_LINE_SIZE;
const DEFAULT_RING_BUFFER_CAPACITY: usize = 1024;
use moirai_core::channel::{ChannelConfig, ChannelError, UnifiedReceiver, UnifiedSender};
use moirai_core::memory::MemoryPool;
use std::marker::PhantomData;
use std::sync::Arc;
pub struct StreamingIterator<T> {
receiver: UnifiedReceiver<T>,
buffer: std::vec::IntoIter<T>,
batch_size: usize,
finished: bool,
}
impl<T> StreamingIterator<T> {
pub fn new(receiver: UnifiedReceiver<T>, batch_size: usize) -> Self {
Self {
receiver,
buffer: Vec::new().into_iter(),
batch_size,
finished: false,
}
}
fn fill_buffer(&mut self) -> bool {
if self.finished {
return false;
}
let new_items = self.receiver.recv_batch(self.batch_size);
if new_items.is_empty() {
if self.receiver.is_closed() {
self.finished = true;
return false;
}
match self.receiver.try_recv() {
Ok(item) => {
self.buffer = vec![item].into_iter();
true
}
Err(_) => false,
}
} else {
self.buffer = new_items.into_iter();
true
}
}
}
impl<T> Iterator for StreamingIterator<T> {
type Item = T;
fn next(&mut self) -> Option<Self::Item> {
if let Some(item) = self.buffer.next() {
return Some(item);
}
if self.fill_buffer() {
self.buffer.next()
} else {
None
}
}
fn size_hint(&self) -> (usize, Option<usize>) {
(self.buffer.len(), None)
}
}
pub struct ProducerConsumerPair<T> {
sender: UnifiedSender<T>,
receiver: UnifiedReceiver<T>,
config: ChannelConfig,
}
impl<T> ProducerConsumerPair<T> {
pub fn new(capacity: usize) -> Result<Self, ChannelError> {
let config = ChannelConfig {
capacity,
enable_batching: true,
batch_size: capacity.min(64),
..Default::default()
};
let (sender, receiver) = moirai_core::channel::unified_channel_with_config(config.clone())?;
Ok(Self {
sender,
receiver,
config,
})
}
pub fn producer(&self) -> &UnifiedSender<T> {
&self.sender
}
pub fn consumer(&self) -> &UnifiedReceiver<T> {
&self.receiver
}
pub fn into_streaming_iter(self) -> StreamingIterator<T> {
StreamingIterator::new(self.receiver, self.config.batch_size)
}
pub fn split(self) -> (UnifiedSender<T>, StreamingIterator<T>) {
let iter = StreamingIterator::new(self.receiver, self.config.batch_size);
(self.sender, iter)
}
}
pub trait PipelineStage<Input, Output>: Send + Sync {
fn process_batch(&self, inputs: Vec<Input>) -> Vec<Output>;
fn preferred_batch_size(&self) -> usize {
64
}
}
pub struct MapStage<F, Input, Output> {
func: F,
batch_size: usize,
_phantom: PhantomData<(Input, Output)>,
}
impl<F, Input, Output> MapStage<F, Input, Output>
where
F: Fn(Input) -> Output + Send + Sync,
{
pub fn new(func: F) -> Self {
Self {
func,
batch_size: 64,
_phantom: PhantomData,
}
}
pub fn with_batch_size(mut self, batch_size: usize) -> Self {
self.batch_size = batch_size;
self
}
}
impl<F, Input, Output> PipelineStage<Input, Output> for MapStage<F, Input, Output>
where
F: Fn(Input) -> Output + Send + Sync,
Input: Send + Sync,
Output: Send + Sync,
{
fn process_batch(&self, inputs: Vec<Input>) -> Vec<Output> {
inputs.into_iter().map(&self.func).collect()
}
fn preferred_batch_size(&self) -> usize {
self.batch_size
}
}
pub struct FilterStage<F, T> {
predicate: F,
batch_size: usize,
_phantom: PhantomData<T>,
}
impl<F, T> FilterStage<F, T>
where
F: Fn(&T) -> bool + Send + Sync,
{
pub fn new(predicate: F) -> Self {
Self {
predicate,
batch_size: 64,
_phantom: PhantomData,
}
}
pub fn with_batch_size(mut self, batch_size: usize) -> Self {
self.batch_size = batch_size;
self
}
}
impl<F, T> PipelineStage<T, T> for FilterStage<F, T>
where
F: Fn(&T) -> bool + Send + Sync,
T: Send + Sync,
{
fn process_batch(&self, inputs: Vec<T>) -> Vec<T> {
inputs.into_iter().filter(|x| (self.predicate)(x)).collect()
}
fn preferred_batch_size(&self) -> usize {
self.batch_size
}
}
pub struct IteratorPipeline<T> {
source: Option<Vec<T>>,
channel_capacity: usize,
memory_pool: Option<Arc<MemoryPool<T>>>,
}
impl<T> IteratorPipeline<T> {
pub fn from_vec(data: Vec<T>) -> Self {
Self {
source: Some(data),
channel_capacity: DEFAULT_RING_BUFFER_CAPACITY,
memory_pool: None,
}
}
pub fn with_channel_capacity(mut self, capacity: usize) -> Self {
self.channel_capacity = capacity;
self
}
pub fn with_memory_pool(mut self, pool: Arc<MemoryPool<T>>) -> Self {
self.memory_pool = Some(pool);
self
}
pub fn map<F, R>(self, func: F) -> MappedPipeline<T, R, F>
where
F: Fn(T) -> R + Send + Sync + 'static,
T: Send + Sync + 'static,
R: Send + Sync + 'static,
{
MappedPipeline {
source: self.source.unwrap_or_default(),
stage: MapStage::new(func),
channel_capacity: self.channel_capacity,
memory_pool: self.memory_pool,
}
}
pub fn filter<F>(self, predicate: F) -> FilteredPipeline<T, F>
where
F: Fn(&T) -> bool + Send + Sync + 'static,
T: Send + Sync + 'static,
{
FilteredPipeline {
source: self.source.unwrap_or_default(),
stage: FilterStage::new(predicate),
channel_capacity: self.channel_capacity,
memory_pool: self.memory_pool,
}
}
pub fn collect(self) -> Vec<T> {
self.source.unwrap_or_default()
}
}
pub struct MappedPipeline<Input, Output, F> {
source: Vec<Input>,
stage: MapStage<F, Input, Output>,
channel_capacity: usize,
memory_pool: Option<Arc<MemoryPool<Input>>>,
}
impl<Input, Output, F> MappedPipeline<Input, Output, F>
where
F: Fn(Input) -> Output + Send + Sync + 'static,
Input: Send + Sync + 'static,
Output: Send + Sync + 'static,
{
pub fn map<G, R>(self, func: G) -> MappedPipeline<Output, R, G>
where
G: Fn(Output) -> R + Send + Sync + 'static,
R: Send + 'static,
{
let intermediate_results = self.stage.process_batch(self.source);
MappedPipeline {
source: intermediate_results,
stage: MapStage::new(func),
channel_capacity: self.channel_capacity,
memory_pool: self.memory_pool.map(|_| {
Arc::new(MemoryPool::new(256))
}),
}
}
pub fn filter<G>(self, predicate: G) -> FilteredPipeline<Output, G>
where
G: Fn(&Output) -> bool + Send + Sync + 'static,
{
let intermediate_results = self.stage.process_batch(self.source);
FilteredPipeline {
source: intermediate_results,
stage: FilterStage::new(predicate),
channel_capacity: self.channel_capacity,
memory_pool: None, }
}
pub fn collect(self) -> Vec<Output> {
if let Some(_pool) = &self.memory_pool {
}
self.stage.process_batch(self.source)
}
pub async fn collect_parallel(self) -> Vec<Output> {
self.collect()
}
}
pub struct FilteredPipeline<T, F> {
source: Vec<T>,
stage: FilterStage<F, T>,
channel_capacity: usize,
memory_pool: Option<Arc<MemoryPool<T>>>,
}
impl<T, F> FilteredPipeline<T, F>
where
F: Fn(&T) -> bool + Send + Sync + 'static,
T: Send + Sync + 'static,
{
pub fn map<G, R>(self, func: G) -> MappedPipeline<T, R, G>
where
G: Fn(T) -> R + Send + Sync + 'static,
R: Send + 'static,
{
let intermediate_results = self.stage.process_batch(self.source);
MappedPipeline {
source: intermediate_results,
stage: MapStage::new(func),
channel_capacity: self.channel_capacity,
memory_pool: self.memory_pool,
}
}
pub fn filter<G>(self, predicate: G) -> FilteredPipeline<T, G>
where
G: Fn(&T) -> bool + Send + Sync + 'static,
{
let intermediate_results = self.stage.process_batch(self.source);
FilteredPipeline {
source: intermediate_results,
stage: FilterStage::new(predicate),
channel_capacity: self.channel_capacity,
memory_pool: self.memory_pool,
}
}
pub fn collect(self) -> Vec<T> {
self.stage.process_batch(self.source)
}
pub async fn collect_parallel(self) -> Vec<T> {
self.collect()
}
}
pub struct CacheAwareIterator<T> {
data: Vec<T>,
chunk_size: usize,
current_chunk: usize,
current_pos: usize,
}
impl<T: Clone> CacheAwareIterator<T> {
pub fn new(data: Vec<T>) -> Self {
let chunk_size = CACHE_LINE_SIZE / std::mem::size_of::<T>().max(1);
Self {
data,
chunk_size,
current_chunk: 0,
current_pos: 0,
}
}
pub fn process_chunks<F, R>(self, processor: F) -> Vec<R>
where
F: FnMut(&[T]) -> R,
{
self.data.chunks(self.chunk_size).map(processor).collect()
}
pub fn map_with_prefetch<F, R>(self, func: F) -> Vec<R>
where
F: Fn(T) -> R + Send + Sync,
T: Send,
R: Send,
{
self.data.into_iter().map(func).collect()
}
}
impl<T: Clone> Iterator for CacheAwareIterator<T> {
type Item = T;
fn next(&mut self) -> Option<Self::Item> {
if self.current_pos >= self.data.len() {
return None;
}
let item = self.data[self.current_pos].clone();
self.current_pos += 1;
if self.current_pos.is_multiple_of(self.chunk_size) {
self.current_chunk += 1;
}
Some(item)
}
fn size_hint(&self) -> (usize, Option<usize>) {
let remaining = self.data.len() - self.current_pos;
(remaining, Some(remaining))
}
}
impl<T: Clone> ExactSizeIterator for CacheAwareIterator<T> {}
pub fn streaming_from_channel<T>(receiver: UnifiedReceiver<T>) -> StreamingIterator<T> {
StreamingIterator::new(receiver, 64)
}
pub fn producer_consumer_channel<T>(
capacity: usize,
) -> Result<ProducerConsumerPair<T>, ChannelError> {
ProducerConsumerPair::new(capacity)
}
pub fn cache_aware_iter<T: Clone>(data: Vec<T>) -> CacheAwareIterator<T> {
CacheAwareIterator::new(data)
}
pub fn pipeline<T>(data: Vec<T>) -> IteratorPipeline<T> {
IteratorPipeline::from_vec(data)
}
#[cfg(test)]
mod tests {
use super::*;
use moirai_core::channel::unified;
#[test]
fn test_streaming_iterator() {
let (sender, receiver) = unified::unified_channel::<i32>(16).unwrap();
for i in 0..10 {
sender.send(i).unwrap();
}
let mut iter = StreamingIterator::new(receiver, 5);
let mut collected = Vec::new();
for _ in 0..5 {
if let Some(item) = iter.next() {
collected.push(item);
}
}
assert_eq!(collected.len(), 5);
}
#[test]
fn test_producer_consumer_pair() {
let pair = ProducerConsumerPair::<i32>::new(32).unwrap();
let producer = pair.producer();
for i in 0..5 {
producer.send(i).unwrap();
}
let iter = pair.into_streaming_iter();
let collected: Vec<_> = iter.take(5).collect();
assert_eq!(collected, vec![0, 1, 2, 3, 4]);
}
#[test]
fn test_pipeline_basic() {
let data = vec![1, 2, 3, 4, 5];
let result = pipeline(data).map(|x| x * 2).filter(|&x| x > 4).collect();
assert_eq!(result, vec![6, 8, 10]);
}
#[test]
fn test_cache_aware_iterator() {
let data = (0..100).collect::<Vec<i32>>();
let iter = CacheAwareIterator::new(data.clone());
let collected: Vec<_> = iter.take(10).collect();
assert_eq!(collected, (0..10).collect::<Vec<_>>());
}
#[test]
fn test_pipeline_chaining() {
let data = vec![1, 2, 3, 4, 5, 6, 7, 8, 9, 10];
let result = pipeline(data)
.filter(|&x| x % 2 == 0) .map(|x| x * x) .filter(|&x| x < 50) .collect();
assert_eq!(result, vec![4, 16, 36]);
}
}