use mongodb::Database;
use std::{
collections::HashMap,
future::Future,
pin::Pin,
sync::Arc,
task::{Context, Poll},
};
use tokio::sync::mpsc;
use tokio_stream::{Stream, StreamExt};
use crate::error::MongoStreamError;
use crate::event::Event;
use mongodb::bson::Document;
type CallbackFn = dyn Fn(&Document) -> Pin<Box<dyn Future<Output = ()> + Send>> + Send + Sync;
pub type Callbacks = HashMap<Event, Arc<CallbackFn>>;
pub struct MongoStream {
db: Database,
collection_callbacks: HashMap<String, Callbacks>,
active_streams: HashMap<String, tokio::sync::mpsc::Sender<()>>,
}
impl MongoStream {
pub fn new(db: Database) -> Self {
Self {
db,
collection_callbacks: HashMap::new(),
active_streams: HashMap::new(),
}
}
pub fn add_callback<F>(&mut self, collection_name: impl Into<String>, event: Event, callback: F)
where
F: Fn(&Document) -> Pin<Box<dyn Future<Output = ()> + Send>> + 'static + Send + Sync,
{
let collection_name = collection_name.into();
let callback_arc = Arc::new(callback);
let default_callbacks = self.create_default_callbacks();
self.collection_callbacks
.entry(collection_name)
.or_insert_with(|| default_callbacks)
.insert(event, callback_arc);
}
fn create_default_callbacks(&self) -> Callbacks {
let mut callbacks: Callbacks = HashMap::new();
for event in [Event::Insert, Event::Update, Event::Delete] {
let event_name = event.event_type_str().to_string();
callbacks.insert(
event,
Arc::new(move |_doc: &Document| {
let event_name = event_name.clone();
Box::pin(async move {
println!("{} event received", event_name);
})
}),
);
}
callbacks
}
fn get_collection_callbacks(&self, collection_name: &str) -> Callbacks {
match self.collection_callbacks.get(collection_name) {
Some(callbacks) => {
callbacks
.iter()
.map(|(event, callback)| (*event, Arc::clone(callback)))
.collect()
}
None => self.create_default_callbacks(),
}
}
pub async fn start_stream(&mut self, collection_name: &str) -> Result<(), MongoStreamError> {
if self.active_streams.contains_key(collection_name) {
return Err(MongoStreamError::new(format!(
"Stream for collection '{}' is already running",
collection_name
)));
}
let collection = self
.db
.collection::<mongodb::bson::Document>(collection_name);
let (tx, mut rx) = mpsc::channel::<()>(1);
self.active_streams.insert(collection_name.to_string(), tx);
let callbacks = self.get_collection_callbacks(collection_name);
let collection_name = collection_name.to_string();
let db = self.db.clone();
tokio::spawn(async move {
let mut stream = match collection.watch(None, None).await {
Ok(s) => s,
Err(e) => {
eprintln!(
"Failed to start stream for collection '{}': {}",
collection_name, e
);
return;
}
};
loop {
tokio::select! {
_ = rx.recv() => {
println!("Closing stream for collection '{}'", collection_name);
break;
}
next_event = stream.next() => {
match next_event {
Some(result) => {
match result {
Ok(change_stream_event) => {
let event_type = Event::from(change_stream_event.operation_type);
if let Some(callback) = callbacks.get(&event_type) {
if let Some(doc) = change_stream_event.full_document {
callback(&doc).await;
}
}
}
Err(e) => {
eprintln!("Error in MongoDB stream for collection '{}': {}. Reconnecting...", collection_name, e);
match db.collection::<mongodb::bson::Document>(&collection_name).watch(None, None).await {
Ok(new_stream) => stream = new_stream,
Err(reconnect_err) => {
eprintln!("Failed to reconnect to collection '{}': {}", collection_name, reconnect_err);
break;
}
}
}
}
}
None => {
println!("Stream for collection '{}' has ended", collection_name);
break;
}
}
}
}
}
});
Ok(())
}
pub async fn close_stream(&mut self, collection_name: &str) -> bool {
if let Some(tx) = self.active_streams.remove(collection_name) {
let _ = tx.send(()).await;
true
} else {
false
}
}
pub async fn close_all_streams(&mut self) -> usize {
let count = self.active_streams.len();
for (collection_name, tx) in self.active_streams.drain() {
let _ = tx.send(()).await;
println!(
"Sent close signal to stream for collection '{}'",
collection_name
);
}
count
}
}
impl Stream for MongoStream {
type Item = Result<Event, MongoStreamError>;
fn poll_next(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Poll::Ready(Some(Ok(Event::Insert)))
}
}
#[cfg(test)]
mod tests {
use super::*;
use mongodb::{bson::Document, options::ClientOptions, Client};
use tokio::sync::mpsc;
async fn setup_test_db() -> Database {
let client_options = ClientOptions::parse("mongodb://localhost:27017")
.await
.unwrap();
let client = Client::with_options(client_options).unwrap();
client.database("test_db")
}
#[tokio::test]
async fn test_mongo_stream_creation() {
let db = setup_test_db().await;
let mongo_stream = MongoStream::new(db.clone());
assert!(mongo_stream.collection_callbacks.is_empty());
assert!(mongo_stream.active_streams.is_empty());
}
#[tokio::test]
async fn test_add_callback() {
let db = setup_test_db().await;
let mut mongo_stream = MongoStream::new(db);
let (tx, mut rx) = mpsc::channel(1);
mongo_stream.add_callback("test_collection", Event::Insert, move |_doc| {
let tx = tx.clone();
Box::pin(async move {
let _ = tx.send(()).await;
})
});
let callbacks = mongo_stream.get_collection_callbacks("test_collection");
assert!(callbacks.contains_key(&Event::Insert));
if let Some(callback) = callbacks.get(&Event::Insert) {
let doc = Document::new();
callback(&doc).await;
assert!(rx.recv().await.is_some());
}
}
#[tokio::test]
async fn test_start_and_close_stream() {
let db = setup_test_db().await;
let mut mongo_stream = MongoStream::new(db);
assert!(mongo_stream.start_stream("test_collection").await.is_ok());
assert!(mongo_stream.active_streams.contains_key("test_collection"));
assert!(mongo_stream.close_stream("test_collection").await);
assert!(!mongo_stream.active_streams.contains_key("test_collection"));
}
#[tokio::test]
async fn test_close_all_streams() {
let db = setup_test_db().await;
let mut mongo_stream = MongoStream::new(db);
mongo_stream.start_stream("collection1").await.unwrap();
mongo_stream.start_stream("collection2").await.unwrap();
let closed_count = mongo_stream.close_all_streams().await;
assert_eq!(closed_count, 2);
assert!(mongo_stream.active_streams.is_empty());
}
#[tokio::test]
async fn test_default_callbacks() {
let db = setup_test_db().await;
let mongo_stream = MongoStream::new(db);
let callbacks = mongo_stream.get_collection_callbacks("test_collection");
assert_eq!(callbacks.len(), 3);
assert!(callbacks.contains_key(&Event::Insert));
assert!(callbacks.contains_key(&Event::Update));
assert!(callbacks.contains_key(&Event::Delete));
}
#[tokio::test]
async fn test_double_start_error() {
let db = setup_test_db().await;
let mut mongo_stream = MongoStream::new(db);
mongo_stream.start_stream("test_collection").await.unwrap();
let result = mongo_stream.start_stream("test_collection").await;
assert!(matches!(result, Err(MongoStreamError { .. })));
}
}