1use crate::casting::CAST_ACTOR_MESH_ID;
10use crate::comm::multicast::CAST_ORIGINATING_SENDER;
11use crate::comm::multicast::CastEnvelope;
12use crate::comm::multicast::CastMessageV1;
13use crate::comm::multicast::ForwardMessageV1;
14use crate::mesh_id::ActorMeshId;
15use crate::resource;
16pub mod multicast;
17
18use std::cmp::Ordering;
19use std::collections::HashMap;
20use std::fmt::Debug;
21
22use anyhow::Result;
23use async_trait::async_trait;
24use hyperactor::Actor;
25use hyperactor::ActorAddr;
26use hyperactor::ActorRef;
27use hyperactor::Context;
28use hyperactor::Endpoint as _;
29use hyperactor::Handler;
30use hyperactor::Instance;
31use hyperactor::OncePortRefRepr;
32use hyperactor::PortAddr;
33use hyperactor::PortRef;
34use hyperactor::PortRefRepr;
35use hyperactor::RemoteEndpoint as _;
36use hyperactor::RemoteMessage;
37use hyperactor::accum::ReducerMode;
38use hyperactor::mailbox::MailboxSender;
39use hyperactor::mailbox::Undeliverable;
40use hyperactor::mailbox::UndeliverableMailboxSender;
41use hyperactor::mailbox::UndeliverableMessageError;
42use hyperactor::mailbox::monitored_return_handle;
43use hyperactor::ordering::SEQ_INFO;
44use hyperactor::ordering::SeqInfo;
45use hyperactor_config::CONFIG;
46use hyperactor_config::ConfigAttr;
47use hyperactor_config::Flattrs;
48use hyperactor_config::attrs::declare_attrs;
49use hyperactor_mesh_macros::sel;
50use ndslice::Point;
51use ndslice::Selection;
52use ndslice::View;
53use ndslice::selection::routing::RoutingFrame;
54use serde::Deserialize;
55use serde::Serialize;
56use typeuri::Named;
57
58use crate::comm::multicast::CastMessage;
59use crate::comm::multicast::CastMessageEnvelope;
60use crate::comm::multicast::ForwardMessage;
61use crate::comm::multicast::set_cast_info_on_headers;
62
63declare_attrs! {
64 @meta(CONFIG = ConfigAttr::new(
66 Some("HYPERACTOR_MESH_ENABLE_NATIVE_V1_CASTING".to_string()),
67 Some("enable_native_v1_casting".to_string()),
68 ))
69 pub attr ENABLE_NATIVE_V1_CASTING: bool = true;
70
71 pub attr MULTICAST_FAILURE_PHASE: String;
73
74 pub attr MULTICAST_FAILURE_COMM_ACTOR: ActorAddr;
76
77 pub attr MULTICAST_FAILURE_ORIGIN: ActorAddr;
79
80 pub attr MULTICAST_FAILURE_RETURN_PORT: String;
82}
83
84fn annotate_multicast_failure(
85 envelope: &mut hyperactor::mailbox::MessageEnvelope,
86 comm_actor: &ActorAddr,
87 phase: &str,
88 origin: &ActorAddr,
89 return_port: &PortAddr,
90) {
91 let actor_mesh_id = envelope.headers().get(CAST_ACTOR_MESH_ID);
92 if let Some(failure) = envelope.root_delivery_failure_mut() {
93 failure
94 .attrs
95 .set(MULTICAST_FAILURE_PHASE, phase.to_string());
96 failure
97 .attrs
98 .set(MULTICAST_FAILURE_COMM_ACTOR, comm_actor.clone());
99 failure.attrs.set(MULTICAST_FAILURE_ORIGIN, origin.clone());
100 failure
101 .attrs
102 .set(MULTICAST_FAILURE_RETURN_PORT, return_port.to_string());
103 if let Some(actor_mesh_id) = actor_mesh_id {
104 failure.attrs.set(CAST_ACTOR_MESH_ID, actor_mesh_id);
105 }
106 }
107}
108
109#[derive(Debug, Clone, Serialize, Deserialize, Named, Default)]
111pub struct CommActorParams {}
112wirevalue::register_type!(CommActorParams);
113
114#[derive(Debug)]
116struct Buffered {
117 seq: usize,
119 deliver_here: bool,
121 next_steps: HashMap<usize, Vec<RoutingFrame>>,
123 message: CastMessageEnvelope,
125}
126
127#[derive(Debug, Default)]
130struct ReceiveState {
131 seq: usize,
133 buffer: HashMap<usize, Buffered>,
136 last_seqs: HashMap<usize, usize>,
138}
139
140#[derive(Debug, Default)]
143#[hyperactor::export(
144 CommMeshConfig,
145 CastMessage,
146 ForwardMessage,
147 CastMessageV1,
148 ForwardMessageV1
149)]
150#[hyperactor::spawnable]
151pub struct CommActor {
152 send_seq: HashMap<(ActorMeshId, ActorAddr), usize>,
154 recv_state: HashMap<(ActorMeshId, ActorAddr), ReceiveState>,
156
157 mesh_config: MeshConfigState,
159}
160
161#[derive(Debug)]
162enum PendingMessage {
163 Cast(CastMessage),
164 Forward(ForwardMessage),
165 ForwardV1(ForwardMessageV1),
166}
167
168#[derive(Debug)]
169enum MeshConfigState {
170 NotConfigured(Vec<PendingMessage>),
172 Configured(CommMeshConfig),
174}
175
176impl Default for MeshConfigState {
177 fn default() -> Self {
178 MeshConfigState::NotConfigured(Vec::new())
179 }
180}
181
182#[derive(Debug, Clone, Serialize, Deserialize, Named)]
184pub struct CommMeshConfig {
185 rank: usize,
187 peers: HashMap<usize, ActorRef<CommActor>>,
189}
190wirevalue::register_type!(CommMeshConfig);
191
192impl CommMeshConfig {
193 pub fn new(rank: usize, peers: HashMap<usize, ActorRef<CommActor>>) -> Self {
195 Self { rank, peers }
196 }
197
198 fn peer_for_rank(&self, rank: usize) -> Result<ActorRef<CommActor>> {
200 self.peers
201 .get(&rank)
202 .cloned()
203 .ok_or_else(|| anyhow::anyhow!("no peer for rank {}", rank))
204 }
205
206 fn self_rank(&self) -> usize {
208 self.rank
209 }
210}
211
212#[async_trait]
213impl Actor for CommActor {
214 async fn init(&mut self, this: &Instance<Self>) -> Result<(), anyhow::Error> {
215 this.set_system();
216 Ok(())
217 }
218
219 async fn handle_undeliverable_message(
221 &mut self,
222 cx: &Instance<Self>,
223 _reason: hyperactor::mailbox::UndeliverableReason,
224 undelivered: hyperactor::mailbox::Undeliverable<hyperactor::mailbox::MessageEnvelope>,
225 ) -> Result<(), anyhow::Error> {
226 self.return_delivery_failure_to_origin(cx, undelivered)
227 .await
228 }
229
230 async fn handle_invalid_reference(
231 &mut self,
232 cx: &Instance<Self>,
233 _invalid: hyperactor::mailbox::InvalidReference,
234 undelivered: hyperactor::mailbox::Undeliverable<hyperactor::mailbox::MessageEnvelope>,
235 ) -> Result<(), anyhow::Error> {
236 self.return_delivery_failure_to_origin(cx, undelivered)
237 .await
238 }
239}
240
241impl CommActor {
242 async fn return_delivery_failure_to_origin(
243 &mut self,
244 cx: &Instance<Self>,
245 undelivered: hyperactor::mailbox::Undeliverable<hyperactor::mailbox::MessageEnvelope>,
246 ) -> Result<(), anyhow::Error> {
247 let mut message_envelope = match undelivered {
248 Undeliverable::Returned(message_envelope) => message_envelope,
249 Undeliverable::Report(report) => {
250 anyhow::bail!(UndeliverableMessageError::Report { report });
251 }
252 };
253
254 if let Ok(ForwardMessage { message, .. }) =
256 message_envelope.deserialized::<ForwardMessage>()
257 {
258 let sender = message.sender();
259 let return_port = PortRef::attest_handler_port(sender);
260 annotate_multicast_failure(
261 &mut message_envelope,
262 cx.self_addr(),
263 "forward",
264 sender,
265 return_port.port_addr(),
266 );
267
268 message_envelope.set_header(CAST_ORIGINATING_SENDER, sender.clone());
271
272 return_port.post(cx, Undeliverable::Returned(message_envelope.clone()));
273 return Ok(());
274 }
275
276 if let Some(sender) = message_envelope.headers().get(CAST_ORIGINATING_SENDER) {
278 let return_port = PortRef::attest_handler_port(&sender);
279 annotate_multicast_failure(
280 &mut message_envelope,
281 cx.self_addr(),
282 "deliver_here",
283 &sender,
284 return_port.port_addr(),
285 );
286 return_port.post(cx, Undeliverable::Returned(message_envelope.clone()));
287 return Ok(());
288 }
289
290 UndeliverableMailboxSender
292 .post(message_envelope, monitored_return_handle());
293 Ok(())
294 }
295
296 fn forward<M: RemoteMessage>(
298 cx: &Context<Self>,
299 config: &CommMeshConfig,
300 rank: usize,
301 message: M,
302 ) -> Result<()>
303 where
304 CommActor: hyperactor::RemoteHandles<M>,
305 {
306 let child = config.peer_for_rank(rank)?;
307 if let Some(cast_actor_mesh_id) = cx.headers().get(CAST_ACTOR_MESH_ID) {
309 let mut headers = Flattrs::new();
310 headers.set(CAST_ACTOR_MESH_ID, cast_actor_mesh_id);
311 child.post_with_headers(cx, headers, message);
312 } else {
313 child.post(cx, message);
314 }
315 Ok(())
316 }
317
318 fn handle_message(
319 cx: &Context<Self>,
320 config: &CommMeshConfig,
321 deliver_here: bool,
322 next_steps: HashMap<usize, Vec<RoutingFrame>>,
323 sender: ActorAddr,
324 mut message: CastMessageEnvelope,
325 seq: usize,
326 last_seqs: &mut HashMap<usize, usize>,
327 ) -> Result<()> {
328 split_ports(cx, message.data_mut(), deliver_here, &next_steps)?;
329
330 if deliver_here {
332 let headers = message.headers().clone();
336 Self::deliver_to_dest(cx, headers, &mut message, config)?;
337 }
338
339 next_steps
341 .into_iter()
342 .map(|(peer, dests)| {
343 let last_seq = last_seqs.entry(peer).or_default();
344 Self::forward(
345 cx,
346 config,
347 peer,
348 ForwardMessage {
349 dests,
350 sender: sender.clone(),
351 message: message.clone(),
352 seq,
353 last_seq: *last_seq,
354 },
355 )?;
356 *last_seq = seq;
357 Ok(())
358 })
359 .collect::<Result<Vec<_>>>()?;
360
361 Ok(())
362 }
363
364 fn deliver_to_dest<M: CastEnvelope>(
365 cx: &Context<Self>,
366 mut headers: Flattrs,
367 message: &mut M,
368 config: &CommMeshConfig,
369 ) -> anyhow::Result<()> {
370 let cast_point = message.cast_point(config)?;
371 replace_with_self_ranks(&cast_point, message.data_mut())?;
373
374 set_cast_info_on_headers(&mut headers, cast_point, message.sender().clone());
375
376 let dest = cx
379 .self_addr()
380 .proc_addr()
381 .actor_addr_uid(message.dest_port().actor_uid().clone())
382 .port_addr(hyperactor::Port::handler_id(
383 message.dest_port().port(),
384 None,
385 ));
386
387 if let Some(seq_info) = headers.get(SEQ_INFO) {
394 hyperactor::mailbox::headers::stamp_sender_actor_id(
395 &mut headers,
396 &seq_info,
397 &dest,
398 message.sender(),
399 );
400 }
401
402 cx.post_with_external_seq_info(dest, headers, message.data().clone().erase_encoding());
403
404 Ok(())
405 }
406}
407
408fn split_ports(
412 cx: &Context<CommActor>,
413 data: &mut wirevalue::Any<wirevalue::encoding::Multipart>,
414 deliver_here: bool,
415 next_steps: &HashMap<usize, Vec<RoutingFrame>>,
416) -> anyhow::Result<()> {
417 data.visit_multipart_parts_mut::<PortRefRepr, anyhow::Error>(|port| {
421 if port.unsplit() {
422 return Ok(());
423 }
424
425 let split = port.port_addr().split(
426 cx,
427 port.reducer_spec().clone(),
428 ReducerMode::Streaming(port.streaming_opts().clone()),
429 port.get_return_undeliverable(),
430 )?;
431
432 #[cfg(test)]
433 tests::collect_split_port(port.port_addr(), &split, deliver_here);
434
435 port.update_port_addr(split);
436 Ok(())
437 })?;
438
439 data.visit_multipart_parts_mut::<OncePortRefRepr, anyhow::Error>(|port| {
440 if port.unsplit() || port.reducer_spec().is_none() {
441 return Ok(());
445 }
446
447 let peer_count = next_steps.len() + if deliver_here { 1 } else { 0 };
448 let split = port.port_addr().split(
449 cx,
450 port.reducer_spec().clone(),
451 ReducerMode::Once(peer_count),
452 true,
453 )?;
454
455 #[cfg(test)]
456 tests::collect_split_port(port.port_addr(), &split, deliver_here);
457
458 port.update_port_addr(split);
459 Ok(())
460 })?;
461
462 Ok(())
463}
464
465fn replace_with_self_ranks(
466 cast_point: &Point,
467 data: &mut wirevalue::Any<wirevalue::encoding::Multipart>,
468) -> anyhow::Result<()> {
469 data.visit_multipart_parts_mut::<resource::RankRepr, anyhow::Error>(
470 |resource::RankRepr(rank)| {
471 *rank = Some(cast_point.rank());
472 Ok(())
473 },
474 )
475}
476
477#[async_trait]
478impl Handler<CommMeshConfig> for CommActor {
479 async fn handle(&mut self, cx: &Context<Self>, config: CommMeshConfig) -> Result<()> {
480 let pending =
481 match std::mem::replace(&mut self.mesh_config, MeshConfigState::Configured(config)) {
482 MeshConfigState::NotConfigured(pending) => pending,
483 MeshConfigState::Configured(_) => Vec::new(),
484 };
485 if !pending.is_empty() {
486 tracing::info!(
487 count = pending.len(),
488 "replaying buffered pre-config messages"
489 );
490 }
491 for msg in pending {
492 match msg {
493 PendingMessage::Cast(m) => self.handle(cx, m).await?,
494 PendingMessage::Forward(m) => self.handle(cx, m).await?,
495 PendingMessage::ForwardV1(m) => self.handle(cx, m).await?,
496 }
497 }
498 Ok(())
499 }
500}
501
502#[async_trait]
504impl Handler<CastMessage> for CommActor {
505 #[tracing::instrument(level = "debug", skip_all)]
506 async fn handle(&mut self, cx: &Context<Self>, cast_message: CastMessage) -> Result<()> {
507 let config = match &mut self.mesh_config {
508 MeshConfigState::NotConfigured(pending) => {
509 pending.push(PendingMessage::Cast(cast_message));
510 return Ok(());
511 }
512 MeshConfigState::Configured(config) => config,
513 };
514 let slice = cast_message.dest.slice.clone();
516 let selection = cast_message.dest.selection.clone();
517 let frame = RoutingFrame::root(selection, slice);
518 let rank = frame.slice.location(&frame.here)?;
519 let seq = self
520 .send_seq
521 .entry(cast_message.message.stream_key())
522 .or_default();
523 let last_seq = *seq;
524 *seq += 1;
525
526 let fwd_message = ForwardMessage {
527 dests: vec![frame],
528 sender: cx.self_addr().clone(),
529 message: cast_message.message,
530 seq: *seq,
531 last_seq,
532 };
533
534 if config.self_rank() == rank {
537 Handler::<ForwardMessage>::handle(self, cx, fwd_message).await?;
538 } else {
539 Self::forward(cx, config, rank, fwd_message)?;
540 }
541 Ok(())
542 }
543}
544
545#[async_trait]
546impl Handler<ForwardMessage> for CommActor {
547 #[tracing::instrument(level = "debug", skip_all)]
548 async fn handle(&mut self, cx: &Context<Self>, fwd_message: ForwardMessage) -> Result<()> {
549 let config = match &mut self.mesh_config {
550 MeshConfigState::NotConfigured(pending) => {
551 pending.push(PendingMessage::Forward(fwd_message));
552 return Ok(());
553 }
554 MeshConfigState::Configured(config) => config,
555 };
556
557 let ForwardMessage {
558 sender,
559 dests,
560 message,
561 seq,
562 last_seq,
563 } = fwd_message;
564
565 let rank = config.self_rank();
567 let (deliver_here, next_steps) =
568 ndslice::selection::routing::resolve_routing(rank, dests, &mut |_| {
569 panic!("Choice encountered in CommActor routing")
570 })?;
571
572 let recv_state = self.recv_state.entry(message.stream_key()).or_default();
573 match recv_state.seq.cmp(&last_seq) {
574 Ordering::Equal => {
576 Self::handle_message(
578 cx,
579 config,
580 deliver_here,
581 next_steps,
582 sender.clone(),
583 message,
584 seq,
585 &mut recv_state.last_seqs,
586 )?;
587 recv_state.seq = seq;
588
589 while let Some(Buffered {
592 seq,
593 deliver_here,
594 next_steps,
595 message,
596 }) = recv_state.buffer.remove(&recv_state.seq)
597 {
598 Self::handle_message(
599 cx,
600 config,
601 deliver_here,
602 next_steps,
603 sender.clone(),
604 message,
605 seq,
606 &mut recv_state.last_seqs,
607 )?;
608 recv_state.seq = seq;
609 }
610 }
611 Ordering::Less => {
614 tracing::warn!(
615 "buffering out-of-order message with seq {} (last {}), expected {}: {:?}",
616 seq,
617 last_seq,
618 recv_state.seq,
619 message
620 );
621 recv_state.buffer.insert(
622 last_seq,
623 Buffered {
624 seq,
625 deliver_here,
626 next_steps,
627 message,
628 },
629 );
630 }
631 Ordering::Greater => {
633 tracing::warn!("received duplicate message with seq {}: {:?}", seq, message);
634 }
635 }
636
637 Ok(())
638 }
639}
640
641#[async_trait]
642impl Handler<CastMessageV1> for CommActor {
643 async fn handle(&mut self, cx: &Context<Self>, cast_message: CastMessageV1) -> Result<()> {
644 let slice = cast_message.dest_region.slice().clone();
645 let frame = RoutingFrame::root(sel!(*), slice);
646 let forward_message = ForwardMessageV1 {
647 dests: vec![frame],
648 message: cast_message,
649 };
650 self.handle(cx, forward_message).await
651 }
652}
653
654#[async_trait]
655impl Handler<ForwardMessageV1> for CommActor {
656 async fn handle(&mut self, cx: &Context<Self>, fwd_message: ForwardMessageV1) -> Result<()> {
657 let config = match &mut self.mesh_config {
658 MeshConfigState::NotConfigured(pending) => {
659 pending.push(PendingMessage::ForwardV1(fwd_message));
660 return Ok(());
661 }
662 MeshConfigState::Configured(config) => config,
663 };
664
665 let ForwardMessageV1 { dests, mut message } = fwd_message;
666 let rank_on_root_mesh = config.self_rank();
668 let (deliver_here, next_steps) =
669 ndslice::selection::routing::resolve_routing(rank_on_root_mesh, dests, &mut |_| {
670 panic!("choice encountered in CommActor routing")
671 })?;
672
673 split_ports(cx, &mut message.data, deliver_here, &next_steps)?;
674
675 if deliver_here {
677 let mut headers = message.headers().clone();
678 let seq = message
679 .seqs
680 .get(message.cast_point(config)?.rank())
681 .expect("mismatched seqs and dest_region");
682 headers.set(
683 SEQ_INFO,
684 SeqInfo::Session {
685 session_id: message.session_id,
686 seq,
687 },
688 );
689 Self::deliver_to_dest(cx, headers, &mut message, config)?;
690 }
691
692 for (peer_rank_on_root_mesh, dests) in next_steps {
694 let forward_message = ForwardMessageV1 {
695 dests,
696 message: message.clone(),
697 };
698 Self::forward(cx, config, peer_rank_on_root_mesh, forward_message)?;
699 }
700
701 Ok(())
702 }
703}
704
705pub mod test_utils {
706 use anyhow::Result;
707 use async_trait::async_trait;
708 use hyperactor::Actor;
709 use hyperactor::ActorAddr;
710 use hyperactor::Context;
711 use hyperactor::Handler;
712 use hyperactor::PortRef;
713 use serde::Deserialize;
714 use serde::Serialize;
715 use typeuri::Named;
716
717 use super::*;
718
719 #[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Named)]
720 pub struct MyReply {
721 pub sender: ActorAddr,
722 pub value: u64,
723 }
724
725 #[derive(Debug, Named, Serialize, Deserialize, PartialEq, Clone)]
726 #[expect(
727 clippy::large_enum_variant,
728 reason = "test fixture; CastAndReply carries PortRefs and boxing fields ripples into handler assertions"
729 )]
730 pub enum TestMessage {
731 Forward(String),
732 CastAndReply {
733 arg: String,
734 reply_to0: PortRef<String>,
736 reply_to1: PortRef<u64>,
737 reply_to2: PortRef<MyReply>,
738 },
739 CastAndReplyOnce {
740 arg: String,
741 reply_to: hyperactor::OncePortRef<u64>,
742 },
743 CastWithUnsplitPort {
744 reply_to: PortRef<u64>,
745 },
746 }
747
748 #[derive(Debug)]
749 #[hyperactor::export(TestMessage)]
750 #[hyperactor::spawnable]
751 pub struct TestActor {
752 forward_port: PortRef<TestMessage>,
755 }
756
757 #[derive(Debug, Clone, Named, Serialize, Deserialize)]
758 pub struct TestActorParams {
759 pub forward_port: PortRef<TestMessage>,
760 }
761
762 #[async_trait]
763 impl Actor for TestActor {}
764
765 #[async_trait]
766 impl hyperactor::RemoteSpawn for TestActor {
767 type Params = TestActorParams;
768
769 async fn new(params: Self::Params, _environment: Flattrs) -> Result<Self> {
770 let Self::Params { forward_port } = params;
771 Ok(Self { forward_port })
772 }
773 }
774
775 #[async_trait]
776 impl Handler<TestMessage> for TestActor {
777 async fn handle(&mut self, cx: &Context<Self>, msg: TestMessage) -> anyhow::Result<()> {
778 if let TestMessage::CastWithUnsplitPort { ref reply_to } = msg {
781 reply_to.post(cx, 42);
782 }
783 self.forward_port.post(cx, msg);
784 Ok(())
785 }
786 }
787
788 #[doc(hidden)]
794 #[derive(Debug, Clone, Serialize, Deserialize, Named)]
795 pub struct SenderCaptureMsg();
796
797 #[doc(hidden)]
798 #[derive(Debug, Clone, PartialEq, Serialize, Deserialize, Named)]
799 pub struct CapturedSender(pub Option<ActorAddr>);
800
801 #[doc(hidden)]
802 #[derive(Debug)]
803 #[hyperactor::export(SenderCaptureMsg)]
804 #[hyperactor::spawnable]
805 pub struct SenderCapturingActor {
806 forward_port: PortRef<CapturedSender>,
807 }
808
809 #[doc(hidden)]
810 #[derive(Debug, Clone, Named, Serialize, Deserialize)]
811 pub struct SenderCapturingActorParams {
812 pub forward_port: PortRef<CapturedSender>,
813 }
814
815 #[async_trait]
816 impl Actor for SenderCapturingActor {}
817
818 #[async_trait]
819 impl hyperactor::RemoteSpawn for SenderCapturingActor {
820 type Params = SenderCapturingActorParams;
821
822 async fn new(params: Self::Params, _environment: Flattrs) -> Result<Self> {
823 let Self::Params { forward_port } = params;
824 Ok(Self { forward_port })
825 }
826 }
827
828 #[async_trait]
829 impl Handler<SenderCaptureMsg> for SenderCapturingActor {
830 async fn handle(&mut self, cx: &Context<Self>, _: SenderCaptureMsg) -> anyhow::Result<()> {
831 let sender = cx
832 .headers()
833 .get(hyperactor::mailbox::headers::SENDER_ACTOR_ID);
834 self.forward_port.post(cx, CapturedSender(sender));
835 Ok(())
836 }
837 }
838}
839
840#[cfg(test)]
841mod tests {
842 use std::collections::BTreeMap;
843 use std::collections::HashSet;
844 use std::fmt::Display;
845 use std::hash::Hash;
846 use std::ops::Deref;
847 use std::ops::DerefMut;
848 use std::sync::Mutex;
849 use std::sync::OnceLock;
850
851 use hyperactor::accum;
852
853 async fn buffering_fixture(
858 proc_name: &str,
859 ) -> (
860 hyperactor::Client,
861 hyperactor::mailbox::PortReceiver<TestMessage>,
862 hyperactor::ActorHandle<CommActor>,
863 crate::mesh_id::ActorMeshId,
864 (hyperactor::ActorHandle<TestActor>, ActorRef<TestActor>),
866 ) {
867 use hyperactor::Proc;
868 use hyperactor::RemoteSpawn;
869 use hyperactor::channel::ChannelTransport;
870 use hyperactor::id::Label;
871
872 let proc = Proc::direct(ChannelTransport::Unix.any(), proc_name.to_string()).unwrap();
873 let client = proc.client("client");
874
875 let actor_mesh_id = crate::mesh_id::ActorMeshId::instance(Label::new("test").unwrap());
876
877 let (tx, rx) = open_port(&client);
878 let forward_port = tx.bind();
879 let test_actor = TestActor::new(TestActorParams { forward_port }, Default::default())
880 .await
881 .unwrap();
882 let test_handle = proc
883 .spawn_with_uid(actor_mesh_id.uid().clone(), test_actor)
884 .unwrap();
885 let test_ref: ActorRef<TestActor> = test_handle.bind::<TestActor>();
886
887 let comm_handle = proc.spawn(CommActor::default());
888
889 (
890 client,
891 rx,
892 comm_handle,
893 actor_mesh_id,
894 (test_handle, test_ref),
895 )
896 }
897
898 fn send_config(client: &hyperactor::Client, comm_handle: &hyperactor::ActorHandle<CommActor>) {
900 let comm_ref = comm_handle.bind::<CommActor>();
901 let mut peers = HashMap::new();
902 peers.insert(0, comm_ref);
903 comm_handle.post(client, CommMeshConfig::new(0, peers));
904 }
905
906 async fn assert_buffered_and_replayed<M: hyperactor::Message>(
909 proc_name: &str,
910 mut make_msg: impl FnMut(&hyperactor::Client, &crate::mesh_id::ActorMeshId, &str) -> M,
911 ) where
912 CommActor: hyperactor::Handler<M>,
913 {
914 let (client, mut rx, comm_handle, actor_mesh_id, _guards) =
915 buffering_fixture(proc_name).await;
916
917 comm_handle.post(&client, make_msg(&client, &actor_mesh_id, "buffered"));
918 send_config(&client, &comm_handle);
919 comm_handle.post(&client, make_msg(&client, &actor_mesh_id, "direct"));
920
921 assert_eq!(
922 rx.recv().await.unwrap(),
923 TestMessage::Forward("buffered".to_string()),
924 );
925 assert_eq!(
926 rx.recv().await.unwrap(),
927 TestMessage::Forward("direct".to_string()),
928 );
929 comm_handle.drain_and_stop("test done").ok();
930 }
931
932 #[async_timed_test(timeout_secs = 1)]
933 async fn cast_before_config_is_buffered_and_replayed() {
934 use ndslice::Slice;
935
936 assert_buffered_and_replayed("test_cast", |client, actor_mesh_id, payload| {
937 let actor_mesh_id = actor_mesh_id.clone();
938 let slice = Slice::new_row_major(vec![1]);
939 let shape = ndslice::Shape::new(vec!["rank".to_string()], slice.clone()).unwrap();
940 let envelope = multicast::CastMessageEnvelope::new::<TestActor, TestMessage>(
941 actor_mesh_id,
942 client.self_addr().clone(),
943 shape,
944 hyperactor_config::Flattrs::new(),
945 TestMessage::Forward(payload.to_string()),
946 )
947 .unwrap();
948 multicast::CastMessage {
949 dest: multicast::Uslice {
950 slice,
951 selection: sel!(*),
952 },
953 message: envelope,
954 }
955 })
956 .await;
957 }
958
959 #[async_timed_test(timeout_secs = 1)]
960 async fn forward_before_config_is_buffered_and_replayed() {
961 use ndslice::Slice;
962 use ndslice::selection::routing::RoutingFrame;
963
964 let mut next_seq: usize = 0;
965 assert_buffered_and_replayed("test_fwd", move |client, actor_mesh_id, payload| {
966 let actor_mesh_id = actor_mesh_id.clone();
967 let slice = Slice::new_row_major(vec![1]);
968 let shape = ndslice::Shape::new(vec!["rank".to_string()], slice.clone()).unwrap();
969 let envelope = multicast::CastMessageEnvelope::new::<TestActor, TestMessage>(
970 actor_mesh_id,
971 client.self_addr().clone(),
972 shape,
973 hyperactor_config::Flattrs::new(),
974 TestMessage::Forward(payload.to_string()),
975 )
976 .unwrap();
977 let frame = RoutingFrame::root(sel!(*), slice);
978 let last_seq = next_seq;
979 next_seq += 1;
980 multicast::ForwardMessage {
981 sender: client.self_addr().clone(),
982 dests: vec![frame],
983 seq: next_seq,
984 last_seq,
985 message: envelope,
986 }
987 })
988 .await;
989 }
990
991 #[async_timed_test(timeout_secs = 1)]
992 async fn forward_v1_before_config_is_buffered_and_replayed() {
993 use ndslice::Region;
994 use ndslice::Slice;
995 use ndslice::selection::routing::RoutingFrame;
996
997 assert_buffered_and_replayed("test_fwd_v1", |client, actor_mesh_id, payload| {
998 let slice = Slice::new_row_major(vec![1]);
999 let region = Region::new(vec!["rank".to_string()], slice.clone());
1000 let cast_msg = multicast::CastMessageV1::new::<TestActor, TestMessage>(
1001 client.self_addr().clone(),
1002 actor_mesh_id,
1003 region.clone(),
1004 hyperactor_config::Flattrs::new(),
1005 TestMessage::Forward(payload.to_string()),
1006 uuid::Uuid::new_v4(),
1007 crate::ValueMesh::from_single(region, 1u64),
1008 )
1009 .unwrap();
1010 let frame = RoutingFrame::root(sel!(*), slice);
1011 multicast::ForwardMessageV1 {
1012 dests: vec![frame],
1013 message: cast_msg,
1014 }
1015 })
1016 .await;
1017 }
1018
1019 use hyperactor::ActorAddr;
1020 use hyperactor::ActorRef;
1021 use hyperactor::Endpoint as _;
1022 use hyperactor::Index;
1023 use hyperactor::OncePortRef;
1024 use hyperactor::PortAddr;
1025 use hyperactor::PortRef;
1026 use hyperactor::ProcAddr;
1027 use hyperactor::accum::Accumulator;
1028 use hyperactor::accum::ReducerSpec;
1029 use hyperactor::channel::ChannelAddr;
1030 use hyperactor::context;
1031 use hyperactor::context::Mailbox;
1032 use hyperactor::mailbox::DeliveryFailure;
1033 use hyperactor::mailbox::MessageEnvelope;
1034 use hyperactor::mailbox::PortReceiver;
1035 use hyperactor::mailbox::TransportFailure;
1036 use hyperactor::mailbox::TransportFailureReason;
1037 use hyperactor::mailbox::UndeliverableReason;
1038 use hyperactor::mailbox::open_port;
1039 use hyperactor::port::Port;
1040 use hyperactor_config;
1041 use hyperactor_mesh_macros::sel;
1042 use maplit::btreemap;
1043 use maplit::hashmap;
1044 use ndslice::Extent;
1045 use ndslice::Selection;
1046 use ndslice::ViewExt as _;
1047 use ndslice::extent;
1048 use ndslice::selection::test_utils::collect_commactor_routing_tree;
1049 use test_utils::*;
1050 use timed_test::async_timed_test;
1051 use tokio::time::Duration;
1052
1053 use super::*;
1054 use crate::ActorMesh;
1055 use crate::ProcMesh;
1056 use crate::host_mesh::HostMesh;
1057 use crate::test_utils::local_host_mesh;
1058 use crate::testing;
1059
1060 #[test]
1061 fn annotate_multicast_failure_adds_attrs_to_root_failure() {
1062 let proc_addr = ProcAddr::singleton(ChannelAddr::Local(1), "test");
1063 let origin = proc_addr.actor_addr("origin");
1064 let comm_actor = proc_addr.actor_addr("comm");
1065 let dest = proc_addr.actor_addr("dest").port_addr(Port::from(42));
1066 let return_port = origin.port_addr(Port::from(7));
1067 let mut envelope =
1068 MessageEnvelope::serialize(origin.clone(), dest.clone(), &(), Flattrs::new()).unwrap();
1069 envelope.push_delivery_failure(DeliveryFailure::new(UndeliverableReason::Transport(
1070 TransportFailure::new(dest, TransportFailureReason::NoRoute),
1071 )));
1072
1073 annotate_multicast_failure(
1074 &mut envelope,
1075 &comm_actor,
1076 "deliver_here",
1077 &origin,
1078 &return_port,
1079 );
1080
1081 let root_failure = envelope
1082 .root_delivery_failure()
1083 .expect("expected root delivery failure");
1084 assert_eq!(
1085 root_failure.attrs.get(MULTICAST_FAILURE_PHASE).as_deref(),
1086 Some("deliver_here")
1087 );
1088 assert_eq!(
1089 root_failure.attrs.get(MULTICAST_FAILURE_COMM_ACTOR),
1090 Some(comm_actor)
1091 );
1092 assert_eq!(
1093 root_failure.attrs.get(MULTICAST_FAILURE_ORIGIN),
1094 Some(origin)
1095 );
1096 assert_eq!(
1097 root_failure.attrs.get(MULTICAST_FAILURE_RETURN_PORT),
1098 Some(return_port.to_string())
1099 );
1100 }
1101
1102 #[test]
1103 fn annotate_multicast_failure_records_forward_phase() {
1104 let proc_addr = ProcAddr::singleton(ChannelAddr::Local(1), "test");
1105 let origin = proc_addr.actor_addr("origin");
1106 let comm_actor = proc_addr.actor_addr("comm");
1107 let dest = proc_addr.actor_addr("dest").port_addr(Port::from(42));
1108 let return_port = origin.port_addr(Port::from(7));
1109 let mut envelope =
1110 MessageEnvelope::serialize(origin.clone(), dest.clone(), &(), Flattrs::new()).unwrap();
1111 envelope.push_delivery_failure(DeliveryFailure::new(UndeliverableReason::Transport(
1112 TransportFailure::new(dest, TransportFailureReason::NoRoute),
1113 )));
1114
1115 annotate_multicast_failure(&mut envelope, &comm_actor, "forward", &origin, &return_port);
1116
1117 let root_failure = envelope
1118 .root_delivery_failure()
1119 .expect("expected root delivery failure");
1120 assert_eq!(
1121 root_failure.attrs.get(MULTICAST_FAILURE_PHASE).as_deref(),
1122 Some("forward")
1123 );
1124 }
1125
1126 #[test]
1127 fn annotate_multicast_failure_copies_actor_mesh_id_attr() {
1128 let proc_addr = ProcAddr::singleton(ChannelAddr::Local(1), "test");
1129 let origin = proc_addr.actor_addr("origin");
1130 let comm_actor = proc_addr.actor_addr("comm");
1131 let dest = proc_addr.actor_addr("dest").port_addr(Port::from(42));
1132 let return_port = origin.port_addr(Port::from(7));
1133 let actor_mesh_id =
1134 crate::mesh_id::ActorMeshId::instance(hyperactor::id::Label::new("mesh").unwrap());
1135 let mut headers = Flattrs::new();
1136 headers.set(CAST_ACTOR_MESH_ID, actor_mesh_id.clone());
1137 let mut envelope =
1138 MessageEnvelope::serialize(origin.clone(), dest.clone(), &(), headers).unwrap();
1139 envelope.push_delivery_failure(DeliveryFailure::new(UndeliverableReason::Transport(
1140 TransportFailure::new(dest, TransportFailureReason::NoRoute),
1141 )));
1142
1143 annotate_multicast_failure(
1144 &mut envelope,
1145 &comm_actor,
1146 "deliver_here",
1147 &origin,
1148 &return_port,
1149 );
1150
1151 let root_failure = envelope
1152 .root_delivery_failure()
1153 .expect("expected root delivery failure");
1154 assert_eq!(
1155 root_failure.attrs.get(CAST_ACTOR_MESH_ID),
1156 Some(actor_mesh_id)
1157 );
1158 }
1159
1160 #[test]
1161 fn annotate_multicast_failure_noops_without_root_failure() {
1162 let proc_addr = ProcAddr::singleton(ChannelAddr::Local(1), "test");
1163 let origin = proc_addr.actor_addr("origin");
1164 let comm_actor = proc_addr.actor_addr("comm");
1165 let dest = proc_addr.actor_addr("dest").port_addr(Port::from(42));
1166 let return_port = origin.port_addr(Port::from(7));
1167 let mut envelope =
1168 MessageEnvelope::serialize(origin.clone(), dest, &(), Flattrs::new()).unwrap();
1169
1170 annotate_multicast_failure(
1171 &mut envelope,
1172 &comm_actor,
1173 "deliver_here",
1174 &origin,
1175 &return_port,
1176 );
1177
1178 assert!(envelope.root_delivery_failure().is_none());
1179 }
1180
1181 struct Edge<T> {
1182 from: T,
1183 to: T,
1184 is_leaf: bool,
1185 }
1186
1187 impl<T> From<(T, T, bool)> for Edge<T> {
1188 fn from((from, to, is_leaf): (T, T, bool)) -> Self {
1189 Self { from, to, is_leaf }
1190 }
1191 }
1192
1193 static SPLIT_PORT_TREE: OnceLock<Mutex<Vec<Edge<PortAddr>>>> = OnceLock::new();
1196
1197 pub(crate) fn collect_split_port(original: &PortAddr, split: &PortAddr, deliver_here: bool) {
1200 let mutex = SPLIT_PORT_TREE.get_or_init(|| Mutex::new(vec![]));
1201 let mut tree = mutex.lock().unwrap();
1202
1203 tree.deref_mut().push(Edge {
1204 from: original.clone(),
1205 to: split.clone(),
1206 is_leaf: deliver_here,
1207 });
1208 }
1209
1210 fn clear_collected_tree() {
1214 if let Some(tree) = SPLIT_PORT_TREE.get() {
1215 let mut tree: std::sync::MutexGuard<'_, Vec<Edge<PortAddr>>> = tree.lock().unwrap();
1216 tree.clear();
1217 }
1218 }
1219
1220 #[derive(PartialEq)]
1224 struct PathToLeaves<T>(BTreeMap<T, Vec<T>>);
1225
1226 impl<T: Display> Debug for PathToLeaves<T> {
1228 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1229 fn vec_to_string<T: Display>(v: &[T]) -> String {
1230 v.iter()
1231 .map(ToString::to_string)
1232 .collect::<Vec<String>>()
1233 .join(", ")
1234 }
1235
1236 for (src, path) in &self.0 {
1237 writeln!(f, "{} -> {}", src, vec_to_string(path))?;
1238 }
1239 Ok(())
1240 }
1241 }
1242
1243 fn build_paths<T: Clone + Eq + Hash + Ord>(edges: &[Edge<T>]) -> PathToLeaves<T> {
1244 let mut child_parent_map = HashMap::new();
1245 let mut all_nodes = HashSet::new();
1246 let mut parents = HashSet::new();
1247 let mut children = HashSet::new();
1248 let mut dests = HashSet::new();
1249
1250 for Edge { from, to, is_leaf } in edges {
1252 child_parent_map.insert(to.clone(), from.clone());
1253 all_nodes.insert(from.clone());
1254 all_nodes.insert(to.clone());
1255 parents.insert(from.clone());
1256 children.insert(to.clone());
1257 if *is_leaf {
1258 dests.insert(to.clone());
1259 }
1260 }
1261
1262 let mut result = BTreeMap::new();
1264 for dest in dests {
1265 let mut path = vec![dest.clone()];
1266 let mut current = dest.clone();
1267 while let Some(parent) = child_parent_map.get(¤t) {
1268 path.push(parent.clone());
1269 current = parent.clone();
1270 }
1271 path.reverse();
1272 result.insert(dest, path);
1273 }
1274
1275 PathToLeaves(result)
1276 }
1277
1278 #[test]
1279 fn test_build_paths() {
1280 let edges: Vec<_> = [
1287 (0, 1, false),
1288 (1, 2, true),
1289 (1, 3, true),
1290 (0, 4, true),
1291 (4, 5, true),
1292 ]
1293 .into_iter()
1294 .map(|(from, to, is_leaf)| Edge { from, to, is_leaf })
1295 .collect();
1296
1297 let paths = build_paths(&edges);
1298
1299 let expected = btreemap! {
1300 2 => vec![0, 1, 2],
1301 3 => vec![0, 1, 3],
1302 4 => vec![0, 4],
1303 5 => vec![0, 4, 5],
1304 };
1305
1306 assert_eq!(paths.0, expected);
1307 }
1308
1309 struct NoneAccumulator;
1310
1311 impl Accumulator for NoneAccumulator {
1312 type State = u64;
1313 type Update = u64;
1314
1315 fn accumulate(
1316 &self,
1317 _state: &mut Self::State,
1318 _update: Self::Update,
1319 ) -> anyhow::Result<()> {
1320 unimplemented!()
1321 }
1322
1323 fn reducer_spec(&self) -> Option<ReducerSpec> {
1324 unimplemented!()
1325 }
1326 }
1327
1328 async fn execute_cast_and_reply(
1329 ranks: Vec<ActorRef<TestActor>>,
1330 instance: &impl context::Actor,
1331 mut reply1_rx: PortReceiver<u64>,
1332 mut reply2_rx: PortReceiver<MyReply>,
1333 reply_tos: Vec<(PortRef<u64>, PortRef<MyReply>)>,
1334 ) {
1335 {
1337 for (rank, (dest_actor, (reply_to1, reply_to2))) in
1338 ranks.iter().zip(reply_tos.iter()).enumerate()
1339 {
1340 let rank_u64 = rank as u64;
1341 reply_to1.post(instance, rank_u64);
1342 let my_reply = MyReply {
1343 sender: dest_actor.actor_addr().clone(),
1344 value: rank_u64,
1345 };
1346 reply_to2.post(instance, my_reply.clone());
1347
1348 assert_eq!(reply1_rx.recv().await.unwrap(), rank_u64);
1349 assert_eq!(reply2_rx.recv().await.unwrap(), my_reply);
1350 }
1351 }
1352
1353 tracing::info!("the 1st updates from all dest actors were receivered by client");
1354
1355 {
1359 let n = 100;
1360 let mut expected2: HashMap<ActorAddr, Vec<MyReply>> = hashmap! {};
1361 for (i, (dest_actor, (_reply_to1, reply_to2))) in
1362 ranks.iter().zip(reply_tos.iter()).enumerate()
1363 {
1364 let mut sent2 = vec![];
1365 for j in 0..n {
1366 let value = (i * 100 + j) as u64;
1367 let my_reply = MyReply {
1368 sender: dest_actor.actor_addr().clone(),
1369 value,
1370 };
1371 reply_to2.post(instance, my_reply.clone());
1372 sent2.push(my_reply);
1373 }
1374 assert!(
1375 expected2
1376 .insert(dest_actor.actor_addr().clone(), sent2)
1377 .is_none(),
1378 "duplicate actor_id {} in map",
1379 dest_actor.actor_addr()
1380 );
1381 }
1382
1383 let mut received2: HashMap<ActorAddr, Vec<MyReply>> = hashmap! {};
1384
1385 for _ in 0..(n * ranks.len()) {
1386 let my_reply = reply2_rx.recv().await.unwrap();
1387 received2
1388 .entry(my_reply.sender.clone())
1389 .or_default()
1390 .push(my_reply);
1391 }
1392 assert_eq!(received2, expected2);
1393 }
1394 }
1395
1396 async fn wait_for_with_timeout(
1397 receiver: &mut PortReceiver<u64>,
1398 expected: u64,
1399 dur: Duration,
1400 ) -> anyhow::Result<()> {
1401 tokio::time::timeout(dur, async {
1403 loop {
1404 let msg = receiver.recv().await.unwrap();
1405 if msg == expected {
1406 break;
1407 }
1408 }
1409 })
1410 .await?;
1411 Ok(())
1412 }
1413
1414 async fn execute_cast_and_accum(
1415 ranks: Vec<ActorRef<TestActor>>,
1416 instance: &impl context::Actor,
1417 mut reply1_rx: PortReceiver<u64>,
1418 reply_tos: Vec<(PortRef<u64>, PortRef<MyReply>)>,
1419 ) {
1420 let mut sum = 0;
1424 let n = 100;
1425 for (i, (_dest_actor, (reply_to1, _reply_to2))) in
1426 ranks.iter().zip(reply_tos.iter()).enumerate()
1427 {
1428 for j in 0..n {
1429 let value = (i + j) as u64;
1430 reply_to1.post(instance, value);
1431 sum += value;
1432 }
1433 }
1434 wait_for_with_timeout(&mut reply1_rx, sum, Duration::from_secs(8))
1435 .await
1436 .unwrap();
1437 tokio::time::sleep(Duration::from_secs(2)).await;
1439 let msg = reply1_rx.try_recv().unwrap();
1440 assert_eq!(msg, None);
1441 }
1442
1443 struct MeshSetupV1 {
1444 instance: &'static Instance<testing::TestRootClient>,
1445 actor_mesh_ref: crate::ActorMeshRef<TestActor>,
1446 reply1_rx: PortReceiver<u64>,
1447 reply2_rx: PortReceiver<MyReply>,
1448 reply_tos: Vec<(PortRef<u64>, PortRef<MyReply>)>,
1449 host_mesh: HostMesh,
1451 }
1452
1453 async fn setup_mesh_v1<A>(accum: Option<A>) -> MeshSetupV1
1454 where
1455 A: Accumulator<Update = u64, State = u64> + Send + Sync + 'static,
1456 {
1457 let instance = crate::testing::instance();
1458 let host_mesh = local_host_mesh(8).await;
1461 let proc_mesh = host_mesh
1462 .spawn(instance, "test", extent!(gpu = 8), None, None)
1463 .await
1464 .unwrap();
1465
1466 let (tx, mut rx) = hyperactor::mailbox::open_port(instance);
1467 let params = TestActorParams {
1468 forward_port: tx.bind(),
1469 };
1470 let actor_name =
1471 crate::mesh_id::ActorMeshId::instance(hyperactor::id::Label::new("test").unwrap());
1472 let actor_mesh = proc_mesh
1476 .spawn_with_name(&instance, actor_name, ¶ms, None, true)
1477 .await
1478 .unwrap();
1479 let actor_mesh_ref: crate::ActorMeshRef<TestActor> = actor_mesh.deref().clone();
1480
1481 let (reply_port_handle0, _) = open_port::<String>(instance);
1482 let reply_port_ref0 = reply_port_handle0.bind();
1483 let (reply_port_handle1, reply1_rx) = match accum {
1484 Some(a) => instance.mailbox().open_accum_port(a),
1485 None => open_port(instance),
1486 };
1487 let reply_port_ref1 = reply_port_handle1.bind();
1488 let (reply_port_handle2, reply2_rx) = open_port::<MyReply>(instance);
1489 let reply_port_ref2 = reply_port_handle2.bind();
1490 let message = TestMessage::CastAndReply {
1491 arg: "abc".to_string(),
1492 reply_to0: reply_port_ref0.clone().unsplit(),
1493 reply_to1: reply_port_ref1.clone(),
1494 reply_to2: reply_port_ref2.clone(),
1495 };
1496
1497 clear_collected_tree();
1498 actor_mesh_ref.cast(instance, message).unwrap();
1499
1500 let mut reply_tos = vec![];
1501 for _ in proc_mesh.extent().points() {
1502 let msg = rx.recv().await.expect("missing");
1503 match msg {
1504 TestMessage::CastAndReply {
1505 arg,
1506 reply_to0,
1507 reply_to1,
1508 reply_to2,
1509 } => {
1510 assert_eq!(arg, "abc");
1511 assert_eq!(reply_to0, reply_port_ref0.clone().unsplit());
1513 assert_ne!(reply_to1, reply_port_ref1);
1515 assert_ne!(reply_to2, reply_port_ref2);
1516 reply_tos.push((reply_to1, reply_to2));
1517 }
1518 _ => {
1519 panic!("unexpected message: {:?}", msg);
1520 }
1521 }
1522 }
1523
1524 MeshSetupV1 {
1525 instance,
1526 actor_mesh_ref,
1527 reply1_rx,
1528 reply2_rx,
1529 reply_tos,
1530 host_mesh,
1531 }
1532 }
1533
1534 async fn execute_cast_and_reply_v1() {
1535 let mut setup = setup_mesh_v1::<NoneAccumulator>(None).await;
1536
1537 let ranks = setup.actor_mesh_ref.values().collect::<Vec<_>>();
1538 execute_cast_and_reply(
1539 ranks,
1540 setup.instance,
1541 setup.reply1_rx,
1542 setup.reply2_rx,
1543 setup.reply_tos,
1544 )
1545 .await;
1546
1547 let _ = setup.host_mesh.shutdown(setup.instance).await;
1548 }
1549
1550 #[async_timed_test(timeout_secs = 60)]
1551 async fn test_cast_and_reply_v1_retrofit() {
1552 let config = hyperactor_config::global::lock();
1553 let _guard = config.override_key(ENABLE_NATIVE_V1_CASTING, false);
1554 let _guard2 = config.override_key(
1555 hyperactor::config::ENABLE_DEST_ACTOR_REORDERING_BUFFER,
1556 false,
1557 );
1558 execute_cast_and_reply_v1().await
1559 }
1560
1561 #[async_timed_test(timeout_secs = 60)]
1562 async fn test_cast_and_reply_v1_native() {
1563 let config = hyperactor_config::global::lock();
1564 let _guard = config.override_key(ENABLE_NATIVE_V1_CASTING, true);
1565 let _guard2 = config.override_key(
1566 hyperactor::config::ENABLE_DEST_ACTOR_REORDERING_BUFFER,
1567 true,
1568 );
1569 execute_cast_and_reply_v1().await
1570 }
1571
1572 #[async_timed_test(timeout_secs = 60)]
1573 async fn test_cast_and_reply_v1_native_p2p() {
1574 let config = hyperactor_config::global::lock();
1575 let _guard = config.override_key(ENABLE_NATIVE_V1_CASTING, true);
1576 let _guard2 = config.override_key(
1577 hyperactor::config::ENABLE_DEST_ACTOR_REORDERING_BUFFER,
1578 true,
1579 );
1580 let _guard3 = config.override_key(crate::config::V1_CAST_POINT_TO_POINT_THRESHOLD, 1024);
1581 execute_cast_and_reply_v1().await
1582 }
1583
1584 async fn execute_cast_and_accum_v1(config: &hyperactor_config::global::ConfigLock) {
1585 let _guard1 = config.override_key(hyperactor::config::SPLIT_MAX_BUFFER_SIZE, 1);
1587
1588 let mut setup = setup_mesh_v1(Some(accum::sum::<u64>())).await;
1589
1590 let ranks = setup.actor_mesh_ref.values().collect::<Vec<_>>();
1591 execute_cast_and_accum(ranks, setup.instance, setup.reply1_rx, setup.reply_tos).await;
1592
1593 let _ = setup.host_mesh.shutdown(setup.instance).await;
1594 }
1595
1596 #[async_timed_test(timeout_secs = 60)]
1597 async fn test_cast_and_accum_v1_retrofit() {
1598 let config = hyperactor_config::global::lock();
1599 let _guard = config.override_key(ENABLE_NATIVE_V1_CASTING, false);
1600 let _guard2 = config.override_key(
1601 hyperactor::config::ENABLE_DEST_ACTOR_REORDERING_BUFFER,
1602 false,
1603 );
1604 execute_cast_and_accum_v1(&config).await
1605 }
1606
1607 #[async_timed_test(timeout_secs = 60)]
1608 async fn test_cast_and_accum_v1_native() {
1609 let config = hyperactor_config::global::lock();
1610 let _guard = config.override_key(ENABLE_NATIVE_V1_CASTING, true);
1611 let _guard2 = config.override_key(
1612 hyperactor::config::ENABLE_DEST_ACTOR_REORDERING_BUFFER,
1613 true,
1614 );
1615 execute_cast_and_accum_v1(&config).await
1616 }
1617
1618 #[async_timed_test(timeout_secs = 60)]
1619 async fn test_cast_and_accum_v1_native_p2p() {
1620 let config = hyperactor_config::global::lock();
1621 let _guard = config.override_key(ENABLE_NATIVE_V1_CASTING, true);
1622 let _guard2 = config.override_key(
1623 hyperactor::config::ENABLE_DEST_ACTOR_REORDERING_BUFFER,
1624 true,
1625 );
1626 let _guard3 = config.override_key(crate::config::V1_CAST_POINT_TO_POINT_THRESHOLD, 1024);
1627 execute_cast_and_accum_v1(&config).await
1628 }
1629
1630 struct OncePortMeshSetupV1 {
1631 instance: &'static Instance<testing::TestRootClient>,
1632 reply_rx: hyperactor::mailbox::OncePortReceiver<u64>,
1633 reply_tos: Vec<OncePortRef<u64>>,
1634 _reply_port_ref: OncePortRef<u64>,
1635 host_mesh: HostMesh,
1636 }
1637
1638 async fn setup_once_port_mesh<A>(reducer: Option<A>) -> OncePortMeshSetupV1
1639 where
1640 A: Accumulator<State = u64, Update = u64> + Send + Sync + 'static,
1641 {
1642 let instance = crate::testing::instance();
1643 let host_mesh = local_host_mesh(8).await;
1646 let proc_mesh = host_mesh
1647 .spawn(instance, "test", extent!(gpu = 8), None, None)
1648 .await
1649 .unwrap();
1650
1651 let (tx, mut rx) = hyperactor::mailbox::open_port(instance);
1652 let params = TestActorParams {
1653 forward_port: tx.bind(),
1654 };
1655 let actor_name =
1656 crate::mesh_id::ActorMeshId::instance(hyperactor::id::Label::new("test").unwrap());
1657 let actor_mesh: crate::ActorMesh<TestActor> = proc_mesh
1659 .spawn_with_name(&instance, actor_name, ¶ms, None, true)
1660 .await
1661 .unwrap();
1662 let actor_mesh_ref = actor_mesh.deref().clone();
1663
1664 let has_reducer = reducer.is_some();
1665 let (reply_port_handle, reply_rx) = match reducer {
1666 Some(reducer) => instance.mailbox().open_reduce_port(reducer),
1667 None => instance.mailbox().open_once_port::<u64>(),
1668 };
1669 let reply_port_ref = reply_port_handle.bind();
1670
1671 let message = TestMessage::CastAndReplyOnce {
1672 arg: "abc".to_string(),
1673 reply_to: reply_port_ref.clone(),
1674 };
1675
1676 clear_collected_tree();
1677 actor_mesh_ref.cast(instance, message).unwrap();
1678
1679 let mut reply_tos = vec![];
1680 for _ in proc_mesh.extent().points() {
1681 let msg = rx.recv().await.expect("missing");
1682 match msg {
1683 TestMessage::CastAndReplyOnce { arg, reply_to } => {
1684 assert_eq!(arg, "abc");
1685 if has_reducer {
1686 assert_ne!(reply_to, reply_port_ref);
1688 } else {
1689 assert_eq!(reply_to, reply_port_ref);
1691 }
1692 reply_tos.push(reply_to);
1693 }
1694 _ => {
1695 panic!("unexpected message: {:?}", msg);
1696 }
1697 }
1698 }
1699
1700 OncePortMeshSetupV1 {
1701 instance,
1702 reply_rx,
1703 reply_tos,
1704 _reply_port_ref: reply_port_ref,
1705 host_mesh,
1706 }
1707 }
1708
1709 async fn execute_cast_and_reply_once_v1() {
1710 let mut setup = setup_once_port_mesh::<NoneAccumulator>(None).await;
1714
1715 let num_replies = setup.reply_tos.len();
1718 for (i, reply_to) in setup.reply_tos.into_iter().enumerate() {
1719 reply_to.post(setup.instance, i as u64);
1720 }
1721
1722 let result = setup.reply_rx.recv().await.unwrap();
1724 assert!(result < num_replies as u64);
1726
1727 let _ = setup.host_mesh.shutdown(setup.instance).await;
1728 }
1729
1730 #[async_timed_test(timeout_secs = 60)]
1731 async fn test_cast_and_reply_once_v1() {
1732 execute_cast_and_reply_once_v1().await
1733 }
1734
1735 #[async_timed_test(timeout_secs = 60)]
1736 async fn test_cast_and_reply_once_v1_p2p() {
1737 let config = hyperactor_config::global::lock();
1738 let _guard = config.override_key(ENABLE_NATIVE_V1_CASTING, true);
1739 let _guard2 = config.override_key(
1740 hyperactor::config::ENABLE_DEST_ACTOR_REORDERING_BUFFER,
1741 true,
1742 );
1743 let _guard3 = config.override_key(crate::config::V1_CAST_POINT_TO_POINT_THRESHOLD, 1024);
1744 execute_cast_and_reply_once_v1().await
1745 }
1746
1747 async fn execute_cast_and_accum_once_v1() {
1748 let mut setup = setup_once_port_mesh(Some(accum::sum::<u64>())).await;
1752
1753 let mut expected_sum = 0u64;
1755 for (i, reply_to) in setup.reply_tos.into_iter().enumerate() {
1756 reply_to.post(setup.instance, i as u64);
1757 expected_sum += i as u64;
1758 }
1759
1760 let result = setup.reply_rx.recv().await.unwrap();
1762 assert_eq!(result, expected_sum);
1763
1764 let _ = setup.host_mesh.shutdown(setup.instance).await;
1765 }
1766
1767 #[async_timed_test(timeout_secs = 60)]
1768 async fn test_cast_and_accum_once_v1() {
1769 execute_cast_and_accum_once_v1().await
1770 }
1771
1772 #[async_timed_test(timeout_secs = 60)]
1773 async fn test_cast_and_accum_once_v1_p2p() {
1774 let config = hyperactor_config::global::lock();
1775 let _guard = config.override_key(ENABLE_NATIVE_V1_CASTING, true);
1776 let _guard2 = config.override_key(
1777 hyperactor::config::ENABLE_DEST_ACTOR_REORDERING_BUFFER,
1778 true,
1779 );
1780 let _guard3 = config.override_key(crate::config::V1_CAST_POINT_TO_POINT_THRESHOLD, 1024);
1781 execute_cast_and_accum_once_v1().await
1782 }
1783
1784 #[async_timed_test(timeout_secs = 60)]
1785 async fn test_unsplit_port_not_split() {
1786 let instance = crate::testing::instance();
1787 let mut host_mesh = local_host_mesh(8).await;
1788 let proc_mesh = host_mesh
1789 .spawn(instance, "test", extent!(gpu = 8), None, None)
1790 .await
1791 .unwrap();
1792
1793 let (tx, mut rx) = hyperactor::mailbox::open_port(instance);
1794 let params = TestActorParams {
1795 forward_port: tx.bind(),
1796 };
1797 let actor_name =
1798 crate::mesh_id::ActorMeshId::instance(hyperactor::id::Label::new("test").unwrap());
1799 let actor_mesh: ActorMesh<TestActor> = proc_mesh
1800 .spawn_with_name(&instance, actor_name, ¶ms, None, true)
1801 .await
1802 .unwrap();
1803 let (reply_port_handle, mut reply_rx) = open_port::<u64>(instance);
1804 let reply_port_ref = reply_port_handle.bind().unsplit();
1805
1806 let message = TestMessage::CastWithUnsplitPort {
1807 reply_to: reply_port_ref.clone(),
1808 };
1809
1810 clear_collected_tree();
1811 actor_mesh.cast(instance, message).unwrap();
1812
1813 let num_points = proc_mesh.extent().points().count();
1815 for _ in 0..num_points {
1816 let msg = rx.recv().await.expect("missing");
1817 match msg {
1818 TestMessage::CastWithUnsplitPort { reply_to } => {
1819 assert_eq!(
1820 reply_to.port_addr(),
1821 reply_port_ref.port_addr(),
1822 "unsplit port should not be replaced by a comm actor split port"
1823 );
1824 }
1825 _ => panic!("unexpected message: {:?}", msg),
1826 }
1827 }
1828
1829 for _ in 0..8 {
1832 let val = reply_rx.recv().await.unwrap();
1833 assert_eq!(val, 42);
1834 }
1835 let _ = host_mesh.shutdown(instance).await;
1836 }
1837
1838 struct SenderCaptureSetup {
1845 instance: &'static Instance<testing::TestRootClient>,
1846 host_mesh: HostMesh,
1847 proc_mesh: ProcMesh,
1848 actor_mesh: ActorMesh<SenderCapturingActor>,
1849 rx: hyperactor::mailbox::PortReceiver<CapturedSender>,
1850 }
1851
1852 async fn setup_sender_capture_mesh(num_ranks: usize) -> SenderCaptureSetup {
1853 let instance = crate::testing::instance();
1854 let host_mesh = local_host_mesh(num_ranks).await;
1855 let proc_mesh = host_mesh
1856 .spawn(
1857 instance,
1858 "sender_capture",
1859 extent!(gpu = num_ranks),
1860 None,
1861 None,
1862 )
1863 .await
1864 .unwrap();
1865 let (tx, rx) = hyperactor::mailbox::open_port(instance);
1866 let params = SenderCapturingActorParams {
1867 forward_port: tx.bind(),
1868 };
1869 let actor_name = crate::mesh_id::ActorMeshId::instance(
1870 hyperactor::id::Label::new("sender_capture").unwrap(),
1871 );
1872 let actor_mesh: ActorMesh<SenderCapturingActor> = proc_mesh
1873 .spawn_with_name(&instance, actor_name, ¶ms, None, true)
1874 .await
1875 .unwrap();
1876 SenderCaptureSetup {
1877 instance,
1878 host_mesh,
1879 proc_mesh,
1880 actor_mesh,
1881 rx,
1882 }
1883 }
1884
1885 #[async_timed_test(timeout_secs = 60)]
1888 async fn test_direct_send_sender_actor_id() {
1889 let mut setup = setup_sender_capture_mesh(1).await;
1890
1891 let actor_ref = setup.actor_mesh.values().next().unwrap();
1892 actor_ref.post(setup.instance, SenderCaptureMsg());
1893
1894 let captured = setup.rx.recv().await.unwrap();
1895 assert_eq!(
1896 captured,
1897 CapturedSender(Some(setup.instance.self_addr().clone())),
1898 );
1899
1900 let _ = setup.host_mesh.shutdown(setup.instance).await;
1901 }
1902
1903 #[async_timed_test(timeout_secs = 60)]
1908 async fn test_v1_native_p2p_sender_actor_id() {
1909 let config = hyperactor_config::global::lock();
1910 let _g1 = config.override_key(ENABLE_NATIVE_V1_CASTING, true);
1911 let _g2 = config.override_key(
1912 hyperactor::config::ENABLE_DEST_ACTOR_REORDERING_BUFFER,
1913 true,
1914 );
1915 let _g3 = config.override_key(crate::config::V1_CAST_POINT_TO_POINT_THRESHOLD, 1024);
1916
1917 let mut setup = setup_sender_capture_mesh(1).await;
1918 setup
1919 .actor_mesh
1920 .cast(setup.instance, SenderCaptureMsg())
1921 .unwrap();
1922
1923 let captured = setup.rx.recv().await.unwrap();
1924 assert_eq!(
1925 captured,
1926 CapturedSender(Some(setup.instance.self_addr().clone())),
1927 );
1928
1929 let _ = setup.host_mesh.shutdown(setup.instance).await;
1930 }
1931
1932 #[async_timed_test(timeout_secs = 60)]
1937 async fn test_v1_native_tree_sender_actor_id() {
1938 let config = hyperactor_config::global::lock();
1939 let _g1 = config.override_key(ENABLE_NATIVE_V1_CASTING, true);
1940 let _g2 = config.override_key(
1941 hyperactor::config::ENABLE_DEST_ACTOR_REORDERING_BUFFER,
1942 true,
1943 );
1944 let _g3 = config.override_key(crate::config::V1_CAST_POINT_TO_POINT_THRESHOLD, 0);
1945
1946 let mut setup = setup_sender_capture_mesh(8).await;
1947 setup
1948 .actor_mesh
1949 .cast(setup.instance, SenderCaptureMsg())
1950 .unwrap();
1951
1952 let num_points = setup.proc_mesh.extent().points().count();
1953 let expected = CapturedSender(Some(setup.instance.self_addr().clone()));
1954 for _ in 0..num_points {
1955 let captured = setup.rx.recv().await.unwrap();
1956 assert_eq!(captured, expected);
1957 }
1958
1959 let _ = setup.host_mesh.shutdown(setup.instance).await;
1960 }
1961}