1use 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 errors.extend(served.serve_on(link).await.err());
122 }
123 if errors.len() == links.len() {
124 return Err(PoolError::NotServed(errors));
125 }
126 self.inner.hold_served(&served)?;
127 for link in self.inner.links() {
128 served.attach(link);
129 }
130 Ok(Served { inner: served })
131 }
132}
133
134impl PoolInner {
135 fn hold_served(&self, served: &Arc<ServedInner>) -> Result<(), PoolError> {
138 let mut state = self.lock();
139 if state.closed {
140 return Err(PoolError::Closed);
141 }
142 state.served.insert(served.id, served.clone());
143 Ok(())
144 }
145
146 pub(super) async fn replay(&self, link: &Link) {
149 let (subs, served) = {
150 let state = self.lock();
151 (
152 state.subs.values().cloned().collect::<Vec<_>>(),
153 state.served.values().cloned().collect::<Vec<_>>(),
154 )
155 };
156 for sub in subs {
157 sub.attach(link).await;
158 }
159 for s in served {
160 s.attach(link.clone());
161 }
162 }
163}
164
165impl Served {
166 pub async fn stop(&self) -> Result<(), LinkError> {
168 if let Some(pool) = self.inner.pool.upgrade() {
169 pool.lock().served.remove(&self.inner.id);
170 }
171 let Some(on_links) = self.inner.mark_stopped() else {
172 return Ok(());
173 };
174 let mut result = Ok(());
175 for on_link in on_links.into_values() {
176 result = on_link.stop().await.and(result);
178 }
179 result
180 }
181}
182
183impl ServedInner {
184 fn lock(&self) -> MutexGuard<'_, Held> {
185 self.held.lock().unwrap_or_else(|p| p.into_inner())
186 }
187
188 fn mark_stopped(&self) -> Option<HashMap<u64, station_link::Served>> {
191 let mut held = self.lock();
192 if held.stopped {
193 return None;
194 }
195 held.stopped = true;
196 Some(std::mem::take(&mut held.on_links))
197 }
198
199 fn needs_no_serving(&self, link: &Link) -> bool {
202 let held = self.lock();
203 held.stopped || held.on_links.contains_key(&link.serial())
204 }
205
206 fn hold_on_link(
210 self: &Arc<Self>,
211 link: &Link,
212 on_link: station_link::Served,
213 ) -> Option<station_link::Served> {
214 let mut held = self.lock();
215 if held.stopped {
216 return Some(on_link);
217 }
218 held.on_links.insert(link.serial(), on_link.clone());
219 drop(held);
220 tokio::spawn(watch(self.clone(), link.clone(), on_link));
221 None
222 }
223
224 fn respawn_delay(&self) -> Option<Duration> {
225 self.pool.upgrade().map(|p| p.opts.respawn_delay)
226 }
227
228 async fn serve_on(self: &Arc<Self>, link: &Link) -> Result<(), LinkError> {
230 if self.needs_no_serving(link) {
231 return Ok(());
232 }
233 let on_link = match link.serve(self.offer.clone()).await {
234 Err(LinkError::AlreadyServed) => return Ok(()),
235 other => other?,
236 };
237 let Some(stopped_meanwhile) = self.hold_on_link(link, on_link) else {
238 return Ok(());
239 };
240 stopped_meanwhile.stop().await
241 }
242
243 pub(super) fn attach(self: &Arc<Self>, link: Link) {
246 tokio::spawn(serve_while_linked(self.clone(), link));
247 }
248}
249
250async fn serve_while_linked(served: Arc<ServedInner>, link: Link) {
253 loop {
254 let outcome = tokio::time::timeout(DEFAULT_CALL_TIMEOUT, served.serve_on(&link)).await;
255 if matches!(outcome, Ok(Ok(()))) {
256 return;
257 }
258 let Some(delay) = served.respawn_delay() else {
259 return;
260 };
261 tokio::select! {
262 _ = link.done() => return,
263 _ = tokio::time::sleep(delay) => {}
264 }
265 }
266}
267
268async fn watch(served: Arc<ServedInner>, link: Link, on_link: station_link::Served) {
271 let why = on_link.done().await;
272 let stopped = {
273 let mut held = served.lock();
274 held.on_links.remove(&link.serial());
275 held.stopped
276 };
277 if stopped || why == LinkError::Stopped || link.error().is_some() {
278 return;
279 }
280 let Some(delay) = served.respawn_delay() else {
281 return;
282 };
283 tokio::select! {
284 _ = link.done() => {}
285 _ = tokio::time::sleep(delay) => served.attach(link.clone()),
286 }
287}