Skip to main content

unb_client/
endpoint.rs

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)
322                    .await?;
323            Ok(Pipe::piped_with_streams(pipe, initiator, bodies))
324        }
325        #[cfg(all(feature = "webtransport", target_arch = "wasm32"))]
326        TransportKind::WebTransport => {
327            let (pipe, initiator) =
328                unb_transport::webtransport::dial(&endpoint.address, endpoint.cert_hash)
329                    .await?;
330            Ok(Pipe::Piped { pipe, initiator })
331        }
332        #[cfg(not(feature = "webtransport"))]
333        TransportKind::WebTransport => Err(WsError::Connect(
334            "webtransport support is not compiled into this client".into(),
335        )),
336    }
337}
338
339#[cfg(test)]
340mod tests {
341    use super::*;
342    use crate::pair;
343
344    fn set(kinds: &[TransportKind]) -> EndpointSet {
345        kinds
346            .iter()
347            .map(|kind| Endpoint {
348                kind: *kind,
349                address: String::new(),
350                cert_hash: None,
351            })
352            .collect()
353    }
354
355    #[test]
356    fn webtransport_is_preferred_when_supported() {
357        let ordered = candidates(
358            &set(&[TransportKind::WebSocket, TransportKind::WebTransport]),
359            None,
360            Instant::now(),
361            &DialConfig::default(),
362        );
363        if TransportKind::WebTransport.supported() {
364            assert_eq!(ordered[0].kind, TransportKind::WebTransport);
365            assert_eq!(ordered[1].kind, TransportKind::WebSocket);
366        } else {
367            assert_eq!(ordered.len(), 1);
368            assert_eq!(ordered[0].kind, TransportKind::WebSocket);
369        }
370    }
371
372    #[test]
373    fn a_fresh_non_preferred_winner_is_tried_first() {
374        let now = Instant::now();
375        let ordered = candidates(
376            &set(&[TransportKind::WebTransport, TransportKind::WebSocket]),
377            Some((TransportKind::WebSocket, now)),
378            now,
379            &DialConfig::default(),
380        );
381        assert_eq!(ordered[0].kind, TransportKind::WebSocket);
382    }
383
384    #[test]
385    fn an_expired_non_preferred_winner_falls_back_to_preference_order() {
386        let now = Instant::now();
387        let stale = now.checked_sub(CACHE_TTL + Duration::from_secs(1)).unwrap();
388        let ordered = candidates(
389            &set(&[TransportKind::WebTransport, TransportKind::WebSocket]),
390            Some((TransportKind::WebSocket, stale)),
391            now,
392            &DialConfig::default(),
393        );
394        if TransportKind::WebTransport.supported() {
395            assert_eq!(ordered[0].kind, TransportKind::WebTransport);
396        }
397    }
398
399    #[test]
400    fn a_cached_kind_absent_from_the_set_is_ignored() {
401        let now = Instant::now();
402        let ordered = candidates(
403            &set(&[TransportKind::WebTransport]),
404            Some((TransportKind::WebSocket, now)),
405            now,
406            &DialConfig::default(),
407        );
408        assert!(ordered
409            .iter()
410            .all(|endpoint| endpoint.kind == TransportKind::WebTransport));
411    }
412
413    #[test]
414    fn injected_platform_support_overrides_the_compiled_intersection() {
415        let config = DialConfig {
416            supported: Some(vec![TransportKind::WebSocket]),
417            ..DialConfig::default()
418        };
419        let ordered = candidates(
420            &set(&[TransportKind::WebTransport, TransportKind::WebSocket]),
421            None,
422            Instant::now(),
423            &config,
424        );
425        assert_eq!(ordered.len(), 1);
426        assert_eq!(ordered[0].kind, TransportKind::WebSocket);
427
428        let none_supported = DialConfig {
429            supported: Some(Vec::new()),
430            ..DialConfig::default()
431        };
432        let empty = candidates(
433            &set(&[TransportKind::WebTransport, TransportKind::WebSocket]),
434            None,
435            Instant::now(),
436            &none_supported,
437        );
438        assert!(empty.is_empty());
439    }
440
441    #[test]
442    fn unsupported_kinds_are_never_candidates() {
443        let ordered = candidates(
444            &set(&[TransportKind::WebTransport, TransportKind::WebSocket]),
445            None,
446            Instant::now(),
447            &DialConfig::default(),
448        );
449        assert!(ordered.iter().all(|endpoint| endpoint.kind.supported()));
450    }
451
452    #[cfg(not(target_arch = "wasm32"))]
453    #[tokio::test]
454    async fn an_empty_intersection_is_an_error() {
455        let result = dial_endpoints(&EndpointSet::new()).await;
456        assert!(matches!(result, Err(WsError::Connect(_))));
457    }
458
459    #[cfg(not(target_arch = "wasm32"))]
460    #[tokio::test]
461    async fn a_handshake_failure_falls_back_to_the_next_endpoint() {
462        let set: EndpointSet = vec![
463            Endpoint {
464                kind: TransportKind::WebTransport,
465                address: "failed".into(),
466                cert_hash: None,
467            },
468            Endpoint {
469                kind: TransportKind::WebSocket,
470                address: "ready".into(),
471                cert_hash: None,
472            },
473        ]
474        .into();
475        let peers = Peers::with_config(DialConfig {
476            attempt_timeout: Duration::from_millis(100),
477            supported: Some(vec![TransportKind::WebTransport, TransportKind::WebSocket]),
478        });
479        let session = peers
480            .dial_with("owner", &set, |endpoint| async move {
481                if endpoint.address == "failed" {
482                    return Err(WsError::Connect("failed".into()));
483                }
484                let (client, server) = pair();
485                let _server = Wire::open(server);
486                Ok(client)
487            })
488            .await
489            .unwrap();
490
491        assert_eq!(Arc::strong_count(&session), 1);
492        assert_eq!(
493            peers
494                .cache
495                .lock()
496                .unwrap()
497                .get("owner")
498                .map(|entry| entry.0),
499            Some(TransportKind::WebSocket)
500        );
501    }
502
503    #[cfg(not(target_arch = "wasm32"))]
504    #[tokio::test]
505    async fn handshake_timeout_falls_back_without_caching_the_timed_out_candidate() {
506        let set: EndpointSet = vec![
507            Endpoint {
508                kind: TransportKind::WebTransport,
509                address: "stalled".into(),
510                cert_hash: None,
511            },
512            Endpoint {
513                kind: TransportKind::WebSocket,
514                address: "ready".into(),
515                cert_hash: None,
516            },
517        ]
518        .into();
519        let peers = Peers::with_config(DialConfig {
520            attempt_timeout: Duration::from_millis(25),
521            supported: Some(vec![TransportKind::WebTransport, TransportKind::WebSocket]),
522        });
523        peers
524            .dial_with("owner", &set, |endpoint| async move {
525                let (client, server) = pair();
526                if endpoint.address == "stalled" {
527                    std::mem::forget(server);
528                } else {
529                    let _server = Wire::open(server);
530                }
531                Ok(client)
532            })
533            .await
534            .unwrap();
535
536        assert_eq!(
537            peers
538                .cache
539                .lock()
540                .unwrap()
541                .get("owner")
542                .map(|entry| entry.0),
543            Some(TransportKind::WebSocket)
544        );
545    }
546
547    #[cfg(not(target_arch = "wasm32"))]
548    #[tokio::test]
549    async fn a_ready_winner_is_preferred_on_the_next_dial() {
550        let set = set(&[TransportKind::WebTransport, TransportKind::WebSocket]);
551        let peers = Peers::with_config(DialConfig {
552            attempt_timeout: Duration::from_millis(100),
553            supported: Some(vec![TransportKind::WebTransport, TransportKind::WebSocket]),
554        });
555        peers
556            .dial_with("owner", &set, |endpoint| async move {
557                if endpoint.kind == TransportKind::WebTransport {
558                    return Err(WsError::Connect("failed".into()));
559                }
560                let (client, server) = pair();
561                let _server = Wire::open(server);
562                Ok(client)
563            })
564            .await
565            .unwrap();
566
567        assert_eq!(
568            peers.ordered_candidates("owner", &set)[0].kind,
569            TransportKind::WebSocket
570        );
571    }
572}