Skip to main content

pb_mapper_client/server/
mod.rs

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/// Why one attempt at holding a control connection ended, and therefore how the
54/// worker should treat it.
55///
56/// The three transport variants are interchangeable as far as the retry decision
57/// goes — they differ only in what they report — while the two rejection variants
58/// each mean something the transport ones do not:
59///
60/// * [`Status::Rejected`] is terminal, so the worker stops.
61/// * [`Status::RejectedRetryable`] is not, but the relay answered and refused, so
62///   the worker waits on the much slower [`registration_reject_backoff`] ladder
63///   instead of the transport one.
64#[derive(Debug)]
65enum Status {
66    ReadMsg,
67    SendPing,
68    ConnectRemote,
69    Cancelled,
70    /// The relay refused the registration for a reason reconnecting cannot fix —
71    /// a namespace the credential does not own, a malformed service name. Retrying
72    /// would loop forever while the caller's `wait_ready` never resolves, so this
73    /// ends the worker instead.
74    Rejected(String),
75    /// The relay refused the registration for a condition that can clear on its
76    /// own — a full per-service connection quota, a namespace at its service
77    /// limit. Worth retrying, but slowly: the relay is reachable and answering,
78    /// and reconnecting hard only adds load to what has to drain.
79    RejectedRetryable(String),
80}
81
82/// The two retry ladders a control worker climbs, and the reason they are two.
83///
84/// A transport failure and a retryable rejection are not the same kind of
85/// problem, so they must not share a counter. Sharing one meant a reconnect —
86/// which happens on its own schedule — reset the wait a full quota was serving,
87/// leaving the worker hammering the relay with a request it had just refused.
88#[derive(Debug)]
89struct ControlBackoff {
90    /// Could not reach or hold the relay. Retry quickly; that is how the tunnel
91    /// comes back.
92    transport: RetryBackoff,
93    /// The relay answered and refused, for something that may clear. Retry
94    /// slowly; see [`registration_reject_backoff`].
95    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    /// Clear both ladders. Called once a registration is established: it is the
108    /// only evidence that either condition has passed.
109    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
169// Callback for notifying status changes to external systems
170pub type StatusCallback = Box<dyn Fn(&str) + Send + Sync>;
171
172/// Turns the control pool's per-worker statuses into the one status a caller sees.
173///
174/// A registration is served by [`control_conn_pool_size`] control connections,
175/// any one of which keeps the service reachable. Reporting only the first
176/// worker's view — which is what this code did — meant a caller was told
177/// `retrying` while the rest of the pool was registered and forwarding fine.
178///
179/// So the pool's status is the best of its workers, in this order:
180///
181/// * `connected` — at least one worker is registered.
182/// * `retrying` — none is, but at least one is still trying.
183/// * `failed` — every worker gave up permanently. Only then, because a caller
184///   awaiting readiness with no timeout must not be released while any worker
185///   could still bring the service up.
186///
187/// The reported status is recomputed on every worker transition and forwarded
188/// only when it actually changes, so a caller sees pool-level transitions rather
189/// than one line per worker.
190struct PoolStatus {
191    callback: StatusCallback,
192    workers: Mutex<PoolStatusState>,
193}
194
195/// What each worker last reported, plus what the pool last published.
196struct PoolStatusState {
197    workers: Vec<WorkerStatus>,
198    published: Option<WorkerStatus>,
199}
200
201/// One worker's latest view of its own control connection.
202#[derive(Clone, Debug, Eq, PartialEq)]
203enum WorkerStatus {
204    /// Connecting, or reconnecting after a retryable failure.
205    Retrying,
206    /// Registered: this worker alone makes the service reachable.
207    Connected,
208    /// Permanently rejected; this worker will not try again.
209    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    /// Record `status` for `worker_index`, and publish the pool's status if the
224    /// aggregate changed.
225    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        // Outside the lock: the callback runs caller code, which must not be able
242        // to deadlock a worker reporting its own transition.
243        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    /// The best status among the workers. `Failed` needs every one of them.
251    ///
252    /// `Connected` short-circuits; `Retrying` cannot, because a later worker may
253    /// still be connected and outrank it.
254    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            // A pool that is out of workers reports the first permanent rejection
268            // it collected, which is the one that explains the others.
269            Some(reason) if !retrying => WorkerStatus::Failed(reason.clone()),
270            _ => WorkerStatus::Retrying,
271        }
272    }
273}
274
275/// How a local-server tunnel forwards: the codec, the transport, and whether
276/// its sockets keep alive.
277///
278/// Grouped rather than passed as three adjacent `bool`s, which is a
279/// transposition waiting to happen at a call site — and named fields make the
280/// call sites say which is which.
281#[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
433/// Same as [`run_server_side_cli_with_pinned_credential`], but the retry loop
434/// returns when `shutdown` is cancelled.
435///
436/// Takes both ends already resolved, because its caller — the SDK — resolves them
437/// itself in order to report a bad address as an error rather than a log line.
438pub 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    // Every worker reports here, and the pool's aggregate is what the caller is
538    // told: one healthy control connection is a reachable service, whichever
539    // worker holds it.
540    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    // Forwarded sessions outlive a single control connection, so the set lives
602    // here rather than inside the attempt: a reconnect is not a reason to drop
603    // streams that are still healthy. It is drained before this worker returns.
604    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        // Every non-terminal outcome retries; only the ladder differs.
632        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    // Cancellation is already signalled by the token the sessions hold; this only
686    // waits for them to observe it.
687    stream_tasks.shutdown().await;
688}
689
690// `backoff` is skipped: it is mutable retry state, and recording it would
691// print two ladders' internals on every span the worker enters.
692#[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    // Start registration with a protocol-v2 first frame. The session is reused for all
746    // subsequent control messages on this TCP connection. The credential is pinned
747    // when the worker starts so a later process-key change cannot retarget reconnects.
748    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    // read register resp to indicate that register has finished
800    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            // A rejection is an answer, not a failure to get one, so it never
833            // goes through `snafu_error_get_or_return_ok!`: that would report the
834            // attempt as having "finished without an error" and retry it on the
835            // transport ladder, which is how a refused registration turned into
836            // thousands of reject lines a minute.
837            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        // This worker holds a registered control connection, which is on its own
870        // enough to make the service reachable.
871        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                // Reap finished sessions so the set does not grow for the lifetime
981                // of the registration. `handle_stream` logs its own failures.
982            }
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            // Tracked rather than detached: a cancelled registration must take its
1183            // in-flight forwarded sessions with it, so `stop()` returning means no
1184            // stream of this tunnel is still moving bytes.
1185            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        // got pong response
1195        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    /// Records what the pool published, in order.
1262    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    /// One connected worker is enough: the service is reachable through it, so
1276    /// the pool must not report the other worker's retry as the pool's state.
1277    #[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    /// The guarantee `wait_ready()` depends on: while any worker might still
1286    /// bring the service up, one worker's permanent rejection is not the pool's.
1287    #[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    /// Once every worker has given up, the pool reports failed and says why —
1301    /// otherwise a caller awaiting readiness with no timeout waits forever.
1302    #[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    /// A worker that reconnects repeatedly must not turn into a stream of
1318    /// identical status lines for the caller.
1319    #[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    /// Losing the last connected worker has to be visible: the service is no
1329    /// longer reachable, and the caller's status must say so.
1330    #[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}