Skip to main content

hyperactor/channel/net/
duplex.rs

1/*
2 * Copyright (c) Meta Platforms, Inc. and affiliates.
3 * All rights reserved.
4 *
5 * This source code is licensed under the BSD-style license found in the
6 * LICENSE file in the root directory of this source tree.
7 */
8
9//! Duplex-mode channels over the net link layer.
10//!
11//! A single physical connection carries messages in both directions,
12//! each with independent sequence/ack state.
13//!
14//! ## Wire protocol
15//!
16//! Each connection starts with a unified `LinkInit` header (13 bytes,
17//! unframed) carrying the protocol kind, session id, and stream id:
18//!
19//! ```text
20//! [magic: 4B ("SMP\0" | "DPX\0")] [session_id: 8B u64 BE] [stream_id: 1B u8]
21//! ```
22//!
23//! Duplex servers expect `DPX\0`; mismatched magics are rejected at
24//! handshake. After the init, the standard tagged frame format is used.
25//! The tag byte in the 8-byte header distinguishes logical channels:
26//!
27//! - `INITIATOR_TO_ACCEPTOR = 0x00`
28//! - `ACCEPTOR_TO_INITIATOR = 0x01`
29
30use std::sync::Arc;
31
32use async_trait::async_trait;
33use backoff::ExponentialBackoffBuilder;
34use backoff::backoff::Backoff;
35use dashmap::DashMap;
36use tokio::sync::mpsc;
37use tokio::sync::watch;
38use tokio::time::Instant;
39use tokio_util::sync::CancellationToken;
40
41use super::ClientError;
42use super::Link;
43use super::LinkStatus;
44use super::ServerError;
45use super::SessionId;
46use super::log_send_error;
47#[cfg(test)]
48use super::read_link_init;
49use super::server::AcceptorLink;
50use super::server::ServerHandle;
51use super::session;
52use super::session::Next;
53use super::session::Session;
54use crate::RemoteMessage;
55use crate::channel::ChannelAddr;
56use crate::channel::ChannelError;
57use crate::channel::CloseReason;
58use crate::channel::CompletionSink;
59use crate::channel::Rx;
60use crate::channel::SendError;
61use crate::channel::SendErrorReason;
62use crate::channel::Tx;
63use crate::channel::TxStatus;
64use crate::channel::net::Stream;
65use crate::metrics;
66
67/// Public duplex server that yields `(DuplexRx<In>, DuplexTx<Out>)` pairs.
68pub struct DuplexServer<In: RemoteMessage, Out: RemoteMessage> {
69    accept_rx: mpsc::Receiver<(DuplexRx<In>, DuplexTx<Out>)>,
70    handle: ServerHandle,
71    addr: ChannelAddr,
72}
73
74impl<In: RemoteMessage, Out: RemoteMessage> DuplexServer<In, Out> {
75    /// Accept a new duplex link, returning `(rx, tx)` handles.
76    pub async fn accept(&mut self) -> Result<(DuplexRx<In>, DuplexTx<Out>), ChannelError> {
77        self.accept_rx.recv().await.ok_or(ChannelError::Closed)
78    }
79
80    /// The address this server is listening on.
81    pub fn addr(&self) -> &ChannelAddr {
82        &self.addr
83    }
84
85    /// Signal the duplex server to stop accepting new connections and
86    /// tear down its listener. Returns immediately; await
87    /// [`join`](Self::join) to confirm shutdown has completed. Forwards
88    /// to the underlying [`ServerHandle::stop`]; `join` also stops the
89    /// handle, so an explicit `stop` is only needed to signal teardown
90    /// without (yet) awaiting it.
91    pub fn stop(&self, reason: &str) {
92        self.handle.stop(reason);
93    }
94
95    /// Gracefully shut down the duplex server. Cancels the listener
96    /// and awaits its task; structured concurrency in
97    /// [`dispatch_duplex_stream`] guarantees every in-flight session
98    /// has finished its terminal cleanup (final ack flush + `Closed`
99    /// emit) before this returns.
100    pub async fn join(mut self) {
101        self.handle.stop(&format!(
102            "DuplexServer joined; channel address: {}",
103            self.addr
104        ));
105        let _ = (&mut self.handle).await;
106    }
107
108    /// Build a [`DuplexServer`] from an accept channel, server handle,
109    /// and bind address. Used by the muxed-listener path where the
110    /// accept queue is driven by a shared accept loop.
111    pub(super) fn from_parts(
112        accept_rx: mpsc::Receiver<(DuplexRx<In>, DuplexTx<Out>)>,
113        handle: ServerHandle,
114        addr: ChannelAddr,
115    ) -> Self {
116        Self {
117            accept_rx,
118            handle,
119            addr,
120        }
121    }
122}
123
124impl<In: RemoteMessage, Out: RemoteMessage> Drop for DuplexServer<In, Out> {
125    fn drop(&mut self) {
126        self.handle.stop(&format!(
127            "DuplexServer dropped; channel address: {}",
128            self.addr
129        ));
130    }
131}
132
133/// Receiver half of a duplex channel.
134pub struct DuplexRx<M: RemoteMessage>(mpsc::Receiver<M>, ChannelAddr);
135
136impl<M: RemoteMessage> DuplexRx<M> {
137    pub(super) fn new(rx: mpsc::Receiver<M>, addr: ChannelAddr) -> Self {
138        Self(rx, addr)
139    }
140}
141
142#[async_trait]
143impl<M: RemoteMessage> Rx<M> for DuplexRx<M> {
144    async fn recv(&mut self) -> Result<M, ChannelError> {
145        self.0.recv().await.ok_or(ChannelError::Closed)
146    }
147
148    fn addr(&self) -> ChannelAddr {
149        self.1.clone()
150    }
151
152    async fn join(self) {}
153}
154
155/// A handle to a duplex client session: wraps the send/recv halves
156/// and the spawned task driving the connection. Owns a cancellation
157/// token so callers can deterministically stop the recv/send loop
158/// via [`DuplexClient::join`].
159///
160/// Dropping a `DuplexClient` does *not* cancel — that would tear
161/// down sessions whose tx/rx halves the application has handed off
162/// elsewhere (e.g., into a mailbox). Call [`join`](Self::join) for
163/// orderly shutdown.
164pub struct DuplexClient<Out: RemoteMessage, In: RemoteMessage> {
165    tx: DuplexTx<Out>,
166    rx: Option<DuplexRx<In>>,
167    join_handle: tokio::task::JoinHandle<()>,
168    cancel_token: CancellationToken,
169    addr: ChannelAddr,
170}
171
172impl<Out: RemoteMessage, In: RemoteMessage> DuplexClient<Out, In> {
173    /// Get a new clone of the [`DuplexTx`] for sending messages to
174    /// the peer.
175    pub fn tx(&self) -> DuplexTx<Out> {
176        self.tx.clone()
177    }
178
179    /// Take the [`DuplexRx`] out of the client. Returns `None` on
180    /// subsequent calls — the receiver is single-consumer.
181    pub fn take_rx(&mut self) -> Option<DuplexRx<In>> {
182        self.rx.take()
183    }
184
185    /// The peer address this client dialed.
186    pub fn addr(&self) -> &ChannelAddr {
187        &self.addr
188    }
189
190    /// Gracefully shut down the duplex client session. Cancels the
191    /// recv/send loop's cancellation token (which the spawned task
192    /// observes in its `select!`s) and awaits the spawned task. On
193    /// return, the task has finished its terminal cleanup (final
194    /// ack flush on `ACCEPTOR_TO_INITIATOR`) and dropped its
195    /// [`inbound_tx`](super::session) / outbound receiver halves —
196    /// so any in-progress [`DuplexRx::recv`](super::Rx::recv) on
197    /// the receiver half resolves with [`ChannelError::Closed`].
198    pub async fn join(self) {
199        self.cancel_token.cancel();
200        let _ = self.join_handle.await;
201    }
202}
203
204impl<Out: RemoteMessage, In: RemoteMessage> std::fmt::Debug for DuplexClient<Out, In> {
205    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
206        f.debug_struct("DuplexClient")
207            .field("addr", &self.addr)
208            .field("rx_taken", &self.rx.is_none())
209            .finish()
210    }
211}
212
213/// Sender half of a duplex channel.
214pub struct DuplexTx<M: RemoteMessage> {
215    tx: mpsc::UnboundedSender<(M, CompletionSink<M>, Instant)>,
216    addr: ChannelAddr,
217    status: watch::Receiver<TxStatus>,
218}
219
220impl<M: RemoteMessage> DuplexTx<M> {
221    pub(super) fn new(
222        tx: mpsc::UnboundedSender<(M, CompletionSink<M>, Instant)>,
223        addr: ChannelAddr,
224        status: watch::Receiver<TxStatus>,
225    ) -> Self {
226        Self { tx, addr, status }
227    }
228}
229
230#[async_trait]
231impl<M: RemoteMessage> Tx<M> for DuplexTx<M> {
232    fn do_post(&self, message: M, completion: CompletionSink<M>) {
233        if let Err(mpsc::error::SendError((message, completion, _))) =
234            self.tx
235                .send((message, completion, tokio::time::Instant::now()))
236        {
237            let reason = self
238                .status
239                .borrow()
240                .as_closed()
241                .map(|r| SendErrorReason::Other(r.to_string()));
242            completion.reject(SendError {
243                error: ChannelError::Closed,
244                message,
245                reason,
246            });
247        }
248    }
249
250    fn addr(&self) -> ChannelAddr {
251        self.addr.clone()
252    }
253
254    fn status(&self) -> &watch::Receiver<TxStatus> {
255        &self.status
256    }
257}
258
259impl<M: RemoteMessage> Clone for DuplexTx<M> {
260    fn clone(&self) -> Self {
261        Self {
262            tx: self.tx.clone(),
263            addr: self.addr.clone(),
264            status: self.status.clone(),
265        }
266    }
267}
268
269impl<M: RemoteMessage> std::fmt::Debug for DuplexTx<M> {
270    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
271        f.debug_struct("DuplexTx")
272            .field("addr", &self.addr)
273            .finish()
274    }
275}
276
277/// Start a duplex server on the given address.
278pub fn serve<In: RemoteMessage, Out: RemoteMessage>(
279    addr: ChannelAddr,
280    listener: Option<std::net::TcpListener>,
281) -> Result<DuplexServer<In, Out>, ServerError> {
282    let (mut listener, channel_addr) = super::listen_with_prebound(addr, listener)?;
283
284    let (accept_tx, accept_rx) = mpsc::channel(16);
285    let cancel_token = CancellationToken::new();
286    let child_token = cancel_token.child_token();
287
288    // A duplex server only enforces the duplex protocol kind when a mux
289    // could be demultiplexing simplex traffic off the same address — i.e.
290    // on net (kernel-socket) transports, where `serve_mux` dispatches
291    // simplex and duplex separately and a plain duplex endpoint should
292    // reject stray simplex dials. The in-process `Local` transport has no
293    // mux, yet gateways still route to Local duplex endpoints via simplex
294    // `channel::dial` (e.g. an in-process host's frontend). Accept either
295    // kind there, matching the pre-link-layer behavior; the dispatch reads
296    // sessions by `session_id`, not by kind.
297    let expected_kind = if channel_addr.transport().is_net() {
298        Some(super::ProtocolKind::Duplex)
299    } else {
300        None
301    };
302    let prepare = super::preparer_for(channel_addr.clone(), expected_kind);
303
304    let sessions: Arc<DashMap<SessionId, mpsc::UnboundedSender<Box<dyn Stream>>>> =
305        Arc::new(DashMap::new());
306    let dispatch_dest = channel_addr.clone();
307    let dispatch_cancel = child_token.clone();
308    let dispatch = {
309        let sessions = Arc::clone(&sessions);
310        let accept_tx = accept_tx.clone();
311        let dest = dispatch_dest;
312        move |link_init: super::LinkInit, stream: Box<dyn Stream>| {
313            let sessions = Arc::clone(&sessions);
314            let accept_tx = accept_tx.clone();
315            let cancel = dispatch_cancel.clone();
316            let dest = dest.clone();
317            async move {
318                dispatch_duplex_stream::<In, Out>(
319                    link_init.session_id,
320                    stream,
321                    sessions,
322                    dest,
323                    &accept_tx,
324                    cancel,
325                )
326                .await;
327            }
328        }
329    };
330
331    let ca = channel_addr.clone();
332    let join_handle = tokio::spawn(async move {
333        super::server::accept_loop(&mut listener, &ca, &child_token, prepare, dispatch).await
334    });
335
336    let server_handle = ServerHandle::new(join_handle, cancel_token, channel_addr.clone());
337
338    Ok(DuplexServer {
339        accept_rx,
340        handle: server_handle,
341        addr: channel_addr,
342    })
343}
344
345/// Test-only variant that accepts an arbitrary [`super::Listener`].
346/// Mirrors [`super::server::serve_with_listener`] but for duplex
347/// servers; lets wire-level tests stage `DuplexStream`s via a
348/// custom listener so they can inspect terminal-flush behavior on
349/// the read side.
350#[cfg(test)]
351pub(super) fn serve_with_listener<In, Out, L>(
352    mut listener: L,
353    channel_addr: ChannelAddr,
354) -> Result<DuplexServer<In, Out>, ServerError>
355where
356    In: RemoteMessage,
357    Out: RemoteMessage,
358    L: super::Listener + 'static,
359    L::Stream: Unpin + std::fmt::Debug + 'static,
360{
361    let (accept_tx, accept_rx) = mpsc::channel(16);
362    let cancel_token = CancellationToken::new();
363    let child_token = cancel_token.child_token();
364
365    let prepare = |stream: L::Stream, source: ChannelAddr| async move {
366        let mut boxed: Box<dyn Stream> = Box::new(stream);
367        let link_init = read_link_init(&mut boxed)
368            .await
369            .map_err(|e| anyhow::anyhow!("LinkInit read failed from {}: {}", source, e))?;
370        if link_init.kind != super::ProtocolKind::Duplex {
371            return Err(anyhow::anyhow!(
372                "duplex server received {:?} client from {}",
373                link_init.kind,
374                source
375            ));
376        }
377        Ok((link_init, boxed))
378    };
379
380    let sessions: Arc<DashMap<SessionId, mpsc::UnboundedSender<Box<dyn Stream>>>> =
381        Arc::new(DashMap::new());
382    let dispatch_cancel = child_token.clone();
383    let dispatch = {
384        let sessions = Arc::clone(&sessions);
385        let accept_tx = accept_tx.clone();
386        let dest = channel_addr.clone();
387        move |link_init: super::LinkInit, stream: Box<dyn Stream>| {
388            let sessions = Arc::clone(&sessions);
389            let accept_tx = accept_tx.clone();
390            let cancel = dispatch_cancel.clone();
391            let dest = dest.clone();
392            async move {
393                dispatch_duplex_stream::<In, Out>(
394                    link_init.session_id,
395                    stream,
396                    sessions,
397                    dest,
398                    &accept_tx,
399                    cancel,
400                )
401                .await;
402            }
403        }
404    };
405
406    let ca = channel_addr.clone();
407    let join_handle = tokio::spawn(async move {
408        super::server::accept_loop(&mut listener, &ca, &child_token, prepare, dispatch).await
409    });
410
411    let server_handle = ServerHandle::new(join_handle, cancel_token, channel_addr.clone());
412
413    Ok(DuplexServer {
414        accept_rx,
415        handle: server_handle,
416        addr: channel_addr,
417    })
418}
419
420/// Helper to distinguish send errors from recv errors in duplex select.
421enum Either {
422    Send(session::SendLoopError),
423    Recv(session::RecvLoopError),
424}
425
426/// Dispatch a stream to the appropriate duplex session, creating one
427/// if this is the first connection for the given session ID.
428///
429/// Structured concurrency: the first dispatch for a session runs the
430/// recv/send loop inline and only returns after its terminal cleanup
431/// (flush any pending recv ack, emit `Closed` on cancellation).
432/// Reconnects hand the connection off via the per-session channel
433/// and return immediately. [`accept_loop`](super::server::accept_loop)
434/// joins every dispatch in its `connections` `JoinSet`, so it
435/// finishes only after every recv/send loop has finished — same
436/// contract as the simplex [`dispatch_stream`](super::server::dispatch_stream).
437#[tracing::instrument(level = "debug", skip_all)]
438pub(super) async fn dispatch_duplex_stream<In: RemoteMessage, Out: RemoteMessage>(
439    session_id: SessionId,
440    stream: Box<dyn Stream>,
441    sessions: Arc<DashMap<SessionId, mpsc::UnboundedSender<Box<dyn Stream>>>>,
442    addr: ChannelAddr,
443    accept_tx: &mpsc::Sender<(DuplexRx<In>, DuplexTx<Out>)>,
444    cancel: CancellationToken,
445) {
446    // Insert into the session map and drop the DashMap shard guard
447    // before any await. Vacant inserts the sender side; occupied
448    // returns the existing sender so this dispatch acts as a feeder.
449    // Unbounded so this task can publish the sender before draining
450    // the first conn without deadlocking a concurrent feeder for
451    // the same session_id.
452    let entry_result = {
453        let entry = sessions.entry(session_id);
454        match entry {
455            dashmap::mapref::entry::Entry::Occupied(e) => Err(e.get().clone()),
456            dashmap::mapref::entry::Entry::Vacant(e) => {
457                let (sender, receiver) = mpsc::unbounded_channel::<Box<dyn Stream>>();
458                e.insert(sender.clone());
459                Ok((sender, receiver))
460            }
461        }
462    };
463
464    let (sender, receiver) = match entry_result {
465        Err(sender) => {
466            // Feeder: forward the conn through the existing channel.
467            // Send returns Err if the processor task exited and
468            // dropped the receiver — drop the conn in that case.
469            let _ = sender.send(stream);
470            return;
471        }
472        Ok(pair) => pair,
473    };
474
475    // First dispatch for this session_id: set up the duplex
476    // application-facing handles and yield them to the server.
477    let (inbound_tx, inbound_rx) = mpsc::channel::<In>(1024);
478    let (outbound_tx, mut outbound_rx) =
479        mpsc::unbounded_channel::<(Out, CompletionSink<Out>, Instant)>();
480    let (notify, status) = watch::channel(TxStatus::Active);
481    let net_rx = DuplexRx(inbound_rx, addr.clone());
482    let net_tx = DuplexTx {
483        tx: outbound_tx,
484        addr: addr.clone(),
485        status,
486    };
487    let _ = accept_tx.send((net_rx, net_tx)).await;
488
489    // Hand the first connection off through the channel; the loop
490    // below picks it up via `session.connect().await`.
491    let _ = sender.send(stream);
492    drop(sender);
493
494    let link = AcceptorLink {
495        dest: addr.clone(),
496        session_id,
497        stream: receiver,
498        cancel: cancel.clone(),
499    };
500    let session_ct = cancel;
501    let dest = addr;
502    let log_id = format!("duplex server {:016x}", session_id.0);
503    let mut deliveries = session::Deliveries {
504        outbox: session::Outbox::new(log_id.clone(), dest.clone(), session_id.0),
505        unacked: session::Unacked::new(None, log_id),
506    };
507    let mut session = Session::new(link);
508    let mut recv_next = Next { seq: 0, ack: 0 };
509
510    loop {
511        let connected = match session.connect().await {
512            Ok(s) => s,
513            Err(_) => break,
514        };
515        deliveries.requeue_unacked();
516        let result = {
517            let recv_stream = connected.stream(super::INITIATOR_TO_ACCEPTOR);
518            let send_stream = connected.stream(super::ACCEPTOR_TO_INITIATOR);
519            tokio::select! {
520                r = session::recv_connected::<In, _, _>(
521                    &recv_stream,
522                    &inbound_tx,
523                    &mut recv_next,
524                ) => r.map_err(Either::Recv),
525                r = session::send_connected(
526                    &send_stream,
527                    &mut deliveries,
528                    &mut outbound_rx,
529                ) => r.map_err(Either::Send),
530                _ = session_ct.cancelled() => Err(Either::Recv(session::RecvLoopError::Cancelled)),
531            }
532        };
533
534        let terminal = match &result {
535            Ok(()) => {
536                tracing::info!(
537                    session_id = session_id.0,
538                    "duplex recv_connected returned EOF, awaiting reconnect"
539                );
540                false
541            }
542            Err(Either::Send(session::SendLoopError::Io(err))) => {
543                tracing::info!(
544                    session_id = session_id.0,
545                    error = %err,
546                    "duplex send error (recoverable)",
547                );
548                false
549            }
550            Err(Either::Recv(session::RecvLoopError::Io(err))) => {
551                tracing::info!(
552                    session_id = session_id.0,
553                    error = %err,
554                    "duplex recv error (recoverable)",
555                );
556                false
557            }
558            Err(Either::Send(e)) => {
559                tracing::info!(
560                    session_id = session_id.0,
561                    error = %e,
562                    "duplex send terminal error"
563                );
564                true
565            }
566            Err(Either::Recv(e)) => {
567                tracing::info!(
568                    session_id = session_id.0,
569                    error = %e,
570                    "duplex recv terminal error"
571                );
572                true
573            }
574        };
575
576        // Flush any pending recv ack so the peer's
577        // unacked queue clears cleanly before this
578        // connection goes away. Mirrors the simplex
579        // server's drain logic (see
580        // `dispatch_stream`); without it, peers
581        // retry-loop until `MESSAGE_DELIVERY_TIMEOUT`.
582        if recv_next.ack < recv_next.seq {
583            let recv_stream = connected.stream(super::INITIATOR_TO_ACCEPTOR);
584            let ack = super::serialize_response(super::NetRxResponse::Ack(recv_next.seq - 1))
585                .expect("serialize ack");
586            let mut completion = recv_stream.write(ack);
587            match completion.drive().await {
588                Ok(()) => {
589                    recv_next.ack = recv_next.seq;
590                }
591                Err(e) => {
592                    tracing::debug!(
593                        session_id = session_id.0,
594                        error = %e,
595                        "duplex: failed to flush acks during cleanup"
596                    );
597                }
598            }
599        }
600
601        // On terminal exit, tell the peer we're
602        // closing so it stops trying to reconnect.
603        let terminal_response = match &result {
604            Err(Either::Recv(session::RecvLoopError::SequenceError(reason))) => {
605                Some(super::NetRxResponse::Reject(reason.clone()))
606            }
607            Err(Either::Recv(session::RecvLoopError::Cancelled))
608            | Err(Either::Send(session::SendLoopError::AppClosed)) => {
609                Some(super::NetRxResponse::Closed)
610            }
611            _ => None,
612        };
613        if let Some(rsp) = terminal_response {
614            let recv_stream = connected.stream(super::INITIATOR_TO_ACCEPTOR);
615            let data = super::serialize_response(rsp).expect("serialize terminal response");
616            let mut completion = recv_stream.write(data);
617            let _ = completion.drive().await;
618        }
619
620        session = connected.release();
621        if terminal {
622            break;
623        }
624    }
625
626    // Recv/send loop is finished — drop the session entry so a later
627    // reconnect for the same session_id starts a fresh dispatch task
628    // instead of feeding a dead channel; any in-flight feeder's send
629    // fails after the link's receiver above is dropped.
630    sessions.remove(&session_id);
631
632    let _ = notify.send(TxStatus::Closed(CloseReason::Other(
633        "duplex session ended".into(),
634    )));
635}
636
637/// Establish a duplex (bidirectional) session over the given link.
638/// Returns a [`DuplexClient`] wrapping the send/recv halves and the
639/// spawned recv/send task; the client owns a cancellation token so
640/// callers can deterministically tear the session down via
641/// [`DuplexClient::stop`] followed by [`DuplexClient::join`].
642#[tracing::instrument(level = "debug", skip_all)]
643pub(crate) fn spawn<Out: RemoteMessage, In: RemoteMessage>(
644    link: impl Link,
645) -> DuplexClient<Out, In> {
646    let addr = link.dest();
647    let session_id = link.link_id();
648    let (outbound_tx, outbound_rx) = tokio::sync::mpsc::unbounded_channel();
649    let (inbound_tx, inbound_rx) = tokio::sync::mpsc::channel::<In>(1024);
650    let (notify, status) = watch::channel(TxStatus::Active);
651    let cancel_token = CancellationToken::new();
652    let task_cancel = cancel_token.clone();
653    let dest = addr.clone();
654    let join_handle = crate::init::get_runtime().spawn(async move {
655        let mut session = Session::new(link);
656        let log_id = format!("session {}.{:016x}", dest, session_id.0);
657        let mut deliveries = session::Deliveries {
658            outbox: session::Outbox::new(log_id.clone(), dest.clone(), session_id.0),
659            unacked: session::Unacked::new(None, log_id),
660        };
661        let mut outbound_rx = outbound_rx;
662        let mut recv_next = Next { seq: 0, ack: 0 };
663        let mut reconnect_backoff = ExponentialBackoffBuilder::new()
664            .with_initial_interval(std::time::Duration::from_millis(10))
665            .with_multiplier(2.0)
666            .with_randomization_factor(0.1)
667            .with_max_interval(std::time::Duration::from_secs(5))
668            .with_max_elapsed_time(None)
669            .build();
670
671        let mut link_status = LinkStatus::NeverConnected;
672
673        loop {
674            // Race connect against cancel so a `DuplexClient::stop`
675            // call mid-dial doesn't have to wait for the dial backoff
676            // to elapse.
677            let connected = tokio::select! {
678                result = session.connect() => match result {
679                    Ok(s) => s,
680                    Err(_) => break,
681                },
682                _ = task_cancel.cancelled() => break,
683            };
684
685            metrics::CHANNEL_CONNECTIONS.add(
686                1,
687                hyperactor_telemetry::kv_pairs!(
688                    "transport" => dest.transport().to_string(),
689                    "mode" => "duplex",
690                    "reason" => "link connected",
691                ),
692            );
693
694            if !deliveries.unacked.is_empty() {
695                metrics::CHANNEL_RECONNECTIONS.add(
696                    1,
697                    hyperactor_telemetry::kv_pairs!(
698                        "dest" => dest.to_string(),
699                        "transport" => dest.transport().to_string(),
700                        "mode" => "duplex",
701                        "reason" => "reconnect_with_unacked",
702                    ),
703                );
704            }
705            deliveries.requeue_unacked();
706
707            link_status.connected();
708            let connected_at = tokio::time::Instant::now();
709
710            let result = {
711                let send_stream = connected.stream(super::INITIATOR_TO_ACCEPTOR);
712                let recv_stream = connected.stream(super::ACCEPTOR_TO_INITIATOR);
713                tokio::select! {
714                    r = session::send_connected(
715                        &send_stream, &mut deliveries, &mut outbound_rx,
716                    ) => r.map_err(Either::Send),
717                    r = session::recv_connected::<In, _, _>(
718                        &recv_stream, &inbound_tx, &mut recv_next,
719                    ) => r.map_err(Either::Recv),
720                    _ = task_cancel.cancelled() => Err(Either::Recv(session::RecvLoopError::Cancelled)),
721                }
722            };
723
724            link_status.disconnected();
725
726            if connected_at.elapsed() > tokio::time::Duration::from_secs(1) {
727                reconnect_backoff.reset();
728            }
729
730            let terminal = match &result {
731                Ok(()) => {
732                    if let Some(delay) = reconnect_backoff.next_backoff() {
733                        tracing::info!(
734                            dest = %dest,
735                            session_id = session_id.0,
736                            delay_ms = delay.as_millis() as u64,
737                            "duplex send_connected returned EOF, reconnecting after backoff; {link_status}"
738                        );
739                        tokio::time::sleep(delay).await;
740                    }
741                    false
742                }
743                Err(Either::Send(e)) => {
744                    let terminal = log_send_error(e, &dest, session_id.0, "duplex", &link_status);
745                    if !terminal {
746                        // Recoverable send error — reconnect after backoff.
747                        if let Some(delay) = reconnect_backoff.next_backoff() {
748                            tracing::info!(
749                                dest = %dest,
750                                session_id = session_id.0,
751                                error = %e,
752                                delay_ms = delay.as_millis() as u64,
753                                mode = "duplex",
754                                "send error (recoverable), reconnecting after backoff; {link_status}",
755                            );
756                            tokio::time::sleep(delay).await;
757                        }
758                    }
759                    terminal
760                }
761                Err(Either::Recv(session::RecvLoopError::Io(err))) => {
762                    if let Some(delay) = reconnect_backoff.next_backoff() {
763                        tracing::info!(
764                            dest = %dest,
765                            session_id = session_id.0,
766                            error = %err,
767                            delay_ms = delay.as_millis() as u64,
768                            mode = "duplex",
769                            "recv error (recoverable), reconnecting after backoff; {link_status}",
770                        );
771                        tokio::time::sleep(delay).await;
772                    }
773                    metrics::CHANNEL_ERRORS.add(
774                        1,
775                        hyperactor_telemetry::kv_pairs!(
776                            "dest" => dest.to_string(),
777                            "session_id" => session_id.0.to_string(),
778                            "error_type" => metrics::ChannelErrorType::SendError.as_str(),
779                            "mode" => "duplex",
780                        ),
781                    );
782                    false
783                }
784                Err(Either::Recv(e)) => {
785                    tracing::info!(
786                        dest = %dest,
787                        session_id = session_id.0,
788                        error = %e,
789                        "duplex recv terminal error; {link_status}"
790                    );
791                    true
792                }
793            };
794
795            // Flush any pending recv ack so the peer's send-side
796            // unacked queue clears cleanly before this connection
797            // goes away. Mirrors the cleanup in
798            // `dispatch_duplex_stream` but on the other tag — the
799            // initiator reads data on `ACCEPTOR_TO_INITIATOR`, so
800            // its acks travel back on that same tag.
801            if recv_next.ack < recv_next.seq {
802                let recv_stream = connected.stream(super::ACCEPTOR_TO_INITIATOR);
803                let ack = super::serialize_response(super::NetRxResponse::Ack(recv_next.seq - 1))
804                    .expect("serialize ack");
805                let mut completion = recv_stream.write(ack);
806                match completion.drive().await {
807                    Ok(()) => {
808                        recv_next.ack = recv_next.seq;
809                    }
810                    Err(e) => {
811                        tracing::debug!(
812                            dest = %dest,
813                            session_id = session_id.0,
814                            error = %e,
815                            "duplex client: failed to flush acks during cleanup"
816                        );
817                    }
818                }
819            }
820
821            // On terminal exit, tell the peer we're closing (or
822            // rejecting) on the same tag we read data on. Mirrors
823            // `dispatch_duplex_stream`; without it, the peer's
824            // server-side dispatch keeps awaiting a reconnect that
825            // will never come.
826            let terminal_response = match &result {
827                Err(Either::Recv(session::RecvLoopError::SequenceError(reason))) => {
828                    Some(super::NetRxResponse::Reject(reason.clone()))
829                }
830                Err(Either::Recv(session::RecvLoopError::Cancelled))
831                | Err(Either::Send(session::SendLoopError::AppClosed)) => {
832                    Some(super::NetRxResponse::Closed)
833                }
834                _ => None,
835            };
836            if let Some(rsp) = terminal_response {
837                let recv_stream = connected.stream(super::ACCEPTOR_TO_INITIATOR);
838                let data =
839                    super::serialize_response(rsp).expect("serialize terminal response");
840                let mut completion = recv_stream.write(data);
841                let _ = completion.drive().await;
842            }
843
844            session = connected.release();
845            if terminal {
846                break;
847            }
848        }
849
850        let _ = notify.send(TxStatus::Closed(CloseReason::Other(
851            "duplex session ended".into(),
852        )));
853    });
854    let tx = DuplexTx::new(outbound_tx, addr.clone(), status);
855    let rx = DuplexRx::new(inbound_rx, addr.clone());
856    DuplexClient {
857        tx,
858        rx: Some(rx),
859        join_handle,
860        cancel_token,
861        addr,
862    }
863}
864
865/// Connect to a duplex server. Returns a [`DuplexClient`] wrapping
866/// the send/recv halves and the spawned recv/send task; callers use
867/// [`DuplexClient::tx`] / [`DuplexClient::take_rx`] to extract the
868/// halves and [`DuplexClient::stop`] followed by [`DuplexClient::join`]
869/// to deterministically shut the session down.
870pub fn dial<Out: RemoteMessage, In: RemoteMessage>(
871    addr: ChannelAddr,
872) -> Result<DuplexClient<Out, In>, ClientError> {
873    let addr = addr.into_dial_addr();
874    Ok(spawn(super::link(
875        addr,
876        super::SessionId::random(),
877        0,
878        super::ProtocolKind::Duplex,
879    )?))
880}
881
882#[cfg(test)]
883mod tests {
884    use timed_test::async_timed_test;
885
886    use super::*;
887    use crate::channel::ChannelTransport;
888
889    #[async_timed_test(timeout_secs = 30)]
890    // TODO: OSS: called `Result::unwrap()` on an `Err` value: Listen(Tcp([::1]:0), Os { code: 99, kind: AddrNotAvailable, message: "Cannot assign requested address" })
891    #[cfg_attr(not(fbcode_build), ignore)]
892    async fn test_duplex_basic() {
893        let mut server =
894            serve::<u64, String>(ChannelAddr::Tcp("[::1]:0".parse().unwrap()), None).unwrap();
895        let server_addr = server.addr().clone();
896
897        // Client: sends u64, receives String.
898        let mut client = dial::<u64, String>(server_addr).unwrap();
899        let client_tx = client.tx();
900        let mut client_rx = client.take_rx().unwrap();
901
902        // Server: receives u64, sends String.
903        let (mut server_rx, server_tx) = server.accept().await.unwrap();
904
905        // Client sends to server.
906        client_tx.post(42);
907        let received = server_rx.recv().await.unwrap();
908        assert_eq!(received, 42);
909
910        // Server sends to client.
911        server_tx.post("hello".to_string());
912        let received = client_rx.recv().await.unwrap();
913        assert_eq!(received, "hello");
914
915        // Multiple messages both ways.
916        for i in 0..10u64 {
917            client_tx.post(i);
918            assert_eq!(server_rx.recv().await.unwrap(), i);
919
920            server_tx.post(format!("msg-{}", i));
921            assert_eq!(client_rx.recv().await.unwrap(), format!("msg-{}", i));
922        }
923    }
924
925    #[async_timed_test(timeout_secs = 30)]
926    #[cfg_attr(not(fbcode_build), ignore)]
927    async fn test_duplex_multiple_links() {
928        let mut server =
929            serve::<u64, u64>(ChannelAddr::Tcp("[::1]:0".parse().unwrap()), None).unwrap();
930        let server_addr = server.addr().clone();
931
932        // Two independent clients.
933        let mut client1 = dial::<u64, u64>(server_addr.clone()).unwrap();
934        let tx1 = client1.tx();
935        let mut rx1 = client1.take_rx().unwrap();
936        let (mut srx1, stx1) = server.accept().await.unwrap();
937
938        let mut client2 = dial::<u64, u64>(server_addr).unwrap();
939        let tx2 = client2.tx();
940        let mut rx2 = client2.take_rx().unwrap();
941        let (mut srx2, stx2) = server.accept().await.unwrap();
942
943        // Send on link 1.
944        tx1.post(100);
945        assert_eq!(srx1.recv().await.unwrap(), 100);
946        stx1.post(200);
947        assert_eq!(rx1.recv().await.unwrap(), 200);
948
949        // Send on link 2.
950        tx2.post(300);
951        assert_eq!(srx2.recv().await.unwrap(), 300);
952        stx2.post(400);
953        assert_eq!(rx2.recv().await.unwrap(), 400);
954    }
955
956    /// Ping-pong helper: server echoes back each message it receives.
957    /// Returns elapsed time for `iterations` round-trips.
958    async fn duplex_ping_pong(
959        addr: ChannelAddr,
960        iterations: usize,
961    ) -> anyhow::Result<std::time::Duration> {
962        let mut server = serve::<u64, u64>(addr, None)?;
963        let server_addr = server.addr().clone();
964
965        let server_handle = tokio::spawn(async move {
966            let (mut rx, tx) = server.accept().await.unwrap();
967            while let Ok(msg) = rx.recv().await {
968                tx.post(msg);
969            }
970        });
971
972        let mut client = dial::<u64, u64>(server_addr).unwrap();
973        let client_tx = client.tx();
974        let mut client_rx = client.take_rx().unwrap();
975
976        // Warmup.
977        for i in 0..10u64 {
978            client_tx.post(i);
979            assert_eq!(client_rx.recv().await?, i);
980        }
981
982        let start = std::time::Instant::now();
983        for i in 0..iterations as u64 {
984            client_tx.post(i);
985            assert_eq!(client_rx.recv().await?, i);
986        }
987        let elapsed = start.elapsed();
988
989        server_handle.abort();
990        Ok(elapsed)
991    }
992
993    #[async_timed_test(timeout_secs = 30)]
994    #[cfg_attr(not(fbcode_build), ignore)]
995    async fn test_duplex_ping_pong_tcp() {
996        let elapsed = duplex_ping_pong(ChannelAddr::Tcp("[::1]:0".parse().unwrap()), 100)
997            .await
998            .unwrap();
999        println!("TCP duplex: 100 round-trips in {elapsed:?}");
1000    }
1001
1002    #[async_timed_test(timeout_secs = 30)]
1003    async fn test_duplex_ping_pong_unix() {
1004        let elapsed = duplex_ping_pong(ChannelAddr::any(ChannelTransport::Unix), 100)
1005            .await
1006            .unwrap();
1007        println!("Unix duplex: 100 round-trips in {elapsed:?}");
1008    }
1009}