1use 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 register(&self.inner, &sub)?;
73 for link in self.inner.links() {
74 sub.attach(&link).await;
75 }
76 Ok(Subscription {
77 inner: sub,
78 events,
79 unsubscribed: false,
80 })
81 }
82
83 pub async fn publish(&self, p: Publication) -> Result<(), PoolError> {
87 let links = self.inner.links();
88 if links.is_empty() {
89 return Err(PoolError::NoLink(Vec::new()));
90 }
91 let signed =
92 SignedPublication::sign(&self.inner.opts.identity, &self.inner.publication_seq, p)?;
93 let mut errors = Vec::new();
94 let mut sent = 0;
95 for link in links.iter().take(self.inner.opts.replication_factor) {
96 match link.publish_signed(&signed).await {
97 Ok(()) => sent += 1,
98 Err(e) => errors.push(e),
99 }
100 }
101 if sent == 0 {
102 return Err(PoolError::NoLink(errors));
103 }
104 Ok(())
105 }
106}
107
108impl Subscription {
109 pub async fn recv(&mut self) -> Option<Event> {
112 self.events.recv().await
113 }
114
115 pub fn dropped(&self) -> u64 {
117 self.inner.dropped.load(Ordering::Relaxed)
118 }
119
120 pub async fn unsubscribe(&mut self) -> Result<(), LinkError> {
122 self.unsubscribed = true;
123 if let Some(pool) = self.inner.pool.upgrade() {
124 pool.lock().subs.remove(&self.inner.id);
125 }
126 self.inner.end().await
127 }
128}
129
130impl Drop for Subscription {
131 fn drop(&mut self) {
133 if self.unsubscribed {
134 return;
135 }
136 if let Some(pool) = self.inner.pool.upgrade() {
137 pool.lock().subs.remove(&self.inner.id);
138 }
139 let inner = self.inner.clone();
140 let Ok(runtime) = tokio::runtime::Handle::try_current() else {
141 return;
142 };
143 runtime.spawn(async move {
144 let _ = inner.end().await;
145 });
146 }
147}
148
149fn register(pool: &PoolInner, sub: &Arc<SubInner>) -> Result<(), PoolError> {
152 let mut state = pool.lock();
153 if state.closed {
154 return Err(PoolError::Closed);
155 }
156 state.subs.insert(sub.id, sub.clone());
157 Ok(())
158}
159
160impl SubInner {
161 fn lock(&self) -> MutexGuard<'_, Held> {
162 self.held.lock().unwrap_or_else(|p| p.into_inner())
163 }
164
165 pub(super) async fn end(&self) -> Result<(), LinkError> {
168 let Some(forwarders) = self.take_forwarders() else {
169 return Ok(());
170 };
171 stop_forwarders(forwarders).await
172 }
173
174 fn take_forwarders(&self) -> Option<HashMap<u64, Forwarder>> {
177 let mut held = self.lock();
178 held.events.take()?;
179 Some(std::mem::take(&mut held.on_links))
180 }
181
182 fn keep_forwarder(&self, serial: u64, forwarder: Forwarder) -> bool {
185 let mut held = self.lock();
186 let wanted = held.events.is_some() && !held.on_links.contains_key(&serial);
187 if wanted {
188 held.on_links.insert(serial, forwarder);
189 }
190 wanted
191 }
192
193 pub(super) async fn attach(self: &Arc<Self>, link: &Link) {
196 let events = {
197 let held = self.lock();
198 match &held.events {
199 Some(events) if !held.on_links.contains_key(&link.serial()) => events.clone(),
200 _ => return,
201 }
202 };
203 let Ok(on_link) = link.subscribe(&self.realm, &self.topic).await else {
204 return;
205 };
206 let (stop, stopped) = oneshot::channel();
207 let (unsubscribed_tx, unsubscribed) = oneshot::channel();
208 let kept = self.keep_forwarder(link.serial(), Forwarder { stop, unsubscribed });
209 if !kept {
210 let _ = on_link.unsubscribe().await;
211 return;
212 }
213 tokio::spawn(forward(
214 self.clone(),
215 link.serial(),
216 on_link,
217 events,
218 stopped,
219 unsubscribed_tx,
220 ));
221 }
222}
223
224async fn stop_forwarders(forwarders: HashMap<u64, Forwarder>) -> Result<(), LinkError> {
227 let mut result = Ok(());
228 for (_, f) in forwarders {
229 let _ = f.stop.send(());
230 if let Ok(Err(e)) = f.unsubscribed.await {
231 result = Err(e);
232 }
233 }
234 result
235}
236
237async fn forward(
238 sub: Arc<SubInner>,
239 serial: u64,
240 mut on_link: station_link::Subscription,
241 events: mpsc::Sender<Event>,
242 mut stop: oneshot::Receiver<()>,
243 unsubscribed: oneshot::Sender<Result<(), LinkError>>,
244) {
245 loop {
246 tokio::select! {
247 _ = &mut stop => {
248 let _ = unsubscribed.send(on_link.unsubscribe().await);
249 return;
250 }
251 event = on_link.recv() => match event {
252 Some(event) => {
253 if events.try_send(event).is_err() {
254 sub.dropped.fetch_add(1, Ordering::Relaxed);
255 }
256 }
257 None => break,
258 },
259 }
260 }
261 sub.lock().on_links.remove(&serial);
263 let _ = unsubscribed.send(Ok(()));
264}