Skip to main content

monarch_rdma/backend/tcp/
manager_actor.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//! TCP manager actor for RDMA fallback transport.
10//!
11//! Transfers buffer data over the default hyperactor channel transport
12//! in chunks controlled by
13//! [`RDMA_MAX_CHUNK_SIZE_MB`](crate::config::RDMA_MAX_CHUNK_SIZE_MB).
14
15use std::collections::HashMap;
16use std::sync::Arc;
17use std::sync::OnceLock;
18use std::time::Duration;
19use std::time::Instant;
20
21use anyhow::Result;
22use async_trait::async_trait;
23use bytes::Bytes;
24use bytes::BytesMut;
25use dashmap::DashMap;
26use hyperactor::Actor;
27use hyperactor::ActorHandle;
28use hyperactor::ActorRef;
29use hyperactor::Context;
30use hyperactor::Endpoint as _;
31use hyperactor::HandleClient;
32use hyperactor::Handler;
33use hyperactor::Instance;
34use hyperactor::OncePortHandle;
35use hyperactor::OncePortRef;
36use hyperactor::PortHandle;
37use hyperactor::RefClient;
38use hyperactor::actor::ActorError;
39use hyperactor::channel;
40use hyperactor::channel::ChannelAddr;
41use hyperactor::channel::ChannelTx;
42use hyperactor::channel::Rx;
43use hyperactor::channel::Tx;
44use hyperactor::context;
45use hyperactor::context::Actor as _;
46use hyperactor_mesh::transport::default_bind_spec;
47use serde::Deserialize;
48use serde::Serialize;
49use serde_multipart::Part;
50use tokio::time::timeout as tokio_timeout;
51use tokio_util::sync::CancellationToken;
52use typeuri::Named;
53
54use super::TcpOp;
55use crate::RdmaOp;
56use crate::RdmaOpType;
57use crate::RdmaTransportLevel;
58use crate::backend::RdmaBackend;
59use crate::backend::RdmaConfig;
60use crate::local_memory::KeepaliveLocalMemory;
61use crate::rdma_manager_actor::GetTcpActorRefClient;
62use crate::rdma_manager_actor::RdmaManagerActor;
63use crate::rdma_manager_actor::RdmaManagerMessageClient;
64
65/// [`Named`] wrapper around [`Part`] for use as a reply type.
66///
67/// [`Part`] itself does not implement [`Named`], which is required by
68/// [`OncePortRef`]. This newtype adds the missing trait.
69#[derive(Debug, Clone, Serialize, Deserialize, Named)]
70pub struct TcpChunk(Part);
71wirevalue::register_type!(TcpChunk);
72
73/// Data chunk sent over direct parallel channels.
74#[derive(Debug, Clone, Serialize, Deserialize, Named)]
75struct TcpDataChunk {
76    // Which specific transfer this chunk is associated with.
77    transfer_id: usize,
78    // Offset into the buffer for this chunk.
79    offset: usize,
80    data: Part,
81}
82wirevalue::register_type!(TcpDataChunk);
83
84/// Tracks the progress of a single parallel transfer (read or write).
85///
86/// Shared between channel receive loops and actor message handlers
87/// via a [`DashMap`].
88#[derive(Debug)]
89struct TransferState {
90    /// Buffer backing this transfer, provided at construction.
91    local_memory: KeepaliveLocalMemory,
92
93    /// Number of chunks received so far.
94    chunks_received: usize,
95
96    /// Total chunks expected for this transfer.
97    total_chunks: usize,
98
99    /// Completion reply port. Fired when all chunks arrive
100    /// or an error occurs.
101    done: OncePortRef<Result<(), String>>,
102}
103
104impl TransferState {
105    fn new(
106        total_chunks: usize,
107        local_memory: KeepaliveLocalMemory,
108        done: OncePortRef<Result<(), String>>,
109    ) -> Self {
110        Self {
111            local_memory,
112            chunks_received: 0,
113            total_chunks,
114            done,
115        }
116    }
117}
118
119/// Sends the result of a completed transfer to the caller's reply port.
120///
121/// Sending an actor message from the spawned receiver task requires the
122/// loop to own a dummy [`context::Actor`] impl. If the task sent directly
123/// to a remote [`OncePortRef`], the message would appear to come from
124/// this dummy context, and undeliverable messages wouldn't be handled
125/// properly. This intermediate message lets us use a [`PortHandle`]
126/// whose message cannot be undeliverable; the handler then forwards the
127/// result using the real actor's context.
128#[derive(Debug, Serialize, Deserialize, Named)]
129struct SendTransferResult {
130    done: OncePortRef<Result<(), String>>,
131    result: Result<(), String>,
132}
133
134/// Fatal error from the receive loop.
135///
136/// The handler logs the error and returns `Err`, which triggers a
137/// supervision event and crashes the actor.
138#[derive(Debug, Serialize, Deserialize, Named)]
139struct TransferError {
140    message: String,
141}
142
143/// Set up the local TcpManagerActor to receive a parallel transfer from
144/// a remote TcpManagerActor.
145#[derive(Debug)]
146struct RegisterTransferLocal {
147    local_memory: KeepaliveLocalMemory,
148    total_chunks: usize,
149    done: OncePortRef<Result<(), String>>,
150    // The transfer ID
151    reply: OncePortHandle<usize>,
152}
153
154/// Tell the local TcpManagerActor to read local memory and push
155/// chunks to `dest_addr`.
156#[derive(Debug)]
157struct ExecuteTransferLocal {
158    transfer_id: usize,
159    local_memory: KeepaliveLocalMemory,
160    chunk_size: usize,
161    dest_addr: ChannelAddr,
162}
163
164/// Serializable messages for the [`TcpManagerActor`].
165///
166/// These travel over the wire between processes. The [`Part`] payload
167/// is transferred via the multipart codec without an extra copy.
168#[derive(Handler, HandleClient, RefClient, Debug, Serialize, Deserialize, Named)]
169enum TcpManagerMessage {
170    /// Write a chunk of data into a registered buffer at the given offset.
171    WriteChunk {
172        buf_id: usize,
173        offset: usize,
174        data: Part,
175        #[reply]
176        reply: OncePortRef<Result<(), String>>,
177    },
178    /// Read a chunk of data from a registered buffer at the given offset.
179    ReadChunk {
180        buf_id: usize,
181        offset: usize,
182        size: usize,
183        #[reply]
184        reply: OncePortRef<Result<TcpChunk, String>>,
185    },
186    /// Return the channel address served by this actor for parallel transfers.
187    /// `None` when parallelism is 1.
188    GetChannelAddress {
189        #[reply]
190        reply: OncePortRef<Option<ChannelAddr>>,
191    },
192    /// Set up a remote TcpManagerActor to receive a parallel transfer from
193    /// the sender.
194    RegisterTransferRemote {
195        buf_id: usize,
196        total_chunks: usize,
197        done: OncePortRef<Result<(), String>>,
198        #[reply]
199        reply: OncePortRef<Result<usize, String>>,
200    },
201    /// Tell the remote TcpManagerActor to read its local memory and push
202    /// chunks to the dest_addr provided by the sender.
203    ExecuteTransferRemote {
204        transfer_id: usize,
205        buf_id: usize,
206        chunk_size: usize,
207        dest_addr: ChannelAddr,
208        #[reply]
209        reply: OncePortRef<Result<(), String>>,
210    },
211}
212wirevalue::register_type!(TcpManagerMessage);
213
214/// TCP fallback RDMA backend actor.
215///
216/// Spawned as a child of [`RdmaManagerActor`]. Transfers buffer data
217/// over the default hyperactor channel transport in chunks.
218#[derive(Debug)]
219#[hyperactor::export(
220    handlers = [TcpManagerMessage],
221)]
222pub struct TcpManagerActor {
223    owner: OnceLock<ActorHandle<RdmaManagerActor>>,
224    next_transfer_id: usize,
225    transfers: Arc<DashMap<usize, TransferState>>,
226    /// Address of the direct channel served for parallel transfers.
227    /// `None` when parallelism is 1 (default).
228    channel_addr: Option<ChannelAddr>,
229    /// Cached outbound connections keyed by remote channel address.
230    outbound: HashMap<ChannelAddr, Vec<Arc<ChannelTx<TcpDataChunk>>>>,
231    /// Cancellation token for spawned tasks.
232    cancel: CancellationToken,
233    /// Signaled when the parallel receive loop exits.
234    receiver_done: Option<tokio::sync::oneshot::Receiver<()>>,
235}
236
237impl TcpManagerActor {
238    pub fn new() -> Self {
239        Self {
240            owner: OnceLock::new(),
241            next_transfer_id: 0,
242            transfers: Arc::new(DashMap::new()),
243            channel_addr: None,
244            outbound: HashMap::new(),
245            cancel: CancellationToken::new(),
246            receiver_done: None,
247        }
248    }
249
250    fn register_transfer(
251        &mut self,
252        local_memory: KeepaliveLocalMemory,
253        total_chunks: usize,
254        done: OncePortRef<Result<(), String>>,
255    ) -> usize {
256        let transfer_id = self.next_transfer_id;
257        self.next_transfer_id += 1;
258        self.transfers.insert(
259            transfer_id,
260            TransferState::new(total_chunks, local_memory, done),
261        );
262        transfer_id
263    }
264
265    fn execute_transfer(
266        &mut self,
267        cx: &Context<Self>,
268        transfer_id: usize,
269        local_memory: KeepaliveLocalMemory,
270        chunk_size: usize,
271        dest_addr: ChannelAddr,
272    ) -> Result<()> {
273        let parallelism =
274            hyperactor_config::global::get(crate::config::RDMA_TCP_FALLBACK_PARALLELISM);
275
276        if !self.outbound.contains_key(&dest_addr) {
277            let conns = (0..parallelism)
278                .map(|_| {
279                    channel::dial::<TcpDataChunk>(dest_addr.clone())
280                        .map(Arc::new)
281                        .map_err(anyhow::Error::from)
282                })
283                .collect::<Result<Vec<_>>>()?;
284            self.outbound.insert(dest_addr.clone(), conns);
285        }
286        let conns = self.outbound.get(&dest_addr).unwrap();
287
288        let size = local_memory.size();
289        let total_chunks = size.div_ceil(chunk_size);
290
291        let chunk_index = Arc::new(std::sync::atomic::AtomicUsize::new(0));
292        let error_port: PortHandle<TransferError> = cx.port();
293        let proc = cx.instance().proc().clone();
294        let cancel = self.cancel.clone();
295
296        for conn in conns.clone() {
297            let mem = local_memory.clone();
298            let chunk_index = chunk_index.clone();
299            let error_port = error_port.clone();
300            let proc = proc.clone();
301            let cancel = cancel.clone();
302
303            tokio::spawn(async move {
304                let sender_name = format!(
305                    "tcp_chunk_sender_{}",
306                    hyperactor_mesh::shortuuid::ShortUuid::generate()
307                );
308                let instance = proc.client(&sender_name);
309
310                loop {
311                    if cancel.is_cancelled() {
312                        return;
313                    }
314
315                    let idx = chunk_index.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
316                    if idx >= total_chunks {
317                        break;
318                    }
319
320                    let offset = idx * chunk_size;
321                    let len = std::cmp::min(chunk_size, size - offset);
322                    let mut buf = BytesMut::zeroed(len);
323                    // SAFETY: the caller is responsible for ensuring that no other
324                    // component writes the target byte range concurrently.
325                    if let Err(e) = unsafe { mem.read_at(offset, &mut buf) } {
326                        error_port.post(
327                            &instance,
328                            TransferError {
329                                message: format!("read_at failed at offset {offset}: {e}"),
330                            },
331                        );
332                        return;
333                    }
334
335                    let chunk = TcpDataChunk {
336                        transfer_id,
337                        offset,
338                        data: Part::from(buf.freeze()),
339                    };
340
341                    if let Err(e) = conn.send(chunk).await {
342                        error_port.post(
343                            &instance,
344                            TransferError {
345                                message: format!("failed to send chunk at offset {offset}: {e}"),
346                            },
347                        );
348                        return;
349                    }
350                }
351            });
352        }
353
354        Ok(())
355    }
356
357    /// Construct an [`ActorHandle`] for the local [`TcpManagerActor`]
358    /// by querying the local [`RdmaManagerActor`].
359    pub async fn local_handle(
360        client: &(impl context::Actor + Send + Sync),
361    ) -> Result<ActorHandle<Self>, anyhow::Error> {
362        let rdma_handle = RdmaManagerActor::local_handle(client);
363        let tcp_ref = rdma_handle.get_tcp_actor_ref(client).await?;
364        tcp_ref
365            .downcast_handle(client)
366            .ok_or_else(|| anyhow::anyhow!("TcpManagerActor is not in the local process"))
367    }
368}
369
370#[async_trait]
371impl Actor for TcpManagerActor {
372    async fn init(&mut self, this: &Instance<Self>) -> Result<(), anyhow::Error> {
373        let owner = this.parent_handle().ok_or_else(|| {
374            anyhow::anyhow!("RdmaManagerActor not found as parent of TcpManagerActor")
375        })?;
376        self.owner
377            .set(owner)
378            .map_err(|_| anyhow::anyhow!("TcpManagerActor owner already set"))?;
379
380        let parallelism =
381            hyperactor_config::global::get(crate::config::RDMA_TCP_FALLBACK_PARALLELISM);
382        if parallelism > 1 {
383            let addr = match default_bind_spec() {
384                channel::BindSpec::Any(transport) => ChannelAddr::any(transport),
385                channel::BindSpec::Addr(addr) => addr,
386            };
387            let (bound_addr, mut rx) = channel::serve::<TcpDataChunk>(addr)?;
388            self.channel_addr = Some(bound_addr);
389
390            let transfers = self.transfers.clone();
391            let proc = this.proc().clone();
392            let result_port: PortHandle<SendTransferResult> = this.port();
393            let error_port: PortHandle<TransferError> = this.port();
394            let cancel = self.cancel.clone();
395            let receiver_name = format!(
396                "tcp_chunk_receiver_{}",
397                hyperactor_mesh::shortuuid::ShortUuid::generate()
398            );
399
400            let (done_tx, done_rx) = tokio::sync::oneshot::channel::<()>();
401            self.receiver_done = Some(done_rx);
402
403            tokio::spawn(async move {
404                let instance = proc.client(&receiver_name.to_string());
405
406                loop {
407                    let chunk = tokio::select! {
408                        _ = cancel.cancelled() => break,
409                        result = rx.recv() => match result {
410                            Ok(chunk) => chunk,
411                            Err(e) => {
412                                error_port
413                                    .post(
414                                        &instance,
415                                        TransferError {
416                                            message: format!(
417                                                "parallel channel receive error: {e}"
418                                            ),
419                                        },
420                                    );
421                                break;
422                            }
423                        },
424                    };
425
426                    let mut entry = match transfers.get_mut(&chunk.transfer_id) {
427                        Some(entry) => entry,
428                        None => {
429                            tracing::warn!(
430                                "received chunk for unknown transfer {:?}",
431                                chunk.transfer_id,
432                            );
433                            continue;
434                        }
435                    };
436
437                    let mut write_offset = chunk.offset;
438                    let fragments = chunk.data.into_fragments();
439                    let write_err = fragments.iter().find_map(|fragment| {
440                        // SAFETY: the caller is responsible for ensuring that no other
441                        // component reads or writes the target byte range concurrently.
442                        let result = unsafe { entry.local_memory.write_at(write_offset, fragment) };
443                        write_offset += fragment.len();
444                        result.err()
445                    });
446                    if let Some(e) = write_err {
447                        let transfer_id = chunk.transfer_id;
448                        drop(entry);
449                        let (_, state) = transfers.remove(&transfer_id).unwrap();
450                        result_port.post(
451                            &instance,
452                            SendTransferResult {
453                                done: state.done,
454                                result: Err(e.to_string()),
455                            },
456                        );
457                        continue;
458                    }
459
460                    entry.chunks_received += 1;
461                    if entry.chunks_received == entry.total_chunks {
462                        let transfer_id = chunk.transfer_id;
463                        drop(entry);
464                        let (_, state) = transfers.remove(&transfer_id).unwrap();
465                        result_port.post(
466                            &instance,
467                            SendTransferResult {
468                                done: state.done,
469                                result: Ok(()),
470                            },
471                        );
472                    }
473                }
474                rx.join().await;
475                done_tx.send(()).unwrap();
476            });
477        }
478
479        Ok(())
480    }
481
482    async fn cleanup(
483        &mut self,
484        _this: &Instance<Self>,
485        _err: Option<&ActorError>,
486    ) -> Result<(), anyhow::Error> {
487        self.cancel.cancel();
488        if let Some(done_rx) = self.receiver_done.take() {
489            done_rx.await?;
490        }
491        Ok(())
492    }
493}
494
495#[async_trait]
496#[hyperactor::handle(TcpManagerMessage)]
497impl TcpManagerMessageHandler for TcpManagerActor {
498    async fn write_chunk(
499        &mut self,
500        cx: &Context<Self>,
501        buf_id: usize,
502        offset: usize,
503        data: Part,
504    ) -> Result<Result<(), String>, anyhow::Error> {
505        let owner = self.owner.get().expect("TcpManagerActor owner not set");
506        let mem = match owner.request_local_memory(cx, buf_id).await {
507            Ok(Some(mem)) => mem,
508            Ok(None) => return Ok(Err(format!("buffer {buf_id} not found"))),
509            Err(e) => return Ok(Err(e.to_string())),
510        };
511
512        let bytes = data.into_bytes();
513        // SAFETY: the remote peer that issued this `WriteChunk` had to
514        // register the buffer locally first; its caller is responsible
515        // for ensuring no other component reads or writes the target byte
516        // range concurrently.
517        if let Err(e) = unsafe { mem.write_at(offset, &bytes) } {
518            return Ok(Err(e.to_string()));
519        }
520
521        Ok(Ok(()))
522    }
523
524    async fn read_chunk(
525        &mut self,
526        cx: &Context<Self>,
527        buf_id: usize,
528        offset: usize,
529        size: usize,
530    ) -> Result<Result<TcpChunk, String>, anyhow::Error> {
531        let owner = self.owner.get().expect("TcpManagerActor owner not set");
532        let mem = match owner.request_local_memory(cx, buf_id).await {
533            Ok(Some(mem)) => mem,
534            Ok(None) => return Ok(Err(format!("buffer {buf_id} not found"))),
535            Err(e) => return Ok(Err(e.to_string())),
536        };
537
538        let mut buf = BytesMut::zeroed(size);
539        // SAFETY: the remote peer that issued this `ReadChunk` had to
540        // register the buffer locally first; its caller is responsible
541        // for ensuring no other component writes the target byte range
542        // concurrently.
543        if let Err(e) = unsafe { mem.read_at(offset, &mut buf) } {
544            return Ok(Err(e.to_string()));
545        }
546        Ok(Ok(TcpChunk(Part::from(buf.freeze()))))
547    }
548
549    async fn get_channel_address(
550        &mut self,
551        _cx: &Context<Self>,
552    ) -> Result<Option<ChannelAddr>, anyhow::Error> {
553        Ok(self.channel_addr.clone())
554    }
555
556    async fn register_transfer_remote(
557        &mut self,
558        cx: &Context<Self>,
559        buf_id: usize,
560        total_chunks: usize,
561        done: OncePortRef<Result<(), String>>,
562    ) -> Result<Result<usize, String>, anyhow::Error> {
563        let owner = self.owner.get().expect("TcpManagerActor owner not set");
564        let mem = match owner.request_local_memory(cx, buf_id).await {
565            Ok(Some(mem)) => mem,
566            Ok(None) => return Ok(Err(format!("buffer {buf_id} not found"))),
567            Err(e) => return Ok(Err(e.to_string())),
568        };
569        let transfer_id = self.register_transfer(mem, total_chunks, done);
570        Ok(Ok(transfer_id))
571    }
572
573    async fn execute_transfer_remote(
574        &mut self,
575        cx: &Context<Self>,
576        transfer_id: usize,
577        buf_id: usize,
578        chunk_size: usize,
579        dest_addr: ChannelAddr,
580    ) -> Result<Result<(), String>, anyhow::Error> {
581        let owner = self.owner.get().expect("TcpManagerActor owner not set");
582        let mem = match owner.request_local_memory(cx, buf_id).await {
583            Ok(Some(mem)) => mem,
584            Ok(None) => return Ok(Err(format!("buffer {buf_id} not found"))),
585            Err(e) => return Ok(Err(e.to_string())),
586        };
587        self.execute_transfer(cx, transfer_id, mem, chunk_size, dest_addr)?;
588        Ok(Ok(()))
589    }
590}
591
592#[async_trait]
593impl Handler<RegisterTransferLocal> for TcpManagerActor {
594    async fn handle(
595        &mut self,
596        cx: &Context<Self>,
597        message: RegisterTransferLocal,
598    ) -> Result<(), anyhow::Error> {
599        let transfer_id =
600            self.register_transfer(message.local_memory, message.total_chunks, message.done);
601        message.reply.post(cx, transfer_id);
602        Ok(())
603    }
604}
605
606#[async_trait]
607impl Handler<ExecuteTransferLocal> for TcpManagerActor {
608    async fn handle(
609        &mut self,
610        cx: &Context<Self>,
611        message: ExecuteTransferLocal,
612    ) -> Result<(), anyhow::Error> {
613        self.execute_transfer(
614            cx,
615            message.transfer_id,
616            message.local_memory,
617            message.chunk_size,
618            message.dest_addr,
619        )
620    }
621}
622
623#[async_trait]
624impl Handler<SendTransferResult> for TcpManagerActor {
625    async fn handle(
626        &mut self,
627        cx: &Context<Self>,
628        message: SendTransferResult,
629    ) -> Result<(), anyhow::Error> {
630        message.done.post(cx, message.result);
631        Ok(())
632    }
633}
634
635#[async_trait]
636impl Handler<TransferError> for TcpManagerActor {
637    async fn handle(
638        &mut self,
639        _cx: &Context<Self>,
640        message: TransferError,
641    ) -> Result<(), anyhow::Error> {
642        tracing::error!("fatal transfer error: {}", message.message);
643        Err(anyhow::anyhow!(message.message))
644    }
645}
646
647/// Wrapper around [`ActorHandle<TcpManagerActor>`] that moves the TCP
648/// data-plane (chunked reads/writes) off the actor loop while keeping
649/// buffer resolution serialized through actor messages.
650///
651/// Because submit logic now runs outside the actor loop, same-process
652/// messages no longer deadlock — the actor loop is free to handle
653/// `WriteChunk`/`ReadChunk` messages.
654#[derive(Debug, Clone)]
655pub struct TcpBackend(pub ActorHandle<TcpManagerActor>);
656
657impl std::ops::Deref for TcpBackend {
658    type Target = ActorHandle<TcpManagerActor>;
659    fn deref(&self) -> &Self::Target {
660        &self.0
661    }
662}
663
664impl TcpBackend {
665    /// Execute a parallel write: register the transfer on the remote
666    /// side, then execute locally to push chunks over direct channels.
667    async fn execute_parallel_write(
668        &self,
669        cx: &(impl context::Actor + Send + Sync),
670        op: &TcpOp,
671        chunk_size: usize,
672        deadline: Instant,
673    ) -> Result<()> {
674        // A write transfers the whole local buffer into the remote prefix; the
675        // remote buffer may be larger, so its tail is left untouched.
676        if op.local_memory.size() > op.remote_size {
677            anyhow::bail!(
678                "remote buffer size ({}) is smaller than local buffer size ({})",
679                op.remote_size,
680                op.local_memory.size(),
681            );
682        }
683        let size = op.local_memory.size();
684        let total_chunks = size.div_ceil(chunk_size);
685
686        let (done_handle, done_rx) = hyperactor::mailbox::open_once_port::<Result<(), String>>(cx);
687        let done_ref = done_handle.bind();
688
689        let remaining = deadline.saturating_duration_since(Instant::now());
690        let transfer_id = tokio_timeout(
691            remaining,
692            op.remote_tcp_manager.register_transfer_remote(
693                cx,
694                op.remote_buf_id,
695                total_chunks,
696                done_ref,
697            ),
698        )
699        .await
700        .map_err(|_| anyhow::anyhow!("register_transfer_remote timed out"))??
701        .map_err(|e| anyhow::anyhow!(e))?;
702
703        let dest_addr = tokio_timeout(
704            deadline.saturating_duration_since(Instant::now()),
705            op.remote_tcp_manager.get_channel_address(cx),
706        )
707        .await
708        .map_err(|_| anyhow::anyhow!("get_channel_address timed out"))??
709        .ok_or_else(|| anyhow::anyhow!("remote does not have parallel channels enabled"))?;
710
711        self.0.post(
712            cx,
713            ExecuteTransferLocal {
714                transfer_id,
715                local_memory: op.local_memory.clone(),
716                chunk_size,
717                dest_addr,
718            },
719        );
720
721        let remaining = deadline.saturating_duration_since(Instant::now());
722        let result = tokio_timeout(remaining, done_rx.recv())
723            .await
724            .map_err(|_| anyhow::anyhow!("parallel write timed out"))?
725            .map_err(|e| anyhow::anyhow!(e))?;
726        result.map_err(|e| anyhow::anyhow!(e))
727    }
728
729    /// Execute a parallel read: register the transfer locally, then
730    /// ask the remote side to push chunks to our channel.
731    async fn execute_parallel_read(
732        &self,
733        cx: &(impl context::Actor + Send + Sync),
734        op: &TcpOp,
735        chunk_size: usize,
736        deadline: Instant,
737    ) -> Result<()> {
738        // A read transfers the whole remote buffer into the local prefix; the
739        // local buffer may be larger, so its tail is left untouched.
740        if op.remote_size > op.local_memory.size() {
741            anyhow::bail!(
742                "remote buffer size ({}) is larger than local buffer size ({})",
743                op.remote_size,
744                op.local_memory.size(),
745            );
746        }
747        let size = op.remote_size;
748        let total_chunks = size.div_ceil(chunk_size);
749
750        let (done_handle, done_rx) = hyperactor::mailbox::open_once_port::<Result<(), String>>(cx);
751        let done_ref = done_handle.bind();
752
753        let (id_handle, id_rx) = hyperactor::mailbox::open_once_port::<usize>(cx);
754
755        self.0.post(
756            cx,
757            RegisterTransferLocal {
758                local_memory: op.local_memory.clone(),
759                total_chunks,
760                done: done_ref,
761                reply: id_handle,
762            },
763        );
764
765        let transfer_id = id_rx
766            .recv()
767            .await
768            .map_err(|e| anyhow::anyhow!("failed to receive transfer id: {e}"))?;
769
770        let my_channel_addr = self
771            .0
772            .get_channel_address(cx)
773            .await?
774            .ok_or_else(|| anyhow::anyhow!("local parallel channels not enabled"))?;
775
776        let remaining = deadline.saturating_duration_since(Instant::now());
777        tokio_timeout(
778            remaining,
779            op.remote_tcp_manager.execute_transfer_remote(
780                cx,
781                transfer_id,
782                op.remote_buf_id,
783                chunk_size,
784                my_channel_addr,
785            ),
786        )
787        .await
788        .map_err(|_| anyhow::anyhow!("execute_transfer_remote timed out"))??
789        .map_err(|e| anyhow::anyhow!(e))?;
790
791        let remaining = deadline.saturating_duration_since(Instant::now());
792        let result = tokio_timeout(remaining, done_rx.recv())
793            .await
794            .map_err(|_| anyhow::anyhow!("parallel read timed out"))?
795            .map_err(|e| anyhow::anyhow!(e))?;
796        result.map_err(|e| anyhow::anyhow!(e))
797    }
798
799    /// Execute a write operation: read local memory in chunks and write
800    /// them into the remote buffer via actor messages.
801    async fn execute_write(
802        &self,
803        cx: &(impl context::Actor + Send + Sync),
804        op: &TcpOp,
805        chunk_size: usize,
806        deadline: Instant,
807    ) -> Result<()> {
808        // A write transfers the whole local buffer into the remote prefix; the
809        // remote buffer may be larger, so its tail is left untouched.
810        if op.local_memory.size() > op.remote_size {
811            anyhow::bail!(
812                "remote buffer size ({}) is smaller than local buffer size ({})",
813                op.remote_size,
814                op.local_memory.size(),
815            );
816        }
817        let size = op.local_memory.size();
818        let mut offset = 0;
819
820        while offset < size {
821            let remaining = deadline.saturating_duration_since(Instant::now());
822            if remaining.is_zero() {
823                anyhow::bail!("tcp write timed out");
824            }
825
826            let len = std::cmp::min(chunk_size, size - offset);
827
828            let mut buf = vec![0u8; len];
829            // SAFETY: `op.local_memory` is the caller's buffer; that
830            // caller is responsible for excluding external writers
831            // while the `write_from_local` operation is in flight.
832            unsafe { op.local_memory.read_at(offset, &mut buf) }?;
833            let data = Part::from(Bytes::from(buf));
834
835            tokio_timeout(
836                remaining,
837                op.remote_tcp_manager
838                    .write_chunk(cx, op.remote_buf_id, offset, data),
839            )
840            .await
841            .map_err(|_| anyhow::anyhow!("tcp write chunk timed out"))??
842            .map_err(|e| anyhow::anyhow!(e))?;
843
844            offset += len;
845        }
846
847        Ok(())
848    }
849
850    /// Execute a read operation: request chunks from the remote buffer
851    /// and write them into local memory via actor messages.
852    async fn execute_read(
853        &self,
854        cx: &(impl context::Actor + Send + Sync),
855        op: &TcpOp,
856        chunk_size: usize,
857        deadline: Instant,
858    ) -> Result<()> {
859        // A read transfers the whole remote buffer into the local prefix; the
860        // local buffer may be larger, so its tail is left untouched.
861        if op.remote_size > op.local_memory.size() {
862            anyhow::bail!(
863                "remote buffer size ({}) is larger than local buffer size ({})",
864                op.remote_size,
865                op.local_memory.size(),
866            );
867        }
868        let size = op.remote_size;
869        let mut offset = 0;
870
871        while offset < size {
872            let remaining = deadline.saturating_duration_since(Instant::now());
873            if remaining.is_zero() {
874                anyhow::bail!("tcp read timed out");
875            }
876
877            let len = std::cmp::min(chunk_size, size - offset);
878
879            let chunk = tokio_timeout(
880                remaining,
881                op.remote_tcp_manager
882                    .read_chunk(cx, op.remote_buf_id, offset, len),
883            )
884            .await
885            .map_err(|_| anyhow::anyhow!("tcp read chunk timed out"))??
886            .map_err(|e| anyhow::anyhow!(e))?;
887            let data = chunk.0.into_bytes();
888
889            anyhow::ensure!(
890                data.len() == len,
891                "tcp read chunk size mismatch: expected {len}, got {}",
892                data.len()
893            );
894
895            // SAFETY: `op.local_memory` is the caller's buffer; that
896            // caller is responsible for excluding external readers
897            // and writers while the `read_into_local` operation is in
898            // flight.
899            unsafe { op.local_memory.write_at(offset, &data) }?;
900
901            offset += len;
902        }
903
904        Ok(())
905    }
906}
907
908#[async_trait]
909impl RdmaBackend for TcpBackend {
910    type RemoteBackendContext = ActorRef<TcpManagerActor>;
911    type TransportInfo = ();
912
913    /// TCP is available when fallback is enabled.
914    fn available() -> bool {
915        hyperactor_config::global::get(crate::config::RDMA_ALLOW_TCP_FALLBACK)
916    }
917
918    fn transport_level(&self) -> RdmaTransportLevel {
919        RdmaTransportLevel::Tcp
920    }
921
922    fn transport_info(&self) -> Option<Self::TransportInfo> {
923        None
924    }
925
926    async fn spawn(cx: &(impl context::Actor + Send + Sync), _config: &RdmaConfig) -> Result<Self> {
927        Ok(TcpBackend(cx.spawn(TcpManagerActor::new())))
928    }
929
930    /// TCP needs no per-buffer registration.
931    async fn register_remote_buffer(
932        &self,
933        _cx: &(impl context::Actor + Send + Sync),
934        _remote_buf_id: usize,
935        _local: KeepaliveLocalMemory,
936    ) -> Result<ActorRef<TcpManagerActor>> {
937        Ok(self.0.bind())
938    }
939
940    async fn release_buffer(
941        &self,
942        _cx: &(impl context::Actor + Send + Sync),
943        _remote_buf_id: usize,
944    ) -> Result<()> {
945        Ok(())
946    }
947
948    /// Submit a batch of RDMA operations over TCP.
949    ///
950    /// Each op is executed directly — sending chunked write/read messages
951    /// to the remote [`TcpManagerActor`].
952    async fn submit(
953        &self,
954        cx: &(impl context::Actor + Send + Sync),
955        ops: Vec<RdmaOp>,
956        timeout: Duration,
957    ) -> Result<()> {
958        let chunk_size =
959            hyperactor_config::global::get(crate::config::RDMA_MAX_CHUNK_SIZE_MB) * 1024 * 1024;
960        let parallelism =
961            hyperactor_config::global::get(crate::config::RDMA_TCP_FALLBACK_PARALLELISM);
962        let deadline = Instant::now() + timeout;
963
964        for op in ops {
965            let remaining = deadline.saturating_duration_since(Instant::now());
966            if remaining.is_zero() {
967                anyhow::bail!("tcp submit timed out");
968            }
969
970            let remote_tcp_manager = op
971                .remote
972                .resolve_tcp()
973                .expect("op routed to incompatible backend");
974            let tcp_op = TcpOp {
975                op_type: op.op_type,
976                remote_buf_id: op.remote.id,
977                remote_size: op.remote.size,
978                local_memory: op.local,
979                remote_tcp_manager,
980            };
981
982            if parallelism > 1 {
983                match tcp_op.op_type {
984                    RdmaOpType::WriteFromLocal => {
985                        self.execute_parallel_write(cx, &tcp_op, chunk_size, deadline)
986                            .await?;
987                    }
988                    RdmaOpType::ReadIntoLocal => {
989                        self.execute_parallel_read(cx, &tcp_op, chunk_size, deadline)
990                            .await?;
991                    }
992                }
993            } else {
994                match tcp_op.op_type {
995                    RdmaOpType::WriteFromLocal => {
996                        self.execute_write(cx, &tcp_op, chunk_size, deadline)
997                            .await?;
998                    }
999                    RdmaOpType::ReadIntoLocal => {
1000                        self.execute_read(cx, &tcp_op, chunk_size, deadline).await?;
1001                    }
1002                }
1003            }
1004        }
1005
1006        Ok(())
1007    }
1008}
1009
1010#[cfg(test)]
1011mod tests {
1012    use std::sync::Arc;
1013    use std::sync::atomic::AtomicUsize;
1014    use std::sync::atomic::Ordering;
1015    use std::time::Duration;
1016
1017    use hyperactor::ActorHandle;
1018    use hyperactor::Proc;
1019    use hyperactor::RemoteSpawn;
1020    use hyperactor::channel::ChannelAddr;
1021    use hyperactor_config::Flattrs;
1022
1023    use super::TcpBackend;
1024    use super::TcpManagerActor;
1025    use crate::RdmaManagerMessageClient;
1026    use crate::RdmaOp;
1027    use crate::RdmaOpType;
1028    use crate::backend::RdmaBackend;
1029    use crate::local_memory::KeepaliveLocalMemory;
1030    use crate::rdma_manager_actor::GetTcpActorRefClient;
1031    use crate::rdma_manager_actor::RdmaManagerActor;
1032
1033    static COUNTER: AtomicUsize = AtomicUsize::new(0);
1034
1035    struct TcpTestProcEnv {
1036        proc: Proc,
1037        rdma_handle: ActorHandle<RdmaManagerActor>,
1038        instance: hyperactor::Client,
1039        tcp_backend: TcpBackend,
1040        rdma_remote_buf: crate::RdmaRemoteBuffer,
1041        local_memory: KeepaliveLocalMemory,
1042    }
1043
1044    impl Drop for TcpTestProcEnv {
1045        fn drop(&mut self) {
1046            use crate::rdma_manager_actor::ReleaseBufferClient;
1047            // Release the buffer so the actor drops its local_memory
1048            // clone while the CUDA runtime is still alive.
1049            tokio::task::block_in_place(|| {
1050                tokio::runtime::Handle::current()
1051                    .block_on(
1052                        self.rdma_remote_buf
1053                            .owner
1054                            .release_buffer(&self.instance, self.rdma_remote_buf.id),
1055                    )
1056                    .expect("failed to release buffer in TcpTestProcEnv drop");
1057            });
1058        }
1059    }
1060
1061    impl TcpTestProcEnv {
1062        /// Create a standalone test environment with its own proc and rdma manager.
1063        async fn new(buffer_size: usize) -> anyhow::Result<Self> {
1064            let id = COUNTER.fetch_add(1, Ordering::Relaxed);
1065            let proc = Proc::direct(
1066                ChannelAddr::any(hyperactor::channel::ChannelTransport::Unix),
1067                format!("tcp_test_{id}"),
1068            )?;
1069            let instance = proc.client("client");
1070
1071            let rdma_actor = RdmaManagerActor::new(None, Flattrs::default()).await?;
1072            let rdma_handle = proc.spawn(rdma_actor);
1073
1074            let tcp_ref = rdma_handle.get_tcp_actor_ref(&instance).await?;
1075            let tcp_backend = TcpBackend(
1076                tcp_ref
1077                    .downcast_handle(&instance)
1078                    .ok_or_else(|| anyhow::anyhow!("tcp actor not local"))?,
1079            );
1080
1081            let (local_memory, rdma_remote_buf) =
1082                Self::alloc_cpu_buffer(&instance, &rdma_handle, buffer_size).await?;
1083
1084            Ok(Self {
1085                proc,
1086                rdma_handle,
1087                instance,
1088                tcp_backend,
1089                rdma_remote_buf,
1090                local_memory,
1091            })
1092        }
1093
1094        /// Create a buffer on an existing proc's rdma manager.
1095        async fn on_proc(
1096            proc: &Proc,
1097            rdma_handle: &ActorHandle<RdmaManagerActor>,
1098            tcp_backend: TcpBackend,
1099            buffer_size: usize,
1100        ) -> anyhow::Result<Self> {
1101            let id = COUNTER.fetch_add(1, Ordering::Relaxed);
1102            let instance = proc.client(&format!("client_{id}"));
1103
1104            let (local_memory, rdma_remote_buf) =
1105                Self::alloc_cpu_buffer(&instance, rdma_handle, buffer_size).await?;
1106
1107            Ok(Self {
1108                proc: proc.clone(),
1109                rdma_handle: rdma_handle.clone(),
1110                instance,
1111                tcp_backend,
1112                rdma_remote_buf,
1113                local_memory,
1114            })
1115        }
1116
1117        async fn alloc_cpu_buffer(
1118            instance: &hyperactor::Client,
1119            rdma_handle: &ActorHandle<RdmaManagerActor>,
1120            buffer_size: usize,
1121        ) -> anyhow::Result<(KeepaliveLocalMemory, crate::RdmaRemoteBuffer)> {
1122            let cpu_buf = vec![0u8; buffer_size].into_boxed_slice();
1123            let local_memory = KeepaliveLocalMemory::new(Arc::new(cpu_buf));
1124            let rdma_remote_buf = rdma_handle
1125                .request_buffer(instance, local_memory.clone())
1126                .await?;
1127            Ok((local_memory, rdma_remote_buf))
1128        }
1129    }
1130
1131    /// Two separate procs, one buffer each.
1132    async fn setup_tcp_env(buf_size: usize) -> anyhow::Result<Vec<TcpTestProcEnv>> {
1133        Ok(vec![
1134            TcpTestProcEnv::new(buf_size).await?,
1135            TcpTestProcEnv::new(buf_size).await?,
1136        ])
1137    }
1138
1139    /// Single proc, two buffers.
1140    async fn setup_same_proc_tcp_env(buf_size: usize) -> anyhow::Result<Vec<TcpTestProcEnv>> {
1141        let first = TcpTestProcEnv::new(buf_size).await?;
1142        let second = TcpTestProcEnv::on_proc(
1143            &first.proc,
1144            &first.rdma_handle,
1145            first.tcp_backend.clone(),
1146            buf_size,
1147        )
1148        .await?;
1149        Ok(vec![first, second])
1150    }
1151
1152    /// Two procs, two buffers each (4 total). For concurrent tests that
1153    /// need independent source/dest pairs.
1154    async fn setup_tcp_env_pairs(buf_size: usize) -> anyhow::Result<Vec<TcpTestProcEnv>> {
1155        let e0 = TcpTestProcEnv::new(buf_size).await?;
1156        let e1 = TcpTestProcEnv::new(buf_size).await?;
1157        let e2 =
1158            TcpTestProcEnv::on_proc(&e0.proc, &e0.rdma_handle, e0.tcp_backend.clone(), buf_size)
1159                .await?;
1160        let e3 =
1161            TcpTestProcEnv::on_proc(&e1.proc, &e1.rdma_handle, e1.tcp_backend.clone(), buf_size)
1162                .await?;
1163        Ok(vec![e0, e1, e2, e3])
1164    }
1165
1166    // --- Shared test helpers ---
1167
1168    /// Test-only wrapper around [`KeepaliveLocalMemory::write_at`].
1169    ///
1170    /// Every [`TcpTestProcEnv`] owns a distinct CPU buffer that no
1171    /// other thread accesses outside of explicit, serialized test
1172    /// operations, so the safety obligation of `write_at` is trivially
1173    /// satisfied across the whole module.
1174    fn test_write(mem: &KeepaliveLocalMemory, offset: usize, src: &[u8]) -> anyhow::Result<()> {
1175        // SAFETY: see the function-level comment.
1176        unsafe { mem.write_at(offset, src) }
1177    }
1178
1179    /// Test-only wrapper around [`KeepaliveLocalMemory::read_at`]. See
1180    /// [`test_write`] for the safety rationale.
1181    fn test_read(mem: &KeepaliveLocalMemory, offset: usize, dst: &mut [u8]) -> anyhow::Result<()> {
1182        // SAFETY: see the function-level comment.
1183        unsafe { mem.read_at(offset, dst) }
1184    }
1185
1186    /// Fill envs[0], write to envs[1], verify.
1187    async fn do_write_test(
1188        envs: &mut [TcpTestProcEnv],
1189        buf_size: usize,
1190        timeout: Duration,
1191    ) -> anyhow::Result<()> {
1192        let mut src = vec![0u8; buf_size];
1193        for (i, byte) in src.iter_mut().enumerate() {
1194            *byte = (i % 256) as u8;
1195        }
1196        test_write(&envs[0].local_memory, 0, &src)?;
1197
1198        let remote = envs[1].rdma_remote_buf.clone();
1199        let env = &mut envs[0];
1200        env.tcp_backend
1201            .submit(
1202                &env.instance,
1203                vec![RdmaOp {
1204                    op_type: RdmaOpType::WriteFromLocal,
1205                    local: env.local_memory.clone(),
1206                    remote,
1207                }],
1208                timeout,
1209            )
1210            .await?;
1211
1212        let mut dst = vec![0u8; buf_size];
1213        test_read(&envs[1].local_memory, 0, &mut dst)?;
1214        for (i, byte) in dst.iter().enumerate() {
1215            assert_eq!(*byte, (i % 256) as u8, "mismatch at offset {i} after write");
1216        }
1217        Ok(())
1218    }
1219
1220    /// Fill envs[1], read into envs[0], verify.
1221    async fn do_read_test(
1222        envs: &mut [TcpTestProcEnv],
1223        buf_size: usize,
1224        timeout: Duration,
1225    ) -> anyhow::Result<()> {
1226        let mut src = vec![0u8; buf_size];
1227        for (i, byte) in src.iter_mut().enumerate() {
1228            *byte = ((i * 7 + 3) % 256) as u8;
1229        }
1230        test_write(&envs[1].local_memory, 0, &src)?;
1231
1232        let remote = envs[1].rdma_remote_buf.clone();
1233        let env = &mut envs[0];
1234        env.tcp_backend
1235            .submit(
1236                &env.instance,
1237                vec![RdmaOp {
1238                    op_type: RdmaOpType::ReadIntoLocal,
1239                    local: env.local_memory.clone(),
1240                    remote,
1241                }],
1242                timeout,
1243            )
1244            .await?;
1245
1246        let mut dst = vec![0u8; buf_size];
1247        test_read(&envs[0].local_memory, 0, &mut dst)?;
1248        for (i, byte) in dst.iter().enumerate() {
1249            assert_eq!(
1250                *byte,
1251                ((i * 7 + 3) % 256) as u8,
1252                "mismatch at offset {i} after read"
1253            );
1254        }
1255        Ok(())
1256    }
1257
1258    /// Write, clear, read-back, verify round-trip.
1259    async fn do_round_trip_test(
1260        envs: &mut [TcpTestProcEnv],
1261        buf_size: usize,
1262        timeout: Duration,
1263    ) -> anyhow::Result<()> {
1264        let mut src = vec![0u8; buf_size];
1265        for (i, byte) in src.iter_mut().enumerate() {
1266            *byte = ((i * 13 + 5) % 256) as u8;
1267        }
1268        test_write(&envs[0].local_memory, 0, &src)?;
1269
1270        let remote = envs[1].rdma_remote_buf.clone();
1271        let env = &mut envs[0];
1272        env.tcp_backend
1273            .submit(
1274                &env.instance,
1275                vec![RdmaOp {
1276                    op_type: RdmaOpType::WriteFromLocal,
1277                    local: env.local_memory.clone(),
1278                    remote: remote.clone(),
1279                }],
1280                timeout,
1281            )
1282            .await?;
1283
1284        test_write(&envs[0].local_memory, 0, &vec![0u8; buf_size])?;
1285
1286        let env = &mut envs[0];
1287        env.tcp_backend
1288            .submit(
1289                &env.instance,
1290                vec![RdmaOp {
1291                    op_type: RdmaOpType::ReadIntoLocal,
1292                    local: env.local_memory.clone(),
1293                    remote,
1294                }],
1295                timeout,
1296            )
1297            .await?;
1298
1299        let mut dst = vec![0u8; buf_size];
1300        test_read(&envs[0].local_memory, 0, &mut dst)?;
1301        for (i, byte) in dst.iter().enumerate() {
1302            assert_eq!(
1303                *byte,
1304                ((i * 13 + 5) % 256) as u8,
1305                "mismatch at offset {i} after round-trip"
1306            );
1307        }
1308        Ok(())
1309    }
1310
1311    // --- Non-parallel two-proc tests ---
1312
1313    /// Write from local buffer 0 into remote buffer 1.
1314    #[timed_test::async_timed_test(timeout_secs = 30)]
1315    async fn test_tcp_write_from_local() -> anyhow::Result<()> {
1316        let config = hyperactor_config::global::lock();
1317        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1318
1319        let mut envs = setup_tcp_env(4096).await?;
1320        do_write_test(&mut envs, 4096, Duration::from_secs(30)).await
1321    }
1322
1323    /// Read from remote buffer 1 into local buffer 0.
1324    #[timed_test::async_timed_test(timeout_secs = 30)]
1325    async fn test_tcp_read_into_local() -> anyhow::Result<()> {
1326        let config = hyperactor_config::global::lock();
1327        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1328
1329        let mut envs = setup_tcp_env(2048).await?;
1330        do_read_test(&mut envs, 2048, Duration::from_secs(30)).await
1331    }
1332
1333    /// Write, clear, read-back, verify round-trip.
1334    #[timed_test::async_timed_test(timeout_secs = 30)]
1335    async fn test_tcp_write_then_read_back() -> anyhow::Result<()> {
1336        let config = hyperactor_config::global::lock();
1337        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1338
1339        let mut envs = setup_tcp_env(4096).await?;
1340        do_round_trip_test(&mut envs, 4096, Duration::from_secs(30)).await
1341    }
1342
1343    /// Multi-chunk write (1 MiB chunks, 1.5 MiB buffer).
1344    #[timed_test::async_timed_test(timeout_secs = 30)]
1345    async fn test_tcp_multi_chunk_write() -> anyhow::Result<()> {
1346        let config = hyperactor_config::global::lock();
1347        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1348        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1349
1350        let buf_size = 3 * 1024 * 512;
1351        let mut envs = setup_tcp_env(buf_size).await?;
1352
1353        let mut src = vec![0u8; buf_size];
1354        for (i, byte) in src.iter_mut().enumerate() {
1355            *byte = (i % 251) as u8;
1356        }
1357        test_write(&envs[0].local_memory, 0, &src)?;
1358
1359        let remote = envs[1].rdma_remote_buf.clone();
1360        let env = &mut envs[0];
1361        env.tcp_backend
1362            .submit(
1363                &env.instance,
1364                vec![RdmaOp {
1365                    op_type: RdmaOpType::WriteFromLocal,
1366                    local: env.local_memory.clone(),
1367                    remote,
1368                }],
1369                Duration::from_secs(30),
1370            )
1371            .await?;
1372
1373        let mut dst = vec![0u8; buf_size];
1374        test_read(&envs[1].local_memory, 0, &mut dst)?;
1375        for (i, byte) in dst.iter().enumerate() {
1376            assert_eq!(*byte, (i % 251) as u8, "mismatch at offset {i}");
1377        }
1378
1379        Ok(())
1380    }
1381
1382    /// Multi-chunk read.
1383    #[timed_test::async_timed_test(timeout_secs = 30)]
1384    async fn test_tcp_multi_chunk_read() -> anyhow::Result<()> {
1385        let config = hyperactor_config::global::lock();
1386        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1387        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1388
1389        let buf_size = 3 * 1024 * 512;
1390        let mut envs = setup_tcp_env(buf_size).await?;
1391
1392        let mut src = vec![0u8; buf_size];
1393        for (i, byte) in src.iter_mut().enumerate() {
1394            *byte = ((i * 3 + 17) % 256) as u8;
1395        }
1396        test_write(&envs[1].local_memory, 0, &src)?;
1397
1398        let remote = envs[1].rdma_remote_buf.clone();
1399        let env = &mut envs[0];
1400        env.tcp_backend
1401            .submit(
1402                &env.instance,
1403                vec![RdmaOp {
1404                    op_type: RdmaOpType::ReadIntoLocal,
1405                    local: env.local_memory.clone(),
1406                    remote,
1407                }],
1408                Duration::from_secs(30),
1409            )
1410            .await?;
1411
1412        let mut dst = vec![0u8; buf_size];
1413        test_read(&envs[0].local_memory, 0, &mut dst)?;
1414        for (i, byte) in dst.iter().enumerate() {
1415            assert_eq!(*byte, ((i * 3 + 17) % 256) as u8, "mismatch at offset {i}");
1416        }
1417
1418        Ok(())
1419    }
1420
1421    /// Multi-chunk write-then-read round-trip.
1422    #[timed_test::async_timed_test(timeout_secs = 30)]
1423    async fn test_tcp_multi_chunk_round_trip() -> anyhow::Result<()> {
1424        let config = hyperactor_config::global::lock();
1425        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1426        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1427
1428        let buf_size = 5 * 1024 * 512;
1429        let mut envs = setup_tcp_env(buf_size).await?;
1430
1431        let mut src = vec![0u8; buf_size];
1432        for (i, byte) in src.iter_mut().enumerate() {
1433            *byte = ((i * 41 + 7) % 256) as u8;
1434        }
1435        test_write(&envs[0].local_memory, 0, &src)?;
1436
1437        let remote = envs[1].rdma_remote_buf.clone();
1438        let env = &mut envs[0];
1439        env.tcp_backend
1440            .submit(
1441                &env.instance,
1442                vec![RdmaOp {
1443                    op_type: RdmaOpType::WriteFromLocal,
1444                    local: env.local_memory.clone(),
1445                    remote: remote.clone(),
1446                }],
1447                Duration::from_secs(30),
1448            )
1449            .await?;
1450
1451        test_write(&envs[0].local_memory, 0, &vec![0u8; buf_size])?;
1452
1453        let env = &mut envs[0];
1454        env.tcp_backend
1455            .submit(
1456                &env.instance,
1457                vec![RdmaOp {
1458                    op_type: RdmaOpType::ReadIntoLocal,
1459                    local: env.local_memory.clone(),
1460                    remote,
1461                }],
1462                Duration::from_secs(30),
1463            )
1464            .await?;
1465
1466        let mut dst = vec![0u8; buf_size];
1467        test_read(&envs[0].local_memory, 0, &mut dst)?;
1468        for (i, byte) in dst.iter().enumerate() {
1469            assert_eq!(
1470                *byte,
1471                ((i * 41 + 7) % 256) as u8,
1472                "mismatch at offset {i} after multi-chunk round-trip"
1473            );
1474        }
1475
1476        Ok(())
1477    }
1478
1479    /// resolve_tcp finds the Tcp backend context in a buffer.
1480    #[timed_test::async_timed_test(timeout_secs = 30)]
1481    async fn test_tcp_resolve_tcp() -> anyhow::Result<()> {
1482        let config = hyperactor_config::global::lock();
1483        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1484
1485        let envs = setup_tcp_env(64).await?;
1486
1487        for (i, env) in envs.iter().enumerate() {
1488            let tcp_ref = env
1489                .rdma_remote_buf
1490                .resolve_tcp()
1491                .unwrap_or_else(|| panic!("tcp backend not found for env {i}"));
1492            let expected: hyperactor::ActorRef<TcpManagerActor> = env.tcp_backend.bind();
1493            assert_eq!(tcp_ref.actor_addr(), expected.actor_addr());
1494        }
1495
1496        Ok(())
1497    }
1498
1499    /// Write to a released buffer returns an error without crashing.
1500    #[timed_test::async_timed_test(timeout_secs = 30)]
1501    async fn test_tcp_write_to_released_buffer() -> anyhow::Result<()> {
1502        let config = hyperactor_config::global::lock();
1503        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1504
1505        let buf_size = 64;
1506        let mut envs = setup_tcp_env(buf_size).await?;
1507
1508        let mut src = vec![0u8; buf_size];
1509        for (i, byte) in src.iter_mut().enumerate() {
1510            *byte = (i % 256) as u8;
1511        }
1512        test_write(&envs[0].local_memory, 0, &src)?;
1513
1514        // Normal write should succeed.
1515        let remote = envs[1].rdma_remote_buf.clone();
1516        let env = &mut envs[0];
1517        env.tcp_backend
1518            .submit(
1519                &env.instance,
1520                vec![RdmaOp {
1521                    op_type: RdmaOpType::WriteFromLocal,
1522                    local: env.local_memory.clone(),
1523                    remote: remote.clone(),
1524                }],
1525                Duration::from_secs(10),
1526            )
1527            .await?;
1528
1529        // Release the remote buffer.
1530        use crate::rdma_manager_actor::ReleaseBufferClient;
1531        let owner_ref = envs[1].rdma_remote_buf.owner.clone();
1532        owner_ref
1533            .release_buffer(&envs[0].instance, envs[1].rdma_remote_buf.id)
1534            .await?;
1535
1536        // Writing to the released buffer should fail.
1537        let env = &mut envs[0];
1538        let result = env
1539            .tcp_backend
1540            .submit(
1541                &env.instance,
1542                vec![RdmaOp {
1543                    op_type: RdmaOpType::WriteFromLocal,
1544                    local: env.local_memory.clone(),
1545                    remote: remote.clone(),
1546                }],
1547                Duration::from_secs(10),
1548            )
1549            .await;
1550        assert!(result.is_err(), "expected error writing to released buffer");
1551
1552        Ok(())
1553    }
1554
1555    /// Read from a released buffer returns an error without crashing.
1556    #[timed_test::async_timed_test(timeout_secs = 30)]
1557    async fn test_tcp_read_from_released_buffer() -> anyhow::Result<()> {
1558        let config = hyperactor_config::global::lock();
1559        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1560
1561        let buf_size = 64;
1562        let mut envs = setup_tcp_env(buf_size).await?;
1563
1564        // Release the remote buffer.
1565        use crate::rdma_manager_actor::ReleaseBufferClient;
1566        let owner_ref = envs[1].rdma_remote_buf.owner.clone();
1567        owner_ref
1568            .release_buffer(&envs[0].instance, envs[1].rdma_remote_buf.id)
1569            .await?;
1570
1571        // Reading from the released buffer should fail.
1572        let remote = envs[1].rdma_remote_buf.clone();
1573        let env = &mut envs[0];
1574        let result = env
1575            .tcp_backend
1576            .submit(
1577                &env.instance,
1578                vec![RdmaOp {
1579                    op_type: RdmaOpType::ReadIntoLocal,
1580                    local: env.local_memory.clone(),
1581                    remote,
1582                }],
1583                Duration::from_secs(10),
1584            )
1585            .await;
1586        assert!(
1587            result.is_err(),
1588            "expected error reading from released buffer"
1589        );
1590
1591        Ok(())
1592    }
1593
1594    // --- Non-parallel same-proc tests ---
1595
1596    /// Same-process write.
1597    #[timed_test::async_timed_test(timeout_secs = 30)]
1598    async fn test_tcp_same_process_write() -> anyhow::Result<()> {
1599        let config = hyperactor_config::global::lock();
1600        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1601
1602        let mut envs = setup_same_proc_tcp_env(4096).await?;
1603        do_write_test(&mut envs, 4096, Duration::from_secs(10)).await
1604    }
1605
1606    /// Same-process read.
1607    #[timed_test::async_timed_test(timeout_secs = 30)]
1608    async fn test_tcp_same_process_read() -> anyhow::Result<()> {
1609        let config = hyperactor_config::global::lock();
1610        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1611
1612        let mut envs = setup_same_proc_tcp_env(2048).await?;
1613        do_read_test(&mut envs, 2048, Duration::from_secs(10)).await
1614    }
1615
1616    /// Same-process write-then-read round-trip.
1617    #[timed_test::async_timed_test(timeout_secs = 30)]
1618    async fn test_tcp_same_process_round_trip() -> anyhow::Result<()> {
1619        let config = hyperactor_config::global::lock();
1620        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1621
1622        let mut envs = setup_same_proc_tcp_env(4096).await?;
1623        do_round_trip_test(&mut envs, 4096, Duration::from_secs(10)).await
1624    }
1625
1626    // --- Multi-GPU TCP fallback tests ---
1627
1628    use crate::backend::cuda_test_utils::CudaAllocator;
1629    use crate::backend::cuda_test_utils::cuda_device_count;
1630
1631    impl TcpTestProcEnv {
1632        /// Create a test environment backed by CUDA device memory.
1633        async fn new_gpu(device: i32, buffer_size: usize) -> anyhow::Result<Self> {
1634            let id = COUNTER.fetch_add(1, Ordering::Relaxed);
1635            let proc = Proc::direct(
1636                ChannelAddr::any(hyperactor::channel::ChannelTransport::Unix),
1637                format!("tcp_gpu_test_{id}"),
1638            )?;
1639            let instance = proc.client("client");
1640
1641            let rdma_actor = RdmaManagerActor::new(None, Flattrs::default()).await?;
1642            let rdma_handle = proc.spawn(rdma_actor);
1643
1644            let tcp_ref = rdma_handle.get_tcp_actor_ref(&instance).await?;
1645            let tcp_backend = TcpBackend(
1646                tcp_ref
1647                    .downcast_handle(&instance)
1648                    .ok_or_else(|| anyhow::anyhow!("tcp actor not local"))?,
1649            );
1650
1651            let alloc = CudaAllocator::get().allocate(device, buffer_size, buffer_size);
1652            let local_memory = KeepaliveLocalMemory::new(Arc::new(alloc));
1653            let rdma_remote_buf = rdma_handle
1654                .request_buffer(&instance, local_memory.clone())
1655                .await?;
1656
1657            Ok(Self {
1658                proc,
1659                rdma_handle,
1660                instance,
1661                tcp_backend,
1662                rdma_remote_buf,
1663                local_memory,
1664            })
1665        }
1666    }
1667
1668    /// TCP write from GPU on cuda:0 to GPU on cuda:1.
1669    #[timed_test::async_timed_test(timeout_secs = 60)]
1670    async fn test_tcp_write_multi_gpu() -> anyhow::Result<()> {
1671        if cuda_device_count() < 2 {
1672            println!("Skipping: need at least 2 CUDA devices");
1673            return Ok(());
1674        }
1675
1676        let config = hyperactor_config::global::lock();
1677        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1678
1679        let buf_size = 2 * 1024 * 1024;
1680        let mut envs = vec![
1681            TcpTestProcEnv::new_gpu(0, buf_size).await?,
1682            TcpTestProcEnv::new_gpu(1, buf_size).await?,
1683        ];
1684        do_write_test(&mut envs, buf_size, Duration::from_secs(30)).await
1685    }
1686
1687    /// TCP read from GPU on cuda:1 into GPU on cuda:0.
1688    #[timed_test::async_timed_test(timeout_secs = 60)]
1689    async fn test_tcp_read_multi_gpu() -> anyhow::Result<()> {
1690        if cuda_device_count() < 2 {
1691            println!("Skipping: need at least 2 CUDA devices");
1692            return Ok(());
1693        }
1694
1695        let config = hyperactor_config::global::lock();
1696        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1697
1698        let buf_size = 2 * 1024 * 1024;
1699        let mut envs = vec![
1700            TcpTestProcEnv::new_gpu(0, buf_size).await?,
1701            TcpTestProcEnv::new_gpu(1, buf_size).await?,
1702        ];
1703        do_read_test(&mut envs, buf_size, Duration::from_secs(30)).await
1704    }
1705
1706    /// TCP write-then-read round-trip between cuda:0 and cuda:1.
1707    #[timed_test::async_timed_test(timeout_secs = 60)]
1708    async fn test_tcp_round_trip_multi_gpu() -> anyhow::Result<()> {
1709        if cuda_device_count() < 2 {
1710            println!("Skipping: need at least 2 CUDA devices");
1711            return Ok(());
1712        }
1713
1714        let config = hyperactor_config::global::lock();
1715        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1716
1717        let buf_size = 2 * 1024 * 1024;
1718        let mut envs = vec![
1719            TcpTestProcEnv::new_gpu(0, buf_size).await?,
1720            TcpTestProcEnv::new_gpu(1, buf_size).await?,
1721        ];
1722        do_round_trip_test(&mut envs, buf_size, Duration::from_secs(30)).await
1723    }
1724
1725    /// Stopping the RdmaManagerActor with parallelism enabled cleanly
1726    /// shuts down the TcpManagerActor's receive loop without hanging.
1727    #[timed_test::async_timed_test(timeout_secs = 30)]
1728    async fn test_tcp_parallel_clean_shutdown() -> anyhow::Result<()> {
1729        let config = hyperactor_config::global::lock();
1730        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1731        let _par_guard = config.override_key(crate::config::RDMA_TCP_FALLBACK_PARALLELISM, 2);
1732        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1733
1734        let buf_size = 3 * 1024 * 1024;
1735        let mut envs = setup_tcp_env(buf_size).await?;
1736
1737        // Do a transfer so the receive loop and outbound connections are live.
1738        let mut src = vec![0u8; buf_size];
1739        for (i, byte) in src.iter_mut().enumerate() {
1740            *byte = (i % 256) as u8;
1741        }
1742        test_write(&envs[0].local_memory, 0, &src)?;
1743        let remote = envs[1].rdma_remote_buf.clone();
1744        let env = &mut envs[0];
1745        env.tcp_backend
1746            .submit(
1747                &env.instance,
1748                vec![RdmaOp {
1749                    op_type: RdmaOpType::WriteFromLocal,
1750                    local: env.local_memory.clone(),
1751                    remote: remote.clone(),
1752                }],
1753                Duration::from_secs(30),
1754            )
1755            .await?;
1756
1757        // Stop the RdmaManagerActor, which cascades to TcpManagerActor.
1758        // The test timeout ensures we detect hangs in the cleanup path.
1759        envs[0].rdma_handle.drain_and_stop("test")?;
1760        envs[0].rdma_handle.clone().await;
1761        envs[1].rdma_handle.drain_and_stop("test")?;
1762        envs[1].rdma_handle.clone().await;
1763
1764        Ok(())
1765    }
1766
1767    // --- Parallel transfer tests ---
1768
1769    /// Parallel write via direct channels.
1770    #[timed_test::async_timed_test(timeout_secs = 30)]
1771    async fn test_tcp_parallel_write() -> anyhow::Result<()> {
1772        let config = hyperactor_config::global::lock();
1773        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1774        let _par_guard = config.override_key(crate::config::RDMA_TCP_FALLBACK_PARALLELISM, 2);
1775        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1776
1777        // 3 MiB, 3 chunks spread across 2 workers.
1778        let buf_size = 3 * 1024 * 1024;
1779        let mut envs = setup_tcp_env(buf_size).await?;
1780        do_write_test(&mut envs, buf_size, Duration::from_secs(30)).await
1781    }
1782
1783    /// Parallel read via direct channels.
1784    #[timed_test::async_timed_test(timeout_secs = 30)]
1785    async fn test_tcp_parallel_read() -> anyhow::Result<()> {
1786        let config = hyperactor_config::global::lock();
1787        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1788        let _par_guard = config.override_key(crate::config::RDMA_TCP_FALLBACK_PARALLELISM, 2);
1789        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1790
1791        let buf_size = 3 * 1024 * 1024;
1792        let mut envs = setup_tcp_env(buf_size).await?;
1793        do_read_test(&mut envs, buf_size, Duration::from_secs(30)).await
1794    }
1795
1796    /// Parallel write-then-read round-trip.
1797    #[timed_test::async_timed_test(timeout_secs = 30)]
1798    async fn test_tcp_parallel_round_trip() -> anyhow::Result<()> {
1799        let config = hyperactor_config::global::lock();
1800        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1801        let _par_guard = config.override_key(crate::config::RDMA_TCP_FALLBACK_PARALLELISM, 2);
1802        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1803
1804        let buf_size = 3 * 1024 * 1024;
1805        let mut envs = setup_tcp_env(buf_size).await?;
1806        do_round_trip_test(&mut envs, buf_size, Duration::from_secs(30)).await
1807    }
1808
1809    /// Same-process parallel write.
1810    #[timed_test::async_timed_test(timeout_secs = 30)]
1811    async fn test_tcp_parallel_same_process_write() -> anyhow::Result<()> {
1812        let config = hyperactor_config::global::lock();
1813        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1814        let _par_guard = config.override_key(crate::config::RDMA_TCP_FALLBACK_PARALLELISM, 2);
1815        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1816
1817        let buf_size = 3 * 1024 * 1024;
1818        let mut envs = setup_same_proc_tcp_env(buf_size).await?;
1819        do_write_test(&mut envs, buf_size, Duration::from_secs(10)).await
1820    }
1821
1822    /// Same-process parallel read.
1823    #[timed_test::async_timed_test(timeout_secs = 30)]
1824    async fn test_tcp_parallel_same_process_read() -> anyhow::Result<()> {
1825        let config = hyperactor_config::global::lock();
1826        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1827        let _par_guard = config.override_key(crate::config::RDMA_TCP_FALLBACK_PARALLELISM, 2);
1828        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1829
1830        let buf_size = 3 * 1024 * 1024;
1831        let mut envs = setup_same_proc_tcp_env(buf_size).await?;
1832        do_read_test(&mut envs, buf_size, Duration::from_secs(10)).await
1833    }
1834
1835    // --- Concurrent parallel tests (4 envs, 2 independent pairs) ---
1836
1837    /// Two concurrent parallel writes to independent buffer pairs.
1838    #[timed_test::async_timed_test(timeout_secs = 30)]
1839    async fn test_tcp_parallel_concurrent_writes() -> anyhow::Result<()> {
1840        let config = hyperactor_config::global::lock();
1841        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1842        let _par_guard = config.override_key(crate::config::RDMA_TCP_FALLBACK_PARALLELISM, 2);
1843        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1844
1845        let buf_size = 3 * 1024 * 1024;
1846        let envs = setup_tcp_env_pairs(buf_size).await?;
1847
1848        // Fill source buffers with distinct patterns.
1849        let mut src0 = vec![0u8; buf_size];
1850        for (i, byte) in src0.iter_mut().enumerate() {
1851            *byte = (i % 256) as u8;
1852        }
1853        test_write(&envs[0].local_memory, 0, &src0)?;
1854        let mut src2 = vec![0u8; buf_size];
1855        for (i, byte) in src2.iter_mut().enumerate() {
1856            *byte = ((i * 3 + 7) % 256) as u8;
1857        }
1858        test_write(&envs[2].local_memory, 0, &src2)?;
1859
1860        // Pair 1: envs[0] -> envs[1], Pair 2: envs[2] -> envs[3].
1861        let remote_1 = envs[1].rdma_remote_buf.clone();
1862        let remote_3 = envs[3].rdma_remote_buf.clone();
1863        let h0 = envs[0].tcp_backend.clone();
1864        let h2 = envs[2].tcp_backend.clone();
1865        let inst_0 = &envs[0].instance;
1866        let inst_2 = &envs[2].instance;
1867        let mem_0 = envs[0].local_memory.clone();
1868        let mem_2 = envs[2].local_memory.clone();
1869        let (r1, r2) = tokio::join!(
1870            h0.submit(
1871                inst_0,
1872                vec![RdmaOp {
1873                    op_type: RdmaOpType::WriteFromLocal,
1874                    local: mem_0,
1875                    remote: remote_1,
1876                }],
1877                Duration::from_secs(30),
1878            ),
1879            h2.submit(
1880                inst_2,
1881                vec![RdmaOp {
1882                    op_type: RdmaOpType::WriteFromLocal,
1883                    local: mem_2,
1884                    remote: remote_3,
1885                }],
1886                Duration::from_secs(30),
1887            ),
1888        );
1889        r1?;
1890        r2?;
1891
1892        let mut dst1 = vec![0u8; buf_size];
1893        test_read(&envs[1].local_memory, 0, &mut dst1)?;
1894        for (i, byte) in dst1.iter().enumerate() {
1895            assert_eq!(*byte, (i % 256) as u8, "pair 1 mismatch at offset {i}");
1896        }
1897        let mut dst3 = vec![0u8; buf_size];
1898        test_read(&envs[3].local_memory, 0, &mut dst3)?;
1899        for (i, byte) in dst3.iter().enumerate() {
1900            assert_eq!(
1901                *byte,
1902                ((i * 3 + 7) % 256) as u8,
1903                "pair 2 mismatch at offset {i}"
1904            );
1905        }
1906
1907        Ok(())
1908    }
1909
1910    /// Two concurrent parallel reads from independent buffer pairs.
1911    #[timed_test::async_timed_test(timeout_secs = 30)]
1912    async fn test_tcp_parallel_concurrent_reads() -> anyhow::Result<()> {
1913        let config = hyperactor_config::global::lock();
1914        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1915        let _par_guard = config.override_key(crate::config::RDMA_TCP_FALLBACK_PARALLELISM, 2);
1916        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1917
1918        let buf_size = 3 * 1024 * 1024;
1919        let envs = setup_tcp_env_pairs(buf_size).await?;
1920
1921        // Fill remote buffers with distinct patterns.
1922        let mut src1 = vec![0u8; buf_size];
1923        for (i, byte) in src1.iter_mut().enumerate() {
1924            *byte = ((i * 11 + 3) % 256) as u8;
1925        }
1926        test_write(&envs[1].local_memory, 0, &src1)?;
1927        let mut src3 = vec![0u8; buf_size];
1928        for (i, byte) in src3.iter_mut().enumerate() {
1929            *byte = ((i * 5 + 13) % 256) as u8;
1930        }
1931        test_write(&envs[3].local_memory, 0, &src3)?;
1932
1933        // Pair 1: envs[0] <- envs[1], Pair 2: envs[2] <- envs[3].
1934        let remote_1 = envs[1].rdma_remote_buf.clone();
1935        let remote_3 = envs[3].rdma_remote_buf.clone();
1936        let h0 = envs[0].tcp_backend.clone();
1937        let h2 = envs[2].tcp_backend.clone();
1938        let inst_0 = &envs[0].instance;
1939        let inst_2 = &envs[2].instance;
1940        let mem_0 = envs[0].local_memory.clone();
1941        let mem_2 = envs[2].local_memory.clone();
1942        let (r1, r2) = tokio::join!(
1943            h0.submit(
1944                inst_0,
1945                vec![RdmaOp {
1946                    op_type: RdmaOpType::ReadIntoLocal,
1947                    local: mem_0,
1948                    remote: remote_1,
1949                }],
1950                Duration::from_secs(30),
1951            ),
1952            h2.submit(
1953                inst_2,
1954                vec![RdmaOp {
1955                    op_type: RdmaOpType::ReadIntoLocal,
1956                    local: mem_2,
1957                    remote: remote_3,
1958                }],
1959                Duration::from_secs(30),
1960            ),
1961        );
1962        r1?;
1963        r2?;
1964
1965        let mut dst0 = vec![0u8; buf_size];
1966        test_read(&envs[0].local_memory, 0, &mut dst0)?;
1967        for (i, byte) in dst0.iter().enumerate() {
1968            assert_eq!(
1969                *byte,
1970                ((i * 11 + 3) % 256) as u8,
1971                "pair 1 mismatch at offset {i}"
1972            );
1973        }
1974        let mut dst2 = vec![0u8; buf_size];
1975        test_read(&envs[2].local_memory, 0, &mut dst2)?;
1976        for (i, byte) in dst2.iter().enumerate() {
1977            assert_eq!(
1978                *byte,
1979                ((i * 5 + 13) % 256) as u8,
1980                "pair 2 mismatch at offset {i}"
1981            );
1982        }
1983
1984        Ok(())
1985    }
1986
1987    /// Concurrent parallel write and read on independent buffer pairs.
1988    #[timed_test::async_timed_test(timeout_secs = 30)]
1989    async fn test_tcp_parallel_concurrent_write_and_read() -> anyhow::Result<()> {
1990        let config = hyperactor_config::global::lock();
1991        let _guard = config.override_key(crate::config::RDMA_ALLOW_TCP_FALLBACK, true);
1992        let _par_guard = config.override_key(crate::config::RDMA_TCP_FALLBACK_PARALLELISM, 2);
1993        let _chunk_guard = config.override_key(crate::config::RDMA_MAX_CHUNK_SIZE_MB, 1);
1994
1995        let buf_size = 3 * 1024 * 1024;
1996        let envs = setup_tcp_env_pairs(buf_size).await?;
1997
1998        // Fill source buffers.
1999        let mut src0 = vec![0u8; buf_size];
2000        for (i, byte) in src0.iter_mut().enumerate() {
2001            *byte = (i % 256) as u8;
2002        }
2003        test_write(&envs[0].local_memory, 0, &src0)?;
2004        let mut src3 = vec![0u8; buf_size];
2005        for (i, byte) in src3.iter_mut().enumerate() {
2006            *byte = ((i * 7 + 13) % 256) as u8;
2007        }
2008        test_write(&envs[3].local_memory, 0, &src3)?;
2009
2010        // Write envs[0] -> envs[1], read envs[2] <- envs[3] concurrently.
2011        let remote_1 = envs[1].rdma_remote_buf.clone();
2012        let remote_3 = envs[3].rdma_remote_buf.clone();
2013        let h0 = envs[0].tcp_backend.clone();
2014        let h2 = envs[2].tcp_backend.clone();
2015        let inst_0 = &envs[0].instance;
2016        let inst_2 = &envs[2].instance;
2017        let mem_0 = envs[0].local_memory.clone();
2018        let mem_2 = envs[2].local_memory.clone();
2019        let (write_result, read_result) = tokio::join!(
2020            h0.submit(
2021                inst_0,
2022                vec![RdmaOp {
2023                    op_type: RdmaOpType::WriteFromLocal,
2024                    local: mem_0,
2025                    remote: remote_1,
2026                }],
2027                Duration::from_secs(30),
2028            ),
2029            h2.submit(
2030                inst_2,
2031                vec![RdmaOp {
2032                    op_type: RdmaOpType::ReadIntoLocal,
2033                    local: mem_2,
2034                    remote: remote_3,
2035                }],
2036                Duration::from_secs(30),
2037            ),
2038        );
2039        write_result?;
2040        read_result?;
2041
2042        let mut dst1 = vec![0u8; buf_size];
2043        test_read(&envs[1].local_memory, 0, &mut dst1)?;
2044        for (i, byte) in dst1.iter().enumerate() {
2045            assert_eq!(*byte, (i % 256) as u8, "write mismatch at offset {i}");
2046        }
2047        let mut dst2 = vec![0u8; buf_size];
2048        test_read(&envs[2].local_memory, 0, &mut dst2)?;
2049        for (i, byte) in dst2.iter().enumerate() {
2050            assert_eq!(
2051                *byte,
2052                ((i * 7 + 13) % 256) as u8,
2053                "read mismatch at offset {i}"
2054            );
2055        }
2056
2057        Ok(())
2058    }
2059}