macula_rust/pool/
pubsub.rs1use std::collections::HashMap;
7use std::sync::atomic::{AtomicU64, Ordering};
8use std::sync::{Arc, Mutex, MutexGuard, Weak};
9
10use tokio::sync::{mpsc, oneshot};
11
12use crate::station_link::{self, Event, Link, LinkError, Publication, SignedPublication};
13
14use super::{Pool, PoolError, PoolInner};
15
16const SUBSCRIPTION_BUFFER: usize = 256;
19
20static NEXT_SUBSCRIPTION: AtomicU64 = AtomicU64::new(1);
21
22pub struct Subscription {
26 inner: Arc<SubInner>,
27 events: mpsc::Receiver<Event>,
28 unsubscribed: bool,
29}
30
31pub(super) struct SubInner {
32 id: u64,
33 pool: Weak<PoolInner>,
34 realm: [u8; 32],
35 topic: String,
36 held: Mutex<Held>,
37 dropped: AtomicU64,
38}
39
40struct Held {
41 events: Option<mpsc::Sender<Event>>,
43 on_links: HashMap<u64, Forwarder>,
46}
47
48struct Forwarder {
49 stop: oneshot::Sender<()>,
50 unsubscribed: oneshot::Receiver<Result<(), LinkError>>,
51}
52
53impl Pool {
54 pub async fn subscribe(
56 &self,
57 realm: &[u8; 32],
58 topic: &str,
59 ) -> Result<Subscription, PoolError> {
60 let (events_tx, events) = mpsc::channel(SUBSCRIPTION_BUFFER);
61 let sub = Arc::new(SubInner {
62 id: NEXT_SUBSCRIPTION.fetch_add(1, Ordering::Relaxed),
63 pool: Arc::downgrade(&self.inner),
64 realm: *realm,
65 topic: topic.to_string(),
66 held: Mutex::new(Held {
67 events: Some(events_tx),
68 on_links: HashMap::new(),
69 }),
70 dropped: AtomicU64::new(0),
71 });
72 {
73 let mut state = self.inner.lock();
74 if state.closed {
75 return Err(PoolError::Closed);
76 }
77 state.subs.insert(sub.id, sub.clone());
78 }
79 for link in self.inner.links() {
80 sub.attach(&link).await;
81 }
82 Ok(Subscription {
83 inner: sub,
84 events,
85 unsubscribed: false,
86 })
87 }
88
89 pub async fn publish(&self, p: Publication) -> Result<(), PoolError> {
93 let links = self.inner.links();
94 if links.is_empty() {
95 return Err(PoolError::NoLink(Vec::new()));
96 }
97 let signed =
98 SignedPublication::sign(&self.inner.opts.identity, &self.inner.publication_seq, p)?;
99 let mut errors = Vec::new();
100 let mut sent = 0;
101 for link in links.iter().take(self.inner.opts.replication_factor) {
102 match link.publish_signed(&signed).await {
103 Ok(()) => sent += 1,
104 Err(e) => errors.push(e),
105 }
106 }
107 if sent == 0 {
108 return Err(PoolError::NoLink(errors));
109 }
110 Ok(())
111 }
112}
113
114impl Subscription {
115 pub async fn recv(&mut self) -> Option<Event> {
118 self.events.recv().await
119 }
120
121 pub fn dropped(&self) -> u64 {
123 self.inner.dropped.load(Ordering::Relaxed)
124 }
125
126 pub async fn unsubscribe(&mut self) -> Result<(), LinkError> {
128 self.unsubscribed = true;
129 if let Some(pool) = self.inner.pool.upgrade() {
130 pool.lock().subs.remove(&self.inner.id);
131 }
132 self.inner.end().await
133 }
134}
135
136impl Drop for Subscription {
137 fn drop(&mut self) {
139 if self.unsubscribed {
140 return;
141 }
142 if let Some(pool) = self.inner.pool.upgrade() {
143 pool.lock().subs.remove(&self.inner.id);
144 }
145 let inner = self.inner.clone();
146 if let Ok(runtime) = tokio::runtime::Handle::try_current() {
147 runtime.spawn(async move {
148 let _ = inner.end().await;
149 });
150 }
151 }
152}
153
154impl SubInner {
155 fn lock(&self) -> MutexGuard<'_, Held> {
156 self.held.lock().unwrap_or_else(|p| p.into_inner())
157 }
158
159 pub(super) async fn end(&self) -> Result<(), LinkError> {
162 let forwarders = {
163 let mut held = self.lock();
164 if held.events.take().is_none() {
165 return Ok(());
166 }
167 std::mem::take(&mut held.on_links)
168 };
169 let mut result = Ok(());
170 for (_, f) in forwarders {
171 let _ = f.stop.send(());
172 if let Ok(Err(e)) = f.unsubscribed.await {
173 result = Err(e);
174 }
175 }
176 result
177 }
178
179 pub(super) async fn attach(self: &Arc<Self>, link: &Link) {
182 let events = {
183 let held = self.lock();
184 match &held.events {
185 Some(events) if !held.on_links.contains_key(&link.serial()) => events.clone(),
186 _ => return,
187 }
188 };
189 let Ok(on_link) = link.subscribe(&self.realm, &self.topic).await else {
190 return;
191 };
192 let (stop, stopped) = oneshot::channel();
193 let (unsubscribed_tx, unsubscribed) = oneshot::channel();
194 let kept = {
195 let mut held = self.lock();
196 let wanted = held.events.is_some() && !held.on_links.contains_key(&link.serial());
197 if wanted {
198 held.on_links
199 .insert(link.serial(), Forwarder { stop, unsubscribed });
200 }
201 wanted
202 };
203 if !kept {
204 let _ = on_link.unsubscribe().await;
205 return;
206 }
207 tokio::spawn(forward(
208 self.clone(),
209 link.serial(),
210 on_link,
211 events,
212 stopped,
213 unsubscribed_tx,
214 ));
215 }
216}
217
218async fn forward(
219 sub: Arc<SubInner>,
220 serial: u64,
221 mut on_link: station_link::Subscription,
222 events: mpsc::Sender<Event>,
223 mut stop: oneshot::Receiver<()>,
224 unsubscribed: oneshot::Sender<Result<(), LinkError>>,
225) {
226 loop {
227 tokio::select! {
228 _ = &mut stop => {
229 let _ = unsubscribed.send(on_link.unsubscribe().await);
230 return;
231 }
232 event = on_link.recv() => match event {
233 Some(event) => {
234 if events.try_send(event).is_err() {
235 sub.dropped.fetch_add(1, Ordering::Relaxed);
236 }
237 }
238 None => break,
239 },
240 }
241 }
242 sub.lock().on_links.remove(&serial);
244 let _ = unsubscribed.send(Ok(()));
245}