Skip to main content

moirai_transport/
router.rs

1#![cfg_attr(test, allow(clippy::unwrap_used, reason = "test scope"))]
2
3use crate::{Transport, TransportResult, lock_mutex, transport::Address};
4use std::{
5    collections::HashMap,
6    fmt,
7    sync::{Arc, Mutex},
8};
9
10/// Remote address for cross-machine communication
11#[derive(Debug, Clone, PartialEq, Eq, Hash)]
12pub struct RemoteAddress {
13    /// Remote host name or IP address.
14    pub host: String,
15    /// Remote TCP port.
16    pub port: u16,
17    /// Service label carried in the address display form.
18    pub service: String,
19}
20
21impl fmt::Display for RemoteAddress {
22    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
23        write!(f, "{}://{}:{}", self.service, self.host, self.port)
24    }
25}
26
27/// Topic-based pub/sub router that delivers published messages to every
28/// subscribed [`Address`] over a shared transport.
29///
30/// The router is generic over the backing [`Transport`] so delivery is
31/// monomorphized and zero-cost; the transport must be the *same instance* the
32/// subscribers receive from (e.g. one `Arc<crate::InMemoryTransport>`), since
33/// in-memory channels are keyed by address within a single transport instance.
34/// The prior implementation constructed a throwaway `InMemoryTransport` per
35/// send and so silently discarded every message.
36pub struct MessageRouter<T: Transport> {
37    transport: Arc<T>,
38    subscriptions: Mutex<HashMap<String, Vec<Address>>>,
39}
40
41impl<T: Transport> MessageRouter<T> {
42    /// Create a router that delivers over `transport`.
43    pub fn new(transport: Arc<T>) -> Self {
44        Self {
45            transport,
46            subscriptions: Mutex::new(HashMap::new()),
47        }
48    }
49
50    /// Subscribe `address` to `topic`. Duplicate (topic, address) pairs are
51    /// ignored so a message is delivered to each subscriber exactly once.
52    pub fn subscribe(&self, topic: &str, address: Address) {
53        let mut subs = lock_mutex(&self.subscriptions);
54        let entry = subs.entry(topic.to_string()).or_default();
55        if !entry.contains(&address) {
56            entry.push(address);
57        }
58    }
59
60    /// Remove `address` from `topic`. Returns `true` if a subscription was
61    /// removed.
62    pub fn unsubscribe(&self, topic: &str, address: &Address) -> bool {
63        let mut subs = lock_mutex(&self.subscriptions);
64        if let Some(entry) = subs.get_mut(topic) {
65            let before = entry.len();
66            entry.retain(|a| a != address);
67            let removed = entry.len() != before;
68            if entry.is_empty() {
69                subs.remove(topic);
70            }
71            return removed;
72        }
73        false
74    }
75
76    /// Publish `data` to every subscriber of `topic` via the shared transport.
77    ///
78    /// Returns the number of subscribers the message was delivered to. Delivery
79    /// is fail-fast: the first transport error is propagated (after the
80    /// subscribers ahead of it have already received the message).
81    ///
82    /// # Errors
83    /// Propagates the first per-subscriber transport send error.
84    pub fn publish(&self, topic: &str, data: Vec<u8>) -> TransportResult<usize> {
85        // Snapshot the subscriber list so the transport sends happen without the
86        // subscriptions lock held (a subscriber's send must not block resubscribe).
87        let targets: Vec<Address> = {
88            let subs = lock_mutex(&self.subscriptions);
89            match subs.get(topic) {
90                Some(addresses) => addresses.clone(),
91                None => return Ok(0),
92            }
93        };
94
95        // `Transport::send` takes ownership (each in-memory subscriber channel
96        // stores its own `Vec<u8>`), so N subscribers need N owned buffers —
97        // but only N-1 copies: the caller's original buffer is moved to the
98        // final subscriber instead of being cloned and dropped.
99        let mut delivered = 0;
100        let Some(last) = targets.len().checked_sub(1) else {
101            return Ok(0);
102        };
103        let mut data = Some(data);
104        for (index, addr) in targets.iter().enumerate() {
105            let payload = if index == last {
106                data.take()
107                    .expect("invariant: original buffer moved exactly once, at the last subscriber")
108            } else {
109                data.as_ref()
110                    .expect("invariant: original buffer present until the last subscriber")
111                    .clone()
112            };
113            self.transport.send(addr, payload)?;
114            delivered += 1;
115        }
116        Ok(delivered)
117    }
118
119    /// Number of distinct addresses subscribed to `topic`.
120    pub fn subscriber_count(&self, topic: &str) -> usize {
121        lock_mutex(&self.subscriptions)
122            .get(topic)
123            .map_or(0, Vec::len)
124    }
125}
126
127#[cfg(test)]
128mod tests {
129    use super::*;
130    use crate::{InMemoryTransport, Transport};
131
132    #[test]
133    fn message_router_delivers_to_each_subscriber_once() {
134        let transport = Arc::new(InMemoryTransport::new());
135        let router = MessageRouter::new(Arc::clone(&transport));
136        let sub_a = Address::Local("sub_a".to_string());
137        let sub_b = Address::Local("sub_b".to_string());
138
139        router.subscribe("topic", sub_a.clone());
140        router.subscribe("topic", sub_b.clone());
141        router.subscribe("topic", sub_a.clone()); // duplicate ignored
142        assert_eq!(router.subscriber_count("topic"), 2);
143
144        let delivered = router.publish("topic", vec![1, 2, 3]).unwrap();
145        assert_eq!(delivered, 2);
146
147        assert_eq!(transport.recv(&sub_a).unwrap(), vec![1, 2, 3]);
148        assert_eq!(transport.recv(&sub_b).unwrap(), vec![1, 2, 3]);
149    }
150
151    #[test]
152    fn message_router_single_subscriber_receives_moved_buffer() {
153        let transport = Arc::new(InMemoryTransport::new());
154        let router = MessageRouter::new(Arc::clone(&transport));
155        let sub = Address::Local("solo".to_string());
156        router.subscribe("t", sub.clone());
157
158        let delivered = router.publish("t", vec![7, 8, 9]).unwrap();
159        assert_eq!(delivered, 1);
160        assert_eq!(transport.recv(&sub).unwrap(), vec![7, 8, 9]);
161    }
162
163    #[test]
164    fn message_router_unknown_topic_delivers_nothing() {
165        let transport = Arc::new(InMemoryTransport::new());
166        let router = MessageRouter::new(transport);
167        assert_eq!(router.publish("absent", vec![0]).unwrap(), 0);
168    }
169
170    #[test]
171    fn message_router_unsubscribe_stops_delivery() {
172        let transport = Arc::new(InMemoryTransport::new());
173        let router = MessageRouter::new(Arc::clone(&transport));
174        let sub = Address::Local("s".to_string());
175
176        router.subscribe("t", sub.clone());
177        assert!(router.unsubscribe("t", &sub));
178        assert!(
179            !router.unsubscribe("t", &sub),
180            "second unsubscribe is a no-op"
181        );
182        assert_eq!(router.subscriber_count("t"), 0);
183        assert_eq!(router.publish("t", vec![9]).unwrap(), 0);
184    }
185}