Skip to main content

kvbm_engine/pubsub/
stub.rs

1// SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4//! In-memory stub implementation of the PubSub traits for testing.
5
6use 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/// Shared state for stub publisher/subscriber pairs.
21#[derive(Clone)]
22pub struct StubBus {
23    inner: Arc<StubBusInner>,
24}
25
26struct StubBusInner {
27    /// Map of subject patterns to broadcast channels.
28    channels: RwLock<HashMap<String, broadcast::Sender<Message>>>,
29    /// Channel capacity for new subscriptions.
30    capacity: usize,
31}
32
33impl Default for StubBus {
34    fn default() -> Self {
35        Self::new(256)
36    }
37}
38
39impl StubBus {
40    /// Create a new stub bus with the specified channel capacity.
41    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    /// Create a publisher for this bus.
51    pub fn publisher(&self) -> StubPublisher {
52        StubPublisher { bus: self.clone() }
53    }
54
55    /// Create a subscriber for this bus.
56    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        // Double-check after acquiring write lock
69        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
79/// Stub implementation of the [`Publisher`] trait for testing.
80pub struct StubPublisher {
81    bus: StubBus,
82}
83
84impl StubPublisher {
85    /// Create a new stub publisher with a dedicated bus.
86    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        // Ignore send errors (no receivers is ok)
100        let _ = tx.send(msg);
101        Ok(())
102    }
103
104    fn flush(&self) -> BoxFuture<'static, Result<()>> {
105        // In-memory delivery is synchronous, nothing to flush
106        async { Ok(()) }.boxed()
107    }
108}
109
110/// Stub implementation of the [`Subscriber`] trait for testing.
111pub 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        // Subscribe first
140        let mut sub = subscriber.subscribe("test.subject").await.unwrap();
141
142        // Publish a message
143        publisher
144            .publish("test.subject", Bytes::from("hello"))
145            .unwrap();
146
147        // Receive the message
148        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}