macula_rust/pool/
serve.rs1use std::collections::HashMap;
7use std::sync::atomic::{AtomicU64, Ordering};
8use std::sync::{Arc, Mutex, MutexGuard, Weak};
9use std::time::Duration;
10
11use crate::frame::StreamMode;
12use crate::station_link::{
13 self, Confidentiality, Handler, Link, LinkError, StreamHandler, StreamOffer,
14 DEFAULT_CALL_TIMEOUT,
15};
16
17use super::{Pool, PoolError, PoolInner};
18
19static NEXT_SERVED: AtomicU64 = AtomicU64::new(1);
20
21#[derive(Clone)]
28pub struct Offer {
29 pub realm: [u8; 32],
30 pub procedure: String,
31 pub handler: Option<Handler>,
32 pub stream: Option<StreamOffer>,
33 pub confidential: Confidentiality,
34}
35
36impl Offer {
37 pub fn unary(realm: [u8; 32], procedure: &str, handler: Handler) -> Offer {
39 Offer {
40 realm,
41 procedure: procedure.to_string(),
42 handler: Some(handler),
43 stream: None,
44 confidential: Confidentiality::Preferred,
45 }
46 }
47
48 pub fn stream(
50 realm: [u8; 32],
51 procedure: &str,
52 mode: StreamMode,
53 handler: StreamHandler,
54 ) -> Offer {
55 Offer {
56 realm,
57 procedure: procedure.to_string(),
58 handler: None,
59 stream: Some(StreamOffer { mode, handler }),
60 confidential: Confidentiality::Preferred,
61 }
62 }
63}
64
65pub struct Served {
68 inner: Arc<ServedInner>,
69}
70
71pub(super) struct ServedInner {
72 id: u64,
73 pool: Weak<PoolInner>,
74 offer: station_link::Offer,
75 held: Mutex<Held>,
76}
77
78struct Held {
79 stopped: bool,
80 on_links: HashMap<u64, station_link::Served>,
81}
82
83impl Pool {
84 pub async fn serve(&self, o: Offer) -> Result<Served, PoolError> {
89 let realm_key = self.inner.realm_key_for(&o.realm, &o.procedure)?;
90 let kem_advertise = self.inner.opts.kem_advertise;
91 if o.confidential == Confidentiality::Required && !kem_advertise {
92 return Err(PoolError::Link(LinkError::KemAdvertiseDisabled));
93 }
94 let keyed_since_ms = (kem_advertise && o.confidential != Confidentiality::Off)
97 .then(|| crate::uuid_v7::now_ms() as i64);
98 let served = Arc::new(ServedInner {
99 id: NEXT_SERVED.fetch_add(1, Ordering::Relaxed),
100 pool: Arc::downgrade(&self.inner),
101 offer: station_link::Offer {
102 realm: o.realm,
103 procedure: o.procedure,
104 handler: o.handler,
105 stream: o.stream,
106 realm_key,
107 confidential: o.confidential,
108 keyed_since_ms,
109 },
110 held: Mutex::new(Held {
111 stopped: false,
112 on_links: HashMap::new(),
113 }),
114 });
115 let links = self.inner.links();
116 if links.is_empty() {
117 return Err(PoolError::NoLink(Vec::new()));
118 }
119 let mut errors = Vec::new();
120 for link in &links {
121 if let Err(e) = served.serve_on(link).await {
122 errors.push(e);
123 }
124 }
125 if errors.len() == links.len() {
126 return Err(PoolError::NotServed(errors));
127 }
128 {
129 let mut state = self.inner.lock();
130 if state.closed {
131 return Err(PoolError::Closed);
132 }
133 state.served.insert(served.id, served.clone());
134 }
135 for link in self.inner.links() {
136 served.attach(link);
137 }
138 Ok(Served { inner: served })
139 }
140}
141
142impl PoolInner {
143 pub(super) async fn replay(&self, link: &Link) {
146 let (subs, served) = {
147 let state = self.lock();
148 (
149 state.subs.values().cloned().collect::<Vec<_>>(),
150 state.served.values().cloned().collect::<Vec<_>>(),
151 )
152 };
153 for sub in subs {
154 sub.attach(link).await;
155 }
156 for s in served {
157 s.attach(link.clone());
158 }
159 }
160}
161
162impl Served {
163 pub async fn stop(&self) -> Result<(), LinkError> {
165 if let Some(pool) = self.inner.pool.upgrade() {
166 pool.lock().served.remove(&self.inner.id);
167 }
168 let on_links = {
169 let mut held = self.inner.lock();
170 if held.stopped {
171 return Ok(());
172 }
173 held.stopped = true;
174 std::mem::take(&mut held.on_links)
175 };
176 let mut result = Ok(());
177 for on_link in on_links.into_values() {
178 if let Err(e) = on_link.stop().await {
179 result = Err(e);
180 }
181 }
182 result
183 }
184}
185
186impl ServedInner {
187 fn lock(&self) -> MutexGuard<'_, Held> {
188 self.held.lock().unwrap_or_else(|p| p.into_inner())
189 }
190
191 fn respawn_delay(&self) -> Option<Duration> {
192 self.pool.upgrade().map(|p| p.opts.respawn_delay)
193 }
194
195 async fn serve_on(self: &Arc<Self>, link: &Link) -> Result<(), LinkError> {
197 {
198 let held = self.lock();
199 if held.stopped || held.on_links.contains_key(&link.serial()) {
200 return Ok(());
201 }
202 }
203 let on_link = match link.serve(self.offer.clone()).await {
204 Err(LinkError::AlreadyServed) => return Ok(()),
205 other => other?,
206 };
207 {
208 let mut held = self.lock();
209 if !held.stopped {
210 held.on_links.insert(link.serial(), on_link.clone());
211 drop(held);
212 tokio::spawn(watch(self.clone(), link.clone(), on_link));
213 return Ok(());
214 }
215 }
216 on_link.stop().await
217 }
218
219 pub(super) fn attach(self: &Arc<Self>, link: Link) {
222 let served = self.clone();
223 tokio::spawn(async move {
224 loop {
225 let outcome =
226 tokio::time::timeout(DEFAULT_CALL_TIMEOUT, served.serve_on(&link)).await;
227 if matches!(outcome, Ok(Ok(()))) {
228 return;
229 }
230 let Some(delay) = served.respawn_delay() else {
231 return;
232 };
233 tokio::select! {
234 _ = link.done() => return,
235 _ = tokio::time::sleep(delay) => {}
236 }
237 }
238 });
239 }
240}
241
242async fn watch(served: Arc<ServedInner>, link: Link, on_link: station_link::Served) {
245 let why = on_link.done().await;
246 let stopped = {
247 let mut held = served.lock();
248 held.on_links.remove(&link.serial());
249 held.stopped
250 };
251 if stopped || why == LinkError::Stopped || link.error().is_some() {
252 return;
253 }
254 let Some(delay) = served.respawn_delay() else {
255 return;
256 };
257 tokio::select! {
258 _ = link.done() => {}
259 _ = tokio::time::sleep(delay) => served.attach(link.clone()),
260 }
261}