Skip to main content

pb_mapper_client/client/
mod.rs

1pub mod error;
2pub mod status;
3mod stream;
4
5use crate::diagnostics::{Diagnostics, RecoveryFailure, RecoveryPhase};
6use crate::endpoint::RelayEndpoint;
7use std::fmt::Debug;
8use std::sync::Arc;
9use std::time::Duration;
10
11use snafu::ResultExt;
12use tokio::sync::{Mutex, Semaphore};
13use tokio::task::JoinSet;
14use tokio::time::Instant;
15use tokio_util::sync::CancellationToken;
16use uni_stream::udp::set_custom_timeout;
17
18use self::error::{AcceptLocalStreamSnafu, BindLocalListenerSnafu};
19use self::status::{get_status_scoped, get_status_with_credential};
20use self::stream::{StreamSetup, handle_local_stream};
21use crate::addr::{resolve_all, resolve_tunnel_ends};
22use crate::recovery::{RecoveryTiming, jitter};
23use pb_mapper_core::checksum::{Credential, get_process_credential};
24use pb_mapper_core::config::ResolvedAddrs;
25use pb_mapper_core::config::{
26    StatusOp, client_health_check_interval, client_health_check_timeout,
27    client_health_failure_threshold,
28};
29use pb_mapper_core::timeout::RetryBackoff;
30use pb_mapper_protocol::command::{PbConnStatusReq, PbConnStatusResp};
31use pb_mapper_protocol::forward::StreamForward;
32use uni_stream::addr::ToSocketAddrs;
33use uni_stream::stream::{ListenerProvider, StreamAccept};
34
35// Callback for notifying status changes to external systems
36pub type ClientStatusCallback = Box<dyn Fn(&str) + Send + Sync>;
37
38/// Resolve both ends, reporting a bad address through `status_callback`.
39///
40/// The callback is the only channel a caller has here — these entry points return
41/// `()` — so an address that never resolves has to settle the tunnel as failed
42/// rather than just leaving a log line behind.
43async fn resolve_ends_or_fail<A: ToSocketAddrs>(
44    local_addr: A,
45    remote_addr: A,
46    status_callback: Option<&ClientStatusCallback>,
47) -> Option<(ResolvedAddrs, ResolvedAddrs)> {
48    let resolved = resolve_tunnel_ends(local_addr, remote_addr).await;
49    if let (None, Some(callback)) = (&resolved, status_callback) {
50        callback("failed");
51    }
52    resolved
53}
54
55pub async fn run_client_side_cli<LocalListener: ListenerProvider, A: ToSocketAddrs>(
56    local_addr: A,
57    remote_addr: A,
58    key: Arc<str>,
59    keep_alive: bool,
60) where
61    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
62{
63    run_client_side_cli_with_callback::<LocalListener, A>(
64        local_addr,
65        remote_addr,
66        key,
67        keep_alive,
68        None,
69    )
70    .await
71}
72
73pub async fn run_client_side_cli_with_callback<LocalListener: ListenerProvider, A: ToSocketAddrs>(
74    local_addr: A,
75    remote_addr: A,
76    key: Arc<str>,
77    keep_alive: bool,
78    status_callback: Option<ClientStatusCallback>,
79) where
80    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
81{
82    run_client_side_cli_with_callback_scoped::<LocalListener, A>(
83        local_addr,
84        remote_addr,
85        key,
86        keep_alive,
87        None,
88        status_callback,
89        None,
90    )
91    .await
92}
93
94pub async fn run_client_side_cli_with_pinned_credential<
95    LocalListener: ListenerProvider,
96    A: ToSocketAddrs,
97>(
98    local_addr: A,
99    remote_addr: A,
100    key: Arc<str>,
101    keep_alive: bool,
102    status_callback: Option<ClientStatusCallback>,
103    credential: pb_mapper_core::checksum::Credential,
104) where
105    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
106{
107    run_client_side_cli_with_callback_scoped::<LocalListener, A>(
108        local_addr,
109        remote_addr,
110        key,
111        keep_alive,
112        None,
113        status_callback,
114        Some(credential),
115    )
116    .await
117}
118
119pub async fn run_client_side_cli_with_callback_scoped<
120    LocalListener: ListenerProvider,
121    A: ToSocketAddrs,
122>(
123    local_addr: A,
124    remote_addr: A,
125    key: Arc<str>,
126    keep_alive: bool,
127    namespace: Option<u64>,
128    status_callback: Option<ClientStatusCallback>,
129    pinned_credential: Option<pb_mapper_core::checksum::Credential>,
130) where
131    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
132{
133    let Some((local_addr, remote_addr)) =
134        resolve_ends_or_fail(local_addr, remote_addr, status_callback.as_ref()).await
135    else {
136        return;
137    };
138    run_client_side_cli_loop::<LocalListener>(
139        local_addr,
140        remote_addr,
141        key,
142        keep_alive,
143        namespace,
144        status_callback,
145        pinned_credential,
146        CancellationToken::new(),
147    )
148    .await
149}
150
151/// Same as [`run_client_side_cli_with_callback_scoped`], but the retry loop
152/// returns when `shutdown` is cancelled.
153///
154/// Takes both ends already resolved, because its caller — the SDK — resolves them
155/// itself in order to report a bad address as an error rather than a log line.
156#[allow(clippy::too_many_arguments)]
157pub async fn run_client_side_cli_with_shutdown<LocalListener: ListenerProvider>(
158    local_addr: ResolvedAddrs,
159    remote_addr: ResolvedAddrs,
160    key: Arc<str>,
161    keep_alive: bool,
162    namespace: Option<u64>,
163    status_callback: Option<ClientStatusCallback>,
164    pinned_credential: Option<pb_mapper_core::checksum::Credential>,
165    shutdown: CancellationToken,
166) where
167    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
168{
169    run_client_side_cli_loop::<LocalListener>(
170        local_addr,
171        remote_addr,
172        key,
173        keep_alive,
174        namespace,
175        status_callback,
176        pinned_credential,
177        shutdown,
178    )
179    .await
180}
181
182#[allow(clippy::too_many_arguments)]
183async fn run_client_side_cli_loop<LocalListener: ListenerProvider>(
184    local_addr: ResolvedAddrs,
185    remote_addr: ResolvedAddrs,
186    key: Arc<str>,
187    keep_alive: bool,
188    namespace: Option<u64>,
189    status_callback: Option<ClientStatusCallback>,
190    pinned_credential: Option<pb_mapper_core::checksum::Credential>,
191    shutdown: CancellationToken,
192) where
193    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
194{
195    run_client_side_cli_recovering::<LocalListener>(
196        local_addr,
197        RelayEndpoint::fixed(remote_addr),
198        key,
199        keep_alive,
200        namespace,
201        status_callback,
202        pinned_credential,
203        shutdown,
204        Diagnostics::default(),
205    )
206    .await;
207}
208
209#[allow(clippy::too_many_arguments)]
210pub(crate) async fn run_client_side_cli_recovering<LocalListener: ListenerProvider>(
211    local_addr: ResolvedAddrs,
212    remote_addr: RelayEndpoint,
213    key: Arc<str>,
214    keep_alive: bool,
215    namespace: Option<u64>,
216    status_callback: Option<ClientStatusCallback>,
217    pinned_credential: Option<pb_mapper_core::checksum::Credential>,
218    shutdown: CancellationToken,
219    diagnostics: Diagnostics,
220) where
221    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
222{
223    remote_addr.start();
224    let mut wake = remote_addr.wake();
225    set_custom_timeout(Duration::from_secs(120));
226
227    let credential = match pinned_credential {
228        Some(credential) => credential,
229        None => match get_process_credential() {
230            Ok(credential) => credential,
231            Err(e) => {
232                tracing::error!("load client credential failed: {e}");
233                if let Some(ref callback) = status_callback {
234                    callback("failed");
235                }
236                return;
237            }
238        },
239    };
240
241    let mut retry_backoff = RetryBackoff::new(Duration::from_millis(100), Duration::from_secs(2));
242    let mut stream_tasks = JoinSet::new();
243    let setup_slots = Arc::new(Semaphore::new(64));
244    let timing = Arc::new(Mutex::new(RecoveryTiming::default()));
245
246    'outer: loop {
247        // Listener lifetime follows the local endpoint, not a remote health
248        // snapshot. Only a local bind/accept error requires replacing it.
249        let listener = tokio::select! {
250            () = shutdown.cancelled() => break,
251            result = LocalListener::bind(local_addr.as_slice()) => match result.context(BindLocalListenerSnafu) {
252                Ok(listener) => listener,
253                Err(error) => {
254                    tracing::warn!(event = "client_local_bind_failed", %key, %local_addr, %error);
255                    if let Some(callback) = &status_callback { callback("retrying"); }
256                    tokio::select! {
257                        () = shutdown.cancelled() => break,
258                        () = tokio::time::sleep(jitter(retry_backoff.next_delay())) => {}
259                    }
260                    continue;
261                }
262            }
263        };
264        tracing::info!(event = "client_local_listener_bound", %key, %local_addr, "local listener bound; checking remote service");
265        let mut probes = JoinSet::new();
266        let mut next_probe = Instant::now();
267        let mut last_success = None;
268        let mut connected = false;
269        let mut consecutive_health_failures = 0usize;
270        let (stream_event_tx, mut stream_event_rx) = tokio::sync::mpsc::channel(1);
271
272        loop {
273            tokio::select! {
274                () = shutdown.cancelled() => break 'outer,
275                () = wake.changed() => {
276                    probes.abort_all();
277                    next_probe = Instant::now();
278                    retry_backoff.reset();
279                },
280                accepted = async {
281                    // Backpressure stays in the listener backlog while setup is
282                    // saturated. Established forwarding releases its permit.
283                    let permit = setup_slots.clone().acquire_owned().await;
284                    (permit, listener.accept().await)
285                } => {
286                    let (Ok(permit), accepted) = accepted else { break 'outer; };
287                    let (stream, _) = match accepted.context(AcceptLocalStreamSnafu) {
288                        Ok(accepted) => accepted,
289                        Err(error) => {
290                            tracing::warn!(event = "client_local_accept_failed", %key, %local_addr, %error);
291                            break;
292                        }
293                    };
294                    let stream_key = key.clone();
295                    let stream_remote = remote_addr.clone();
296                    let stream_shutdown = shutdown.clone();
297                    let event_tx = stream_event_tx.clone();
298                    let setup = StreamSetup { permit, timing: timing.clone(), events: event_tx.clone(), diagnostics: diagnostics.clone() };
299                    stream_tasks.spawn(async move {
300                        let result = tokio::select! {
301                            () = stream_shutdown.cancelled() => return,
302                            result = handle_local_stream(stream, stream_key, stream_remote, keep_alive, namespace, credential, setup) => result,
303                        };
304                        if let Err(error) = result {
305                            tracing::warn!(event = "client_local_stream_failed_before_forward", reason = %snafu::Report::from_error(error));
306                            let _ = event_tx.try_send(false);
307                        }
308                    });
309                }
310                () = tokio::time::sleep_until(next_probe), if probes.is_empty() => {
311                    diagnostics.attempt();
312                    diagnostics.phase(RecoveryPhase::Handshake);
313                    let probe_remote = remote_addr.clone();
314                    let probe_key = key.clone();
315                    let budget = timing.lock().await.timeout().min(client_health_check_timeout());
316                    probes.spawn(async move {
317                        let _permit = probe_remote.control_permit().await;
318                        let started = Instant::now();
319                        let result = probe_remote_key(&probe_remote, &probe_key, namespace, credential, budget).await;
320                        (started, started.elapsed(), result)
321                    });
322                }
323                Some(result) = probes.join_next() => {
324                    let (started, elapsed, result) = match result {
325                        Ok(result) => result,
326                        Err(error) if error.is_cancelled() => continue,
327                        Err(error) => (Instant::now(), Duration::ZERO, Err(ProbeFailure::transient(error.to_string()))),
328                    };
329                    match result {
330                        Ok(()) => {
331                            timing.lock().await.record(elapsed);
332                            remote_addr.protocol_succeeded();
333                            diagnostics.succeeded(elapsed);
334                            last_success = Some(Instant::now());
335                            consecutive_health_failures = 0;
336                            retry_backoff.reset();
337                            next_probe = Instant::now() + client_health_check_interval();
338                            if !connected {
339                                connected = true;
340                                if let Some(callback) = &status_callback { callback("connected"); }
341                            }
342                        }
343                        Err(failure) if failure.permanent => {
344                            if failure.protocol_replied {
345                                remote_addr.protocol_succeeded();
346                                diagnostics.responded(elapsed);
347                            }
348                            diagnostics.failed(failure.kind, Duration::ZERO);
349                            if let Some(callback) = &status_callback { callback(&format!("failed: {failure}")); }
350                            break 'outer;
351                        }
352                        Err(_) if last_success.is_some_and(|success| success > started) => {
353                            // Actual traffic completed after this probe began.
354                            next_probe = Instant::now() + client_health_check_interval();
355                        }
356                        Err(failure) => {
357                            failure.update_timing(&mut *timing.lock().await, elapsed);
358                            if failure.protocol_replied {
359                                remote_addr.protocol_succeeded();
360                                diagnostics.responded(elapsed);
361                            } else { remote_addr.transport_failed(); }
362                            consecutive_health_failures = consecutive_health_failures.saturating_add(1);
363                            if !connected || consecutive_health_failures >= client_health_failure_threshold() {
364                                connected = false;
365                                if let Some(callback) = &status_callback { callback("retrying"); }
366                            }
367                            let delay = jitter(retry_backoff.next_delay());
368                            next_probe = Instant::now() + delay;
369                            if diagnostics.failed(failure.kind, delay) {
370                            tracing::warn!(event = "client_remote_probe_failed", %key, reason = %failure, consecutive_health_failures, retry_delay = ?delay, "remote probe failed; listener remains active");
371                            }
372                        }
373                    }
374                }
375                Some(success) = stream_event_rx.recv() => {
376                    if success {
377                        last_success = Some(Instant::now());
378                        consecutive_health_failures = 0;
379                        retry_backoff.reset();
380                        if !connected {
381                            connected = true;
382                            if let Some(callback) = &status_callback { callback("connected"); }
383                        }
384                    } else if probes.is_empty() {
385                        // Coalesce failures and rate limit probes independently
386                        // of the number of callers arriving during an outage.
387                        next_probe = next_probe.min(Instant::now() + Duration::from_millis(100));
388                    }
389                }
390                Some(_) = stream_tasks.join_next() => {}
391            }
392        }
393        drop(probes);
394        if let Some(callback) = &status_callback {
395            callback("retrying");
396        }
397        tokio::select! {
398            () = shutdown.cancelled() => break,
399            () = tokio::time::sleep(jitter(retry_backoff.next_delay())) => {}
400        }
401    }
402
403    // Wait for the forwarding sessions to observe the cancellation, so returning
404    // from here means no stream of this tunnel is still moving bytes.
405    stream_tasks.shutdown().await;
406}
407
408/// Why a remote probe did not confirm the service, and whether waiting could
409/// change that.
410///
411/// The distinction is what keeps the retry loop from spinning: a service that has
412/// not registered yet may still appear, but a namespace the credential does not
413/// own never will.
414#[derive(Debug)]
415struct ProbeFailure {
416    reason: String,
417    permanent: bool,
418    kind: RecoveryFailure,
419    protocol_replied: bool,
420}
421
422impl ProbeFailure {
423    fn update_timing(&self, timing: &mut RecoveryTiming, elapsed: Duration) {
424        if self.kind == RecoveryFailure::Timeout {
425            timing.timed_out();
426        } else if self.protocol_replied {
427            timing.record(elapsed);
428        }
429    }
430
431    /// A failure worth retrying: a transport error, or a service that is simply
432    /// not registered yet.
433    fn transient(reason: impl Into<String>) -> Self {
434        Self {
435            reason: reason.into(),
436            permanent: false,
437            kind: RecoveryFailure::Transport,
438            protocol_replied: false,
439        }
440    }
441
442    /// A status failure classified by the relay's own `retryable` verdict, which
443    /// is the only party that knows whether the refusal is final.
444    fn from_status_error(context: &str, error: crate::client::error::Error) -> Self {
445        let permanent = error.remote_retryable() == Some(false);
446        let protocol_replied = error.remote_retryable().is_some();
447        Self {
448            reason: format!("{context}: {}", snafu::Report::from_error(error)),
449            permanent,
450            kind: if protocol_replied {
451                RecoveryFailure::Rejected
452            } else {
453                RecoveryFailure::Transport
454            },
455            protocol_replied,
456        }
457    }
458}
459
460impl std::fmt::Display for ProbeFailure {
461    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
462        formatter.write_str(&self.reason)
463    }
464}
465
466async fn probe_remote_key(
467    remote_addr: &RelayEndpoint,
468    key: &str,
469    namespace: Option<u64>,
470    credential: Credential,
471    timeout: Duration,
472) -> std::result::Result<(), ProbeFailure> {
473    match tokio::time::timeout(timeout, async {
474        let addresses = remote_addr
475            .addresses()
476            .await
477            .map_err(|error| ProbeFailure {
478                reason: error.to_string(),
479                permanent: false,
480                kind: RecoveryFailure::Dns,
481                protocol_replied: false,
482            })?;
483        probe_remote_key_once(&addresses, key, namespace, credential).await
484    })
485    .await
486    {
487        Ok(result) => result,
488        Err(_) => Err(ProbeFailure {
489            reason: format!("remote key probe timed out after {timeout:?}"),
490            permanent: false,
491            kind: RecoveryFailure::Timeout,
492            protocol_replied: false,
493        }),
494    }
495}
496
497async fn probe_remote_key_once(
498    remote_addr: &ResolvedAddrs,
499    key: &str,
500    namespace: Option<u64>,
501    credential: Credential,
502) -> std::result::Result<(), ProbeFailure> {
503    match fetch_remote_status(
504        remote_addr,
505        PbConnStatusReq::Service {
506            key: key.to_string(),
507        },
508        namespace,
509        credential,
510    )
511    .await
512    {
513        Ok(PbConnStatusResp::Service { connections, .. }) => {
514            if connections.iter().any(|conn| conn.healthy) {
515                return Ok(());
516            }
517            return Err(ProbeFailure {
518                reason: format!("client key `{key}` has no healthy remote server connections"),
519                permanent: false,
520                kind: RecoveryFailure::ServiceUnavailable,
521                protocol_replied: true,
522            });
523        }
524        Ok(status_resp) => {
525            return Err(ProbeFailure::transient(format!(
526                "expected service status response, got {status_resp:?}"
527            )));
528        }
529        // A refusal the relay calls final applies to the keys probe just as much,
530        // so there is nothing to fall back to.
531        Err(failure) if failure.permanent => return Err(failure),
532        Err(service_failure) => {
533            tracing::debug!(
534                event = "client_remote_service_probe_failed",
535                key = %key,
536                remote_addr = %remote_addr,
537                reason = %service_failure,
538                "service status probe failed; falling back to key status"
539            );
540        }
541    }
542
543    let status_resp =
544        fetch_remote_status(remote_addr, PbConnStatusReq::Keys, namespace, credential).await?;
545    let PbConnStatusResp::Keys(keys) = status_resp else {
546        return Err(ProbeFailure::transient(format!(
547            "expected keys status response, got {status_resp:?}"
548        )));
549    };
550    if keys.iter().any(|candidate| candidate == key) {
551        Ok(())
552    } else {
553        Err(ProbeFailure {
554            reason: format!("client key `{key}` is not registered on remote server"),
555            permanent: false,
556            kind: RecoveryFailure::ServiceUnavailable,
557            protocol_replied: true,
558        })
559    }
560}
561
562async fn fetch_remote_status(
563    remote_addr: &ResolvedAddrs,
564    req: PbConnStatusReq,
565    namespace: Option<u64>,
566    credential: Credential,
567) -> std::result::Result<PbConnStatusResp, ProbeFailure> {
568    let mut stream = crate::addr::connect_tcp(remote_addr)
569        .await
570        .map_err(|error| {
571            ProbeFailure::transient(format!("connect remote stream failed: {error}"))
572        })?;
573    get_status_with_credential(&mut stream, req, namespace, &credential)
574        .await
575        .map_err(|error| ProbeFailure::from_status_error("get status failed", error))
576}
577
578pub async fn handle_status_cli_scoped<A: ToSocketAddrs>(
579    op: StatusOp,
580    addr: A,
581    namespace: Option<u64>,
582) -> Result<(), Box<dyn std::error::Error>> {
583    match op {
584        StatusOp::RemoteId => show_status_scoped(addr, PbConnStatusReq::RemoteId, namespace).await,
585        StatusOp::Keys => show_status_scoped(addr, PbConnStatusReq::Keys, namespace).await,
586    }
587}
588
589pub async fn show_status_scoped<A: ToSocketAddrs>(
590    remote_addr: A,
591    req: PbConnStatusReq,
592    namespace: Option<u64>,
593) -> Result<(), Box<dyn std::error::Error>> {
594    let remote_addr = resolve_all(remote_addr).await?;
595    let mut stream = crate::addr::connect_tcp(&remote_addr)
596        .await
597        .map_err(|error| format!("get status stream: {error}"))?;
598    let status = get_status_scoped(&mut stream, req, namespace).await?;
599    let status = serde_json::to_string_pretty(&status)?;
600    println!("Status:{status}");
601    Ok(())
602}
603
604#[cfg(test)]
605#[allow(clippy::unwrap_used, clippy::expect_used)]
606mod cancellation_audit {
607    use std::{
608        future::Future,
609        sync::{
610            Arc,
611            atomic::{AtomicBool, Ordering},
612        },
613        task::{Context, Wake, Waker},
614        time::Duration,
615    };
616    struct Notified(AtomicBool);
617    impl Wake for Notified {
618        fn wake(self: Arc<Self>) {
619            self.0.store(true, Ordering::SeqCst);
620        }
621        fn wake_by_ref(self: &Arc<Self>) {
622            self.0.store(true, Ordering::SeqCst);
623        }
624    }
625    #[tokio::test]
626    async fn cancelled_udp_accept_must_preserve_the_delivered_peer() {
627        let listener = uni_stream::udp::UdpListener::bind("127.0.0.1:0")
628            .await
629            .unwrap();
630        let addr = listener.local_addr().unwrap();
631        let client = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
632        let notified = Arc::new(Notified(AtomicBool::new(false)));
633        let waker = Waker::from(notified.clone());
634        let mut accept = Box::pin(listener.accept());
635        assert!(
636            accept
637                .as_mut()
638                .poll(&mut Context::from_waker(&waker))
639                .is_pending()
640        );
641        client.send_to(b"first datagram", addr).await.unwrap();
642        tokio::time::timeout(Duration::from_secs(1), async {
643            while !notified.0.load(Ordering::SeqCst) {
644                tokio::task::yield_now().await;
645            }
646        })
647        .await
648        .unwrap();
649        drop(accept);
650        let delivered = tokio::time::timeout(Duration::from_millis(100), listener.accept()).await;
651        assert!(
652            delivered.is_ok(),
653            "UDP peer/first datagram disappeared when a competing select branch cancelled accept"
654        );
655    }
656}
657
658#[cfg(test)]
659mod recovery_tests {
660    use super::*;
661    #[test]
662    fn absent_service_learns_response_latency_without_inflating_timeout() {
663        let mut timing = RecoveryTiming::default();
664        let missing = ProbeFailure {
665            reason: String::new(),
666            permanent: false,
667            kind: RecoveryFailure::ServiceUnavailable,
668            protocol_replied: true,
669        };
670        for _ in 0..20 {
671            missing.update_timing(&mut timing, Duration::from_millis(40));
672        }
673        assert_eq!(timing.timeout(), Duration::from_secs(1));
674        let refusal = ProbeFailure::transient("connection refused");
675        refusal.update_timing(&mut timing, Duration::from_millis(1));
676        assert_eq!(timing.timeout(), Duration::from_secs(1));
677        let timeout = ProbeFailure {
678            kind: RecoveryFailure::Timeout,
679            ..refusal
680        };
681        timeout.update_timing(&mut timing, Duration::from_secs(1));
682        assert_eq!(timing.timeout(), Duration::from_secs(2));
683    }
684}