kvbm_engine/pubsub/
stub.rs1use std::collections::HashMap;
7use std::sync::Arc;
8
9use anyhow::Result;
10use bytes::Bytes;
11use futures::future::BoxFuture;
12use futures::stream::BoxStream;
13use futures::{FutureExt, StreamExt};
14use parking_lot::RwLock;
15use tokio::sync::broadcast;
16use tokio_stream::wrappers::BroadcastStream;
17
18use super::{Message, Publisher, Subscriber, Subscription};
19
20#[derive(Clone)]
22pub struct StubBus {
23 inner: Arc<StubBusInner>,
24}
25
26struct StubBusInner {
27 channels: RwLock<HashMap<String, broadcast::Sender<Message>>>,
29 capacity: usize,
31}
32
33impl Default for StubBus {
34 fn default() -> Self {
35 Self::new(256)
36 }
37}
38
39impl StubBus {
40 pub fn new(capacity: usize) -> Self {
42 Self {
43 inner: Arc::new(StubBusInner {
44 channels: RwLock::new(HashMap::new()),
45 capacity,
46 }),
47 }
48 }
49
50 pub fn publisher(&self) -> StubPublisher {
52 StubPublisher { bus: self.clone() }
53 }
54
55 pub fn subscriber(&self) -> StubSubscriber {
57 StubSubscriber { bus: self.clone() }
58 }
59
60 fn get_or_create_channel(&self, subject: &str) -> broadcast::Sender<Message> {
61 let channels = self.inner.channels.read();
62 if let Some(tx) = channels.get(subject) {
63 return tx.clone();
64 }
65 drop(channels);
66
67 let mut channels = self.inner.channels.write();
68 if let Some(tx) = channels.get(subject) {
70 return tx.clone();
71 }
72
73 let (tx, _) = broadcast::channel(self.inner.capacity);
74 channels.insert(subject.to_string(), tx.clone());
75 tx
76 }
77}
78
79pub struct StubPublisher {
81 bus: StubBus,
82}
83
84impl StubPublisher {
85 pub fn new() -> (Self, StubSubscriber) {
87 let bus = StubBus::default();
88 (bus.publisher(), bus.subscriber())
89 }
90}
91
92impl Publisher for StubPublisher {
93 fn publish(&self, subject: &str, payload: Bytes) -> Result<()> {
94 let tx = self.bus.get_or_create_channel(subject);
95 let msg = Message {
96 subject: subject.to_string(),
97 payload,
98 };
99 let _ = tx.send(msg);
101 Ok(())
102 }
103
104 fn flush(&self) -> BoxFuture<'static, Result<()>> {
105 async { Ok(()) }.boxed()
107 }
108}
109
110pub struct StubSubscriber {
112 bus: StubBus,
113}
114
115impl Subscriber for StubSubscriber {
116 fn subscribe(&self, subject: &str) -> BoxFuture<'static, Result<Subscription>> {
117 let tx = self.bus.get_or_create_channel(subject);
118 let rx = tx.subscribe();
119
120 let stream: BoxStream<'static, Message> = BroadcastStream::new(rx)
121 .filter_map(|result| async move { result.ok() })
122 .boxed();
123
124 async move { Ok(stream) }.boxed()
125 }
126}
127
128#[cfg(test)]
129mod tests {
130 use super::*;
131 use futures::StreamExt;
132
133 #[tokio::test]
134 async fn test_stub_pubsub() {
135 let bus = StubBus::default();
136 let publisher = bus.publisher();
137 let subscriber = bus.subscriber();
138
139 let mut sub = subscriber.subscribe("test.subject").await.unwrap();
141
142 publisher
144 .publish("test.subject", Bytes::from("hello"))
145 .unwrap();
146
147 let msg = sub.next().await.unwrap();
149 assert_eq!(msg.subject, "test.subject");
150 assert_eq!(msg.payload.as_ref(), b"hello");
151 }
152
153 #[tokio::test]
154 async fn test_stub_multiple_subscribers() {
155 let bus = StubBus::default();
156 let publisher = bus.publisher();
157
158 let mut sub1 = bus.subscriber().subscribe("multi").await.unwrap();
159 let mut sub2 = bus.subscriber().subscribe("multi").await.unwrap();
160
161 publisher
162 .publish("multi", Bytes::from("broadcast"))
163 .unwrap();
164
165 let msg1 = sub1.next().await.unwrap();
166 let msg2 = sub2.next().await.unwrap();
167
168 assert_eq!(msg1.payload.as_ref(), b"broadcast");
169 assert_eq!(msg2.payload.as_ref(), b"broadcast");
170 }
171}