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};
16use crate::ucan::Policy;
17
18use super::{Pool, PoolError, PoolInner};
19
20static NEXT_SERVED: AtomicU64 = AtomicU64::new(1);
21
22#[derive(Clone)]
30pub struct Offer {
31 pub realm: [u8; 32],
32 pub procedure: String,
33 pub handler: Option<Handler>,
34 pub stream: Option<StreamOffer>,
35 pub confidential: Confidentiality,
36 pub policy: Option<Policy>,
37}
38
39impl Offer {
40 pub fn unary(realm: [u8; 32], procedure: &str, handler: Handler) -> Offer {
42 Offer {
43 realm,
44 procedure: procedure.to_string(),
45 handler: Some(handler),
46 stream: None,
47 confidential: Confidentiality::Preferred,
48 policy: None,
49 }
50 }
51
52 pub fn stream(
54 realm: [u8; 32],
55 procedure: &str,
56 mode: StreamMode,
57 handler: StreamHandler,
58 ) -> Offer {
59 Offer {
60 realm,
61 procedure: procedure.to_string(),
62 handler: None,
63 stream: Some(StreamOffer { mode, handler }),
64 confidential: Confidentiality::Preferred,
65 policy: None,
66 }
67 }
68}
69
70pub struct Served {
73 inner: Arc<ServedInner>,
74}
75
76pub(super) struct ServedInner {
77 id: u64,
78 pool: Weak<PoolInner>,
79 offer: station_link::Offer,
80 held: Mutex<Held>,
81}
82
83struct Held {
84 stopped: bool,
85 on_links: HashMap<u64, station_link::Served>,
86}
87
88impl Pool {
89 pub async fn serve(&self, o: Offer) -> Result<Served, PoolError> {
94 if o.policy.as_ref().is_some_and(|p| !p.valid()) {
95 return Err(PoolError::Link(LinkError::InvalidOffer));
96 }
97 let realm_key = self.inner.realm_key_for(&o.realm, &o.procedure)?;
98 let kem_advertise = self.inner.opts.kem_advertise;
99 if o.confidential == Confidentiality::Required && !kem_advertise {
100 return Err(PoolError::Link(LinkError::KemAdvertiseDisabled));
101 }
102 let keyed_since_ms = (kem_advertise && o.confidential != Confidentiality::Off)
105 .then(|| crate::uuid_v7::now_ms() as i64);
106 let served = Arc::new(ServedInner {
107 id: NEXT_SERVED.fetch_add(1, Ordering::Relaxed),
108 pool: Arc::downgrade(&self.inner),
109 offer: station_link::Offer {
110 realm: o.realm,
111 procedure: o.procedure,
112 handler: o.handler,
113 stream: o.stream,
114 realm_key,
115 confidential: o.confidential,
116 keyed_since_ms,
117 policy: o.policy,
118 },
119 held: Mutex::new(Held {
120 stopped: false,
121 on_links: HashMap::new(),
122 }),
123 });
124 let links = self.inner.links();
125 if links.is_empty() {
126 return Err(PoolError::NoLink(Vec::new()));
127 }
128 let mut errors = Vec::new();
129 for link in &links {
130 errors.extend(served.serve_on(link).await.err());
131 }
132 if errors.len() == links.len() {
133 return Err(PoolError::NotServed(errors));
134 }
135 self.inner.hold_served(&served)?;
136 for link in self.inner.links() {
137 served.attach(link);
138 }
139 Ok(Served { inner: served })
140 }
141}
142
143impl PoolInner {
144 fn hold_served(&self, served: &Arc<ServedInner>) -> Result<(), PoolError> {
147 let mut state = self.lock();
148 if state.closed {
149 return Err(PoolError::Closed);
150 }
151 state.served.insert(served.id, served.clone());
152 Ok(())
153 }
154
155 pub(super) async fn replay(&self, link: &Link) {
158 let (subs, served) = {
159 let state = self.lock();
160 (
161 state.subs.values().cloned().collect::<Vec<_>>(),
162 state.served.values().cloned().collect::<Vec<_>>(),
163 )
164 };
165 for sub in subs {
166 sub.attach(link).await;
167 }
168 for s in served {
169 s.attach(link.clone());
170 }
171 }
172}
173
174impl Served {
175 pub async fn stop(&self) -> Result<(), LinkError> {
177 if let Some(pool) = self.inner.pool.upgrade() {
178 pool.lock().served.remove(&self.inner.id);
179 }
180 let Some(on_links) = self.inner.mark_stopped() else {
181 return Ok(());
182 };
183 let mut result = Ok(());
184 for on_link in on_links.into_values() {
185 result = on_link.stop().await.and(result);
187 }
188 result
189 }
190}
191
192impl ServedInner {
193 fn lock(&self) -> MutexGuard<'_, Held> {
194 self.held.lock().unwrap_or_else(|p| p.into_inner())
195 }
196
197 fn mark_stopped(&self) -> Option<HashMap<u64, station_link::Served>> {
200 let mut held = self.lock();
201 if held.stopped {
202 return None;
203 }
204 held.stopped = true;
205 Some(std::mem::take(&mut held.on_links))
206 }
207
208 fn needs_no_serving(&self, link: &Link) -> bool {
211 let held = self.lock();
212 held.stopped || held.on_links.contains_key(&link.serial())
213 }
214
215 fn hold_on_link(
219 self: &Arc<Self>,
220 link: &Link,
221 on_link: station_link::Served,
222 ) -> Option<station_link::Served> {
223 let mut held = self.lock();
224 if held.stopped {
225 return Some(on_link);
226 }
227 held.on_links.insert(link.serial(), on_link.clone());
228 drop(held);
229 tokio::spawn(watch(self.clone(), link.clone(), on_link));
230 None
231 }
232
233 fn respawn_delay(&self) -> Option<Duration> {
234 self.pool.upgrade().map(|p| p.opts.respawn_delay)
235 }
236
237 async fn serve_on(self: &Arc<Self>, link: &Link) -> Result<(), LinkError> {
239 if self.needs_no_serving(link) {
240 return Ok(());
241 }
242 let on_link = match link.serve(self.offer.clone()).await {
243 Err(LinkError::AlreadyServed) => return Ok(()),
244 other => other?,
245 };
246 let Some(stopped_meanwhile) = self.hold_on_link(link, on_link) else {
247 return Ok(());
248 };
249 stopped_meanwhile.stop().await
250 }
251
252 pub(super) fn attach(self: &Arc<Self>, link: Link) {
255 tokio::spawn(serve_while_linked(self.clone(), link));
256 }
257}
258
259async fn serve_while_linked(served: Arc<ServedInner>, link: Link) {
262 loop {
263 let outcome = tokio::time::timeout(DEFAULT_CALL_TIMEOUT, served.serve_on(&link)).await;
264 if matches!(outcome, Ok(Ok(()))) {
265 return;
266 }
267 let Some(delay) = served.respawn_delay() else {
268 return;
269 };
270 tokio::select! {
271 _ = link.done() => return,
272 _ = tokio::time::sleep(delay) => {}
273 }
274 }
275}
276
277async fn watch(served: Arc<ServedInner>, link: Link, on_link: station_link::Served) {
280 let why = on_link.done().await;
281 let stopped = {
282 let mut held = served.lock();
283 held.on_links.remove(&link.serial());
284 held.stopped
285 };
286 if stopped || why == LinkError::Stopped || link.error().is_some() {
287 return;
288 }
289 let Some(delay) = served.respawn_delay() else {
290 return;
291 };
292 tokio::select! {
293 _ = link.done() => {}
294 _ = tokio::time::sleep(delay) => served.attach(link.clone()),
295 }
296}