1use crate::diagnostics::{Diagnostics, RecoveryFailure, RecoveryPhase};
2use crate::endpoint::{RecoveryWake, RelayEndpoint};
3pub mod error;
4mod stream;
5
6use std::fmt::Debug;
7use std::sync::{Arc, Mutex};
8use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
9
10use snafu::ResultExt;
11use tokio::task::JoinSet;
12use tokio::time::MissedTickBehavior;
13use tokio_util::sync::CancellationToken;
14use tracing::instrument;
15
16use self::error::{
17 ControlIoTimeoutSnafu, DecodeRegisterRespSnafu, DecodeStreamReqSnafu, EncodePingMsgSnafu,
18 EncodeRegisterReqSnafu, EncodeStreamAckMsgSnafu, ReadRegisterRespSnafu, ReadStreamReqSnafu,
19 RegisterRespNotMatchSnafu, SendRegisterReqSnafu, WritePingMsgSnafu, WriteStreamAckMsgSnafu,
20};
21use self::stream::{StreamConnect, handle_stream};
22use crate::addr::resolve_tunnel_ends;
23use crate::recovery::{RecoveryTiming, jitter};
24use pb_mapper_core::checksum::{Credential, get_process_credential};
25use pb_mapper_core::config::{
26 ResolvedAddrs, control_conn_pool_size, control_heartbeat_interval, control_heartbeat_tolerance,
27 control_io_timeout, control_suspect_grace, registration_probe_timeout,
28 registration_reject_backoff,
29};
30use pb_mapper_core::timeout::RetryBackoff;
31use pb_mapper_core::{
32 snafu_error_get_or_continue, snafu_error_get_or_return, snafu_error_get_or_return_ok,
33 snafu_error_handle,
34};
35use pb_mapper_protocol::command::{
36 CONTROL_PROTOCOL_V2, LocalServer, MessageSerializer, PbConnRequest, PbConnResponse,
37 PbConnStatusReq, PbConnStatusResp, PbServerRequest,
38};
39use pb_mapper_protocol::forward::StreamForward;
40use pb_mapper_protocol::secure::ClientHeaderSession;
41use pb_mapper_protocol::{MessageReader, MessageWriter};
42use uni_stream::addr::ToSocketAddrs;
43use uni_stream::stream::{StreamProvider, set_tcp_keep_alive, set_tcp_nodelay};
44
45fn get_ping_message(protocol_version: u16, seq: u64) -> error::Result<Vec<u8>> {
46 if protocol_version >= CONTROL_PROTOCOL_V2 {
47 PbServerRequest::PingV2 { seq }
48 .encode()
49 .context(EncodePingMsgSnafu)
50 } else {
51 PbServerRequest::Ping.encode().context(EncodePingMsgSnafu)
52 }
53}
54
55#[derive(Debug)]
67enum Status {
68 ReadMsg,
69 SendPing,
70 ConnectRemote,
71 Resolve,
72 Timeout,
73 NetworkChanged,
74 Cancelled,
75 Rejected(String),
80 RejectedRetryable(String),
85}
86
87#[derive(Debug)]
94struct ControlBackoff {
95 transport: RetryBackoff,
98 reject: RetryBackoff,
101 timing: RecoveryTiming,
102}
103
104impl ControlBackoff {
105 fn new() -> Self {
106 let (reject_min, reject_max) = registration_reject_backoff();
107 Self {
108 transport: RetryBackoff::new(Duration::from_millis(100), Duration::from_secs(2)),
109 timing: RecoveryTiming::default(),
110 reject: RetryBackoff::new(reject_min, reject_max),
111 }
112 }
113
114 fn reset(&mut self) {
117 self.transport.reset();
118 self.reject.reset();
119 }
120}
121
122enum LocalControlWrite {
123 Ping {
124 seq: u64,
125 },
126 StreamAck {
127 client_id: u32,
128 server_generation: u64,
129 },
130}
131
132#[derive(Debug, Clone, Copy)]
133struct ControlRegistration {
134 conn_id: u32,
135 generation: u64,
136 protocol_version: u16,
137 lease_ttl_ms: u64,
138}
139
140#[derive(Debug)]
141struct ControlLeaseState {
142 last_rx_at: Instant,
143 last_pong_at: Option<Instant>,
144}
145
146impl ControlLeaseState {
147 fn new() -> Self {
148 Self {
149 last_rx_at: Instant::now(),
150 last_pong_at: None,
151 }
152 }
153
154 fn record_rx(&mut self) {
155 self.last_rx_at = Instant::now();
156 }
157
158 fn record_pong(&mut self) {
159 let now = Instant::now();
160 self.last_rx_at = now;
161 self.last_pong_at = Some(now);
162 }
163
164 fn last_rx_age(&self) -> Duration {
165 self.last_rx_at.elapsed()
166 }
167}
168
169#[derive(Debug)]
170enum RegistrationProbeResult {
171 Present,
172 Missing,
173 Failed(String),
174}
175
176pub type StatusCallback = Box<dyn Fn(&str) + Send + Sync>;
178
179struct PoolStatus {
198 callback: StatusCallback,
199 workers: Mutex<PoolStatusState>,
200}
201
202struct PoolStatusState {
204 workers: Vec<WorkerStatus>,
205 published: Option<WorkerStatus>,
206}
207
208#[derive(Clone, Debug, Eq, PartialEq)]
210enum WorkerStatus {
211 Retrying,
213 Connected,
215 Failed(String),
217}
218
219impl PoolStatus {
220 fn new(callback: StatusCallback, pool_size: usize) -> Self {
221 Self {
222 callback,
223 workers: Mutex::new(PoolStatusState {
224 workers: vec![WorkerStatus::Retrying; pool_size],
225 published: None,
226 }),
227 }
228 }
229
230 fn report(&self, worker_index: usize, status: WorkerStatus) {
233 let next = {
234 let mut state = self
235 .workers
236 .lock()
237 .unwrap_or_else(|poisoned| poisoned.into_inner());
238 if let Some(slot) = state.workers.get_mut(worker_index) {
239 *slot = status;
240 }
241 let aggregate = Self::aggregate(&state.workers);
242 if state.published.as_ref() == Some(&aggregate) {
243 return;
244 }
245 state.published = Some(aggregate.clone());
246 aggregate
247 };
248 match next {
251 WorkerStatus::Connected => (self.callback)("connected"),
252 WorkerStatus::Retrying => (self.callback)("retrying"),
253 WorkerStatus::Failed(reason) => (self.callback)(&format!("failed: {reason}")),
254 }
255 }
256
257 fn aggregate(workers: &[WorkerStatus]) -> WorkerStatus {
262 let mut retrying = false;
263 let mut first_failure = None;
264 for status in workers {
265 match status {
266 WorkerStatus::Connected => return WorkerStatus::Connected,
267 WorkerStatus::Retrying => retrying = true,
268 WorkerStatus::Failed(reason) => {
269 first_failure.get_or_insert(reason);
270 }
271 }
272 }
273 match first_failure {
274 Some(reason) if !retrying => WorkerStatus::Failed(reason.clone()),
277 _ => WorkerStatus::Retrying,
278 }
279 }
280}
281
282#[derive(Clone, Copy, Debug)]
289pub struct ServerTunnelOptions {
290 pub need_codec: bool,
291 pub is_datagram: bool,
292 pub keep_alive: bool,
293 pub namespace: Option<u64>,
294 pub force_namespace: bool,
295}
296
297#[derive(Clone)]
298struct ServerCliRunConfig {
299 local_addr: ResolvedAddrs,
300 remote_addr: RelayEndpoint,
301 diagnostics: Diagnostics,
302 key: Arc<str>,
303 options: ServerTunnelOptions,
304 worker_index: usize,
305 credential: Credential,
306}
307
308impl Debug for ServerCliRunConfig {
309 fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
310 formatter
311 .debug_struct("ServerCliRunConfig")
312 .field("local_addr", &self.local_addr)
313 .field("remote_addr", &self.remote_addr)
314 .field("key", &self.key)
315 .field("options", &self.options)
316 .field("worker_index", &self.worker_index)
317 .field("credential_key_id", &self.credential.key_id())
318 .finish()
319 }
320}
321
322fn duration_to_millis(duration: Duration) -> u64 {
323 duration.as_millis().min(u128::from(u64::MAX)) as u64
324}
325
326fn new_client_instance_id(worker_index: usize) -> String {
327 let now_ms = SystemTime::now()
328 .duration_since(UNIX_EPOCH)
329 .map(duration_to_millis)
330 .unwrap_or_default();
331 format!("{}-{worker_index}-{now_ms}", std::process::id())
332}
333
334async fn probe_remote_registration(
335 remote_addr: ResolvedAddrs,
336 key: Arc<str>,
337 registration: ControlRegistration,
338 namespace: Option<u64>,
339 credential: Credential,
340) -> RegistrationProbeResult {
341 let timeout = registration_probe_timeout();
342 let result = tokio::time::timeout(timeout, async {
343 let mut stream = crate::addr::connect_tcp(&remote_addr)
344 .await
345 .map_err(|e| format!("connect remote status stream failed: {e}"))?;
346 crate::client::status::get_status_with_credential(
347 &mut stream,
348 PbConnStatusReq::Service {
349 key: key.to_string(),
350 },
351 namespace,
352 &credential,
353 )
354 .await
355 .map_err(|e| {
356 format!(
357 "get service status failed: {}",
358 snafu::Report::from_error(e)
359 )
360 })
361 })
362 .await;
363
364 let status = match result {
365 Ok(Ok(status)) => status,
366 Ok(Err(reason)) => return RegistrationProbeResult::Failed(reason),
367 Err(_) => {
368 return RegistrationProbeResult::Failed(format!(
369 "status probe timed out after {timeout:?}"
370 ));
371 }
372 };
373
374 match status {
375 PbConnStatusResp::Service { connections, .. } => {
376 let present = connections.iter().any(|conn| {
377 conn.conn_id == registration.conn_id
378 && conn.generation == registration.generation
379 && conn.healthy
380 });
381 if present {
382 RegistrationProbeResult::Present
383 } else {
384 RegistrationProbeResult::Missing
385 }
386 }
387 PbConnStatusResp::Keys(keys) => {
388 if keys.iter().any(|candidate| candidate == key.as_ref()) {
389 RegistrationProbeResult::Present
390 } else {
391 RegistrationProbeResult::Missing
392 }
393 }
394 other => RegistrationProbeResult::Failed(format!(
395 "unexpected status response while probing registration: {other:?}"
396 )),
397 }
398}
399
400pub async fn run_server_side_cli<LocalStream, A>(
401 local_addr: A,
402 remote_addr: A,
403 key: Arc<str>,
404 options: ServerTunnelOptions,
405) where
406 LocalStream: StreamProvider + Send + 'static,
407 LocalStream::Item: StreamForward,
408 A: ToSocketAddrs,
409{
410 run_server_side_cli_with_callback::<LocalStream, A>(local_addr, remote_addr, key, options, None)
411 .await
412}
413
414pub async fn run_server_side_cli_with_pinned_credential<LocalStream, A>(
415 local_addr: A,
416 remote_addr: A,
417 key: Arc<str>,
418 options: ServerTunnelOptions,
419 status_callback: Option<StatusCallback>,
420 credential: Credential,
421) where
422 LocalStream: StreamProvider + Send + 'static,
423 LocalStream::Item: StreamForward,
424 A: ToSocketAddrs,
425{
426 let Some((local_addr, remote_addr)) = resolve_tunnel_ends(local_addr, remote_addr).await else {
427 return;
428 };
429 run_server_side_cli_pool::<LocalStream>(
430 local_addr,
431 remote_addr,
432 key,
433 options,
434 status_callback,
435 Some(credential),
436 CancellationToken::new(),
437 )
438 .await;
439}
440
441pub async fn run_server_side_cli_with_shutdown<LocalStream>(
447 local_addr: ResolvedAddrs,
448 remote_addr: ResolvedAddrs,
449 key: Arc<str>,
450 options: ServerTunnelOptions,
451 status_callback: Option<StatusCallback>,
452 credential: Credential,
453 shutdown: CancellationToken,
454) where
455 LocalStream: StreamProvider + Send + 'static,
456 LocalStream::Item: StreamForward,
457{
458 run_server_side_cli_pool::<LocalStream>(
459 local_addr,
460 remote_addr,
461 key,
462 options,
463 status_callback,
464 Some(credential),
465 shutdown,
466 )
467 .await;
468}
469
470pub async fn run_server_side_cli_with_callback<LocalStream, A>(
471 local_addr: A,
472 remote_addr: A,
473 key: Arc<str>,
474 options: ServerTunnelOptions,
475 status_callback: Option<StatusCallback>,
476) where
477 LocalStream: StreamProvider + Send + 'static,
478 LocalStream::Item: StreamForward,
479 A: ToSocketAddrs,
480{
481 let Some((local_addr, remote_addr)) = resolve_tunnel_ends(local_addr, remote_addr).await else {
482 return;
483 };
484 run_server_side_cli_pool::<LocalStream>(
485 local_addr,
486 remote_addr,
487 key,
488 options,
489 status_callback,
490 None,
491 CancellationToken::new(),
492 )
493 .await;
494}
495
496async fn resolve_registration_credential(
497 pinned: Option<Credential>,
498 shutdown: &CancellationToken,
499) -> Option<Credential> {
500 if let Some(credential) = pinned {
501 return Some(credential);
502 }
503 let mut retry_backoff = RetryBackoff::default();
504 loop {
505 if shutdown.is_cancelled() {
506 return None;
507 }
508 match get_process_credential() {
509 Ok(credential) => return Some(credential),
510 Err(error) => {
511 tracing::error!("load registration credential failed: {error}");
512 tokio::select! {
513 () = shutdown.cancelled() => return None,
514 () = tokio::time::sleep(retry_backoff.next_delay()) => {}
515 }
516 }
517 }
518 }
519}
520
521async fn run_server_side_cli_pool<LocalStream>(
522 local_addr: ResolvedAddrs,
523 remote_addr: ResolvedAddrs,
524 key: Arc<str>,
525 options: ServerTunnelOptions,
526 status_callback: Option<StatusCallback>,
527 pinned_credential: Option<Credential>,
528 shutdown: CancellationToken,
529) where
530 LocalStream: StreamProvider + Send + 'static,
531 LocalStream::Item: StreamForward,
532{
533 run_server_side_cli_recovering::<LocalStream>(
534 local_addr,
535 RelayEndpoint::fixed(remote_addr),
536 key,
537 options,
538 status_callback,
539 pinned_credential,
540 shutdown,
541 Diagnostics::default(),
542 )
543 .await;
544}
545
546#[allow(clippy::too_many_arguments)]
547pub(crate) async fn run_server_side_cli_recovering<LocalStream>(
548 local_addr: ResolvedAddrs,
549 remote_addr: RelayEndpoint,
550 key: Arc<str>,
551 options: ServerTunnelOptions,
552 status_callback: Option<StatusCallback>,
553 pinned_credential: Option<Credential>,
554 shutdown: CancellationToken,
555 diagnostics: Diagnostics,
556) where
557 LocalStream: StreamProvider + Send + 'static,
558 LocalStream::Item: StreamForward,
559{
560 let Some(credential) = resolve_registration_credential(pinned_credential, &shutdown).await
561 else {
562 return;
563 };
564 remote_addr.start();
565 let pool_size = control_conn_pool_size().max(1);
566 tracing::info!(
567 event = "local_server_control_pool_starting",
568 key = %key,
569 pool_size,
570 "starting local server control connection pool"
571 );
572 let mut workers = JoinSet::new();
573 let pool_status =
577 status_callback.map(|callback| Arc::new(PoolStatus::new(callback, pool_size)));
578 for worker_index in 0..pool_size {
579 let worker_key = key.clone();
580 let worker_status = pool_status.clone();
581 let worker_shutdown = shutdown.clone();
582 let worker_local = local_addr.clone();
583 let worker_remote = remote_addr.clone();
584 let worker_diagnostics = diagnostics.clone();
585 workers.spawn(async move {
586 run_server_side_cli_worker::<LocalStream>(
587 worker_local,
588 worker_remote,
589 worker_key,
590 options,
591 worker_status,
592 worker_index,
593 credential,
594 worker_shutdown,
595 worker_diagnostics,
596 )
597 .await;
598 });
599 }
600 while let Some(result) = workers.join_next().await {
601 if let Err(e) = result {
602 tracing::warn!(
603 event = "local_server_control_worker_join_failed",
604 error = %e,
605 "local server control worker join failed"
606 );
607 }
608 }
609}
610
611#[allow(clippy::too_many_arguments)]
612async fn run_server_side_cli_worker<LocalStream>(
613 local_addr: ResolvedAddrs,
614 remote_addr: RelayEndpoint,
615 key: Arc<str>,
616 options: ServerTunnelOptions,
617 pool_status: Option<Arc<PoolStatus>>,
618 worker_index: usize,
619 credential: Credential,
620 shutdown: CancellationToken,
621 diagnostics: Diagnostics,
622) where
623 LocalStream: StreamProvider + Send + 'static,
624 LocalStream::Item: StreamForward,
625{
626 let report = |status: WorkerStatus| {
627 if let Some(ref pool_status) = pool_status {
628 pool_status.report(worker_index, status);
629 }
630 };
631 let mut backoff = ControlBackoff::new();
632 let run_config = ServerCliRunConfig {
633 local_addr: local_addr.clone(),
634 remote_addr: remote_addr.clone(),
635 diagnostics: diagnostics.clone(),
636 key: key.clone(),
637 options,
638 worker_index,
639 credential,
640 };
641 let mut stream_tasks = JoinSet::new();
645 let setup_slots = remote_addr.data_slots();
646 let mut wake = remote_addr.wake();
647 'outer: loop {
648 if shutdown.is_cancelled() {
649 break 'outer;
650 }
651 wake.acknowledge();
652 diagnostics.attempt();
653 let status = if let Err(status) = run_server_side_cli_inner::<LocalStream>(
654 &mut backoff,
655 run_config.clone(),
656 &report,
657 shutdown.clone(),
658 &mut stream_tasks,
659 &setup_slots,
660 &mut wake,
661 )
662 .await
663 {
664 status
665 } else {
666 if shutdown.is_cancelled() {
667 break 'outer;
668 }
669 tracing::warn!(
670 event = "local_server_control_worker_finished",
671 key = %key,
672 worker_index,
673 "local server control worker finished without an error; reconnecting"
674 );
675 Status::ReadMsg
676 };
677 if matches!(status, Status::Cancelled) {
678 break 'outer;
679 }
680 let failure = match &status {
682 Status::Resolve => RecoveryFailure::Dns,
683 Status::Timeout => RecoveryFailure::Timeout,
684 Status::Rejected(_) | Status::RejectedRetryable(_) => RecoveryFailure::Rejected,
685 Status::NetworkChanged => RecoveryFailure::NetworkChanged,
686 _ => RecoveryFailure::Transport,
687 };
688 if matches!(
689 failure,
690 RecoveryFailure::Transport | RecoveryFailure::Timeout | RecoveryFailure::Dns
691 ) {
692 remote_addr.transport_failed();
693 }
694 let (retry_interval, retry_count) = match status {
695 Status::Cancelled => break 'outer,
696 Status::Rejected(reason) => {
697 diagnostics.failed(RecoveryFailure::Rejected, Duration::ZERO);
698 tracing::error!(
699 event = "local_server_registration_rejected_permanently",
700 key = %key,
701 worker_index,
702 reason = %reason,
703 "pb server permanently rejected this registration; not reconnecting"
704 );
705 report(WorkerStatus::Failed(reason));
706 break 'outer;
707 }
708 Status::RejectedRetryable(ref reason) => {
709 let interval = backoff.reject.next_delay();
710 tracing::debug!(
711 event = "local_server_registration_rejected_retryable",
712 key = %key,
713 worker_index,
714 reason = %reason,
715 retry_delay = ?interval,
716 retry_count = backoff.reject.failures(),
717 "pb server rejected this registration for a condition that may clear"
718 );
719 (interval, backoff.reject.failures())
720 }
721 Status::NetworkChanged => {
722 backoff.transport.reset();
723 (Duration::ZERO, 0)
724 }
725 Status::ReadMsg
726 | Status::SendPing
727 | Status::ConnectRemote
728 | Status::Resolve
729 | Status::Timeout => (
730 jitter(backoff.transport.next_delay()),
731 backoff.transport.failures(),
732 ),
733 };
734 if diagnostics.failed(failure, retry_interval) {
735 tracing::info!(
736 event = "local_server_control_reconnect_scheduled",
737 key = %key,
738 worker_index,
739 local_addr = ?local_addr,
740 remote_addr = ?remote_addr,
741 status = ?status,
742 retry_delay = ?retry_interval,
743 retry_count,
744 "local server control connection will reconnect"
745 );
746 }
747 report(WorkerStatus::Retrying);
748
749 tokio::select! {
750 () = shutdown.cancelled() => break 'outer,
751 () = tokio::time::sleep(retry_interval) => {},
752 () = wake.changed(), if !matches!(failure, RecoveryFailure::Rejected) => { backoff.transport.reset(); }
753 }
754 if shutdown.is_cancelled() {
755 break 'outer;
756 }
757 }
758
759 stream_tasks.shutdown().await;
762}
763
764#[instrument(skip(backoff, report, shutdown, stream_tasks, setup_slots, wake))]
767async fn run_server_side_cli_inner<LocalStream: StreamProvider>(
768 backoff: &mut ControlBackoff,
769 config: ServerCliRunConfig,
770 report: &impl Fn(WorkerStatus),
771 shutdown: CancellationToken,
772 stream_tasks: &mut JoinSet<()>,
773 setup_slots: &Arc<tokio::sync::Semaphore>,
774 wake: &mut RecoveryWake,
775) -> std::result::Result<(), Status>
776where
777 LocalStream::Item: StreamForward,
778{
779 let ServerCliRunConfig {
780 local_addr,
781 remote_addr,
782 diagnostics,
783 key,
784 options:
785 ServerTunnelOptions {
786 need_codec,
787 is_datagram,
788 keep_alive,
789 namespace,
790 force_namespace,
791 },
792 worker_index,
793 credential,
794 } = config;
795 let control_permit = tokio::select! {
796 () = shutdown.cancelled() => return Err(Status::Cancelled),
797 () = wake.network_changed() => return Err(Status::NetworkChanged),
798 permit = remote_addr.control_permit() => permit,
799 };
800 let started = tokio::time::Instant::now();
801 let timeout = backoff.timing.timeout();
802 let deadline = started + timeout;
803 diagnostics.phase(RecoveryPhase::Resolving);
804 let addresses = tokio::select! {
805 () = shutdown.cancelled() => return Err(Status::Cancelled),
806 () = wake.network_changed() => return Err(Status::NetworkChanged),
807 result = tokio::time::timeout_at(deadline, remote_addr.addresses()) => match result {
808 Ok(Ok(addresses)) => addresses,
809 _ => return Err(Status::Resolve),
810 },
811 };
812 diagnostics.phase(RecoveryPhase::Connecting);
813 let mut manager_stream = tokio::select! {
814 () = shutdown.cancelled() => return Err(Status::Cancelled),
815 () = wake.network_changed() => return Err(Status::NetworkChanged),
816 result = tokio::time::timeout_at(deadline, crate::addr::connect_tcp(&addresses)) => {
817 match result {
818 Ok(Ok(stream)) => stream,
819 Ok(Err(error)) => {
820 tracing::debug!(event = "local_server_dial_failed", %error, %key, worker_index);
821 return Err(Status::ConnectRemote);
822 }
823 Err(_) => {
824 backoff.timing.timed_out();
825 tracing::debug!(event = "local_server_setup_timeout", %key, worker_index, ?timeout, phase = "dial");
826 return Err(Status::Timeout);
827 }
828 }
829 }
830 };
831 tracing::debug!(
832 event = "local_server_connected_remote",
833 key = %key,
834 worker_index,
835 local_addr = %local_addr,
836 remote_addr = %remote_addr,
837 need_codec,
838 is_datagram,
839 "local server connected to pb server"
840 );
841
842 if keep_alive {
843 snafu_error_handle!(
844 set_tcp_keep_alive(&manager_stream),
845 "manager stream set tcp keep alive"
846 );
847 }
848 snafu_error_handle!(
849 set_tcp_nodelay(&manager_stream),
850 "manager stream set tcp nodelay"
851 );
852
853 let registered_addr = manager_stream
854 .peer_addr()
855 .map(ResolvedAddrs::from)
856 .unwrap_or(addresses);
857 diagnostics.phase(RecoveryPhase::Handshake);
858 let session = match ClientHeaderSession::new_v2(&credential) {
862 Ok(session) => session,
863 Err(error) => {
864 tracing::error!("create manager protocol-v2 session failed: {error}");
865 return Err(Status::ConnectRemote);
866 }
867 };
868 let heartbeat_interval = control_heartbeat_interval();
869 let heartbeat_tolerance = control_heartbeat_tolerance();
870 let request = match namespace {
871 Some(namespace) => PbConnRequest::RegisterScoped {
872 key: key.to_string(),
873 namespace,
874 force_namespace,
875 need_codec,
876 is_datagram,
877 protocol_version: Some(CONTROL_PROTOCOL_V2),
878 client_instance_id: Some(new_client_instance_id(worker_index)),
879 heartbeat_interval_ms: Some(duration_to_millis(heartbeat_interval)),
880 heartbeat_tolerance_ms: Some(duration_to_millis(heartbeat_tolerance)),
881 },
882 None => PbConnRequest::Register {
883 key: key.to_string(),
884 need_codec,
885 is_datagram,
886 protocol_version: Some(CONTROL_PROTOCOL_V2),
887 client_instance_id: Some(new_client_instance_id(worker_index)),
888 heartbeat_interval_ms: Some(duration_to_millis(heartbeat_interval)),
889 heartbeat_tolerance_ms: Some(duration_to_millis(heartbeat_tolerance)),
890 },
891 };
892 let msg = snafu_error_get_or_return_ok!(request.encode().context(EncodeRegisterReqSnafu));
893 tokio::select! {
894 () = shutdown.cancelled() => return Err(Status::Cancelled),
895 () = wake.network_changed() => return Err(Status::NetworkChanged),
896 result = tokio::time::timeout_at(deadline, session.write_initial(&mut manager_stream, &msg)) => {
897 match result {
898 Ok(result) => snafu_error_get_or_return_ok!(result.context(SendRegisterReqSnafu)),
899 Err(_) => {
900 backoff.timing.timed_out();
901 tracing::debug!(event = "local_server_setup_timeout", %key, worker_index, ?timeout, phase = "write");
902 return Err(Status::Timeout);
903 }
904 }
905 }
906 }
907 let (mut reader, mut writer) = manager_stream.into_split();
908 let mut msg_reader = match session.response_reader(&mut reader) {
909 Ok(reader) => reader,
910 Err(e) => {
911 tracing::error!("create manager header reader failed: {e}");
912 return Err(Status::ReadMsg);
913 }
914 };
915 let (key, registration) = {
917 let msg = tokio::select! {
918 () = shutdown.cancelled() => return Err(Status::Cancelled),
919 () = wake.network_changed() => return Err(Status::NetworkChanged),
920 result = tokio::time::timeout_at(deadline, msg_reader.read_msg()) => {
921 match result {
922 Ok(result) => snafu_error_get_or_return_ok!(result.context(ReadRegisterRespSnafu)),
923 Err(_) => {
924 backoff.timing.timed_out();
925 tracing::debug!(event = "local_server_setup_timeout", %key, worker_index, ?timeout, phase = "response");
926 return Err(Status::Timeout);
927 }
928 }
929 }
930 };
931 let resp = snafu_error_get_or_return_ok!(
932 PbConnResponse::decode(msg).context(DecodeRegisterRespSnafu)
933 );
934 remote_addr.protocol_succeeded();
935 diagnostics.responded(started.elapsed());
936 backoff.timing.record(started.elapsed());
937 let registration = match resp {
938 PbConnResponse::RegisterV2 {
939 conn_id,
940 generation,
941 lease_ttl_ms,
942 } => ControlRegistration {
943 conn_id,
944 generation,
945 protocol_version: CONTROL_PROTOCOL_V2,
946 lease_ttl_ms,
947 },
948 PbConnResponse::Register(conn_id) => ControlRegistration {
949 conn_id,
950 generation: 0,
951 protocol_version: 1,
952 lease_ttl_ms: 0,
953 },
954 PbConnResponse::Error(error) => {
960 tracing::debug!(
961 event = "local_server_registration_rejected",
962 key = %key,
963 worker_index,
964 reason = %error.code,
965 retryable = error.retryable,
966 message = %error.message,
967 "pb server rejected service registration"
968 );
969 let reason = format!("{}: {}", error.code, error.message);
970 return Err(if error.retryable {
971 Status::RejectedRetryable(reason)
972 } else {
973 Status::Rejected(reason)
974 });
975 }
976 _ => snafu_error_get_or_return_ok!(RegisterRespNotMatchSnafu {}.fail()),
977 };
978 tracing::info!(
979 event = "local_server_registered",
980 key = %key,
981 conn_id = %registration.conn_id,
982 generation = registration.generation,
983 protocol_version = registration.protocol_version,
984 lease_ttl_ms = registration.lease_ttl_ms,
985 worker_index,
986 local_addr = %local_addr,
987 remote_addr = %remote_addr,
988 "local server registered with pb server"
989 );
990
991 report(WorkerStatus::Connected);
994 (key, registration)
995 };
996
997 drop(control_permit);
998 diagnostics.succeeded(started.elapsed());
999 remote_addr.protocol_succeeded();
1000 backoff.timing.record(started.elapsed());
1001 tracing::debug!(event = "local_server_setup_latency", elapsed_ms = duration_to_millis(started.elapsed()), next_timeout_ms = duration_to_millis(backoff.timing.timeout()), %key, worker_index);
1002 backoff.reset();
1003 let (write_tx, mut write_rx) = tokio::sync::mpsc::channel::<LocalControlWrite>(64);
1004 let lease_state = Arc::new(tokio::sync::Mutex::new(ControlLeaseState::new()));
1005 let writer_key = key.clone();
1006 let writer_registration = registration;
1007 let mut writer_handle = tokio_util::task::AbortOnDropHandle::new(tokio::spawn(async move {
1008 let mut msg_writer = match session.continuation_writer(&mut writer) {
1009 Ok(writer) => writer,
1010 Err(e) => {
1011 tracing::error!("create manager header writer failed: {e}");
1012 return Err(Status::SendPing);
1013 }
1014 };
1015 loop {
1016 let Some(cmd) = write_rx.recv().await else {
1017 return Ok(());
1018 };
1019 match cmd {
1020 LocalControlWrite::Ping { seq } => {
1021 snafu_error_get_or_return!(
1022 handle_ping_interval(
1023 &mut msg_writer,
1024 writer_key.clone(),
1025 writer_registration,
1026 seq
1027 )
1028 .await,
1029 "[send ping]",
1030 Err(Status::SendPing)
1031 );
1032 tracing::debug!(
1033 event = "local_server_heartbeat_sent",
1034 key = %writer_key,
1035 conn_id = %writer_registration.conn_id,
1036 generation = writer_registration.generation,
1037 protocol_version = writer_registration.protocol_version,
1038 seq,
1039 interval = ?control_heartbeat_interval(),
1040 "local server heartbeat sent"
1041 );
1042 }
1043 LocalControlWrite::StreamAck {
1044 client_id,
1045 server_generation,
1046 } => {
1047 snafu_error_get_or_return!(
1048 write_stream_ack(&mut msg_writer, client_id, server_generation).await,
1049 "[send stream ack]",
1050 Err(Status::SendPing)
1051 );
1052 }
1053 }
1054 }
1055 }));
1056
1057 let heartbeat_interval = control_heartbeat_interval();
1058 let heartbeat_tolerance = control_heartbeat_tolerance();
1059 let suspect_grace = control_suspect_grace();
1060 let mut heartbeat = tokio::time::interval(heartbeat_interval);
1061 heartbeat.set_missed_tick_behavior(MissedTickBehavior::Delay);
1062 heartbeat.tick().await;
1063 let mut probes = JoinSet::new();
1064 let mut ping_seq = 0_u64;
1065 let mut control_deadline = tokio::time::Instant::now() + heartbeat_tolerance + suspect_grace;
1066
1067 let result = loop {
1068 tokio::select! {
1069 msg = msg_reader.read_msg() => {
1070 let msg = match msg.context(ReadStreamReqSnafu) {
1071 Ok(msg) => msg,
1072 Err(e) => {
1073 tracing::error!(
1074 event = "local_server_control_read_failed",
1075 key = %key,
1076 conn_id = %registration.conn_id,
1077 generation = registration.generation,
1078 error = %snafu::Report::from_error(e),
1079 "local server control read failed"
1080 );
1081 break Err(Status::ReadMsg);
1082 }
1083 };
1084 remote_addr.protocol_succeeded();
1085 diagnostics.heard();
1086 lease_state.lock().await.record_rx();
1087 control_deadline = tokio::time::Instant::now() + heartbeat_tolerance + suspect_grace;
1088 snafu_error_get_or_continue!(
1089 handle_request::<LocalStream>(
1090 msg,
1091 StreamConnect {
1092 local_addr: local_addr.clone(),
1093 remote_addr: registered_addr.clone(),
1094 keep_alive,
1095 namespace,
1096 credential,
1097 },
1098 key.clone(),
1099 registration.conn_id,
1100 &write_tx,
1101 lease_state.clone(),
1102 stream_tasks,
1103 &shutdown,
1104 setup_slots,
1105 )
1106 .await
1107 );
1108 }
1109 Some(_) = stream_tasks.join_next() => {
1110 }
1113 result = &mut writer_handle => {
1114 break match result {
1115 Ok(result) => result,
1116 Err(e) => {
1117 tracing::error!(
1118 event = "local_server_control_writer_join_failed",
1119 key = %key,
1120 conn_id = %registration.conn_id,
1121 generation = registration.generation,
1122 error = %e,
1123 "local server control writer task failed"
1124 );
1125 Err(Status::SendPing)
1126 }
1127 };
1128 }
1129 _ = heartbeat.tick() => {
1130 ping_seq = ping_seq.wrapping_add(1);
1131 if write_tx.try_send(LocalControlWrite::Ping { seq: ping_seq }).is_err() {
1132 break Err(Status::SendPing);
1133 }
1134
1135 let last_rx_age = lease_state.lock().await.last_rx_age();
1136 if registration.protocol_version >= CONTROL_PROTOCOL_V2
1137 && last_rx_age >= heartbeat_tolerance
1138 && probes.is_empty()
1139 {
1140 tracing::warn!(
1141 event = "local_server_lease_suspect",
1142 key = %key,
1143 conn_id = %registration.conn_id,
1144 generation = registration.generation,
1145 worker_index,
1146 last_rx_age_ms = duration_to_millis(last_rx_age),
1147 heartbeat_tolerance_ms = duration_to_millis(heartbeat_tolerance),
1148 "local server control lease is suspect; probing remote registration"
1149 );
1150 let probe_key = key.clone();
1151 let probe_remote = registered_addr.clone();
1152 let probe_endpoint = remote_addr.clone();
1153 probes.spawn(async move {
1154 let _permit = probe_endpoint.control_permit().await;
1155 probe_remote_registration(
1156 probe_remote,
1157 probe_key,
1158 registration,
1159 namespace,
1160 credential,
1161 )
1162 .await
1163 });
1164 }
1165 }
1166 () = shutdown.cancelled() => {
1167 tracing::info!(
1168 event = "local_server_control_cancelled",
1169 key = %key,
1170 conn_id = %registration.conn_id,
1171 generation = registration.generation,
1172 worker_index,
1173 "local server control loop cancelled"
1174 );
1175 break Err(Status::Cancelled);
1176 }
1177 () = tokio::time::sleep_until(control_deadline) => {
1178 tracing::warn!(event = "local_server_control_unresponsive", %key, worker_index, "control socket received no reply past grace; reconnecting");
1179 break Err(Status::ReadMsg);
1180 }
1181 Some(probe_result) = probes.join_next() => {
1182 let probe_result = probe_result.unwrap_or_else(|error| RegistrationProbeResult::Failed(error.to_string()));
1183 let last_rx_age = lease_state.lock().await.last_rx_age();
1184 match probe_result {
1185 RegistrationProbeResult::Present => {
1186 tracing::debug!(
1187 event = "local_server_registration_probe_ok",
1188 key = %key,
1189 conn_id = %registration.conn_id,
1190 generation = registration.generation,
1191 worker_index,
1192 last_rx_age_ms = duration_to_millis(last_rx_age),
1193 "remote registration still contains this control connection"
1194 );
1195 }
1196 RegistrationProbeResult::Missing if last_rx_age >= heartbeat_tolerance => {
1197 tracing::warn!(
1198 event = "local_server_registration_missing",
1199 key = %key,
1200 conn_id = %registration.conn_id,
1201 generation = registration.generation,
1202 worker_index,
1203 last_rx_age_ms = duration_to_millis(last_rx_age),
1204 "remote registration no longer contains this control connection; reconnecting"
1205 );
1206 break Err(Status::ReadMsg);
1207 }
1208 RegistrationProbeResult::Missing => {
1209 tracing::debug!(
1210 event = "local_server_registration_missing_ignored_after_recent_activity",
1211 key = %key,
1212 conn_id = %registration.conn_id,
1213 generation = registration.generation,
1214 worker_index,
1215 last_rx_age_ms = duration_to_millis(last_rx_age),
1216 "remote registration probe was stale after recent control activity"
1217 );
1218 }
1219 RegistrationProbeResult::Failed(reason)
1220 if last_rx_age >= heartbeat_tolerance + suspect_grace =>
1221 {
1222 tracing::warn!(
1223 event = "local_server_status_probe_failed",
1224 key = %key,
1225 conn_id = %registration.conn_id,
1226 generation = registration.generation,
1227 worker_index,
1228 last_rx_age_ms = duration_to_millis(last_rx_age),
1229 reason = %reason,
1230 "registration probe failed past suspect grace; reconnecting"
1231 );
1232 break Err(Status::ReadMsg);
1233 }
1234 RegistrationProbeResult::Failed(reason) => {
1235 tracing::warn!(
1236 event = "local_server_status_probe_failed",
1237 key = %key,
1238 conn_id = %registration.conn_id,
1239 generation = registration.generation,
1240 worker_index,
1241 last_rx_age_ms = duration_to_millis(last_rx_age),
1242 reason = %reason,
1243 "registration probe failed; waiting inside suspect grace"
1244 );
1245 }
1246 }
1247 }
1248 }
1249 };
1250 if !writer_handle.is_finished() {
1251 writer_handle.abort();
1252 }
1253 result
1254}
1255
1256#[instrument(skip(writer))]
1257async fn handle_ping_interval<T: MessageWriter>(
1258 writer: &mut T,
1259 _key: Arc<str>,
1260 registration: ControlRegistration,
1261 seq: u64,
1262) -> error::Result<()> {
1263 let ping_msg = get_ping_message(registration.protocol_version, seq)?;
1264 let timeout = control_io_timeout();
1265 match tokio::time::timeout(timeout, writer.write_msg(&ping_msg)).await {
1266 Ok(result) => result.context(WritePingMsgSnafu),
1267 Err(_) => ControlIoTimeoutSnafu {
1268 action: "write ping message",
1269 timeout,
1270 }
1271 .fail(),
1272 }
1273}
1274
1275#[instrument(skip(
1276 msg,
1277 target,
1278 write_tx,
1279 lease_state,
1280 stream_tasks,
1281 shutdown,
1282 setup_slots
1283))]
1284#[allow(clippy::too_many_arguments)]
1285async fn handle_request<LocalStream: StreamProvider>(
1286 msg: &[u8],
1287 target: StreamConnect,
1288 key: Arc<str>,
1289 conn_id: u32,
1290 write_tx: &tokio::sync::mpsc::Sender<LocalControlWrite>,
1291 lease_state: Arc<tokio::sync::Mutex<ControlLeaseState>>,
1292 stream_tasks: &mut JoinSet<()>,
1293 shutdown: &CancellationToken,
1294 setup_slots: &Arc<tokio::sync::Semaphore>,
1295) -> error::Result<()>
1296where
1297 LocalStream::Item: StreamForward,
1298{
1299 let req = LocalServer::decode(msg).context(DecodeStreamReqSnafu)?;
1300
1301 match req {
1302 LocalServer::Stream {
1303 client_id,
1304 server_generation,
1305 } => {
1306 tracing::debug!(
1307 event = "local_server_stream_request_received",
1308 key = %key,
1309 server_conn_id = %conn_id,
1310 client_conn_id = client_id,
1311 server_generation,
1312 "local server received stream request"
1313 );
1314 let Ok(permit) = setup_slots.clone().try_acquire_owned() else {
1315 tracing::warn!(event = "local_server_setup_saturated", %key, client_id, "stream setup capacity reached; relay can select another worker");
1316 return Ok(());
1317 };
1318 write_tx
1319 .try_send(LocalControlWrite::StreamAck {
1320 client_id,
1321 server_generation,
1322 })
1323 .map_err(|_| error::Error::ControlWriterClosed {
1324 action: "stream ack message",
1325 })?;
1326 let key = key.clone();
1327 let stream_shutdown = shutdown.clone();
1328 stream_tasks.spawn(async move {
1332 let forward =
1333 handle_stream::<LocalStream>(key, client_id, server_generation, target, permit);
1334 tokio::select! {
1335 () = stream_shutdown.cancelled() => {}
1336 result = forward => snafu_error_handle!(result),
1337 }
1338 });
1339 }
1340 LocalServer::Pong => {
1342 lease_state.lock().await.record_pong();
1343 tracing::debug!(
1344 event = "local_server_pong_received",
1345 key = %key,
1346 server_conn_id = %conn_id,
1347 "local server received pong"
1348 );
1349 }
1350 LocalServer::PongV2 { seq } => {
1351 lease_state.lock().await.record_pong();
1352 tracing::debug!(
1353 event = "local_server_pong_received",
1354 key = %key,
1355 server_conn_id = %conn_id,
1356 seq,
1357 "local server received pong v2"
1358 );
1359 }
1360 LocalServer::Retire {
1361 reason,
1362 conn_id: retired_conn_id,
1363 server_generation,
1364 } => {
1365 tracing::warn!(
1366 event = "local_server_control_retired",
1367 key = %key,
1368 server_conn_id = %conn_id,
1369 retired_conn_id,
1370 server_generation,
1371 reason = %reason,
1372 "remote server retired this local control connection"
1373 );
1374 }
1375 }
1376 Ok(())
1377}
1378
1379async fn write_stream_ack<T: MessageWriter>(
1380 writer: &mut T,
1381 client_id: u32,
1382 server_generation: u64,
1383) -> error::Result<()> {
1384 let ack = PbServerRequest::StreamAck {
1385 client_id,
1386 server_generation,
1387 }
1388 .encode()
1389 .context(EncodeStreamAckMsgSnafu)?;
1390 let timeout = control_io_timeout();
1391 match tokio::time::timeout(timeout, writer.write_msg(&ack)).await {
1392 Ok(result) => result.context(WriteStreamAckMsgSnafu),
1393 Err(_) => ControlIoTimeoutSnafu {
1394 action: "write stream ack message",
1395 timeout,
1396 }
1397 .fail(),
1398 }
1399}
1400
1401#[cfg(test)]
1402mod pool_status_tests {
1403 use std::sync::Mutex as StdMutex;
1404
1405 use super::*;
1406
1407 fn recording_pool(pool_size: usize) -> (PoolStatus, Arc<StdMutex<Vec<String>>>) {
1409 let seen = Arc::new(StdMutex::new(Vec::new()));
1410 let sink = seen.clone();
1411 let callback: StatusCallback = Box::new(move |status: &str| {
1412 sink.lock().unwrap().push(status.to_string());
1413 });
1414 (PoolStatus::new(callback, pool_size), seen)
1415 }
1416
1417 fn published(seen: &Arc<StdMutex<Vec<String>>>) -> Vec<String> {
1418 seen.lock().unwrap().clone()
1419 }
1420
1421 #[test]
1424 fn one_connected_worker_makes_the_pool_connected() {
1425 let (pool, seen) = recording_pool(2);
1426 pool.report(0, WorkerStatus::Connected);
1427 pool.report(1, WorkerStatus::Retrying);
1428 assert_eq!(published(&seen), vec!["connected".to_string()]);
1429 }
1430
1431 #[test]
1434 fn a_single_failure_does_not_fail_the_pool() {
1435 let (pool, seen) = recording_pool(2);
1436 pool.report(0, WorkerStatus::Failed("service_transport_mismatch".into()));
1437 assert_eq!(published(&seen), vec!["retrying".to_string()]);
1438
1439 pool.report(1, WorkerStatus::Connected);
1440 assert_eq!(
1441 published(&seen),
1442 vec!["retrying".to_string(), "connected".to_string()]
1443 );
1444 }
1445
1446 #[test]
1449 fn the_pool_fails_only_when_every_worker_has() {
1450 let (pool, seen) = recording_pool(2);
1451 pool.report(0, WorkerStatus::Failed("service_transport_mismatch".into()));
1452 pool.report(1, WorkerStatus::Failed("namespace_access_denied".into()));
1453 assert_eq!(
1454 published(&seen),
1455 vec![
1456 "retrying".to_string(),
1457 "failed: service_transport_mismatch".to_string(),
1458 ],
1459 "the first permanent rejection is the one that explains the pool"
1460 );
1461 }
1462
1463 #[test]
1466 fn unchanged_aggregates_are_not_republished() {
1467 let (pool, seen) = recording_pool(2);
1468 pool.report(0, WorkerStatus::Connected);
1469 pool.report(1, WorkerStatus::Connected);
1470 pool.report(0, WorkerStatus::Connected);
1471 assert_eq!(published(&seen), vec!["connected".to_string()]);
1472 }
1473
1474 #[test]
1477 fn losing_the_last_connection_publishes_retrying() {
1478 let (pool, seen) = recording_pool(1);
1479 pool.report(0, WorkerStatus::Connected);
1480 pool.report(0, WorkerStatus::Retrying);
1481 assert_eq!(
1482 published(&seen),
1483 vec!["connected".to_string(), "retrying".to_string()]
1484 );
1485 }
1486}
1487
1488#[cfg(test)]
1489mod shutdown_tests {
1490 use std::time::Duration;
1491
1492 use super::*;
1493 use uni_stream::stream::TcpStreamProvider;
1494
1495 #[tokio::test]
1496 async fn register_loop_stops_on_cancel() {
1497 let shutdown = CancellationToken::new();
1498 let token = shutdown.clone();
1499 let addr = ResolvedAddrs::from("127.0.0.1:1".parse::<std::net::SocketAddr>().unwrap());
1500 let task = tokio::spawn(async move {
1501 run_server_side_cli_with_shutdown::<TcpStreamProvider>(
1502 addr.clone(),
1503 addr,
1504 "k".into(),
1505 ServerTunnelOptions {
1506 need_codec: false,
1507 is_datagram: false,
1508 keep_alive: false,
1509 namespace: None,
1510 force_namespace: false,
1511 },
1512 None,
1513 Credential::Admin(*b"0123456789abcdefghijklmnopqrstuv"),
1514 token,
1515 )
1516 .await;
1517 });
1518 tokio::time::sleep(Duration::from_millis(80)).await;
1519 shutdown.cancel();
1520 tokio::time::timeout(Duration::from_secs(3), task)
1521 .await
1522 .expect("register loop did not stop after cancel")
1523 .expect("join");
1524 }
1525}