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).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}