Skip to main content

pb_mapper_client/client/
mod.rs

1pub mod error;
2pub mod status;
3mod stream;
4
5use std::fmt::Debug;
6use std::sync::Arc;
7use std::time::Duration;
8
9use snafu::ResultExt;
10use tokio::net::TcpStream;
11use tokio::task::JoinSet;
12use tokio::time::MissedTickBehavior;
13use tokio_util::sync::CancellationToken;
14use uni_stream::udp::set_custom_timeout;
15
16use self::error::{AcceptLocalStreamSnafu, BindLocalListenerSnafu};
17use self::status::{get_status_scoped, get_status_with_credential};
18use self::stream::handle_local_stream;
19use crate::addr::{resolve_all, resolve_tunnel_ends};
20use pb_mapper_core::checksum::{Credential, get_process_credential};
21use pb_mapper_core::config::ResolvedAddrs;
22use pb_mapper_core::config::{
23    StatusOp, client_health_check_interval, client_health_check_timeout,
24    client_health_failure_threshold,
25};
26use pb_mapper_core::timeout::RetryBackoff;
27use pb_mapper_protocol::command::{PbConnStatusReq, PbConnStatusResp};
28use pb_mapper_protocol::forward::StreamForward;
29use uni_stream::addr::{ToSocketAddrs, each_addr};
30use uni_stream::stream::{ListenerProvider, StreamAccept};
31
32// Callback for notifying status changes to external systems
33pub type ClientStatusCallback = Box<dyn Fn(&str) + Send + Sync>;
34
35/// Resolve both ends, reporting a bad address through `status_callback`.
36///
37/// The callback is the only channel a caller has here — these entry points return
38/// `()` — so an address that never resolves has to settle the tunnel as failed
39/// rather than just leaving a log line behind.
40async fn resolve_ends_or_fail<A: ToSocketAddrs>(
41    local_addr: A,
42    remote_addr: A,
43    status_callback: Option<&ClientStatusCallback>,
44) -> Option<(ResolvedAddrs, ResolvedAddrs)> {
45    let resolved = resolve_tunnel_ends(local_addr, remote_addr).await;
46    if let (None, Some(callback)) = (&resolved, status_callback) {
47        callback("failed");
48    }
49    resolved
50}
51
52pub async fn run_client_side_cli<LocalListener: ListenerProvider, A: ToSocketAddrs>(
53    local_addr: A,
54    remote_addr: A,
55    key: Arc<str>,
56    keep_alive: bool,
57) where
58    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
59{
60    run_client_side_cli_with_callback::<LocalListener, A>(
61        local_addr,
62        remote_addr,
63        key,
64        keep_alive,
65        None,
66    )
67    .await
68}
69
70pub async fn run_client_side_cli_with_callback<LocalListener: ListenerProvider, A: ToSocketAddrs>(
71    local_addr: A,
72    remote_addr: A,
73    key: Arc<str>,
74    keep_alive: bool,
75    status_callback: Option<ClientStatusCallback>,
76) where
77    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
78{
79    run_client_side_cli_with_callback_scoped::<LocalListener, A>(
80        local_addr,
81        remote_addr,
82        key,
83        keep_alive,
84        None,
85        status_callback,
86        None,
87    )
88    .await
89}
90
91pub async fn run_client_side_cli_with_pinned_credential<
92    LocalListener: ListenerProvider,
93    A: ToSocketAddrs,
94>(
95    local_addr: A,
96    remote_addr: A,
97    key: Arc<str>,
98    keep_alive: bool,
99    status_callback: Option<ClientStatusCallback>,
100    credential: pb_mapper_core::checksum::Credential,
101) where
102    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
103{
104    run_client_side_cli_with_callback_scoped::<LocalListener, A>(
105        local_addr,
106        remote_addr,
107        key,
108        keep_alive,
109        None,
110        status_callback,
111        Some(credential),
112    )
113    .await
114}
115
116pub async fn run_client_side_cli_with_callback_scoped<
117    LocalListener: ListenerProvider,
118    A: ToSocketAddrs,
119>(
120    local_addr: A,
121    remote_addr: A,
122    key: Arc<str>,
123    keep_alive: bool,
124    namespace: Option<u64>,
125    status_callback: Option<ClientStatusCallback>,
126    pinned_credential: Option<pb_mapper_core::checksum::Credential>,
127) where
128    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
129{
130    let Some((local_addr, remote_addr)) =
131        resolve_ends_or_fail(local_addr, remote_addr, status_callback.as_ref()).await
132    else {
133        return;
134    };
135    run_client_side_cli_loop::<LocalListener>(
136        local_addr,
137        remote_addr,
138        key,
139        keep_alive,
140        namespace,
141        status_callback,
142        pinned_credential,
143        CancellationToken::new(),
144    )
145    .await
146}
147
148/// Same as [`run_client_side_cli_with_callback_scoped`], but the retry loop
149/// returns when `shutdown` is cancelled.
150///
151/// Takes both ends already resolved, because its caller — the SDK — resolves them
152/// itself in order to report a bad address as an error rather than a log line.
153#[allow(clippy::too_many_arguments)]
154pub async fn run_client_side_cli_with_shutdown<LocalListener: ListenerProvider>(
155    local_addr: ResolvedAddrs,
156    remote_addr: ResolvedAddrs,
157    key: Arc<str>,
158    keep_alive: bool,
159    namespace: Option<u64>,
160    status_callback: Option<ClientStatusCallback>,
161    pinned_credential: Option<pb_mapper_core::checksum::Credential>,
162    shutdown: CancellationToken,
163) where
164    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
165{
166    run_client_side_cli_loop::<LocalListener>(
167        local_addr,
168        remote_addr,
169        key,
170        keep_alive,
171        namespace,
172        status_callback,
173        pinned_credential,
174        shutdown,
175    )
176    .await
177}
178
179#[allow(clippy::too_many_arguments)]
180async fn run_client_side_cli_loop<LocalListener: ListenerProvider>(
181    local_addr: ResolvedAddrs,
182    remote_addr: ResolvedAddrs,
183    key: Arc<str>,
184    keep_alive: bool,
185    namespace: Option<u64>,
186    status_callback: Option<ClientStatusCallback>,
187    pinned_credential: Option<pb_mapper_core::checksum::Credential>,
188    shutdown: CancellationToken,
189) where
190    <LocalListener::Listener as StreamAccept>::Item: StreamForward,
191{
192    set_custom_timeout(Duration::from_secs(120));
193
194    let credential = match pinned_credential {
195        Some(credential) => credential,
196        None => match get_process_credential() {
197            Ok(credential) => credential,
198            Err(e) => {
199                tracing::error!("load client credential failed: {e}");
200                if let Some(ref callback) = status_callback {
201                    callback("failed");
202                }
203                return;
204            }
205        },
206    };
207
208    let mut retry_backoff = RetryBackoff::default();
209    // Accepted local streams are tracked rather than detached, so a cancelled
210    // tunnel takes its in-flight forwarding sessions down with it instead of
211    // leaving them forwarding after `stop()` has returned. The set is declared
212    // outside the loop: a listener restart is not a reason to drop live sessions.
213    let mut stream_tasks = JoinSet::new();
214
215    'outer: loop {
216        if shutdown.is_cancelled() {
217            break 'outer;
218        }
219        tracing::debug!(
220            event = "client_probe_start",
221            key = %key,
222            local_addr = %local_addr,
223            remote_addr = %remote_addr,
224            retry_count = retry_backoff.failures(),
225            "client probing remote server"
226        );
227
228        if let Err(failure) =
229            probe_remote_key(&remote_addr, key.as_ref(), namespace, credential).await
230        {
231            // A refusal the relay marks permanent — a namespace this credential
232            // does not own, an invalid service name — cannot be fixed by trying
233            // again. Looping on it would leave the caller's `wait_ready` pending
234            // forever with nothing but "retrying" to show for it, so the reason is
235            // reported and the loop ends. Mirrors the register side's
236            // `Status::Rejected`.
237            if failure.permanent {
238                tracing::error!(
239                    event = "client_remote_probe_rejected_permanently",
240                    key = %key,
241                    local_addr = %local_addr,
242                    remote_addr = %remote_addr,
243                    reason = %failure,
244                    "pb server permanently refused this subscription; not retrying"
245                );
246                if let Some(ref callback) = status_callback {
247                    callback(&format!("failed: {failure}"));
248                }
249                break 'outer;
250            }
251            let retry_delay = retry_backoff.next_delay();
252            tracing::warn!(
253                event = "client_remote_probe_failed",
254                key = %key,
255                local_addr = %local_addr,
256                remote_addr = %remote_addr,
257                reason = %failure,
258                retry_delay = ?retry_delay,
259                retry_count = retry_backoff.failures(),
260                "client remote probe failed; retrying"
261            );
262            if let Some(ref callback) = status_callback {
263                callback("retrying");
264            }
265            tokio::select! {
266                () = shutdown.cancelled() => break 'outer,
267                () = tokio::time::sleep(retry_delay) => {}
268            }
269            continue;
270        }
271
272        tracing::info!(
273            event = "client_key_available",
274            key = %key,
275            local_addr = %local_addr,
276            remote_addr = %remote_addr,
277            "remote server key is available; local listener will start"
278        );
279
280        retry_backoff.reset();
281
282        // The listener binds before "connected" is reported: that status is what
283        // drives readiness for external callers, and a caller told the tunnel is
284        // up must be able to reach the local endpoint. Reporting it on the remote
285        // probe alone would call an occupied local address ready.
286        let listener = match LocalListener::bind(local_addr.as_slice())
287            .await
288            .context(BindLocalListenerSnafu)
289        {
290            Ok(listener) => listener,
291            Err(e) => {
292                tracing::error!(
293                    event = "client_local_bind_failed",
294                    key = %key,
295                    local_addr = %local_addr,
296                    error = %e,
297                    "failed to bind local listener"
298                );
299                if let Some(ref callback) = status_callback {
300                    callback("retrying");
301                }
302                let retry_delay = retry_backoff.next_delay();
303                tokio::select! {
304                    () = shutdown.cancelled() => break 'outer,
305                    () = tokio::time::sleep(retry_delay) => {}
306                }
307                continue;
308            }
309        };
310
311        tracing::info!(
312            event = "client_local_listener_bound",
313            key = %key,
314            local_addr = %local_addr,
315            remote_addr = %remote_addr,
316            "local listener bound; tunnel is ready"
317        );
318
319        if let Some(ref callback) = status_callback {
320            callback("connected");
321        }
322
323        let (stream_failure_tx, mut stream_failure_rx) = tokio::sync::mpsc::unbounded_channel();
324        let mut health_interval = tokio::time::interval(client_health_check_interval());
325        let health_failure_threshold = client_health_failure_threshold();
326        let mut consecutive_health_failures = 0usize;
327        health_interval.set_missed_tick_behavior(MissedTickBehavior::Delay);
328        health_interval.tick().await;
329
330        loop {
331            tokio::select! {
332                () = shutdown.cancelled() => {
333                    tracing::info!(
334                        event = "client_listener_cancelled",
335                        key = %key,
336                        local_addr = %local_addr,
337                        "client listener loop cancelled"
338                    );
339                    break 'outer;
340                }
341                accepted = listener.accept() => {
342                    let (stream, peer_addr) = match accepted.context(AcceptLocalStreamSnafu) {
343                        Ok(result) => result,
344                        Err(e) => {
345                            tracing::error!(
346                                event = "client_local_accept_failed",
347                                key = %key,
348                                local_addr = %local_addr,
349                                error = %e,
350                                "failed to accept local stream"
351                            );
352                            break;
353                        }
354                    };
355                    tracing::debug!(
356                        event = "client_local_stream_accepted",
357                        key = %key,
358                        local_addr = %local_addr,
359                        peer_addr = ?peer_addr,
360                        "accepted local client stream"
361                    );
362                    let key = key.clone();
363                    let failure_tx = stream_failure_tx.clone();
364                    let stream_shutdown = shutdown.clone();
365                    let stream_remote = remote_addr.clone();
366                    stream_tasks.spawn(async move {
367                        let forward = handle_local_stream(stream, key, stream_remote.clone(), keep_alive, namespace, credential);
368                        let forward = tokio::select! {
369                            () = stream_shutdown.cancelled() => return,
370                            result = forward => result,
371                        };
372                        if let Err(e) = forward
373                        {
374                            let reason = snafu::Report::from_error(e).to_string();
375                            tracing::warn!(
376                                event = "client_local_stream_failed_before_forward",
377                                remote_addr = %stream_remote,
378                                reason = %reason,
379                                "local client stream failed before forwarding started"
380                            );
381                            let _ = failure_tx.send(reason);
382                        }
383                    });
384                }
385                _ = health_interval.tick() => {
386                    if let Err(reason) = probe_remote_key(&remote_addr, key.as_ref(), namespace, credential).await {
387                        consecutive_health_failures = consecutive_health_failures.saturating_add(1);
388                        if consecutive_health_failures < health_failure_threshold {
389                            tracing::warn!(
390                                event = "client_remote_health_check_missed",
391                                key = %key,
392                                local_addr = %local_addr,
393                                remote_addr = %remote_addr,
394                                reason = %reason,
395                                consecutive_failures = consecutive_health_failures,
396                                failure_threshold = health_failure_threshold,
397                                "client remote health check failed; listener remains active"
398                            );
399                            continue;
400                        }
401                        tracing::warn!(
402                            event = "client_remote_health_check_failed",
403                            key = %key,
404                            local_addr = %local_addr,
405                            remote_addr = %remote_addr,
406                            reason = %reason,
407                            consecutive_failures = consecutive_health_failures,
408                            failure_threshold = health_failure_threshold,
409                            "client remote health checks failed repeatedly; listener will restart"
410                        );
411                        if let Some(ref callback) = status_callback {
412                            callback("retrying");
413                        }
414                        break;
415                    }
416                    consecutive_health_failures = 0;
417                    retry_backoff.reset();
418                }
419                Some(_) = stream_tasks.join_next() => {
420                    // Reap finished sessions so the set does not grow for the
421                    // lifetime of the process. Failures are already reported
422                    // through `stream_failure_tx`.
423                }
424                Some(stream_failure) = stream_failure_rx.recv() => {
425                    tracing::warn!(
426                        event = "client_stream_failure_reported",
427                        key = %key,
428                        local_addr = %local_addr,
429                        remote_addr = %remote_addr,
430                        stream_failure = %stream_failure,
431                        "local stream failure reported; probing remote key"
432                    );
433                    if let Err(reason) = probe_remote_key(&remote_addr, key.as_ref(), namespace, credential).await {
434                        tracing::warn!(
435                            event = "client_remote_probe_failed_after_stream_error",
436                            key = %key,
437                            local_addr = %local_addr,
438                            remote_addr = %remote_addr,
439                            reason = %reason,
440                            "remote key probe failed after local stream error; listener will restart"
441                        );
442                        if let Some(ref callback) = status_callback {
443                            callback("retrying");
444                        }
445                        break;
446                    }
447                    consecutive_health_failures = 0;
448                    retry_backoff.reset();
449                }
450            }
451        }
452
453        if shutdown.is_cancelled() {
454            break 'outer;
455        }
456        let retry_delay = retry_backoff.next_delay();
457        tracing::info!(
458            event = "client_listener_restart_scheduled",
459            key = %key,
460            local_addr = %local_addr,
461            remote_addr = %remote_addr,
462            retry_delay = ?retry_delay,
463            retry_count = retry_backoff.failures(),
464            "client listener stopped; remote probe will retry"
465        );
466        tokio::select! {
467            () = shutdown.cancelled() => break 'outer,
468            () = tokio::time::sleep(retry_delay) => {}
469        }
470    }
471
472    // Wait for the forwarding sessions to observe the cancellation, so returning
473    // from here means no stream of this tunnel is still moving bytes.
474    stream_tasks.shutdown().await;
475}
476
477/// Why a remote probe did not confirm the service, and whether waiting could
478/// change that.
479///
480/// The distinction is what keeps the retry loop from spinning: a service that has
481/// not registered yet may still appear, but a namespace the credential does not
482/// own never will.
483#[derive(Debug)]
484struct ProbeFailure {
485    reason: String,
486    permanent: bool,
487}
488
489impl ProbeFailure {
490    /// A failure worth retrying: a transport error, or a service that is simply
491    /// not registered yet.
492    fn transient(reason: impl Into<String>) -> Self {
493        Self {
494            reason: reason.into(),
495            permanent: false,
496        }
497    }
498
499    /// A status failure classified by the relay's own `retryable` verdict, which
500    /// is the only party that knows whether the refusal is final.
501    fn from_status_error(context: &str, error: crate::client::error::Error) -> Self {
502        let permanent = error.remote_retryable() == Some(false);
503        Self {
504            reason: format!("{context}: {}", snafu::Report::from_error(error)),
505            permanent,
506        }
507    }
508}
509
510impl std::fmt::Display for ProbeFailure {
511    fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
512        formatter.write_str(&self.reason)
513    }
514}
515
516async fn probe_remote_key(
517    remote_addr: &ResolvedAddrs,
518    key: &str,
519    namespace: Option<u64>,
520    credential: Credential,
521) -> std::result::Result<(), ProbeFailure> {
522    let timeout = client_health_check_timeout();
523    match tokio::time::timeout(
524        timeout,
525        probe_remote_key_once(remote_addr, key, namespace, credential),
526    )
527    .await
528    {
529        Ok(result) => result,
530        Err(_) => Err(ProbeFailure::transient(format!(
531            "remote key probe timed out after {timeout:?}"
532        ))),
533    }
534}
535
536async fn probe_remote_key_once(
537    remote_addr: &ResolvedAddrs,
538    key: &str,
539    namespace: Option<u64>,
540    credential: Credential,
541) -> std::result::Result<(), ProbeFailure> {
542    match fetch_remote_status(
543        remote_addr,
544        PbConnStatusReq::Service {
545            key: key.to_string(),
546        },
547        namespace,
548        credential,
549    )
550    .await
551    {
552        Ok(PbConnStatusResp::Service { connections, .. }) => {
553            if connections.iter().any(|conn| conn.healthy) {
554                return Ok(());
555            }
556            return Err(ProbeFailure::transient(format!(
557                "client key `{key}` has no healthy remote server connections"
558            )));
559        }
560        Ok(status_resp) => {
561            return Err(ProbeFailure::transient(format!(
562                "expected service status response, got {status_resp:?}"
563            )));
564        }
565        // A refusal the relay calls final applies to the keys probe just as much,
566        // so there is nothing to fall back to.
567        Err(failure) if failure.permanent => return Err(failure),
568        Err(service_failure) => {
569            tracing::debug!(
570                event = "client_remote_service_probe_failed",
571                key = %key,
572                remote_addr = %remote_addr,
573                reason = %service_failure,
574                "service status probe failed; falling back to key status"
575            );
576        }
577    }
578
579    let status_resp =
580        fetch_remote_status(remote_addr, PbConnStatusReq::Keys, namespace, credential).await?;
581    let PbConnStatusResp::Keys(keys) = status_resp else {
582        return Err(ProbeFailure::transient(format!(
583            "expected keys status response, got {status_resp:?}"
584        )));
585    };
586    if keys.iter().any(|candidate| candidate == key) {
587        Ok(())
588    } else {
589        Err(ProbeFailure::transient(format!(
590            "client key `{key}` is not registered on remote server; valid keys: {keys:?}"
591        )))
592    }
593}
594
595async fn fetch_remote_status(
596    remote_addr: &ResolvedAddrs,
597    req: PbConnStatusReq,
598    namespace: Option<u64>,
599    credential: Credential,
600) -> std::result::Result<PbConnStatusResp, ProbeFailure> {
601    let mut stream = each_addr(remote_addr.as_slice(), TcpStream::connect)
602        .await
603        .map_err(|error| {
604            ProbeFailure::transient(format!("connect remote stream failed: {error}"))
605        })?;
606    get_status_with_credential(&mut stream, req, namespace, &credential)
607        .await
608        .map_err(|error| ProbeFailure::from_status_error("get status failed", error))
609}
610
611pub async fn handle_status_cli_scoped<A: ToSocketAddrs>(
612    op: StatusOp,
613    addr: A,
614    namespace: Option<u64>,
615) -> Result<(), Box<dyn std::error::Error>> {
616    match op {
617        StatusOp::RemoteId => show_status_scoped(addr, PbConnStatusReq::RemoteId, namespace).await,
618        StatusOp::Keys => show_status_scoped(addr, PbConnStatusReq::Keys, namespace).await,
619    }
620}
621
622pub async fn show_status_scoped<A: ToSocketAddrs>(
623    remote_addr: A,
624    req: PbConnStatusReq,
625    namespace: Option<u64>,
626) -> Result<(), Box<dyn std::error::Error>> {
627    let remote_addr = resolve_all(remote_addr).await?;
628    let mut stream = each_addr(remote_addr.as_slice(), TcpStream::connect)
629        .await
630        .map_err(|error| format!("get status stream: {error}"))?;
631    let status = get_status_scoped(&mut stream, req, namespace).await?;
632    let status = serde_json::to_string_pretty(&status)?;
633    println!("Status:{status}");
634    Ok(())
635}