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
32pub type ClientStatusCallback = Box<dyn Fn(&str) + Send + Sync>;
34
35async 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#[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 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 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 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 }
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 stream_tasks.shutdown().await;
475}
476
477#[derive(Debug)]
484struct ProbeFailure {
485 reason: String,
486 permanent: bool,
487}
488
489impl ProbeFailure {
490 fn transient(reason: impl Into<String>) -> Self {
493 Self {
494 reason: reason.into(),
495 permanent: false,
496 }
497 }
498
499 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 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}