use crate::error::{Error, Result};
use futures::Stream;
use std::collections::VecDeque;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio::sync::mpsc;
use tracing::{debug, info, warn};
#[derive(Debug, Clone)]
pub struct BulkProcessingConfig {
pub chunk_size: usize,
pub max_memory_usage: usize,
pub concurrency: usize,
pub enable_compression: bool,
pub progress_interval: usize,
}
impl Default for BulkProcessingConfig {
fn default() -> Self {
Self {
chunk_size: 1000,
max_memory_usage: 100 * 1024 * 1024, concurrency: 4,
enable_compression: true,
progress_interval: 10000, }
}
}
#[derive(Debug, Clone)]
pub struct BulkProgress {
pub processed_items: usize,
pub total_items: Option<usize>,
pub processing_rate: f64, pub estimated_completion: Option<std::time::Duration>,
pub memory_usage: usize,
}
pub struct ChunkProcessor<T> {
config: BulkProcessingConfig,
buffer: VecDeque<T>,
processed_count: usize,
start_time: std::time::Instant,
last_progress_report: usize,
}
impl<T> ChunkProcessor<T>
where
T: Clone + Send + Sync,
{
pub fn new(config: BulkProcessingConfig) -> Self {
Self {
config,
buffer: VecDeque::new(),
processed_count: 0,
start_time: std::time::Instant::now(),
last_progress_report: 0,
}
}
pub fn add_items(&mut self, items: Vec<T>) -> Result<()> {
let estimated_memory = self.estimate_memory_usage(&items);
if estimated_memory > self.config.max_memory_usage {
return Err(Error::Custom(
"Memory limit exceeded. Consider reducing chunk size.".to_string(),
));
}
self.buffer.extend(items);
Ok(())
}
pub async fn process_chunks<F, Fut, R>(&mut self, mut processor: F) -> Result<Vec<R>>
where
F: FnMut(Vec<T>) -> Fut,
Fut: std::future::Future<Output = Result<R>>,
R: Send,
{
let mut results = Vec::new();
while !self.buffer.is_empty() {
let chunk_size = std::cmp::min(self.config.chunk_size, self.buffer.len());
let chunk: Vec<T> = self.buffer.drain(..chunk_size).collect();
match processor(chunk).await {
Ok(result) => {
results.push(result);
self.processed_count += chunk_size;
if self.processed_count - self.last_progress_report
>= self.config.progress_interval
{
self.report_progress();
self.last_progress_report = self.processed_count;
}
}
Err(e) => {
warn!("Chunk processing failed: {}", e);
return Err(e);
}
}
}
info!("Completed processing {} items", self.processed_count);
Ok(results)
}
fn estimate_memory_usage(&self, items: &[T]) -> usize {
std::mem::size_of::<T>() * (self.buffer.len() + items.len())
+ 1024 * (self.buffer.len() + items.len())
}
fn report_progress(&self) {
let elapsed = self.start_time.elapsed();
let rate = self.processed_count as f64 / elapsed.as_secs_f64();
debug!(
"Processed {} items, rate: {:.2} items/sec, elapsed: {:?}",
self.processed_count, rate, elapsed
);
}
pub fn get_progress(&self) -> BulkProgress {
let elapsed = self.start_time.elapsed();
let rate = if elapsed.as_secs_f64() > 0.0 {
self.processed_count as f64 / elapsed.as_secs_f64()
} else {
0.0
};
BulkProgress {
processed_items: self.processed_count,
total_items: None, processing_rate: rate,
estimated_completion: None, memory_usage: self.estimate_memory_usage(&[]),
}
}
}
pub struct StreamingProcessor<T> {
sender: mpsc::UnboundedSender<T>,
receiver: mpsc::UnboundedReceiver<T>,
config: BulkProcessingConfig,
}
impl<T> StreamingProcessor<T>
where
T: Send + 'static,
{
pub fn new(config: BulkProcessingConfig) -> Self {
let (sender, receiver) = mpsc::unbounded_channel();
Self {
sender,
receiver,
config,
}
}
pub fn get_sender(&self) -> mpsc::UnboundedSender<T> {
self.sender.clone()
}
pub async fn process_stream<F, Fut, R>(&mut self, mut processor: F) -> Result<Vec<R>>
where
F: FnMut(Vec<T>) -> Fut,
Fut: std::future::Future<Output = Result<R>>,
R: Send,
{
let mut results = Vec::new();
let mut buffer = Vec::with_capacity(self.config.chunk_size);
while let Some(item) = self.receiver.recv().await {
buffer.push(item);
if buffer.len() >= self.config.chunk_size {
let chunk = std::mem::take(&mut buffer);
match processor(chunk).await {
Ok(result) => results.push(result),
Err(e) => return Err(e),
}
}
}
if !buffer.is_empty() {
match processor(buffer).await {
Ok(result) => results.push(result),
Err(e) => return Err(e),
}
}
Ok(results)
}
}
pub struct BulkDataStream<T> {
data: VecDeque<T>,
chunk_size: usize,
}
impl<T> BulkDataStream<T> {
pub fn new(data: Vec<T>, chunk_size: usize) -> Self {
Self {
data: VecDeque::from(data),
chunk_size,
}
}
}
impl<T> Stream for BulkDataStream<T>
where
T: Clone + Unpin,
{
type Item = Vec<T>;
fn poll_next(mut self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
if self.data.is_empty() {
return Poll::Ready(None);
}
let chunk_size = std::cmp::min(self.chunk_size, self.data.len());
let chunk: Vec<T> = self.data.drain(..chunk_size).collect();
Poll::Ready(Some(chunk))
}
}
pub struct DataAggregator<T> {
config: BulkProcessingConfig,
temp_storage: Vec<Vec<T>>,
total_items: usize,
}
impl<T> DataAggregator<T>
where
T: Clone + Send + Sync,
{
pub fn new(config: BulkProcessingConfig) -> Self {
Self {
config,
temp_storage: Vec::new(),
total_items: 0,
}
}
pub fn add_batch(&mut self, batch: Vec<T>) -> Result<()> {
let batch_size = batch.len();
let estimated_new_memory =
self.estimate_total_memory() + self.estimate_batch_memory(&batch);
if estimated_new_memory > self.config.max_memory_usage {
if self.config.enable_compression {
self.compress_oldest_batch()?;
} else {
return Err(Error::Custom(
"Memory limit exceeded and compression disabled".to_string(),
));
}
}
self.temp_storage.push(batch);
self.total_items += batch_size;
Ok(())
}
pub fn get_all_data(&mut self) -> Vec<T> {
let mut all_data = Vec::with_capacity(self.total_items);
for batch in self.temp_storage.drain(..) {
all_data.extend(batch);
}
self.total_items = 0;
all_data
}
pub fn drain_chunks(&mut self, chunk_size: usize) -> Vec<Vec<T>> {
let mut chunks = Vec::new();
let mut current_chunk = Vec::with_capacity(chunk_size);
for batch in self.temp_storage.drain(..) {
for item in batch {
current_chunk.push(item);
if current_chunk.len() >= chunk_size {
chunks.push(std::mem::take(&mut current_chunk));
current_chunk = Vec::with_capacity(chunk_size);
}
}
}
if !current_chunk.is_empty() {
chunks.push(current_chunk);
}
self.total_items = 0;
chunks
}
fn estimate_batch_memory(&self, batch: &[T]) -> usize {
std::mem::size_of::<T>() * batch.len() + 1024 * batch.len() }
fn estimate_total_memory(&self) -> usize {
self.temp_storage
.iter()
.map(|batch| self.estimate_batch_memory(batch))
.sum()
}
fn compress_oldest_batch(&mut self) -> Result<()> {
if self.temp_storage.is_empty() {
return Ok(());
}
warn!("Memory limit reached, removing oldest batch");
if !self.temp_storage.is_empty() {
let removed_batch = self.temp_storage.remove(0);
self.total_items -= removed_batch.len();
}
Ok(())
}
pub fn get_stats(&self) -> (usize, usize, usize) {
(
self.total_items,
self.temp_storage.len(),
self.estimate_total_memory(),
)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_chunk_processor() {
let config = BulkProcessingConfig {
chunk_size: 3,
..Default::default()
};
let mut processor = ChunkProcessor::new(config);
let items = vec![1, 2, 3, 4, 5, 6, 7];
processor.add_items(items).unwrap();
let results = processor
.process_chunks(|chunk| async move {
Ok(chunk.len()) })
.await
.unwrap();
assert_eq!(results, vec![3, 3, 1]); }
#[test]
fn test_data_aggregator() {
let config = BulkProcessingConfig::default();
let mut aggregator = DataAggregator::new(config);
aggregator.add_batch(vec![1, 2, 3]).unwrap();
aggregator.add_batch(vec![4, 5]).unwrap();
let all_data = aggregator.get_all_data();
assert_eq!(all_data, vec![1, 2, 3, 4, 5]);
}
#[tokio::test]
async fn test_bulk_data_stream() {
use futures::StreamExt;
let data = vec![1, 2, 3, 4, 5, 6, 7];
let mut stream = BulkDataStream::new(data, 3);
let chunks: Vec<Vec<i32>> = stream.collect().await;
assert_eq!(chunks, vec![vec![1, 2, 3], vec![4, 5, 6], vec![7]]);
}
}