use std::sync::Arc;
use async_trait::async_trait;
use tokio::sync::mpsc;
use super::checkpoint::CdcError;
use super::event::ChangeEvent;
#[async_trait]
pub trait CdcSink: Send + Sync {
async fn handle(&self, event: &ChangeEvent) -> Result<(), CdcError>;
}
pub struct CdcEventDispatcher {
sinks: Vec<Arc<dyn CdcSink>>,
backpressure_capacity: usize,
}
impl CdcEventDispatcher {
pub fn new(sinks: Vec<Arc<dyn CdcSink>>, capacity: usize) -> Self {
Self {
sinks,
backpressure_capacity: capacity,
}
}
pub async fn dispatch(&self, event: ChangeEvent) -> Result<(), CdcError> {
for sink in &self.sinks {
sink.handle(&event).await?;
}
Ok(())
}
pub async fn dispatch_batch(&self, events: Vec<ChangeEvent>) -> Result<usize, CdcError> {
let mut count = 0;
for event in events {
self.dispatch(event).await?;
count += 1;
}
Ok(count)
}
pub async fn dispatch_and_confirm(
&self,
event: ChangeEvent,
checkpoint: &super::checkpoint::SharedCheckpointStore,
) -> Result<bool, CdcError> {
match self.dispatch(event.clone()).await {
Ok(()) => {
checkpoint
.save_checkpoint_with_retry(&event.position, 3)
.await?;
Ok(true)
}
Err(e) => Err(e),
}
}
pub async fn run(&self, mut receiver: mpsc::Receiver<ChangeEvent>) -> Result<usize, CdcError> {
let mut count = 0;
while let Some(event) = receiver.recv().await {
self.dispatch(event).await?;
count += 1;
}
Ok(count)
}
pub fn sink_count(&self) -> usize {
self.sinks.len()
}
pub fn backpressure_capacity(&self) -> usize {
self.backpressure_capacity
}
}
pub struct MemorySink {
events: parking_lot::RwLock<Vec<ChangeEvent>>,
}
impl MemorySink {
pub fn new() -> Self {
Self {
events: parking_lot::RwLock::new(Vec::new()),
}
}
pub fn events(&self) -> Vec<ChangeEvent> {
self.events.read().clone()
}
pub fn count(&self) -> usize {
self.events.read().len()
}
}
impl Default for MemorySink {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl CdcSink for MemorySink {
async fn handle(&self, event: &ChangeEvent) -> Result<(), CdcError> {
self.events.write().push(event.clone());
Ok(())
}
}
pub struct TableFilterSink {
table: String,
inner: Arc<dyn CdcSink>,
}
impl TableFilterSink {
pub fn new(table: &str, inner: Arc<dyn CdcSink>) -> Self {
Self {
table: table.to_string(),
inner,
}
}
}
#[async_trait]
impl CdcSink for TableFilterSink {
async fn handle(&self, event: &ChangeEvent) -> Result<(), CdcError> {
if event.source_table == self.table {
self.inner.handle(event).await?;
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::super::event::ChangeEventType;
use super::*;
fn make_event(table: &str, pos: u64) -> ChangeEvent {
ChangeEvent::new(
ChangeEventType::Insert,
"test_db",
table,
serde_json::json!({"id": pos}),
super::super::event::ChangePosition::MysqlBinlog {
filename: "bin.001".to_string(),
position: pos,
},
1000 + pos,
)
}
#[tokio::test]
async fn dispatch_single_sink() {
let sink = Arc::new(MemorySink::new());
let dispatcher = CdcEventDispatcher::new(vec![sink.clone()], 100);
let event = make_event("users", 1);
dispatcher.dispatch(event).await.unwrap();
assert_eq!(sink.count(), 1);
}
#[tokio::test]
async fn dispatch_multiple_sinks() {
let sink1 = Arc::new(MemorySink::new());
let sink2 = Arc::new(MemorySink::new());
let dispatcher = CdcEventDispatcher::new(vec![sink1.clone(), sink2.clone()], 100);
let event = make_event("users", 1);
dispatcher.dispatch(event).await.unwrap();
assert_eq!(sink1.count(), 1);
assert_eq!(sink2.count(), 1);
}
#[tokio::test]
async fn dispatch_batch_events() {
let sink = Arc::new(MemorySink::new());
let dispatcher = CdcEventDispatcher::new(vec![sink.clone()], 100);
let events: Vec<_> = (0..10).map(|i| make_event("users", i)).collect();
let count = dispatcher.dispatch_batch(events).await.unwrap();
assert_eq!(count, 10);
assert_eq!(sink.count(), 10);
}
#[tokio::test]
async fn table_filter_sink() {
let inner = Arc::new(MemorySink::new());
let filter = TableFilterSink::new("users", inner.clone());
let users_event = make_event("users", 1);
let orders_event = make_event("orders", 2);
filter.handle(&users_event).await.unwrap();
filter.handle(&orders_event).await.unwrap();
assert_eq!(inner.count(), 1);
}
#[tokio::test]
async fn dispatcher_run_from_channel() {
let sink = Arc::new(MemorySink::new());
let dispatcher = CdcEventDispatcher::new(vec![sink.clone()], 100);
let (tx, rx) = mpsc::channel(100);
for i in 0..5 {
tx.send(make_event("users", i)).await.unwrap();
}
drop(tx);
let count = dispatcher.run(rx).await.unwrap();
assert_eq!(count, 5);
assert_eq!(sink.count(), 5);
}
}