use std::collections::HashMap;
use std::fmt;
use std::sync::Mutex;
use std::sync::atomic::{AtomicU64, Ordering};
use bytes::Bytes;
use ruststream::testing::Coordinator;
use ruststream::{Headers, RawMessage};
use tokio::sync::mpsc;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub(crate) struct SubscriptionId(u64);
#[derive(Debug, Clone)]
pub(crate) struct TestDelivery {
pub(crate) payload: Bytes,
pub(crate) headers: Headers,
}
pub(crate) type DeliverySender = mpsc::UnboundedSender<TestDelivery>;
pub(crate) type DeliveryReceiver = mpsc::UnboundedReceiver<TestDelivery>;
#[derive(Debug)]
struct Subscription {
topic: String,
sender: DeliverySender,
}
#[derive(Debug, Default)]
struct RouterState {
subscriptions: HashMap<SubscriptionId, Subscription>,
log: HashMap<String, Vec<RawMessage>>,
}
#[derive(Default)]
pub(crate) struct KeyRouter {
state: Mutex<RouterState>,
next_id: AtomicU64,
}
impl KeyRouter {
pub(crate) fn subscribe_many(
&self,
topics: &[String],
) -> (Vec<SubscriptionId>, DeliverySender, DeliveryReceiver) {
let (sender, receiver) = mpsc::unbounded_channel();
let mut state = self.state.lock().expect("test router mutex poisoned");
let ids = topics
.iter()
.map(|topic| {
let id = SubscriptionId(self.next_id.fetch_add(1, Ordering::Relaxed));
state.subscriptions.insert(
id,
Subscription {
topic: topic.clone(),
sender: sender.clone(),
},
);
id
})
.collect();
(ids, sender, receiver)
}
pub(crate) fn unsubscribe(&self, id: SubscriptionId) {
let mut state = self.state.lock().expect("test router mutex poisoned");
state.subscriptions.remove(&id);
}
pub(crate) fn publish(
&self,
topic: &str,
payload: &Bytes,
headers: &Headers,
coordinator: Option<&Coordinator>,
) {
let senders: Vec<DeliverySender> = {
let mut state = self.state.lock().expect("test router mutex poisoned");
state
.log
.entry(topic.to_owned())
.or_default()
.push(RawMessage::new(topic, payload.clone()).with_headers(headers.clone()));
state
.subscriptions
.values()
.filter(|subscription| subscription.topic == topic)
.map(|subscription| subscription.sender.clone())
.collect()
};
for sender in senders {
let delivery = TestDelivery {
payload: payload.clone(),
headers: headers.clone(),
};
if sender.send(delivery).is_ok()
&& let Some(coordinator) = coordinator
{
coordinator.enqueued();
}
}
}
pub(crate) fn published(&self, topic: &str) -> Vec<RawMessage> {
let state = self.state.lock().expect("test router mutex poisoned");
state.log.get(topic).cloned().unwrap_or_default()
}
pub(crate) fn clear(&self) {
let mut state = self.state.lock().expect("test router mutex poisoned");
state.subscriptions.clear();
state.log.clear();
}
}
impl fmt::Debug for KeyRouter {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("KeyRouter").finish_non_exhaustive()
}
}