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
35pub type ClientStatusCallback = Box<dyn Fn(&str) + Send + Sync>;
37
38async 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#[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 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 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 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 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 stream_tasks.shutdown().await;
406}
407
408#[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 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 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 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}