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