1use crate::prelude::results::PeerConnectSuccess;
68use crate::prelude::*;
69use crate::test_common::wait_for_peers;
70use citadel_io::tokio::sync::mpsc::{Receiver, UnboundedSender};
71use citadel_io::{tokio, Mutex};
72use citadel_proto::prelude::async_trait;
73use citadel_user::hypernode_account::UserIdentifierExt;
74use futures::stream::FuturesUnordered;
75use futures::TryStreamExt;
76use std::collections::HashMap;
77use std::future::Future;
78use std::marker::PhantomData;
79use std::sync::Arc;
80use uuid::Uuid;
81
82pub struct PeerConnectionKernel<'a, F, Fut, R: Ratchet> {
85 inner_kernel: Box<dyn NetKernel<R> + 'a>,
86 shared: Shared,
87 _pd: PhantomData<fn() -> (F, Fut)>,
89}
90
91#[derive(Clone)]
92#[doc(hidden)]
93pub struct Shared {
94 active_peer_conns: Arc<Mutex<HashMap<PeerConnectionType, PeerContext>>>,
95}
96
97struct PeerContext {
98 #[allow(dead_code)]
99 conn_type: PeerConnectionType,
100 send_file_transfer_tx: UnboundedSender<ObjectTransferHandler>,
101}
102
103#[derive(Debug)]
104pub struct FileTransferHandleRx {
105 pub inner: citadel_io::tokio::sync::mpsc::UnboundedReceiver<ObjectTransferHandler>,
106 pub conn_type: VirtualTargetType,
107}
108
109impl FileTransferHandleRx {
110 pub fn accept_all(mut self) {
112 let task = tokio::task::spawn(async move {
113 let rx = &mut self.inner;
114 while let Some(mut handle) = rx.recv().await {
115 let task = tokio::task::spawn(async move {
116 if let Err(err) = handle.exhaust_stream().await {
117 let orientation = handle.orientation;
118 log::warn!(target: "citadel", "Error background handling of file transfer for {orientation:?}: {err:?}");
119 }
120 });
121
122 drop(task);
123 }
124 });
125
126 drop(task);
127 }
128}
129
130impl std::ops::Deref for FileTransferHandleRx {
131 type Target = citadel_io::tokio::sync::mpsc::UnboundedReceiver<ObjectTransferHandler>;
132
133 fn deref(&self) -> &Self::Target {
134 &self.inner
135 }
136}
137
138impl std::ops::DerefMut for FileTransferHandleRx {
139 fn deref_mut(&mut self) -> &mut Self::Target {
140 &mut self.inner
141 }
142}
143
144impl Drop for FileTransferHandleRx {
145 fn drop(&mut self) {
146 log::trace!(target: "citadel", "Dropping file transfer handle receiver {:?}", self.conn_type);
147 }
148}
149
150#[async_trait]
151impl<F, Fut, R: Ratchet> NetKernel<R> for PeerConnectionKernel<'_, F, Fut, R> {
152 fn load_remote(&mut self, server_remote: NodeRemote<R>) -> Result<(), NetworkError> {
153 self.inner_kernel.load_remote(server_remote)
154 }
155
156 async fn on_start(&self) -> Result<(), NetworkError> {
157 self.inner_kernel.on_start().await
158 }
159
160 #[allow(clippy::collapsible_else_if)]
161 async fn on_node_event_received(&self, message: NodeResult<R>) -> Result<(), NetworkError> {
162 match message {
163 NodeResult::ObjectTransferHandle(ObjectTransferHandle {
164 ticket: _,
165 handle,
166 session_cid,
167 }) => {
168 let is_revfs = matches!(
169 handle.metadata.transfer_type,
170 TransferType::RemoteEncryptedVirtualFilesystem { .. }
171 );
172 let active_peers = self.shared.active_peer_conns.lock();
173 let v_conn = if is_revfs {
174 let peer_cid = if session_cid != handle.source {
175 handle.source
176 } else {
177 handle.receiver
178 };
179 PeerConnectionType::LocalGroupPeer {
180 session_cid,
181 peer_cid,
182 }
183 } else {
184 if matches!(
185 handle.orientation,
186 ObjectTransferOrientation::Receiver { .. }
187 ) {
188 PeerConnectionType::LocalGroupPeer {
189 session_cid,
190 peer_cid: handle.source,
191 }
192 } else {
193 PeerConnectionType::LocalGroupPeer {
194 session_cid,
195 peer_cid: handle.receiver,
196 }
197 }
198 };
199
200 if let Some(peer_ctx) = active_peers.get(&v_conn) {
201 if let Err(err) = peer_ctx.send_file_transfer_tx.send(handle) {
202 log::warn!(target: "citadel", "Error forwarding file transfer handle: {:?}", err.to_string());
203 }
204 } else {
205 log::warn!(target: "citadel", "Unable to find key for inbound file transfer handle: {:?}\n Active Peers: {:?} \n handle_source = {}, handle_receiver = {}", v_conn, active_peers.keys().cloned().collect::<Vec<_>>(), handle.source, handle.receiver);
206 }
207
208 Ok(())
209 }
210
211 unprocessed @ NodeResult::Disconnect(..) | unprocessed => {
215 self.inner_kernel.on_node_event_received(unprocessed).await
217 }
218 }
219 }
220
221 async fn on_stop(&mut self) -> Result<(), NetworkError> {
222 self.inner_kernel.on_stop().await
223 }
224}
225
226#[derive(Debug, Default, Clone)]
229pub struct PeerConnectionSetupAggregator {
230 inner: Vec<PeerConnectionSettings>,
231}
232
233#[derive(Debug, Clone)]
234struct PeerConnectionSettings {
235 id: UserIdentifier,
236 session_security_settings: SessionSecuritySettings,
237 udp_mode: UdpMode,
238 ensure_registered: bool,
239 peer_session_password: Option<PreSharedKey>,
240}
241
242pub struct AddedPeer {
243 list: PeerConnectionSetupAggregator,
244 id: UserIdentifier,
245 session_security_settings: Option<SessionSecuritySettings>,
246 ensure_registered: bool,
247 udp_mode: Option<UdpMode>,
248 peer_session_password: Option<PreSharedKey>,
249}
250
251impl AddedPeer {
252 pub fn add(mut self) -> PeerConnectionSetupAggregator {
254 let new = PeerConnectionSettings {
255 id: self.id,
256 session_security_settings: self.session_security_settings.unwrap_or_default(),
257 udp_mode: self.udp_mode.unwrap_or_default(),
258 ensure_registered: self.ensure_registered,
259 peer_session_password: self.peer_session_password,
260 };
261
262 self.list.inner.push(new);
263 self.list
264 }
265
266 pub fn with_udp_mode(mut self, udp_mode: UdpMode) -> Self {
268 self.udp_mode = Some(udp_mode);
269 self
270 }
271
272 pub fn disable_udp(self) -> Self {
274 self.with_udp_mode(UdpMode::Disabled)
275 }
276
277 pub fn enable_udp(self) -> Self {
279 self.with_udp_mode(UdpMode::Enabled)
280 }
281
282 pub fn with_session_security_settings(
284 mut self,
285 session_security_settings: SessionSecuritySettings,
286 ) -> Self {
287 self.session_security_settings = Some(session_security_settings);
288 self
289 }
290
291 pub fn ensure_registered(mut self) -> Self {
293 self.ensure_registered = true;
294 self
295 }
296
297 pub fn with_session_password<T: Into<PreSharedKey>>(mut self, password: T) -> Self {
300 self.peer_session_password = Some(password.into());
301 self
302 }
303}
304
305impl PeerConnectionSetupAggregator {
306 pub fn with_peer<T: Into<UserIdentifier>>(self, peer: T) -> PeerConnectionSetupAggregator {
315 self.with_peer_custom(peer).add()
316 }
317
318 pub fn with_peer_custom<T: Into<UserIdentifier>>(self, peer: T) -> AddedPeer {
334 AddedPeer {
335 list: self,
336 id: peer.into(),
337 ensure_registered: false,
338 session_security_settings: None,
339 udp_mode: None,
340 peer_session_password: None,
341 }
342 }
343}
344
345impl From<PeerConnectionSetupAggregator> for Vec<PeerConnectionSettings> {
346 fn from(this: PeerConnectionSetupAggregator) -> Self {
347 this.inner
348 }
349}
350
351impl From<Vec<UserIdentifier>> for PeerConnectionSetupAggregator {
352 fn from(ids: Vec<UserIdentifier>) -> Self {
353 let mut this = PeerConnectionSetupAggregator::default();
354 for peer in ids {
355 this = this.with_peer(peer);
356 }
357
358 this
359 }
360}
361
362impl From<UserIdentifier> for PeerConnectionSetupAggregator {
363 fn from(this: UserIdentifier) -> Self {
364 Self::from(vec![this])
365 }
366}
367
368impl From<Uuid> for PeerConnectionSetupAggregator {
369 fn from(user: Uuid) -> Self {
370 let user_identifier: UserIdentifier = user.into();
371 user_identifier.into()
372 }
373}
374
375impl From<String> for PeerConnectionSetupAggregator {
376 fn from(this: String) -> Self {
377 let user_identifier: UserIdentifier = this.into();
378 user_identifier.into()
379 }
380}
381
382impl From<&str> for PeerConnectionSetupAggregator {
383 fn from(this: &str) -> Self {
384 let user_identifier: UserIdentifier = this.into();
385 user_identifier.into()
386 }
387}
388
389impl From<u64> for PeerConnectionSetupAggregator {
390 fn from(this: u64) -> Self {
391 let user_identifier: UserIdentifier = this.into();
392 user_identifier.into()
393 }
394}
395
396#[async_trait]
397impl<'a, F, Fut, T: Into<PeerConnectionSetupAggregator> + Send + 'a, R: Ratchet>
398 PrefabFunctions<'a, T, R> for PeerConnectionKernel<'a, F, Fut, R>
399where
400 F: FnOnce(
401 Receiver<Result<PeerConnectSuccess<R>, NetworkError>>,
402 CitadelClientServerConnection<R>,
403 ) -> Fut
404 + Send
405 + 'a,
406 Fut: Future<Output = Result<(), NetworkError>> + Send + 'a,
407{
408 type UserLevelInputFunction = F;
409 type SharedBundle = Shared;
410
411 fn get_shared_bundle(&self) -> Self::SharedBundle {
412 self.shared.clone()
413 }
414
415 #[allow(clippy::blocks_in_conditions)]
416 #[cfg_attr(
417 feature = "localhost-testing",
418 tracing::instrument(level = "trace", target = "citadel", skip_all, ret, err(Debug))
419 )]
420 async fn on_c2s_channel_received(
421 connect_success: CitadelClientServerConnection<R>,
422 peers_to_connect: T,
423 f: Self::UserLevelInputFunction,
424 shared: Shared,
425 ) -> Result<(), NetworkError> {
426 let shared = &shared;
427 let session_cid = connect_success.cid;
428 let mut peers_already_registered = vec![];
429
430 wait_for_peers().await;
431 let peers_to_connect = peers_to_connect.into().inner;
432
433 for peer in &peers_to_connect {
434 peers_already_registered.push(
436 peer.id
437 .search_peer(session_cid, connect_success.account_manager())
438 .await?,
439 )
440 }
441
442 let remote = connect_success.clone();
443 let (tx, rx) = citadel_io::tokio::sync::mpsc::channel(peers_to_connect.len());
444 let requests = FuturesUnordered::new();
445
446 for (mutually_registered, peer_to_connect) in
447 peers_already_registered.into_iter().zip(peers_to_connect)
448 {
449 let remote = remote.clone();
452 let tx = tx.clone();
453 let PeerConnectionSettings {
454 id,
455 session_security_settings,
456 udp_mode,
457 ensure_registered,
458 peer_session_password,
459 } = peer_to_connect;
460
461 let task = async move {
462 let inner_task = async move {
463 let (file_transfer_tx, file_transfer_rx) =
464 citadel_io::tokio::sync::mpsc::unbounded_channel();
465
466 let peer_cid = if let Some(mutual_peer) = &mutually_registered {
468 mutual_peer.cid
469 } else {
470 id.get_cid()
471 };
472
473 let handle = if let Some(_already_registered) = mutually_registered {
474 remote.find_target(session_cid, id).await?
475 } else {
476 log::info!(target: "citadel", "{session_cid} proposing target {id:?} to central node");
478 let handle = remote.propose_target(session_cid, id.clone()).await?;
479 if ensure_registered {
482 loop {
483 if handle.is_peer_registered().await? {
484 break;
485 }
486 citadel_io::time::sleep(std::time::Duration::from_millis(200))
487 .await;
488 }
489 }
490
491 log::info!(target: "citadel", "{session_cid} registering to peer {id:?}");
492 let registration = handle.register_to_peer().await?;
499 if let Some(reason) = registration.refusal_reason() {
500 return Err(NetworkError::generic(format!(
501 "Cannot connect to peer {id:?}: {reason}"
502 )));
503 }
504 log::info!(target: "citadel", "{session_cid} registered to peer {id:?} registered || success -> now connecting");
505 handle
506 };
507
508 let peer_conn = PeerConnectionType::LocalGroupPeer {
511 session_cid,
512 peer_cid,
513 };
514 let peer_context = PeerContext {
515 conn_type: peer_conn,
516 send_file_transfer_tx: file_transfer_tx.clone(),
517 };
518 log::debug!(target: "citadel", "Early registering peer connection: {peer_conn:?}");
519 let _ = shared
520 .active_peer_conns
521 .lock()
522 .insert(peer_conn, peer_context);
523
524 handle
525 .connect_to_peer_custom(
526 session_security_settings,
527 udp_mode,
528 peer_session_password,
529 )
530 .await
531 .map(|mut success| {
532 let actual_peer_conn = success.channel.get_peer_conn_type().unwrap();
533
534 if actual_peer_conn != peer_conn {
537 log::debug!(target: "citadel", "Updating peer connection registration from {peer_conn:?} to {actual_peer_conn:?}");
538 let mut active_peers = shared.active_peer_conns.lock();
539 if let Some(peer_ctx) = active_peers.remove(&peer_conn) {
540 let _ = active_peers.insert(actual_peer_conn, peer_ctx);
541 }
542 }
543 success.incoming_object_transfer_handles = Some(FileTransferHandleRx {
545 inner: file_transfer_rx,
546 conn_type: actual_peer_conn.as_virtual_connection(),
547 });
548 success
549 })
550 .inspect_err(|_err| {
551 let _ = shared.active_peer_conns.lock().remove(&peer_conn);
553 })
554 };
555
556 tx.send(inner_task.await)
557 .await
558 .map_err(|err| NetworkError::generic(err.to_string()))
559 };
560
561 requests.push(Box::pin(task))
562 }
563
564 drop(tx);
567
568 let (collection_result, user_result) =
571 citadel_io::tokio::join!(requests.try_collect::<()>(), f(rx, connect_success));
572
573 collection_result?;
575 user_result
576 }
577
578 fn construct(kernel: Box<dyn NetKernel<R> + 'a>) -> Self {
579 Self {
580 inner_kernel: kernel,
581 shared: Shared {
582 active_peer_conns: Arc::new(Mutex::new(Default::default())),
583 },
584 _pd: Default::default(),
585 }
586 }
587}
588
589#[cfg(all(test, feature = "localhost-testing"))]
590mod tests {
591 use crate::prefabs::client::peer_connection::PeerConnectionKernel;
592 use crate::prefabs::client::DefaultServerConnectionSettingsBuilder;
593 use crate::prelude::*;
594 use crate::remote_ext::results::PeerConnectSuccess;
595 use crate::test_common::{server_info, wait_for_peers, TestBarrier};
596 use citadel_io::tokio;
597 use citadel_io::tokio::sync::mpsc::{Receiver, UnboundedSender};
598 use citadel_user::prelude::UserIdentifierExt;
599 use futures::stream::FuturesUnordered;
600 use futures::TryStreamExt;
601 use rstest::rstest;
602 use std::collections::HashMap;
603 use std::future::Future;
604 use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
605 use std::time::Duration;
606 use uuid::Uuid;
607
608 lazy_static::lazy_static! {
609 pub static ref PEERS: Vec<(String, String, String)> = {
610 ["alpha", "beta", "charlie", "echo", "delta", "epsilon", "foxtrot"]
611 .iter().map(|base| (format!("{base}.username"), format!("{base}.password"), format!("{base}.full_name")))
612 .collect()
613 };
614 }
615
616 #[rstest]
617 #[case(2, UdpMode::Enabled)]
618 #[case(3, UdpMode::Disabled)]
619 #[timeout(Duration::from_secs(90))]
620 #[tokio::test(flavor = "multi_thread")]
621 async fn peer_to_peer_connect(#[case] peer_count: usize, #[case] udp_mode: UdpMode) {
622 assert!(peer_count > 1);
623 citadel_logging::setup_log();
624 TestBarrier::setup(peer_count);
625
626 let client_success = &AtomicUsize::new(0);
627 let (server, server_addr) = server_info::<StackedRatchet>();
628
629 let client_kernels = FuturesUnordered::new();
630 let total_peers = (0..peer_count)
631 .map(|idx| PEERS.get(idx).unwrap().0.clone())
632 .collect::<Vec<String>>();
633
634 for idx in 0..peer_count {
635 let (username, password, full_name) = PEERS.get(idx).unwrap();
636 let peers = total_peers
637 .clone()
638 .into_iter()
639 .filter(|r| r != username)
640 .map(UserIdentifier::Username)
641 .collect::<Vec<UserIdentifier>>();
642
643 let mut agg = PeerConnectionSetupAggregator::default();
644
645 for peer in peers {
646 agg = agg
647 .with_peer_custom(peer)
648 .ensure_registered()
649 .with_udp_mode(udp_mode)
650 .with_session_security_settings(SessionSecuritySettings::default())
651 .add();
652 }
653
654 let server_connection_settings =
655 DefaultServerConnectionSettingsBuilder::credentialed_registration(
656 server_addr,
657 username,
658 full_name,
659 password.as_str(),
660 )
661 .build()
662 .unwrap();
663
664 let username = username.clone();
665
666 let client_kernel = PeerConnectionKernel::new(
667 server_connection_settings,
668 agg.clone(),
669 move |results, connection| async move {
670 log::info!(target: "citadel", "***PEER {username} CONNECTED ***");
671 let session_cid = connection.conn_type.get_session_cid();
672 let check = move |conn: PeerConnectSuccess<_>| async move {
673 let session_cid = conn.channel.get_session_cid();
674 let _mutual_peers = conn
675 .remote
676 .remote()
677 .get_local_group_mutual_peers(session_cid)
678 .await
679 .unwrap();
680 conn
681 };
682 let p2p_remotes = handle_peer_connect_successes(
683 results,
684 session_cid,
685 peer_count,
686 udp_mode,
687 check,
688 )
689 .await
690 .into_iter()
691 .map(|r| (r.channel.get_peer_cid(), r.remote))
692 .collect::<HashMap<_, _>>();
693
694 let network_peers = connection.get_peers(None).await.unwrap();
698 for user in agg.inner {
699 let peer_cid = user.id.get_cid();
700 assert!(network_peers.iter().any(|r| r.cid == peer_cid))
701 }
702
703 let session_cid = connection.conn_type.get_session_cid();
705 let mutual_peers = connection
706 .get_local_group_mutual_peers(session_cid)
707 .await
708 .unwrap();
709 for (peer_cid, _) in p2p_remotes {
710 assert!(mutual_peers.iter().any(|r| r.cid == peer_cid))
711 }
712
713 log::info!(target: "citadel", "***PEER {username} finished all checks***");
714 let _ = client_success.fetch_add(1, Ordering::Relaxed);
715 wait_for_peers().await;
716 connection.shutdown_kernel().await
717 },
718 );
719
720 let client = DefaultNodeBuilder::default().build(client_kernel).unwrap();
721 client_kernels.push(async move { client.await.map(|_| ()) });
722 }
723
724 let clients = Box::pin(async move { client_kernels.try_collect::<()>().await.map(|_| ()) });
725
726 assert!(futures::future::try_select(server, clients).await.is_ok());
727
728 assert_eq!(client_success.load(Ordering::Relaxed), peer_count);
729 }
730
731 #[rstest]
732 #[case(2, HeaderObfuscatorSettings::default())]
733 #[case(2, HeaderObfuscatorSettings::Enabled)]
734 #[case(2, HeaderObfuscatorSettings::EnabledWithKey(12345))]
735 #[case(3, HeaderObfuscatorSettings::default())]
736 #[timeout(Duration::from_secs(90))]
737 #[tokio::test(flavor = "multi_thread")]
738 async fn peer_to_peer_connect_transient(
739 #[case] peer_count: usize,
740 #[case] header_obfuscator_settings: HeaderObfuscatorSettings,
741 ) -> Result<(), Box<dyn std::error::Error>> {
742 assert!(peer_count > 1);
743 citadel_logging::setup_log();
744 TestBarrier::setup(peer_count);
745 let udp_mode = UdpMode::Enabled;
746
747 let do_deregister = peer_count == 2;
748
749 let client_success = &AtomicUsize::new(0);
750 let (server, server_addr) = server_info::<StackedRatchet>();
751
752 let client_kernels = FuturesUnordered::new();
753 let total_peers = (0..peer_count)
754 .map(|_| Uuid::new_v4())
755 .collect::<Vec<Uuid>>();
756
757 for idx in 0..peer_count {
758 let uuid = total_peers.get(idx).cloned().unwrap();
759 let peers = total_peers
760 .clone()
761 .into_iter()
762 .filter(|r| r != &uuid)
763 .map(UserIdentifier::from)
764 .collect::<Vec<UserIdentifier>>();
765
766 let mut agg = PeerConnectionSetupAggregator::default();
767
768 for peer in peers {
769 let security_settings = SessionSecuritySettings {
770 header_obfuscator_settings,
771 ..Default::default()
772 };
773 agg = agg
774 .with_peer_custom(peer)
775 .with_udp_mode(udp_mode)
776 .ensure_registered()
777 .with_session_security_settings(security_settings)
778 .add();
779 }
780
781 let server_connection_settings =
782 DefaultServerConnectionSettingsBuilder::transient_with_id(server_addr, uuid)
783 .build()
784 .unwrap();
785
786 let client_kernel = PeerConnectionKernel::new(
787 server_connection_settings,
788 agg,
789 move |results, remote| async move {
790 log::info!(target: "citadel", "***PEER {uuid} CONNECTED***");
791 let session_cid = remote.conn_type.get_session_cid();
792
793 let check = move |conn: PeerConnectSuccess<_>| async move {
794 if do_deregister {
795 conn.remote
796 .deregister()
797 .await
798 .expect("Deregistration failed");
799 assert!(!conn
800 .remote
801 .inner
802 .account_manager()
803 .get_persistence_handler()
804 .hyperlan_peer_exists(session_cid, conn.channel.get_peer_cid())
805 .await
806 .unwrap());
807 }
808 conn
809 };
810
811 let _ = handle_peer_connect_successes(
812 results,
813 session_cid,
814 peer_count,
815 udp_mode,
816 check,
817 )
818 .await;
819
820 log::info!(target: "citadel", "***PEER {uuid} finished all checks***");
821 let _ = client_success.fetch_add(1, Ordering::Relaxed);
822 wait_for_peers().await;
823 remote.shutdown_kernel().await
824 },
825 );
826
827 let client = DefaultNodeBuilder::default().build(client_kernel)?;
828 client_kernels.push(async move { client.await.map(|_| ()) });
829 }
830
831 let clients = Box::pin(async move { client_kernels.try_collect::<()>().await.map(|_| ()) });
832
833 if let Err(err) = futures::future::try_select(server, clients).await {
834 return match err {
835 futures::future::Either::Left(res) => Err(res.0.into_string().into()),
836 futures::future::Either::Right(res) => Err(res.0.into_string().into()),
837 };
838 }
839
840 assert_eq!(client_success.load(Ordering::Relaxed), peer_count);
841 Ok(())
842 }
843
844 #[rstest]
845 #[case(2)]
846 #[case(3)]
847 #[timeout(std::time::Duration::from_secs(180))]
848 #[tokio::test(flavor = "multi_thread")]
849 async fn test_peer_to_peer_file_transfer(
850 #[case] peer_count: usize,
851 ) -> Result<(), Box<dyn std::error::Error>> {
852 assert!(peer_count > 1);
853 citadel_logging::setup_log();
854 TestBarrier::setup(peer_count);
855 let udp_mode = UdpMode::Enabled;
856
857 let sender_success = &AtomicBool::new(false);
858 let receiver_success = &AtomicBool::new(false);
859
860 let (server, server_addr) = server_info::<StackedRatchet>();
861
862 let client_kernels = FuturesUnordered::new();
863 let total_peers = (0..peer_count)
864 .map(|_| Uuid::new_v4())
865 .collect::<Vec<Uuid>>();
866
867 let sender_uuid = total_peers[0];
868
869 for idx in 0..peer_count {
870 let uuid = total_peers.get(idx).cloned().unwrap();
871 let mut peers = total_peers
872 .clone()
873 .into_iter()
874 .filter(|r| r != &uuid)
875 .map(UserIdentifier::from)
876 .collect::<Vec<UserIdentifier>>();
877 if idx != 0 {
883 peers = vec![sender_uuid.into()];
884 }
885
886 let mut agg = PeerConnectionSetupAggregator::default();
887
888 for peer in peers {
889 agg = agg
890 .with_peer_custom(peer)
891 .ensure_registered()
892 .with_udp_mode(udp_mode)
893 .with_session_security_settings(SessionSecuritySettings::default())
894 .add();
895 }
896
897 let server_connection_settings =
898 DefaultServerConnectionSettingsBuilder::transient_with_id(server_addr, uuid)
899 .build()
900 .unwrap();
901
902 let client_kernel = PeerConnectionKernel::new(
903 server_connection_settings,
904 agg,
905 move |results, remote| async move {
906 log::info!(target: "citadel", "***PEER {uuid} CONNECTED***");
907 wait_for_peers().await;
908 let session_cid = remote.conn_type.get_session_cid();
909 let is_sender = idx == 0; let check = move |mut conn: PeerConnectSuccess<_>| async move {
911 if is_sender {
912 conn.remote
913 .send_file_with_custom_opts(
914 "../resources/TheBridge.pdf",
915 32 * 1024,
916 TransferType::FileTransfer,
917 )
918 .await
919 .expect("Failed to send file");
920 } else {
921 const HANDLE_BUDGET: std::time::Duration =
938 std::time::Duration::from_secs(60);
939 const STATUS_BUDGET: std::time::Duration =
940 std::time::Duration::from_secs(60);
941
942 let mut handles = conn.incoming_object_transfer_handles.take().unwrap();
943 let mut handle = match citadel_io::tokio::time::timeout(
944 HANDLE_BUDGET,
945 handles.recv(),
946 )
947 .await
948 {
949 Ok(Some(handle)) => handle,
950 Ok(None) => panic!(
951 "the incoming-transfer channel closed before any handle arrived"
952 ),
953 Err(_) => panic!(
954 "no incoming file-transfer handle within {HANDLE_BUDGET:?}: the sender never started, or its request never routed here"
955 ),
956 };
957 handle.accept().unwrap();
958
959 use citadel_types::proto::ObjectTransferStatus;
960 use futures::StreamExt;
961 let mut path = None;
962 while let Some(status) = match citadel_io::tokio::time::timeout(
963 STATUS_BUDGET,
964 handle.next(),
965 )
966 .await
967 {
968 Ok(status) => status,
969 Err(_) => panic!(
970 "the transfer stalled: no status within {STATUS_BUDGET:?}, last seen {path:?}"
971 ),
972 } {
973 match status {
974 ObjectTransferStatus::ReceptionComplete => {
975 let cmp =
976 include_bytes!("../../../../resources/TheBridge.pdf");
977 let streamed_data =
978 citadel_io::tokio::fs::read(path.clone().unwrap())
979 .await
980 .unwrap();
981 assert_eq!(
982 cmp,
983 streamed_data.as_slice(),
984 "Original data and streamed data does not match"
985 );
986
987 log::info!(target: "citadel", "Peer has finished receiving and verifying the file!");
988 break;
989 }
990
991 ObjectTransferStatus::ReceptionBeginning(file_path, vfm) => {
992 path = Some(file_path);
993 assert_eq!(vfm.name, "TheBridge.pdf")
994 }
995
996 _ => {}
997 }
998 }
999 }
1000
1001 conn
1002 };
1003 let peer_count = if idx == 0 { peer_count } else { 2 };
1006 let _ = handle_peer_connect_successes(
1007 results,
1008 session_cid,
1009 peer_count,
1010 udp_mode,
1011 check,
1012 )
1013 .await;
1014
1015 if is_sender {
1016 sender_success.store(true, Ordering::Relaxed);
1017 } else {
1018 receiver_success.store(true, Ordering::Relaxed);
1019 }
1020
1021 log::info!(target: "citadel", "***PEER {uuid} (is_sender: {is_sender}) finished all checks***");
1022 wait_for_peers().await;
1023 log::info!(target: "citadel", "***PEER {uuid} (is_sender: {is_sender}) shutting down***");
1024 remote.shutdown_kernel().await
1025 },
1026 );
1027
1028 let client = DefaultNodeBuilder::default().build(client_kernel).unwrap();
1029 client_kernels.push(async move { client.await.map(|_| ()) });
1030 }
1031
1032 let clients = Box::pin(async move { client_kernels.try_collect::<()>().await.map(|_| ()) });
1033
1034 if let Err(err) = futures::future::try_select(server, clients).await {
1035 return match err {
1036 futures::future::Either::Left(res) => Err(res.0.into_string().into()),
1037 futures::future::Either::Right(res) => Err(res.0.into_string().into()),
1038 };
1039 }
1040
1041 assert!(sender_success.load(Ordering::Relaxed));
1042 assert!(receiver_success.load(Ordering::Relaxed));
1043 Ok(())
1044 }
1045
1046 #[rstest]
1047 #[case(2)]
1048 #[timeout(std::time::Duration::from_secs(90))]
1049 #[tokio::test(flavor = "multi_thread")]
1050 async fn test_peer_to_peer_rekey(
1051 #[case] peer_count: usize,
1052 ) -> Result<(), Box<dyn std::error::Error>> {
1053 assert!(peer_count > 1);
1054 citadel_logging::setup_log();
1055 TestBarrier::setup(peer_count);
1056 let udp_mode = UdpMode::Enabled;
1057
1058 let client_success = &AtomicUsize::new(0);
1059 let (server, server_addr) = server_info::<StackedRatchet>();
1060
1061 let client_kernels = FuturesUnordered::new();
1062 let total_peers = (0..peer_count)
1063 .map(|_| Uuid::new_v4())
1064 .collect::<Vec<Uuid>>();
1065
1066 for idx in 0..peer_count {
1067 let uuid = total_peers.get(idx).cloned().unwrap();
1068 let peers = total_peers
1069 .clone()
1070 .into_iter()
1071 .filter(|r| r != &uuid)
1072 .map(UserIdentifier::from)
1073 .collect::<Vec<UserIdentifier>>();
1074
1075 let mut agg = PeerConnectionSetupAggregator::default();
1076
1077 for peer in peers {
1078 agg = agg
1079 .with_peer_custom(peer)
1080 .ensure_registered()
1081 .with_udp_mode(udp_mode)
1082 .with_session_security_settings(SessionSecuritySettings::default())
1083 .add();
1084 }
1085
1086 let server_connection_settings =
1087 DefaultServerConnectionSettingsBuilder::transient_with_id(server_addr, uuid)
1088 .build()
1089 .unwrap();
1090
1091 let client_kernel = PeerConnectionKernel::new(
1092 server_connection_settings,
1093 agg,
1094 move |results, remote| async move {
1095 log::info!(target: "citadel", "***PEER {uuid} CONNECTED***");
1096 let session_cid = remote.conn_type.get_session_cid();
1097
1098 let check = move |conn: PeerConnectSuccess<_>| async move {
1099 if idx == 0 {
1100 for x in 1..10 {
1101 assert_eq!(
1102 conn.remote.rekey().await.expect("Failed to rekey"),
1103 Some(x)
1104 );
1105 }
1106 }
1107
1108 conn
1109 };
1110
1111 let results = handle_peer_connect_successes(
1112 results,
1113 session_cid,
1114 peer_count,
1115 udp_mode,
1116 check,
1117 )
1118 .await;
1119
1120 log::info!(target: "citadel", "***PEER {uuid} finished all check (count: {})s***", results.len());
1121 let _ = client_success.fetch_add(1, Ordering::Relaxed);
1122 wait_for_peers().await;
1123 remote.shutdown_kernel().await
1124 },
1125 );
1126
1127 let client = DefaultNodeBuilder::default().build(client_kernel)?;
1128 client_kernels.push(async move { client.await.map(|_| ()) });
1129 }
1130
1131 let clients = Box::pin(async move { client_kernels.try_collect::<()>().await.map(|_| ()) });
1132
1133 if let Err(err) = futures::future::try_select(server, clients).await {
1134 return match err {
1135 futures::future::Either::Left(res) => Err(res.0.into_string().into()),
1136 futures::future::Either::Right(res) => Err(res.0.into_string().into()),
1137 };
1138 }
1139
1140 assert_eq!(client_success.load(Ordering::Relaxed), peer_count);
1141 Ok(())
1142 }
1143
1144 #[rstest]
1145 #[case(2)]
1146 #[timeout(std::time::Duration::from_secs(90))]
1147 #[tokio::test(flavor = "multi_thread")]
1148 async fn test_peer_to_peer_disconnect(
1149 #[case] peer_count: usize,
1150 ) -> Result<(), Box<dyn std::error::Error>> {
1151 assert!(peer_count > 1);
1152 citadel_logging::setup_log();
1153 TestBarrier::setup(peer_count);
1154 let udp_mode = UdpMode::Enabled;
1155
1156 let client_success = &AtomicUsize::new(0);
1157 let (server, server_addr) = server_info::<StackedRatchet>();
1158
1159 let client_kernels = FuturesUnordered::new();
1160 let total_peers = (0..peer_count)
1161 .map(|_| Uuid::new_v4())
1162 .collect::<Vec<Uuid>>();
1163
1164 for idx in 0..peer_count {
1165 let uuid = total_peers.get(idx).cloned().unwrap();
1166 let peers = total_peers
1167 .clone()
1168 .into_iter()
1169 .filter(|r| r != &uuid)
1170 .map(UserIdentifier::from)
1171 .collect::<Vec<UserIdentifier>>();
1172
1173 let mut agg = PeerConnectionSetupAggregator::default();
1174
1175 for peer in peers {
1176 agg = agg
1177 .with_peer_custom(peer)
1178 .ensure_registered()
1179 .with_udp_mode(udp_mode)
1180 .with_session_security_settings(SessionSecuritySettings::default())
1181 .add();
1182 }
1183
1184 let server_connection_settings =
1185 DefaultServerConnectionSettingsBuilder::transient_with_id(server_addr, uuid)
1186 .build()
1187 .unwrap();
1188
1189 let client_kernel = PeerConnectionKernel::new(
1190 server_connection_settings,
1191 agg,
1192 move |results, remote| async move {
1193 log::info!(target: "citadel", "***PEER {uuid} CONNECTED***");
1194 wait_for_peers().await;
1195 let session_cid = remote.conn_type.get_session_cid();
1196
1197 let check = move |conn: PeerConnectSuccess<_>| async move {
1198 conn.remote
1199 .disconnect()
1200 .await
1201 .expect("Failed to p2p disconnect");
1202 conn
1203 };
1204 let _ = handle_peer_connect_successes(
1205 results,
1206 session_cid,
1207 peer_count,
1208 udp_mode,
1209 check,
1210 )
1211 .await;
1212 log::info!(target: "citadel", "***PEER {uuid} finished all checks***");
1213
1214 let _ = client_success.fetch_add(1, Ordering::Relaxed);
1215 wait_for_peers().await;
1216 remote.shutdown_kernel().await
1217 },
1218 );
1219
1220 let client = DefaultNodeBuilder::default().build(client_kernel)?;
1221 client_kernels.push(async move { client.await.map(|_| ()) });
1222 }
1223
1224 let clients = Box::pin(async move { client_kernels.try_collect::<()>().await.map(|_| ()) });
1225
1226 if let Err(err) = futures::future::try_select(server, clients).await {
1227 return match err {
1228 futures::future::Either::Left(res) => Err(res.0.into_string().into()),
1229 futures::future::Either::Right(res) => Err(res.0.into_string().into()),
1230 };
1231 }
1232
1233 assert_eq!(client_success.load(Ordering::Relaxed), peer_count);
1234 Ok(())
1235 }
1236
1237 #[rstest]
1238 #[case(SecrecyMode::BestEffort, Some("test-p2p-password"))]
1239 #[timeout(std::time::Duration::from_secs(240))]
1240 #[citadel_io::tokio::test(flavor = "multi_thread")]
1241 async fn test_p2p_wrong_session_password(
1242 #[case] secrecy_mode: SecrecyMode,
1243 #[case] p2p_password: Option<&'static str>,
1244 #[values(KemAlgorithm::MlKem)] kem: KemAlgorithm,
1245 #[values(EncryptionAlgorithm::AES_GCM_256)] enx: EncryptionAlgorithm,
1246 ) {
1247 citadel_logging::setup_log_no_panic_hook();
1248 TestBarrier::setup(2);
1249 let (server, server_addr) = server_info::<StackedRatchet>();
1250 let peer_0_error_received = &AtomicBool::new(false);
1251 let peer_1_error_received = &AtomicBool::new(false);
1252
1253 let uuid0 = Uuid::new_v4();
1254 let uuid1 = Uuid::new_v4();
1255 let session_security = SessionSecuritySettingsBuilder::default()
1256 .with_secrecy_mode(secrecy_mode)
1257 .with_crypto_params(kem + enx)
1258 .build()
1259 .unwrap();
1260
1261 let mut peer0_agg = PeerConnectionSetupAggregator::default()
1262 .with_peer_custom(uuid1)
1263 .ensure_registered()
1264 .with_session_security_settings(session_security);
1265
1266 if let Some(password) = p2p_password {
1267 peer0_agg = peer0_agg.with_session_password(password);
1268 }
1269
1270 let peer0_connection = peer0_agg.add();
1271
1272 let mut peer1_agg = PeerConnectionSetupAggregator::default()
1273 .with_peer_custom(uuid0)
1274 .ensure_registered()
1275 .with_session_security_settings(session_security);
1276
1277 if let Some(_password) = p2p_password {
1278 peer1_agg = peer1_agg.with_session_password("wrong password");
1279 }
1280
1281 let peer1_connection = peer1_agg.add();
1282
1283 let server_connection_settings0 =
1284 DefaultServerConnectionSettingsBuilder::transient_with_id(server_addr, uuid0)
1285 .with_udp_mode(UdpMode::Enabled)
1286 .with_session_security_settings(session_security)
1287 .build()
1288 .unwrap();
1289
1290 let server_connection_settings1 =
1291 DefaultServerConnectionSettingsBuilder::transient_with_id(server_addr, uuid1)
1292 .with_udp_mode(UdpMode::Enabled)
1293 .with_session_security_settings(session_security)
1294 .build()
1295 .unwrap();
1296
1297 let client_kernel0 = PeerConnectionKernel::new(
1298 server_connection_settings0,
1299 peer0_connection,
1300 move |mut connection, remote| async move {
1301 wait_for_peers().await;
1302 let conn = connection.recv().await.unwrap();
1303 log::trace!(target: "citadel", "Peer 0 {} received: {:?}", remote.conn_type.get_session_cid(), conn);
1304 if conn.is_ok() {
1305 peer_0_error_received.store(true, Ordering::SeqCst);
1306 }
1307 wait_for_peers().await;
1308 remote.shutdown_kernel().await
1309 },
1310 );
1311
1312 let client_kernel1 = PeerConnectionKernel::new(
1313 server_connection_settings1,
1314 peer1_connection,
1315 move |mut connection, remote| async move {
1316 wait_for_peers().await;
1317 let conn = connection.recv().await.unwrap();
1318 log::trace!(target: "citadel", "Peer 1 {} received: {:?}", remote.conn_type.get_session_cid(), conn);
1319 if conn.is_ok() {
1320 peer_1_error_received.store(true, Ordering::SeqCst);
1321 }
1322 wait_for_peers().await;
1323 remote.shutdown_kernel().await
1324 },
1325 );
1326
1327 let client0 = DefaultNodeBuilder::default().build(client_kernel0).unwrap();
1328 let client1 = DefaultNodeBuilder::default().build(client_kernel1).unwrap();
1329 let clients = futures::future::try_join(client0, client1);
1330
1331 let task = async move {
1332 tokio::select! {
1333 server_res = server => Err(citadel_io::error!(citadel_io::ErrorCode::PrefabServerEndedPrematurely, citadel_io::Dbg(server_res.map(|_| ())))),
1334 client_res = clients => client_res.map(|_| ())
1335 }
1336 };
1337
1338 tokio::time::timeout(Duration::from_secs(120), task)
1339 .await
1340 .unwrap()
1341 .unwrap();
1342
1343 assert!(!peer_0_error_received.load(Ordering::SeqCst));
1344 assert!(!peer_1_error_received.load(Ordering::SeqCst));
1345 }
1346
1347 async fn handle_peer_connect_successes<F, Fut, R: Ratchet>(
1348 mut conn_rx: Receiver<Result<PeerConnectSuccess<R>, NetworkError>>,
1349 session_cid: u64,
1350 peer_count: usize,
1351 udp_mode: UdpMode,
1352 checks: F,
1353 ) -> Vec<PeerConnectSuccess<R>>
1354 where
1355 F: Fn(PeerConnectSuccess<R>) -> Fut + Send + Clone + 'static,
1356 Fut: Future<Output = PeerConnectSuccess<R>> + Send,
1357 {
1358 let (finished_tx, finished_rx) = tokio::sync::oneshot::channel();
1359
1360 let task = async move {
1361 let (done_tx, mut done_rx) = tokio::sync::mpsc::unbounded_channel();
1362 let mut conns = vec![];
1363 while let Some(conn) = conn_rx.recv().await {
1364 conns.push(conn);
1365 if conns.len() == peer_count - 1 {
1366 break;
1367 }
1368 }
1369
1370 log::info!(target: "citadel", "~~~*** Peer {session_cid} has {} connections to other peers ***~~~", conns.len());
1371
1372 for conn in conns {
1373 let conn = conn.expect("Error receiving peer connection");
1374 handle_peer_connect_success(
1375 conn,
1376 done_tx.clone(),
1377 session_cid,
1378 udp_mode,
1379 checks.clone(),
1380 );
1381 }
1382
1383 let mut ret = vec![];
1385 while let Some(done) = done_rx.recv().await {
1386 ret.push(done);
1387 if ret.len() == peer_count - 1 {
1388 break;
1389 }
1390 }
1391
1392 finished_tx
1393 .send(ret)
1394 .expect("Error sending finished signal in handle_peer_connect_successes");
1395 };
1396
1397 drop(tokio::task::spawn(task));
1398 let ret = finished_rx
1399 .await
1400 .expect("Error receiving finished signal in handle_peer_connect_successes");
1401
1402 assert_eq!(ret.len(), peer_count - 1);
1403 ret
1404 }
1405
1406 fn handle_peer_connect_success<F, Fut, R: Ratchet>(
1407 mut conn: PeerConnectSuccess<R>,
1408 done_tx: UnboundedSender<PeerConnectSuccess<R>>,
1409 session_cid: u64,
1410 udp_mode: UdpMode,
1411 checks: F,
1412 ) where
1413 F: Fn(PeerConnectSuccess<R>) -> Fut + Send + Clone + 'static,
1414 Fut: Future<Output = PeerConnectSuccess<R>> + Send,
1415 {
1416 let task = async move {
1417 let chan = conn.udp_channel_rx.take();
1418 crate::test_common::p2p_assertions(session_cid, &conn).await;
1419 crate::test_common::udp_mode_assertions(udp_mode, chan).await;
1420 let conn = checks(conn).await;
1421 done_tx
1422 .send(conn)
1423 .expect("Error sending done signal in handle_peer_connect_success");
1424 };
1425
1426 drop(tokio::task::spawn(task));
1427 }
1428}