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