1use std::collections::{HashMap, VecDeque};
14use std::sync::{Arc, Mutex};
15
16use super::source::{AsyncMessageSource, ReceivedMessage};
17use super::{run_source, Bus, BusConsumer, MessageRouter, RunOptions, TransportError};
18use super::{Message, MessageKind};
19
20type Queues = Arc<Mutex<HashMap<String, VecDeque<Message>>>>;
21type Topics = Arc<Mutex<HashMap<String, Vec<Message>>>>;
22
23fn lock_poisoned(what: &str) -> TransportError {
24 TransportError::permanent(format!("in-memory bus {what} lock poisoned"))
25}
26
27#[derive(Clone, Default)]
32pub struct InMemoryBus {
33 queues: Queues,
34 topics: Topics,
35}
36
37impl InMemoryBus {
38 pub fn new() -> Self {
39 Self::default()
40 }
41
42 fn enqueue(&self, message: Message) -> Result<(), TransportError> {
43 self.queues
44 .lock()
45 .map_err(|_| lock_poisoned("queue"))?
46 .entry(message.name().to_string())
47 .or_default()
48 .push_back(message);
49 Ok(())
50 }
51
52 fn append(&self, message: Message) -> Result<(), TransportError> {
53 self.topics
54 .lock()
55 .map_err(|_| lock_poisoned("topic"))?
56 .entry(message.name().to_string())
57 .or_default()
58 .push(message);
59 Ok(())
60 }
61}
62
63impl Bus for InMemoryBus {
64 async fn send(&self, name: &str, payload: Vec<u8>) -> Result<(), TransportError> {
65 self.enqueue(Message::new(name, MessageKind::Command, payload))
66 }
67
68 async fn publish(&self, name: &str, payload: Vec<u8>) -> Result<(), TransportError> {
69 self.append(Message::new(name, MessageKind::Event, payload))
70 }
71
72 async fn send_message(&self, message: Message) -> Result<(), TransportError> {
73 self.enqueue(message)
74 }
75
76 async fn publish_message(&self, message: Message) -> Result<(), TransportError> {
77 self.append(message)
78 }
79}
80
81impl BusConsumer for InMemoryBus {
82 async fn listen<R: MessageRouter>(
83 &self,
84 router: Arc<R>,
85 options: RunOptions,
86 ) -> Result<(), TransportError> {
87 let names = router.subscription_plan().commands;
88 let source = QueueSource {
89 queues: self.queues.clone(),
90 names,
91 };
92 run_source(router, source, options).await
93 }
94
95 async fn subscribe<R: MessageRouter>(
96 &self,
97 router: Arc<R>,
98 options: RunOptions,
99 ) -> Result<(), TransportError> {
100 let names = router.subscription_plan().events;
101 let source = TopicSource {
102 topics: self.topics.clone(),
103 names,
104 cursors: HashMap::new(),
105 };
106 run_source(router, source, options).await
107 }
108}
109
110struct QueueSource {
112 queues: Queues,
113 names: Vec<String>,
114}
115
116impl AsyncMessageSource for QueueSource {
117 type Received = InMemoryReceived;
118
119 async fn recv(&mut self) -> Result<Option<Self::Received>, TransportError> {
120 let mut queues = self.queues.lock().map_err(|_| lock_poisoned("queue"))?;
121 for name in &self.names {
122 if let Some(message) = queues.get_mut(name).and_then(VecDeque::pop_front) {
123 return Ok(Some(InMemoryReceived { message }));
124 }
125 }
126 Ok(None)
127 }
128}
129
130struct TopicSource {
133 topics: Topics,
134 names: Vec<String>,
135 cursors: HashMap<String, usize>,
136}
137
138impl AsyncMessageSource for TopicSource {
139 type Received = InMemoryReceived;
140
141 async fn recv(&mut self) -> Result<Option<Self::Received>, TransportError> {
142 let topics = self.topics.lock().map_err(|_| lock_poisoned("topic"))?;
143 for name in &self.names {
144 let Some(log) = topics.get(name) else {
145 continue;
146 };
147 let cursor = self.cursors.entry(name.clone()).or_insert(0);
148 if *cursor < log.len() {
149 let message = log[*cursor].clone();
150 *cursor += 1;
151 return Ok(Some(InMemoryReceived { message }));
152 }
153 }
154 Ok(None)
155 }
156}
157
158pub struct InMemoryReceived {
161 message: Message,
162}
163
164impl ReceivedMessage for InMemoryReceived {
165 fn message(&self) -> &Message {
166 &self.message
167 }
168 async fn ack(self) -> Result<(), TransportError> {
169 Ok(())
170 }
171 async fn nack(self, _reason: &str) -> Result<(), TransportError> {
172 Ok(())
173 }
174}
175
176#[cfg(test)]
177mod tests {
178 use super::*;
179 use crate::bus::Handlers;
180 use std::future::Future;
181
182 fn block_on<F: Future>(future: F) -> F::Output {
183 use std::ptr;
184 use std::task::{Context, Poll, RawWaker, RawWakerVTable, Waker};
185 const VTABLE: RawWakerVTable = RawWakerVTable::new(
186 |_| RawWaker::new(ptr::null(), &VTABLE),
187 |_| {},
188 |_| {},
189 |_| {},
190 );
191 let waker = unsafe { Waker::from_raw(RawWaker::new(ptr::null(), &VTABLE)) };
192 let mut cx = Context::from_waker(&waker);
193 let mut future = std::pin::pin!(future);
194 loop {
195 if let Poll::Ready(output) = future.as_mut().poll(&mut cx) {
196 return output;
197 }
198 }
199 }
200
201 fn recorder() -> Arc<Mutex<Vec<String>>> {
202 Arc::new(Mutex::new(Vec::new()))
203 }
204
205 fn command_service(rec: Arc<Mutex<Vec<String>>>) -> Arc<Handlers> {
206 Arc::new(Handlers::new().on_command("work", move |msg: &Message| {
207 let rec = rec.clone();
208 let name = msg.name().to_string();
209 async move {
210 rec.lock().unwrap().push(name);
211 Ok(())
212 }
213 }))
214 }
215
216 fn event_service(rec: Arc<Mutex<Vec<String>>>) -> Arc<Handlers> {
217 Arc::new(Handlers::new().on_event("evt", move |msg: &Message| {
218 let rec = rec.clone();
219 let id = msg.id().unwrap_or("?").to_string();
220 async move {
221 rec.lock().unwrap().push(id);
222 Ok(())
223 }
224 }))
225 }
226
227 #[test]
228 fn send_then_listen_dispatches_each_command() {
229 let bus = InMemoryBus::new();
230 for _ in 0..3 {
231 block_on(bus.send("work", b"{}".to_vec())).unwrap();
232 }
233 let rec = recorder();
234 block_on(bus.listen(command_service(rec.clone()), RunOptions::idempotent())).unwrap();
235 assert_eq!(
236 rec.lock().unwrap().len(),
237 3,
238 "the listener handles all 3 commands"
239 );
240 }
241
242 #[test]
243 fn listen_is_point_to_point_each_message_popped_once() {
244 let bus = InMemoryBus::new();
246 for i in 0..4 {
247 block_on(bus.send_message(
248 Message::new("work", MessageKind::Command, b"{}".to_vec()).with_id(format!("m{i}")),
249 ))
250 .unwrap();
251 }
252 let mut a = QueueSource {
253 queues: bus.queues.clone(),
254 names: vec!["work".to_string()],
255 };
256 let mut b = QueueSource {
257 queues: bus.queues.clone(),
258 names: vec!["work".to_string()],
259 };
260 let mut got = Vec::new();
261 for _ in 0..4 {
263 if let Some(r) = block_on(a.recv()).unwrap() {
264 got.push(r.message().id().unwrap().to_string());
265 }
266 if let Some(r) = block_on(b.recv()).unwrap() {
267 got.push(r.message().id().unwrap().to_string());
268 }
269 }
270 got.sort();
271 assert_eq!(
272 got,
273 vec!["m0", "m1", "m2", "m3"],
274 "each message delivered exactly once"
275 );
276 assert!(block_on(a.recv()).unwrap().is_none());
278 assert!(block_on(b.recv()).unwrap().is_none());
279 }
280
281 #[test]
282 fn publish_then_subscribe_fans_out_to_every_subscriber() {
283 let bus = InMemoryBus::new();
284 for i in 0..3 {
285 block_on(bus.publish_message(
286 Message::new("evt", MessageKind::Event, b"{}".to_vec()).with_id(format!("e{i}")),
287 ))
288 .unwrap();
289 }
290 let a = recorder();
292 let b = recorder();
293 block_on(bus.subscribe(event_service(a.clone()), RunOptions::idempotent())).unwrap();
294 block_on(bus.subscribe(event_service(b.clone()), RunOptions::idempotent())).unwrap();
295 let mut a_ids = a.lock().unwrap().clone();
296 let mut b_ids = b.lock().unwrap().clone();
297 a_ids.sort();
298 b_ids.sort();
299 assert_eq!(a_ids, vec!["e0", "e1", "e2"]);
300 assert_eq!(b_ids, vec!["e0", "e1", "e2"]);
301 }
302
303 #[test]
304 fn unknown_command_is_acked_and_ignored() {
305 let bus = InMemoryBus::new();
307 block_on(bus.send("unrelated", b"{}".to_vec())).unwrap();
308 block_on(bus.send("work", b"{}".to_vec())).unwrap();
309 let rec = recorder();
310 block_on(bus.listen(command_service(rec.clone()), RunOptions::idempotent())).unwrap();
311 assert_eq!(rec.lock().unwrap().clone(), vec!["work"]);
312 }
313
314 #[test]
315 fn handler_error_does_not_panic_the_loop() {
316 let bus = InMemoryBus::new();
317 block_on(bus.send("work", b"{}".to_vec())).unwrap();
318 let handlers: Arc<Handlers> = Arc::new(
319 Handlers::new().on_command("work", |_: &Message| async move {
320 Err(TransportError::permanent("no"))
321 }),
322 );
323 block_on(bus.listen(handlers, RunOptions::idempotent())).unwrap();
326 }
327}