1use 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#[derive(Debug, Clone, Serialize, Deserialize, Named)]
70pub struct TcpChunk(Part);
71wirevalue::register_type!(TcpChunk);
72
73#[derive(Debug, Clone, Serialize, Deserialize, Named)]
75struct TcpDataChunk {
76 transfer_id: usize,
78 offset: usize,
80 data: Part,
81}
82wirevalue::register_type!(TcpDataChunk);
83
84#[derive(Debug)]
89struct TransferState {
90 local_memory: KeepaliveLocalMemory,
92
93 chunks_received: usize,
95
96 total_chunks: usize,
98
99 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#[derive(Debug, Serialize, Deserialize, Named)]
129struct SendTransferResult {
130 done: OncePortRef<Result<(), String>>,
131 result: Result<(), String>,
132}
133
134#[derive(Debug, Serialize, Deserialize, Named)]
139struct TransferError {
140 message: String,
141}
142
143#[derive(Debug)]
146struct RegisterTransferLocal {
147 local_memory: KeepaliveLocalMemory,
148 total_chunks: usize,
149 done: OncePortRef<Result<(), String>>,
150 reply: OncePortHandle<usize>,
152}
153
154#[derive(Debug)]
157struct ExecuteTransferLocal {
158 transfer_id: usize,
159 local_memory: KeepaliveLocalMemory,
160 chunk_size: usize,
161 dest_addr: ChannelAddr,
162}
163
164#[derive(Handler, HandleClient, RefClient, Debug, Serialize, Deserialize, Named)]
169enum TcpManagerMessage {
170 WriteChunk {
172 buf_id: usize,
173 offset: usize,
174 data: Part,
175 #[reply]
176 reply: OncePortRef<Result<(), String>>,
177 },
178 ReadChunk {
180 buf_id: usize,
181 offset: usize,
182 size: usize,
183 #[reply]
184 reply: OncePortRef<Result<TcpChunk, String>>,
185 },
186 GetChannelAddress {
189 #[reply]
190 reply: OncePortRef<Option<ChannelAddr>>,
191 },
192 RegisterTransferRemote {
195 buf_id: usize,
196 total_chunks: usize,
197 done: OncePortRef<Result<(), String>>,
198 #[reply]
199 reply: OncePortRef<Result<usize, String>>,
200 },
201 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#[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 channel_addr: Option<ChannelAddr>,
229 outbound: HashMap<ChannelAddr, Vec<Arc<ChannelTx<TcpDataChunk>>>>,
231 cancel: CancellationToken,
233 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 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 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 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 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 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#[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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 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 fn test_write(mem: &KeepaliveLocalMemory, offset: usize, src: &[u8]) -> anyhow::Result<()> {
1175 unsafe { mem.write_at(offset, src) }
1177 }
1178
1179 fn test_read(mem: &KeepaliveLocalMemory, offset: usize, dst: &mut [u8]) -> anyhow::Result<()> {
1182 unsafe { mem.read_at(offset, dst) }
1184 }
1185
1186 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 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 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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 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 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 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 #[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 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 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 #[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 #[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 #[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 use crate::backend::cuda_test_utils::CudaAllocator;
1629 use crate::backend::cuda_test_utils::cuda_device_count;
1630
1631 impl TcpTestProcEnv {
1632 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 #[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 #[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 #[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 #[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 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 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 #[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 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 #[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 #[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 #[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 #[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 #[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 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 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 #[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 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 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 #[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 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 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}