use crate::metrics;
use crate::models::stream_event::Event;
use crate::pipeline::{
filter::EventFilter, mapper::FieldMapper, soft_delete::SoftDeleteHandler,
transformer::EventTransformer,
};
use std::sync::Arc;
use tokio::sync::{mpsc, RwLock, Semaphore};
use tokio::time::{interval, Duration};
use tracing::{debug, info, warn};
#[derive(Debug, Clone)]
pub struct ParallelConfig {
pub workers_per_table: usize,
pub max_concurrent_events: usize,
pub work_stealing: bool,
pub steal_interval_ms: u64,
}
impl Default for ParallelConfig {
fn default() -> Self {
Self {
workers_per_table: 4,
max_concurrent_events: 1000,
work_stealing: true,
steal_interval_ms: 100,
}
}
}
struct WorkerParams {
worker_id: usize,
table_name: String,
event_queue: Arc<RwLock<Vec<Event>>>,
semaphore: Arc<Semaphore>,
filter: Option<EventFilter>,
transformer: Option<EventTransformer>,
mapper: Option<FieldMapper>,
soft_delete_handler: Option<SoftDeleteHandler>,
destination_tx: mpsc::Sender<Vec<Event>>,
}
pub struct ParallelTableProcessor {
table_name: String,
config: ParallelConfig,
event_queue: Arc<RwLock<Vec<Event>>>,
processing_semaphore: Arc<Semaphore>,
worker_handles: Vec<tokio::task::JoinHandle<()>>,
}
impl ParallelTableProcessor {
pub fn new(table_name: String, config: ParallelConfig) -> Self {
let processing_semaphore = Arc::new(Semaphore::new(config.max_concurrent_events));
Self {
table_name,
config,
event_queue: Arc::new(RwLock::new(Vec::new())),
processing_semaphore,
worker_handles: Vec::new(),
}
}
pub fn start_workers(
&mut self,
filter: Option<EventFilter>,
transformer: Option<EventTransformer>,
mapper: Option<FieldMapper>,
soft_delete_handler: Option<SoftDeleteHandler>,
destination_tx: mpsc::Sender<Vec<Event>>,
shutdown_rx: tokio::sync::watch::Receiver<bool>,
) {
info!(
"Starting {} parallel workers for table '{}'",
self.config.workers_per_table, self.table_name
);
for worker_id in 0..self.config.workers_per_table {
let table_name = self.table_name.clone();
let event_queue = self.event_queue.clone();
let semaphore = self.processing_semaphore.clone();
let filter = filter.clone();
let transformer = transformer.clone();
let mapper = mapper.clone();
let soft_delete_handler = soft_delete_handler.clone();
let destination_tx = destination_tx.clone();
let mut shutdown_rx = shutdown_rx.clone();
let handle = tokio::spawn(async move {
let params = WorkerParams {
worker_id,
table_name,
event_queue,
semaphore,
filter,
transformer,
mapper,
soft_delete_handler,
destination_tx,
};
Self::worker_loop(params, &mut shutdown_rx).await;
});
self.worker_handles.push(handle);
}
}
pub async fn enqueue_events(&self, events: Vec<Event>) {
let mut queue = self.event_queue.write().await;
let event_count = events.len();
queue.extend(events);
let queue_size = queue.len();
debug!(
"Table '{}': Enqueued {} events, queue size: {}",
self.table_name, event_count, queue_size
);
metrics::PARALLEL_QUEUE_SIZE
.with_label_values(&[&self.table_name])
.set(queue_size as f64);
}
pub async fn queue_size(&self) -> usize {
self.event_queue.read().await.len()
}
pub async fn steal_events(&self, count: usize) -> Vec<Event> {
let mut queue = self.event_queue.write().await;
let steal_count = count.min(queue.len() / 2);
if steal_count > 0 {
let mut stolen = Vec::with_capacity(steal_count);
for _ in 0..steal_count {
if let Some(event) = queue.pop() {
stolen.push(event);
}
}
debug!(
"Table '{}': {} events stolen from queue",
self.table_name,
stolen.len()
);
stolen
} else {
Vec::new()
}
}
async fn worker_loop(
params: WorkerParams,
shutdown_rx: &mut tokio::sync::watch::Receiver<bool>,
) {
info!(
"Worker {} for table '{}' started",
params.worker_id, params.table_name
);
let mut processed_count = 0;
let mut check_interval = interval(Duration::from_millis(10));
loop {
tokio::select! {
_ = check_interval.tick() => {
let events_to_process = {
let mut queue = params.event_queue.write().await;
let take_count = 10.min(queue.len()); queue.drain(..take_count).collect::<Vec<_>>()
};
if !events_to_process.is_empty() {
let permits = params.semaphore
.acquire_many(events_to_process.len() as u32)
.await
.unwrap();
let mut processed_events = Vec::new();
for event in events_to_process {
if let Some(processed) = Self::process_single_event(
event,
¶ms.filter,
¶ms.transformer,
¶ms.mapper,
¶ms.soft_delete_handler,
).await {
processed_events.push(processed);
processed_count += 1;
}
}
if !processed_events.is_empty() {
metrics::PARALLEL_WORKER_EVENTS
.with_label_values(&[¶ms.table_name, ¶ms.worker_id.to_string()])
.inc_by(processed_events.len() as f64);
if params.destination_tx.send(processed_events).await.is_err() {
warn!(
"Worker {} for table '{}': Failed to send to destination",
params.worker_id, params.table_name
);
}
}
drop(permits);
let queue_size = params.event_queue.read().await.len();
metrics::PARALLEL_QUEUE_SIZE
.with_label_values(&[¶ms.table_name])
.set(queue_size as f64);
}
}
_ = shutdown_rx.changed() => {
if *shutdown_rx.borrow() {
info!(
"Worker {} for table '{}' shutting down, processed {} events",
params.worker_id, params.table_name, processed_count
);
break;
}
}
}
}
}
async fn process_single_event(
event: Event,
filter: &Option<EventFilter>,
transformer: &Option<EventTransformer>,
mapper: &Option<FieldMapper>,
soft_delete_handler: &Option<SoftDeleteHandler>,
) -> Option<Event> {
if let Some(f) = filter {
match f.should_process(&event) {
Ok(true) => {}
Ok(false) => return None,
Err(e) => {
warn!("Error in filter: {}", e);
return None;
}
}
}
let mut current_event = event;
if let Some(sdh) = soft_delete_handler {
match sdh.transform_event(current_event) {
Some(transformed) => current_event = transformed,
None => return None, }
}
if let Some(t) = transformer {
match t.transform(current_event) {
Ok(Some(transformed)) => current_event = transformed,
Ok(None) => return None,
Err(e) => {
warn!("Error in transformer: {}", e);
return None;
}
}
}
if let Some(m) = mapper {
match m.map_event(current_event) {
Ok(mapped) => current_event = mapped,
Err(e) => {
warn!("Error in mapper: {}", e);
return None;
}
}
}
Some(current_event)
}
pub async fn shutdown(mut self) {
for handle in self.worker_handles.drain(..) {
let _ = handle.await;
}
}
}
type ProcessorRegistry = Arc<RwLock<Vec<(String, Arc<ParallelTableProcessor>)>>>;
pub struct WorkStealingCoordinator {
processors: ProcessorRegistry,
config: ParallelConfig,
}
impl WorkStealingCoordinator {
pub fn new(config: ParallelConfig) -> Self {
Self {
processors: Arc::new(RwLock::new(Vec::new())),
config,
}
}
pub async fn register_processor(
&self,
table_name: String,
processor: Arc<ParallelTableProcessor>,
) {
let mut processors = self.processors.write().await;
processors.push((table_name, processor));
}
pub fn start(
&self,
shutdown_rx: tokio::sync::watch::Receiver<bool>,
) -> tokio::task::JoinHandle<()> {
let processors = self.processors.clone();
let steal_interval = Duration::from_millis(self.config.steal_interval_ms);
tokio::spawn(async move {
let mut interval = interval(steal_interval);
let mut shutdown_rx = shutdown_rx;
loop {
tokio::select! {
_ = interval.tick() => {
Self::balance_work(&processors).await;
}
_ = shutdown_rx.changed() => {
if *shutdown_rx.borrow() {
info!("Work stealing coordinator shutting down");
break;
}
}
}
}
})
}
async fn balance_work(processors: &ProcessorRegistry) {
let processors = processors.read().await;
if processors.len() < 2 {
return; }
let mut queue_info = Vec::new();
for (table_name, processor) in processors.iter() {
let size = processor.queue_size().await;
queue_info.push((table_name.clone(), processor.clone(), size));
}
queue_info.sort_by_key(|&(_, _, size)| size);
if let (
Some((small_table, small_proc, small_size)),
Some((large_table, large_proc, large_size)),
) = (queue_info.first(), queue_info.last())
{
if *large_size > *small_size + 100 {
let steal_count = (large_size - small_size) / 2;
let stolen = large_proc.steal_events(steal_count).await;
if !stolen.is_empty() {
debug!(
"Work stealing: Moved {} events from '{}' to '{}'",
stolen.len(),
large_table,
small_table
);
metrics::WORK_STEALING_OPERATIONS
.with_label_values(&[large_table, small_table])
.inc();
small_proc.enqueue_events(stolen).await;
}
}
}
}
}