1use crate::prelude::*;
66use crate::test_common::wait_for_peers;
67use citadel_io::tokio::sync::Mutex;
68use citadel_user::prelude::UserIdentifierExt;
69use futures::{Future, StreamExt};
70use std::marker::PhantomData;
71use std::pin::Pin;
72use std::sync::atomic::{AtomicBool, Ordering};
73use std::sync::Arc;
74use uuid::Uuid;
75
76pub struct BroadcastKernel<'a, F, Fut, R: Ratchet> {
82 inner_kernel: Box<dyn NetKernel<R> + 'a>,
83 shared: Arc<BroadcastShared>,
84 _pd: PhantomData<fn() -> (F, Fut)>,
85}
86
87pub struct BroadcastShared {
88 route_registers: AtomicBool,
89 register_rx:
90 citadel_io::Mutex<Option<citadel_io::tokio::sync::mpsc::UnboundedReceiver<PeerSignal>>>,
91 register_tx: citadel_io::tokio::sync::mpsc::UnboundedSender<PeerSignal>,
92}
93
94pub enum GroupInitRequestType {
102 Create {
107 local_user: UserIdentifier,
108 invite_list: Vec<UserIdentifier>,
109 group_id: Uuid,
110 accept_registrations: bool,
111 },
112 Join {
120 local_user: UserIdentifier,
121 owner: UserIdentifier,
122 group_id: Uuid,
123 do_peer_register: bool,
124 },
125}
126
127#[async_trait]
128impl<'a, F, Fut, R: Ratchet> PrefabFunctions<'a, GroupInitRequestType, R>
129 for BroadcastKernel<'a, F, Fut, R>
130where
131 F: FnOnce(GroupChannel, CitadelClientServerConnection<R>) -> Fut + Send + 'a,
132 Fut: Future<Output = Result<(), NetworkError>> + Send + 'a,
133{
134 type UserLevelInputFunction = F;
135 type SharedBundle = Arc<BroadcastShared>;
136
137 fn get_shared_bundle(&self) -> Self::SharedBundle {
138 self.shared.clone()
139 }
140
141 #[allow(unreachable_code, clippy::blocks_in_conditions)]
142 #[cfg_attr(
143 feature = "localhost-testing",
144 tracing::instrument(level = "trace", target = "citadel", skip_all, ret, err(Debug))
145 )]
146 async fn on_c2s_channel_received(
147 connect_success: CitadelClientServerConnection<R>,
148 arg: GroupInitRequestType,
149 fx: Self::UserLevelInputFunction,
150 shared: Arc<BroadcastShared>,
151 ) -> Result<(), NetworkError> {
152 let session_cid = connect_success.cid;
153 wait_for_peers().await;
154 let mut creator_only_accept_inbound_registers = false;
155
156 let mut is_owner = false;
157 let request = match arg {
158 GroupInitRequestType::Create {
159 local_user,
160 invite_list,
161 group_id,
162 accept_registrations,
163 } => {
164 is_owner = true;
165 let mut peers_registered = vec![];
167
168 for peer in &invite_list {
169 let peer = peer
170 .search_peer(session_cid, connect_success.account_manager())
171 .await?
172 .ok_or_else(|| {
173 citadel_io::error!(
174 citadel_io::ErrorCode::BroadcastCreateUserNotRegistered,
175 format!("{peer:?}"),
176 format!("{local_user:?}")
177 )
178 })?;
179
180 peers_registered.push(peer.cid)
181 }
182
183 creator_only_accept_inbound_registers = accept_registrations;
184
185 GroupBroadcast::Create {
186 initial_invitees: peers_registered,
187 options: MessageGroupOptions {
188 group_type: GroupType::Public,
189 id: group_id.as_u128(),
190 ..Default::default()
191 },
192 }
193 }
194
195 GroupInitRequestType::Join {
196 local_user,
197 owner,
198 group_id,
199 do_peer_register,
200 } => {
201 let owner_orig = owner;
203 let owner_find = owner_orig
204 .search_peer(session_cid, connect_success.account_manager())
205 .await?;
206
207 let owner = if let Some(owner) = owner_find {
208 Some(owner)
209 } else if do_peer_register {
210 let handle = connect_success
211 .propose_target(local_user.clone(), owner_orig.clone())
212 .await?;
213 let registration = handle.register_to_peer().await?;
217 if let Some(reason) = registration.refusal_reason() {
218 return Err(NetworkError::generic(format!(
219 "Cannot join the group: registering with its owner failed: {reason}"
220 )));
221 }
222 owner_orig
224 .search_peer(session_cid, connect_success.account_manager())
225 .await?
226 } else {
227 None
228 };
229
230 let owner = owner.ok_or_else(|| {
231 citadel_io::error!(
232 citadel_io::ErrorCode::BroadcastJoinUserNotRegistered,
233 format!("{owner_orig:?}"),
234 format!("{local_user:?}")
235 )
236 })?;
237
238 let expected_message_group_key = MessageGroupKey {
239 cid: owner.cid,
240 mgid: group_id.as_u128(),
241 };
242
243 let mut retries = 0;
245 let group_owner_handle = connect_success
246 .propose_target(local_user.clone(), owner.cid)
247 .await?;
248 loop {
249 let owned_groups = group_owner_handle.list_owned_groups().await?;
250 if owned_groups.contains(&expected_message_group_key) {
251 break;
252 } else {
253 citadel_io::time::sleep(std::time::Duration::from_secs(2u64.pow(retries)))
254 .await;
255
256 retries += 1;
257 if retries > 4 {
258 return Err(citadel_io::error!(
259 citadel_io::ErrorCode::BroadcastOwnerGroupMissing,
260 citadel_io::Dbg(owner),
261 citadel_io::Dbg(group_id)
262 ));
263 }
264 }
265 }
266
267 GroupBroadcast::RequestJoin {
268 sender: local_user.get_cid(),
269 key: expected_message_group_key,
270 }
271 }
272 };
273
274 let request = NodeRequest::GroupBroadcastCommand(GroupBroadcastCommand {
275 session_cid,
276 command: request,
277 });
278
279 let subscription = &Mutex::new(Some(
280 connect_success.send_callback_subscription(request).await?,
281 ));
282
283 log::trace!(target: "citadel", "Peer {session_cid} is attempting to join group");
284 let acceptor_task = if creator_only_accept_inbound_registers {
285 shared.route_registers.store(true, Ordering::Relaxed);
286 let mut reg_rx = shared.register_rx.lock().take().unwrap();
287 let remote = connect_success.remote_ref().clone();
288 Box::pin(async move {
289 let mut subscription = subscription.lock().await.take().unwrap();
290 let mut count_registered = 0;
292 loop {
293 let post_register = citadel_io::tokio::select! {
294 reg_request = reg_rx.recv() => {
295 reg_request.ok_or_else(|| citadel_io::error!(citadel_io::ErrorCode::BroadcastStreamEndedUnexpectedly, "reg_rx"))?
296 },
297
298 reg_request2 = subscription.next() => {
299 let signal = reg_request2.ok_or_else(|| citadel_io::error!(citadel_io::ErrorCode::BroadcastStreamEndedUnexpectedly, "subscription"))?;
300 if let NodeResult::PeerEvent(PeerEvent { event: sig @ PeerSignal::PostRegister { .. }, .. }) = &signal {
301 sig.clone()
302 } else {
303 continue;
304 }
305 }
306 };
307
308 log::trace!(target: "citadel", "ACCEPTOR {session_cid} RECV reg_request: {post_register:?}");
309 if let PeerSignal::PostRegister {
310 peer_conn_type: peer_conn,
311 inviter_username: _,
312 invitee_username: _,
313 ticket_opt: _,
314 invitee_response: None,
315 } = &post_register
316 {
317 let cid = peer_conn.get_original_target_cid();
318 if cid != session_cid {
319 log::warn!(target: "citadel", "Received the wrong CID. Will not accept request");
320 continue;
321 }
322
323 let _ = responses::peer_register(post_register, true, &remote).await?;
324 if cfg!(feature = "localhost-testing") {
325 count_registered += 1;
326 if count_registered == crate::test_common::num_local_test_peers() - 1 {
327 break;
329 }
330 }
331 }
332 }
333
334 Ok::<_, NetworkError>(())
335 })
336 as Pin<
337 Box<
338 dyn futures::Future<
339 Output = Result<(), citadel_proto::prelude::NetworkError>,
340 > + Send,
341 >,
342 >
343 } else {
344 Box::pin(async move { Ok::<_, NetworkError>(()) })
345 as Pin<
346 Box<
347 dyn futures::Future<
348 Output = Result<(), citadel_proto::prelude::NetworkError>,
349 > + Send,
350 >,
351 >
352 };
353
354 let mut lock = subscription.lock().await;
355 let subscription = lock.as_mut().unwrap();
356 while let Some(event) = subscription.next().await {
357 match event.into_result()? {
358 NodeResult::PeerEvent(PeerEvent {
359 event: ref ps @ PeerSignal::PostRegister { .. },
360 ..
361 }) => {
362 shared
363 .register_tx
364 .send(ps.clone())
365 .map_err(|err| NetworkError::generic(err.to_string()))?;
366 }
367 NodeResult::GroupChannelCreated(GroupChannelCreated {
368 ticket: _,
369 channel,
370 session_cid: _,
371 }) => {
372 drop(lock);
375 return if is_owner {
376 citadel_io::tokio::try_join!(fx(channel, connect_success), acceptor_task)
377 .map(|_| ())
378 } else {
379 fx(channel, connect_success).await.map(|_| ())
380 };
381 }
382
383 NodeResult::GroupEvent(GroupEvent {
384 session_cid: _,
385 ticket: _,
386 event: GroupBroadcast::CreateResponse { key: None },
387 }) => {
388 return Err(citadel_io::error!(
389 citadel_io::ErrorCode::BroadcastCreateGroupFailed
390 ))
391 }
392
393 _ => {}
394 }
395 }
396
397 Ok(())
398 }
399
400 fn construct(kernel: Box<dyn NetKernel<R> + 'a>) -> Self {
401 let (tx, rx) = citadel_io::tokio::sync::mpsc::unbounded_channel();
402 Self {
403 shared: Arc::new(BroadcastShared {
404 route_registers: AtomicBool::new(false),
405 register_rx: citadel_io::Mutex::new(Some(rx)),
406 register_tx: tx,
407 }),
408 inner_kernel: kernel,
409 _pd: Default::default(),
410 }
411 }
412}
413
414#[async_trait]
415impl<F, Fut, R: Ratchet> NetKernel<R> for BroadcastKernel<'_, F, Fut, R> {
416 fn load_remote(&mut self, node_remote: NodeRemote<R>) -> Result<(), NetworkError> {
417 self.inner_kernel.load_remote(node_remote)
418 }
419
420 async fn on_start(&self) -> Result<(), NetworkError> {
421 self.inner_kernel.on_start().await
422 }
423
424 async fn on_node_event_received(&self, message: NodeResult<R>) -> Result<(), NetworkError> {
425 if let NodeResult::PeerEvent(PeerEvent {
426 event: ps @ PeerSignal::PostRegister { .. },
427 ..
428 }) = &message
429 {
430 if self.shared.route_registers.load(Ordering::Relaxed) {
431 return self
432 .shared
433 .register_tx
434 .send(ps.clone())
435 .map_err(|err| NetworkError::generic(err.to_string()));
436 }
437 }
438
439 self.inner_kernel.on_node_event_received(message).await
440 }
441
442 async fn on_stop(&mut self) -> Result<(), NetworkError> {
443 self.inner_kernel.on_stop().await
444 }
445}
446
447#[cfg(all(test, feature = "localhost-testing"))]
448mod tests {
449 use crate::prefabs::client::broadcast::{BroadcastKernel, GroupInitRequestType};
450 use crate::prefabs::client::peer_connection::PeerConnectionKernel;
451 use crate::prefabs::client::DefaultServerConnectionSettingsBuilder;
452 use crate::prelude::*;
453 use crate::test_common::{server_info, wait_for_peers, TestBarrier};
454 use citadel_io::tokio;
455 use futures::prelude::stream::FuturesUnordered;
456 use futures::TryStreamExt;
457 use rstest::rstest;
458 use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
459 use uuid::Uuid;
460
461 #[citadel_io::tokio::test(flavor = "multi_thread")]
462 async fn group_connect_list_members() -> Result<(), Box<dyn std::error::Error>> {
463 let peer_count = 3;
464 assert!(peer_count > 1);
465 citadel_logging::setup_log();
466 TestBarrier::setup(peer_count);
467
468 let client_success = &AtomicUsize::new(0);
469 let (server, server_addr) = server_info::<StackedRatchet>();
470
471 let client_kernels = FuturesUnordered::new();
472 let total_peers = (0..peer_count)
473 .map(|_| Uuid::new_v4())
474 .collect::<Vec<Uuid>>();
475 let group_id = Uuid::new_v4();
476
477 for idx in 0..peer_count {
478 let uuid = total_peers.get(idx).cloned().unwrap();
479
480 let request = if idx == 0 {
481 GroupInitRequestType::Create {
483 local_user: UserIdentifier::from(uuid),
484 invite_list: vec![],
485 group_id,
486 accept_registrations: true,
487 }
488 } else {
489 GroupInitRequestType::Join {
490 local_user: UserIdentifier::from(uuid),
491 owner: total_peers.first().cloned().unwrap().into(),
492 group_id,
493 do_peer_register: true,
494 }
495 };
496
497 let server_connection_settings =
498 DefaultServerConnectionSettingsBuilder::transient_with_id(server_addr, uuid)
499 .build()
500 .unwrap();
501
502 let client_kernel = BroadcastKernel::new(
503 server_connection_settings,
504 request,
505 move |channel, connection| async move {
506 wait_for_peers().await;
507 log::trace!(target: "citadel", "***GROUP PEER {}={}={} CONNECT SUCCESS***", idx, uuid, connection.conn_type.get_session_cid());
508
509 let owned_groups = connection.list_owned_groups().await.unwrap();
510
511 if idx == 0 {
512 assert_eq!(owned_groups.len(), 1);
513 } else {
514 assert_eq!(owned_groups.len(), 0);
515 }
516
517 log::trace!(target: "citadel", "Peer {idx}={} is COMPLETE!", connection.conn_type.get_session_cid());
518
519 let _ = client_success.fetch_add(1, Ordering::Relaxed);
520 wait_for_peers().await;
521 drop(channel);
522 connection.shutdown_kernel().await
523 },
524 );
525
526 let client = DefaultNodeBuilder::default().build(client_kernel).unwrap();
527
528 client_kernels.push(async move { client.await.map(|_| ()) });
529 }
530
531 let clients = Box::pin(async move { client_kernels.try_collect::<()>().await.map(|_| ()) });
532
533 let res = futures::future::try_select(server, clients).await;
534 if let Err(err) = res {
535 return match err {
536 futures::future::Either::Left(left) => Err(left.0.into_string().into()),
537 futures::future::Either::Right(right) => Err(right.0.into_string().into()),
538 };
539 }
540
541 assert_eq!(client_success.load(Ordering::Relaxed), peer_count);
542 Ok(())
543 }
544
545 #[rstest]
546 #[case(2)]
547 #[timeout(std::time::Duration::from_secs(90))]
548 #[citadel_io::tokio::test(flavor = "multi_thread")]
549 async fn test_manual_group_connect(
550 #[case] peer_count: usize,
551 ) -> Result<(), Box<dyn std::error::Error>> {
552 assert!(peer_count > 1);
557 citadel_logging::setup_log();
558 TestBarrier::setup(peer_count);
559
560 let client_success = &AtomicBool::new(false);
561 let receiver_success = &AtomicBool::new(false);
562
563 let (server, server_addr) = server_info::<StackedRatchet>();
564
565 let client_kernels = FuturesUnordered::new();
566 let total_peers = (0..peer_count)
567 .map(|_| Uuid::new_v4())
568 .collect::<Vec<Uuid>>();
569
570 for idx in 0..peer_count {
571 let uuid = total_peers.get(idx).cloned().unwrap();
572 let peers = total_peers
573 .clone()
574 .into_iter()
575 .filter(|r| r != &uuid)
576 .map(UserIdentifier::from)
577 .collect::<Vec<UserIdentifier>>();
578
579 let server_connection_settings =
580 DefaultServerConnectionSettingsBuilder::transient_with_id(server_addr, uuid)
581 .build()
582 .unwrap();
583
584 let client_kernel = PeerConnectionKernel::new(
585 server_connection_settings,
586 peers,
587 move |mut results, remote| async move {
588 let _sender = remote.conn_type.get_session_cid();
589 let mut signals = remote.get_unprocessed_signals_receiver().unwrap();
590
591 wait_for_peers().await;
592 let conn = results.recv().await.unwrap()?;
593 log::trace!(target: "citadel", "User {uuid} received {conn:?}");
594
595 if idx == 0 {
597 let _channel = remote
598 .create_group(Some(vec![conn.channel.get_peer_cid().into()]))
599 .await?;
600 log::info!(target: "citadel", "The designated node has finished creating a group");
601
602 wait_for_peers().await;
603 client_success.store(true, Ordering::Relaxed);
604 return remote.shutdown_kernel().await;
605 } else {
606 while let Some(evt) = signals.recv().await {
608 log::info!(target: "citadel", "Received unprocessed signal: {evt:?}");
609 match evt {
610 NodeResult::GroupEvent(GroupEvent {
611 session_cid: _,
612 ticket: _,
613 event:
614 GroupBroadcast::Invitation {
615 sender: _,
616 key: _key,
617 },
618 }) => {
619 let _ =
620 crate::responses::group_invite(evt, true, &remote.inner)
621 .await?;
622 }
623
624 NodeResult::GroupChannelCreated(GroupChannelCreated {
625 ticket: _,
626 channel: _chan,
627 session_cid: _,
628 }) => {
629 receiver_success.store(true, Ordering::Relaxed);
630 log::trace!(target: "citadel", "***PEER {uuid} CONNECT***");
631 wait_for_peers().await;
632 return remote.shutdown_kernel().await;
633 }
634
635 val => {
636 log::warn!(target: "citadel", "Unhandled response: {val:?}")
637 }
638 }
639 }
640 }
641
642 Err(citadel_io::error!(
643 citadel_io::ErrorCode::BroadcastStreamEndedUnexpectedly,
644 "signals_recv"
645 ))
646 },
647 );
648
649 let client = DefaultNodeBuilder::default().build(client_kernel).unwrap();
650 client_kernels.push(async move { client.await.map(|_| ()) });
651 }
652
653 let clients = Box::pin(async move { client_kernels.try_collect::<()>().await.map(|_| ()) });
654
655 if let Err(err) = futures::future::try_select(server, clients).await {
656 return match err {
657 futures::future::Either::Left(res) => Err(res.0.into_string().into()),
658 futures::future::Either::Right(res) => Err(res.0.into_string().into()),
659 };
660 }
661
662 assert!(client_success.load(Ordering::Relaxed));
663 assert!(receiver_success.load(Ordering::Relaxed));
664 Ok(())
665 }
666
667 #[citadel_io::tokio::test(flavor = "multi_thread")]
672 async fn group_command_hierarchy_superior_reads_subordinate(
673 ) -> Result<(), Box<dyn std::error::Error>> {
674 use crate::prelude::GroupBroadcastPayload;
675 use citadel_types::crypto::SecBuffer;
676 use citadel_types::proto::{
677 CommandPath, GroupHierarchyMode, MessageGroupOptions, ReadPolicy,
678 };
679 use std::collections::HashMap;
680
681 let peer_count = 2;
682 citadel_logging::setup_log();
683 TestBarrier::setup(peer_count);
684
685 let owner_read = &AtomicBool::new(false);
686 let (server, server_addr) = server_info::<StackedRatchet>();
687 let client_kernels = FuturesUnordered::new();
688 let total_peers = (0..peer_count)
689 .map(|_| Uuid::new_v4())
690 .collect::<Vec<Uuid>>();
691
692 for idx in 0..peer_count {
693 let uuid = total_peers.get(idx).cloned().unwrap();
694 let peers = total_peers
695 .clone()
696 .into_iter()
697 .filter(|r| r != &uuid)
698 .map(UserIdentifier::from)
699 .collect::<Vec<UserIdentifier>>();
700 let server_connection_settings =
701 DefaultServerConnectionSettingsBuilder::transient_with_id(server_addr, uuid)
702 .build()
703 .unwrap();
704
705 let client_kernel = PeerConnectionKernel::new(
706 server_connection_settings,
707 peers,
708 move |mut results, remote| async move {
709 let mut signals = remote.get_unprocessed_signals_receiver().unwrap();
710 wait_for_peers().await;
711 let conn = results.recv().await.unwrap()?;
712
713 if idx == 0 {
714 let sub_cid: u64 = conn.channel.get_peer_cid();
716 let mut ranks = HashMap::new();
717 let _ = ranks.insert(sub_cid, CommandPath::parse("/alpha"));
718 let options = MessageGroupOptions {
719 hierarchy: GroupHierarchyMode::CommandHierarchy {
720 read_policy: ReadPolicy::SuperiorOnly,
721 ranks,
722 },
723 ..Default::default()
724 };
725 let mut channel = remote
726 .create_group_with_options(Some(vec![sub_cid.into()]), options)
727 .await?;
728
729 loop {
731 match channel.recv().await {
732 Some(GroupBroadcastPayload::Message { payload, sender: _ }) => {
733 assert_eq!(payload.as_ref(), b"sitrep from subordinate");
734 owner_read.store(true, Ordering::Relaxed);
735 break;
736 }
737 Some(_) => continue,
738 None => break,
739 }
740 }
741 wait_for_peers().await;
742 return remote.shutdown_kernel().await;
743 }
744
745 while let Some(evt) = signals.recv().await {
747 match evt {
748 NodeResult::GroupEvent(GroupEvent {
749 event: GroupBroadcast::Invitation { .. },
750 ..
751 }) => {
752 let _ = crate::responses::group_invite(evt, true, &remote.inner)
753 .await?;
754 }
755 NodeResult::GroupChannelCreated(GroupChannelCreated {
756 channel,
757 ..
758 }) => {
759 channel
760 .send_message(SecBuffer::from(
761 b"sitrep from subordinate".to_vec(),
762 ))
763 .await?;
764 wait_for_peers().await;
765 return remote.shutdown_kernel().await;
766 }
767 _ => {}
768 }
769 }
770
771 Err(citadel_io::error!(
772 citadel_io::ErrorCode::BroadcastStreamEndedUnexpectedly,
773 "signals"
774 ))
775 },
776 );
777 let client = DefaultNodeBuilder::default().build(client_kernel).unwrap();
778 client_kernels.push(async move { client.await.map(|_| ()) });
779 }
780
781 let clients = Box::pin(async move { client_kernels.try_collect::<()>().await.map(|_| ()) });
782 if let Err(err) = futures::future::try_select(server, clients).await {
783 return match err {
784 futures::future::Either::Left(res) => Err(res.0.into_string().into()),
785 futures::future::Either::Right(res) => Err(res.0.into_string().into()),
786 };
787 }
788
789 assert!(
790 owner_read.load(Ordering::Relaxed),
791 "owner (hierarchy root) must read the subordinate's DHE message"
792 );
793 Ok(())
794 }
795}