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