1use std::collections::HashMap;
2use std::future::Future;
3use std::sync::{Arc, Mutex};
4use std::time::Duration;
5
6use unb_runtime::{ClientSession, Pipe, SessionOutcome, Wire, WsError};
7use web_time::Instant;
8
9pub const DEFAULT_DIAL_TIMEOUT: Duration = Duration::from_secs(5);
10pub const CACHE_TTL: Duration = Duration::from_secs(300);
11
12#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
13pub enum TransportKind {
14 Unix,
15 WebTransport,
16 WebSocket,
17}
18
19impl TransportKind {
20 fn rank(self) -> u8 {
21 match self {
22 TransportKind::Unix => 0,
23 TransportKind::WebTransport => 1,
24 TransportKind::WebSocket => 2,
25 }
26 }
27
28 fn supported(self) -> bool {
29 match self {
30 TransportKind::Unix => cfg!(all(feature = "unix", unix)),
31 TransportKind::WebTransport => cfg!(feature = "webtransport"),
32 TransportKind::WebSocket => true,
33 }
34 }
35}
36
37#[derive(Debug, Clone)]
38pub struct Endpoint {
39 pub kind: TransportKind,
40 pub address: String,
41 pub cert_hash: Option<[u8; 32]>,
42}
43
44#[derive(Debug, Clone, Default)]
45pub struct EndpointSet(Vec<Endpoint>);
46
47impl EndpointSet {
48 pub fn new() -> EndpointSet {
49 EndpointSet(Vec::new())
50 }
51
52 pub fn push(&mut self, endpoint: Endpoint) {
53 self.0.push(endpoint);
54 }
55
56 pub fn is_empty(&self) -> bool {
57 self.0.is_empty()
58 }
59
60 pub fn len(&self) -> usize {
61 self.0.len()
62 }
63
64 pub fn iter(&self) -> impl Iterator<Item = &Endpoint> {
65 self.0.iter()
66 }
67
68 pub fn cache_key(&self) -> String {
69 let mut addresses: Vec<&str> = self.0.iter().map(|e| e.address.as_str()).collect();
70 addresses.sort();
71 addresses.join("|")
72 }
73}
74
75impl From<Endpoint> for EndpointSet {
76 fn from(endpoint: Endpoint) -> EndpointSet {
77 EndpointSet(vec![endpoint])
78 }
79}
80
81impl<const N: usize> From<[Endpoint; N]> for EndpointSet {
82 fn from(endpoints: [Endpoint; N]) -> EndpointSet {
83 EndpointSet(endpoints.into())
84 }
85}
86
87impl From<Vec<Endpoint>> for EndpointSet {
88 fn from(endpoints: Vec<Endpoint>) -> EndpointSet {
89 EndpointSet(endpoints)
90 }
91}
92
93impl FromIterator<Endpoint> for EndpointSet {
94 fn from_iter<I: IntoIterator<Item = Endpoint>>(iter: I) -> EndpointSet {
95 EndpointSet(iter.into_iter().collect())
96 }
97}
98
99#[derive(Debug, Clone)]
100pub struct DialConfig {
101 pub attempt_timeout: Duration,
102 pub supported: Option<Vec<TransportKind>>,
103}
104
105impl Default for DialConfig {
106 fn default() -> DialConfig {
107 DialConfig {
108 attempt_timeout: DEFAULT_DIAL_TIMEOUT,
109 supported: None,
110 }
111 }
112}
113
114impl DialConfig {
115 fn kind_supported(&self, kind: TransportKind) -> bool {
116 match &self.supported {
117 Some(kinds) => kinds.contains(&kind),
118 None => kind.supported(),
119 }
120 }
121}
122
123pub struct Peers {
124 cache: Mutex<HashMap<String, (TransportKind, Instant)>>,
125 config: DialConfig,
126 #[cfg(all(feature = "webtransport", not(target_arch = "wasm32")))]
127 webtransport: unb_transport::webtransport::ClientPool,
128}
129
130impl Peers {
131 pub fn new() -> Peers {
132 Peers::with_config(DialConfig::default())
133 }
134
135 pub fn with_config(config: DialConfig) -> Peers {
136 Peers {
137 cache: Mutex::new(HashMap::new()),
138 config,
139 #[cfg(all(feature = "webtransport", not(target_arch = "wasm32")))]
140 webtransport: unb_transport::webtransport::ClientPool::new(),
141 }
142 }
143
144 #[cfg(not(all(target_arch = "wasm32", target_os = "wasi")))]
145 pub async fn dial(&self, peer: &str, set: &EndpointSet) -> Result<Arc<ClientSession>, WsError> {
146 self.dial_with(peer, set, |endpoint| async move {
147 self.dial_candidate(&endpoint).await
148 })
149 .await
150 }
151
152 pub async fn dial_with<D, Fut>(
153 &self,
154 peer: &str,
155 set: &EndpointSet,
156 dialer: D,
157 ) -> Result<Arc<ClientSession>, WsError>
158 where
159 D: Fn(Endpoint) -> Fut,
160 Fut: Future<Output = Result<Pipe, WsError>>,
161 {
162 let ordered = self.ordered_candidates(peer, set);
163 let session = try_candidates_with(&ordered, &self.config, dialer).await?;
164 self.record_winner(peer, ordered[session.1].kind);
165 Ok(session.0)
166 }
167
168 pub fn ordered_candidates(&self, key: &str, set: &EndpointSet) -> Vec<Endpoint> {
169 let cached = {
170 let cache = self.cache.lock().expect("peer cache lock");
171 cache.get(key).copied()
172 };
173 candidates(set, cached, Instant::now(), &self.config)
174 }
175
176 pub fn record_winner(&self, key: &str, kind: TransportKind) {
177 let mut cache = self.cache.lock().expect("peer cache lock");
178 cache.insert(key.to_string(), (kind, Instant::now()));
179 }
180
181 pub fn attempt_timeout(&self) -> Duration {
182 self.config.attempt_timeout
183 }
184
185 #[cfg(not(all(target_arch = "wasm32", target_os = "wasi")))]
186 pub async fn dial_candidate(&self, endpoint: &Endpoint) -> Result<Pipe, WsError> {
187 match endpoint.kind {
188 TransportKind::Unix => dial_unix(&endpoint.address).await,
189 TransportKind::WebSocket => {
190 let (pipe, initiator) = unb_transport::ws::dial(&endpoint.address).await?;
191 Ok(Pipe::Piped { pipe, initiator })
192 }
193 #[cfg(all(feature = "webtransport", not(target_arch = "wasm32")))]
194 TransportKind::WebTransport => {
195 let (pipe, initiator, bodies) = self
196 .webtransport
197 .dial(&endpoint.address, endpoint.cert_hash)
198 .await?;
199 Ok(Pipe::piped_with_streams(pipe, initiator, bodies))
200 }
201 #[cfg(all(feature = "webtransport", target_arch = "wasm32"))]
202 TransportKind::WebTransport => {
203 let (pipe, initiator) =
204 unb_transport::webtransport::dial(&endpoint.address, endpoint.cert_hash)
205 .await?;
206 Ok(Pipe::Piped { pipe, initiator })
207 }
208 #[cfg(not(feature = "webtransport"))]
209 TransportKind::WebTransport => Err(WsError::Connect(
210 "webtransport support is not compiled into this client".into(),
211 )),
212 }
213 }
214}
215
216impl Default for Peers {
217 fn default() -> Peers {
218 Peers::new()
219 }
220}
221
222#[cfg(not(all(target_arch = "wasm32", target_os = "wasi")))]
223pub async fn dial_endpoints(set: &EndpointSet) -> Result<Arc<ClientSession>, WsError> {
224 let ordered = candidates(set, None, Instant::now(), &DialConfig::default());
225 let dialer = |endpoint: Endpoint| async move { dial_candidate(&endpoint).await };
226 Ok(
227 try_candidates_with(&ordered, &DialConfig::default(), dialer)
228 .await?
229 .0,
230 )
231}
232
233fn candidates(
234 set: &EndpointSet,
235 cached: Option<(TransportKind, Instant)>,
236 now: Instant,
237 config: &DialConfig,
238) -> Vec<Endpoint> {
239 let mut ordered: Vec<Endpoint> = set
240 .0
241 .iter()
242 .filter(|endpoint| config.kind_supported(endpoint.kind))
243 .cloned()
244 .collect();
245 ordered.sort_by_key(|endpoint| endpoint.kind.rank());
246 if let Some((kind, recorded_at)) = cached {
247 let advertised = ordered.iter().any(|endpoint| endpoint.kind == kind);
248 let preferred = ordered
249 .first()
250 .map(|endpoint| endpoint.kind == kind)
251 .unwrap_or(false);
252 let fresh = now.duration_since(recorded_at) < CACHE_TTL;
253 if advertised && (preferred || fresh) {
254 ordered.sort_by_key(|endpoint| (endpoint.kind != kind, endpoint.kind.rank()));
255 }
256 }
257 ordered
258}
259
260async fn try_candidates_with<D, Fut>(
261 ordered: &[Endpoint],
262 config: &DialConfig,
263 dialer: D,
264) -> Result<(Arc<ClientSession>, usize), WsError>
265where
266 D: Fn(Endpoint) -> Fut,
267 Fut: Future<Output = Result<Pipe, WsError>>,
268{
269 let mut last_error = WsError::Connect("no supported endpoint in set".into());
270 for (index, endpoint) in ordered.iter().enumerate() {
271 match n0_future::time::timeout(config.attempt_timeout, async {
272 ready_session(dialer(endpoint.clone()).await?).await
273 })
274 .await
275 {
276 Ok(Ok(session)) => return Ok((session, index)),
277 Ok(Err(error)) => last_error = error,
278 Err(_) => {
279 last_error = WsError::Connect(format!("dial timed out for {:?}", endpoint.kind));
280 }
281 }
282 }
283 Err(last_error)
284}
285
286async fn ready_session(pipe: Pipe) -> Result<Arc<ClientSession>, WsError> {
287 let wire = Wire::open(pipe);
288 match wire.session_outcome().await? {
289 SessionOutcome::Established => Ok(Arc::new(wire.client_session())),
290 SessionOutcome::Retired(reason) => {
291 Err(WsError::Connect(format!("session retired: {reason:?}")))
292 }
293 }
294}
295
296#[cfg(all(feature = "unix", unix))]
297async fn dial_unix(address: &str) -> Result<Pipe, WsError> {
298 let (pipe, bodies) =
299 unb_transport::unix::connect_with_bodies(std::path::Path::new(address)).await?;
300 Ok(Pipe::piped_with_streams(pipe, true, bodies))
301}
302
303#[cfg(not(all(feature = "unix", unix)))]
304async fn dial_unix(_address: &str) -> Result<Pipe, WsError> {
305 Err(WsError::Connect(
306 "unix transport is not compiled into this client".into(),
307 ))
308}
309
310#[cfg(not(all(target_arch = "wasm32", target_os = "wasi")))]
311pub async fn dial_candidate(endpoint: &Endpoint) -> Result<Pipe, WsError> {
312 match endpoint.kind {
313 TransportKind::Unix => dial_unix(&endpoint.address).await,
314 TransportKind::WebSocket => {
315 let (pipe, initiator) = unb_transport::ws::dial(&endpoint.address).await?;
316 Ok(Pipe::Piped { pipe, initiator })
317 }
318 #[cfg(all(feature = "webtransport", not(target_arch = "wasm32")))]
319 TransportKind::WebTransport => {
320 let (pipe, initiator, bodies) =
321 unb_transport::webtransport::dial(&endpoint.address, endpoint.cert_hash).await?;
322 Ok(Pipe::piped_with_streams(pipe, initiator, bodies))
323 }
324 #[cfg(all(feature = "webtransport", target_arch = "wasm32"))]
325 TransportKind::WebTransport => {
326 let (pipe, initiator) =
327 unb_transport::webtransport::dial(&endpoint.address, endpoint.cert_hash).await?;
328 Ok(Pipe::Piped { pipe, initiator })
329 }
330 #[cfg(not(feature = "webtransport"))]
331 TransportKind::WebTransport => Err(WsError::Connect(
332 "webtransport support is not compiled into this client".into(),
333 )),
334 }
335}
336
337#[cfg(test)]
338mod tests {
339 use super::*;
340 use crate::pair;
341
342 fn set(kinds: &[TransportKind]) -> EndpointSet {
343 kinds
344 .iter()
345 .map(|kind| Endpoint {
346 kind: *kind,
347 address: String::new(),
348 cert_hash: None,
349 })
350 .collect()
351 }
352
353 #[test]
354 fn webtransport_is_preferred_when_supported() {
355 let ordered = candidates(
356 &set(&[TransportKind::WebSocket, TransportKind::WebTransport]),
357 None,
358 Instant::now(),
359 &DialConfig::default(),
360 );
361 if TransportKind::WebTransport.supported() {
362 assert_eq!(ordered[0].kind, TransportKind::WebTransport);
363 assert_eq!(ordered[1].kind, TransportKind::WebSocket);
364 } else {
365 assert_eq!(ordered.len(), 1);
366 assert_eq!(ordered[0].kind, TransportKind::WebSocket);
367 }
368 }
369
370 #[test]
371 fn a_fresh_non_preferred_winner_is_tried_first() {
372 let now = Instant::now();
373 let ordered = candidates(
374 &set(&[TransportKind::WebTransport, TransportKind::WebSocket]),
375 Some((TransportKind::WebSocket, now)),
376 now,
377 &DialConfig::default(),
378 );
379 assert_eq!(ordered[0].kind, TransportKind::WebSocket);
380 }
381
382 #[test]
383 fn an_expired_non_preferred_winner_falls_back_to_preference_order() {
384 let now = Instant::now();
385 let stale = now.checked_sub(CACHE_TTL + Duration::from_secs(1)).unwrap();
386 let ordered = candidates(
387 &set(&[TransportKind::WebTransport, TransportKind::WebSocket]),
388 Some((TransportKind::WebSocket, stale)),
389 now,
390 &DialConfig::default(),
391 );
392 if TransportKind::WebTransport.supported() {
393 assert_eq!(ordered[0].kind, TransportKind::WebTransport);
394 }
395 }
396
397 #[test]
398 fn a_cached_kind_absent_from_the_set_is_ignored() {
399 let now = Instant::now();
400 let ordered = candidates(
401 &set(&[TransportKind::WebTransport]),
402 Some((TransportKind::WebSocket, now)),
403 now,
404 &DialConfig::default(),
405 );
406 assert!(ordered
407 .iter()
408 .all(|endpoint| endpoint.kind == TransportKind::WebTransport));
409 }
410
411 #[test]
412 fn injected_platform_support_overrides_the_compiled_intersection() {
413 let config = DialConfig {
414 supported: Some(vec![TransportKind::WebSocket]),
415 ..DialConfig::default()
416 };
417 let ordered = candidates(
418 &set(&[TransportKind::WebTransport, TransportKind::WebSocket]),
419 None,
420 Instant::now(),
421 &config,
422 );
423 assert_eq!(ordered.len(), 1);
424 assert_eq!(ordered[0].kind, TransportKind::WebSocket);
425
426 let none_supported = DialConfig {
427 supported: Some(Vec::new()),
428 ..DialConfig::default()
429 };
430 let empty = candidates(
431 &set(&[TransportKind::WebTransport, TransportKind::WebSocket]),
432 None,
433 Instant::now(),
434 &none_supported,
435 );
436 assert!(empty.is_empty());
437 }
438
439 #[test]
440 fn unsupported_kinds_are_never_candidates() {
441 let ordered = candidates(
442 &set(&[TransportKind::WebTransport, TransportKind::WebSocket]),
443 None,
444 Instant::now(),
445 &DialConfig::default(),
446 );
447 assert!(ordered.iter().all(|endpoint| endpoint.kind.supported()));
448 }
449
450 #[cfg(not(target_arch = "wasm32"))]
451 #[tokio::test]
452 async fn an_empty_intersection_is_an_error() {
453 let result = dial_endpoints(&EndpointSet::new()).await;
454 assert!(matches!(result, Err(WsError::Connect(_))));
455 }
456
457 #[cfg(not(target_arch = "wasm32"))]
458 #[tokio::test]
459 async fn a_handshake_failure_falls_back_to_the_next_endpoint() {
460 let set: EndpointSet = vec![
461 Endpoint {
462 kind: TransportKind::WebTransport,
463 address: "failed".into(),
464 cert_hash: None,
465 },
466 Endpoint {
467 kind: TransportKind::WebSocket,
468 address: "ready".into(),
469 cert_hash: None,
470 },
471 ]
472 .into();
473 let peers = Peers::with_config(DialConfig {
474 attempt_timeout: Duration::from_millis(100),
475 supported: Some(vec![TransportKind::WebTransport, TransportKind::WebSocket]),
476 });
477 let session = peers
478 .dial_with("owner", &set, |endpoint| async move {
479 if endpoint.address == "failed" {
480 return Err(WsError::Connect("failed".into()));
481 }
482 let (client, server) = pair();
483 let _server = Wire::open(server);
484 Ok(client)
485 })
486 .await
487 .unwrap();
488
489 assert_eq!(Arc::strong_count(&session), 1);
490 assert_eq!(
491 peers
492 .cache
493 .lock()
494 .unwrap()
495 .get("owner")
496 .map(|entry| entry.0),
497 Some(TransportKind::WebSocket)
498 );
499 }
500
501 #[cfg(not(target_arch = "wasm32"))]
502 #[tokio::test]
503 async fn handshake_timeout_falls_back_without_caching_the_timed_out_candidate() {
504 let set: EndpointSet = vec![
505 Endpoint {
506 kind: TransportKind::WebTransport,
507 address: "stalled".into(),
508 cert_hash: None,
509 },
510 Endpoint {
511 kind: TransportKind::WebSocket,
512 address: "ready".into(),
513 cert_hash: None,
514 },
515 ]
516 .into();
517 let peers = Peers::with_config(DialConfig {
518 attempt_timeout: Duration::from_millis(25),
519 supported: Some(vec![TransportKind::WebTransport, TransportKind::WebSocket]),
520 });
521 peers
522 .dial_with("owner", &set, |endpoint| async move {
523 let (client, server) = pair();
524 if endpoint.address == "stalled" {
525 std::mem::forget(server);
526 } else {
527 let _server = Wire::open(server);
528 }
529 Ok(client)
530 })
531 .await
532 .unwrap();
533
534 assert_eq!(
535 peers
536 .cache
537 .lock()
538 .unwrap()
539 .get("owner")
540 .map(|entry| entry.0),
541 Some(TransportKind::WebSocket)
542 );
543 }
544
545 #[cfg(not(target_arch = "wasm32"))]
546 #[tokio::test]
547 async fn a_ready_winner_is_preferred_on_the_next_dial() {
548 let set = set(&[TransportKind::WebTransport, TransportKind::WebSocket]);
549 let peers = Peers::with_config(DialConfig {
550 attempt_timeout: Duration::from_millis(100),
551 supported: Some(vec![TransportKind::WebTransport, TransportKind::WebSocket]),
552 });
553 peers
554 .dial_with("owner", &set, |endpoint| async move {
555 if endpoint.kind == TransportKind::WebTransport {
556 return Err(WsError::Connect("failed".into()));
557 }
558 let (client, server) = pair();
559 let _server = Wire::open(server);
560 Ok(client)
561 })
562 .await
563 .unwrap();
564
565 assert_eq!(
566 peers.ordered_candidates("owner", &set)[0].kind,
567 TransportKind::WebSocket
568 );
569 }
570}