1use citadel_proto::prelude::*;
39
40use citadel_io::ServerMode;
41use citadel_proto::kernel::KernelExecutorArguments;
42use citadel_proto::macros::{ContextRequirements, LocalContextRequirements};
43use citadel_types::crypto::{HeaderObfuscatorSettings, PreSharedKey};
44use futures::Future;
45use std::fmt::{Debug, Formatter};
46use std::marker::PhantomData;
47use std::pin::Pin;
48use std::task::{Context, Poll};
49
50pub struct NodeBuilder<R: Ratchet = StackedRatchet, T: PlatformOps = DefaultTransport> {
52 hypernode_type: Option<NodeType>,
53 underlying_protocol: Option<ServerMode<T>>,
54 backend_type: Option<BackendType>,
55 server_argon_settings: Option<ArgonDefaultServerSettings>,
56 #[cfg(feature = "google-services")]
57 services: Option<ServicesConfig>,
58 server_misc_settings: Option<ServerMiscSettings>,
59 client_tls_config: Option<T::ClientConfig>,
60 kernel_executor_settings: Option<KernelExecutorSettings>,
61 stun_servers: Option<Vec<String>>,
62 turn_servers: Option<Vec<TurnServerConfig>>,
63 local_only_server_settings: Option<ServerOnlySessionInitSettings>,
64 websocket_listen_addr: Option<std::net::SocketAddr>,
65 #[cfg(target_family = "wasm")]
66 serverless_config: Option<ServerlessConfig>,
67 _ratchet: PhantomData<R>,
68 _transport: PhantomData<T>,
69}
70
71pub type DefaultNodeBuilder = NodeBuilder<StackedRatchet, DefaultTransport>;
73pub type LightweightNodeBuilder = NodeBuilder<MonoRatchet, DefaultTransport>;
75
76impl<R: Ratchet, T: PlatformOps> Default for NodeBuilder<R, T> {
77 fn default() -> Self {
78 Self {
79 hypernode_type: None,
80 underlying_protocol: None,
81 backend_type: None,
82 server_argon_settings: None,
83 #[cfg(feature = "google-services")]
84 services: None,
85 server_misc_settings: None,
86 client_tls_config: None,
87 kernel_executor_settings: None,
88 stun_servers: None,
89 turn_servers: None,
90 local_only_server_settings: None,
91 websocket_listen_addr: None,
92 #[cfg(target_family = "wasm")]
93 serverless_config: None,
94 _ratchet: Default::default(),
95 _transport: Default::default(),
96 }
97 }
98}
99
100pub struct NodeFuture<'a, K> {
102 inner: Pin<Box<dyn FutureContextRequirements<'a, Result<K, NetworkError>>>>,
103 _pd: PhantomData<fn() -> K>,
104}
105
106#[cfg(feature = "multi-threaded")]
107trait FutureContextRequirements<'a, Output>:
108 Future<Output = Output> + Send + LocalContextRequirements<'a>
109{
110}
111#[cfg(feature = "multi-threaded")]
112impl<'a, T: Future<Output = Output> + Send + LocalContextRequirements<'a>, Output>
113 FutureContextRequirements<'a, Output> for T
114{
115}
116
117#[cfg(not(feature = "multi-threaded"))]
118trait FutureContextRequirements<'a, Output>:
119 Future<Output = Output> + LocalContextRequirements<'a>
120{
121}
122#[cfg(not(feature = "multi-threaded"))]
123impl<'a, T: Future<Output = Output> + LocalContextRequirements<'a>, Output>
124 crate::builder::node_builder::FutureContextRequirements<'a, Output> for T
125{
126}
127
128impl<K> Debug for NodeFuture<'_, K> {
129 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
130 write!(f, "NodeFuture")
131 }
132}
133
134impl<K> Future for NodeFuture<'_, K> {
135 type Output = Result<K, NetworkError>;
136
137 fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
138 self.inner.as_mut().poll(cx)
139 }
140}
141
142impl<R: Ratchet + ContextRequirements, T: PlatformOps> NodeBuilder<R, T> {
143 pub fn build<'a, 'b: 'a, K: NetKernel<R> + 'b>(
145 &'a mut self,
146 kernel: K,
147 ) -> anyhow::Result<NodeFuture<'b, K>> {
148 self.check()?;
149 let hypernode_type = self.hypernode_type.take().unwrap_or_default();
150 let backend_type = self.backend_type.take().unwrap_or_default();
151 let server_argon_settings = self.server_argon_settings.take();
152 #[cfg(feature = "google-services")]
153 let server_services_cfg = self.services.take();
154 #[cfg(not(feature = "google-services"))]
155 let server_services_cfg = None;
156 let server_misc_settings = self.server_misc_settings.take();
157 let client_config = self.client_tls_config.take();
158 let kernel_executor_settings = self.kernel_executor_settings.take().unwrap_or_default();
159 let stun_servers = self.stun_servers.take();
160 let turn_servers = self.turn_servers.take();
161 let underlying_proto = self.underlying_protocol.take();
162 let server_only_session_init_settings = self.local_only_server_settings.take();
163 let websocket_listen_addr = self.websocket_listen_addr.take();
164 #[cfg(target_family = "wasm")]
165 let serverless_config = self.serverless_config.take();
166
167 Ok(NodeFuture {
168 _pd: Default::default(),
169 inner: Box::pin(async move {
170 let underlying_proto = match underlying_proto {
171 Some(proto) => proto,
172 None => T::default_server_config().await.map_err(|err| {
173 citadel_io::error!(
174 citadel_io::ErrorCode::NodeDefaultServerConfigFailed,
175 err.to_string()
176 )
177 })?,
178 };
179
180 T::config_warnings(&underlying_proto);
181
182 #[cfg(target_family = "wasm")]
185 let (pre_built_listener, client_config, hypernode_type) = if let Some(sl_config) =
186 serverless_config
187 {
188 let conn = establish_serverless_connection(
189 sl_config.signaling.as_ref(),
190 &sl_config.room_token,
191 &sl_config.ice_servers,
192 sl_config.poll_interval_ms,
193 sl_config.timeout_ms,
194 )
195 .await
196 .map_err(|e: std::io::Error| NetworkError::generic(e.to_string()))?;
197
198 T::setup_serverless_transport(conn.stream, conn.is_server_role, client_config)
199 } else {
200 (None, client_config, hypernode_type)
201 };
202
203 #[cfg(not(target_family = "wasm"))]
204 let pre_built_listener = None;
205
206 log::trace!(target: "citadel", "[NodeBuilder] Checking Tokio runtime ...");
207 let rt = citadel_io::try_current_runtime().map_err(NetworkError::generic)?;
208 log::trace!(target: "citadel", "[NodeBuilder] Creating account manager ...");
209 let account_manager = AccountManager::new(
210 backend_type,
211 server_argon_settings,
212 server_services_cfg,
213 server_misc_settings,
214 )
215 .await?;
216
217 let args: KernelExecutorArguments<_, _, T> = KernelExecutorArguments {
218 rt,
219 hypernode_type,
220 account_manager,
221 kernel,
222 underlying_proto,
223 client_config,
224 kernel_executor_settings,
225 stun_servers,
226 turn_servers,
227 server_only_session_init_settings,
228 websocket_listen_addr,
229 pre_built_listener,
230 };
231
232 log::trace!(target: "citadel", "[NodeBuilder] Creating KernelExecutor ...");
233 let kernel_executor = KernelExecutor::<_, R>::new(args).await?;
234 log::trace!(target: "citadel", "[NodeBuilder] Executing kernel");
235 kernel_executor.execute().await
236 }),
237 })
238 }
239
240 pub fn with_node_type(&mut self, node_type: NodeType) -> &mut Self {
248 self.hypernode_type = Some(node_type);
249 self
250 }
251
252 pub fn with_backend(&mut self, backend_type: BackendType) -> &mut Self {
256 self.backend_type = Some(backend_type);
257 self
258 }
259
260 pub fn with_kernel_executor_settings(
262 &mut self,
263 kernel_executor_settings: KernelExecutorSettings,
264 ) -> &mut Self {
265 self.kernel_executor_settings = Some(kernel_executor_settings);
266 self
267 }
268
269 pub fn with_server_argon_settings(
271 &mut self,
272 settings: ArgonDefaultServerSettings,
273 ) -> &mut Self {
274 self.server_argon_settings = Some(settings);
275 self
276 }
277
278 #[cfg(feature = "google-services")]
280 pub fn with_google_services_json_path<V: Into<String>>(&mut self, path: V) -> &mut Self {
281 let cfg = self.get_or_create_services();
282 cfg.google_services_json_path = Some(path.into());
283 self
284 }
285
286 pub fn with_server_misc_settings(&mut self, misc_settings: ServerMiscSettings) -> &mut Self {
288 self.server_misc_settings = Some(misc_settings);
289 self
290 }
291
292 #[cfg(feature = "google-services")]
295 pub fn with_google_realtime_database_config<V: Into<String>, W: Into<String>>(
296 &mut self,
297 url: V,
298 api_key: W,
299 ) -> &mut Self {
300 let cfg = self.get_or_create_services();
301 cfg.google_rtdb = Some(RtdbConfig {
302 url: url.into(),
303 api_key: api_key.into(),
304 });
305 self
306 }
307
308 pub fn with_underlying_protocol(&mut self, proto: ServerMode<T>) -> &mut Self {
311 self.underlying_protocol = Some(proto);
312 self
313 }
314
315 pub fn with_client_config(&mut self, config: T::ClientConfig) -> &mut Self {
317 self.client_tls_config = Some(config);
318 self
319 }
320
321 #[cfg(feature = "google-services")]
322 fn get_or_create_services(&mut self) -> &mut ServicesConfig {
323 if self.services.is_some() {
324 self.services.as_mut().unwrap()
325 } else {
326 let cfg = ServicesConfig::default();
327 self.services = Some(cfg);
328 self.services.as_mut().unwrap()
329 }
330 }
331
332 pub fn with_stun_servers<V: Into<String>, S: Into<Vec<V>>>(&mut self, servers: S) -> &mut Self {
334 self.stun_servers = Some(servers.into().into_iter().map(|t| t.into()).collect());
335 self
336 }
337
338 pub fn with_turn_servers<S: Into<Vec<TurnServerConfig>>>(&mut self, servers: S) -> &mut Self {
344 self.turn_servers = Some(servers.into());
345 self
346 }
347
348 pub fn with_websocket_listener(&mut self, addr: std::net::SocketAddr) -> &mut Self {
356 self.websocket_listen_addr = Some(addr);
357 self
358 }
359
360 #[cfg(target_family = "wasm")]
367 pub fn with_no_central_server(&mut self, config: ServerlessConfig) -> &mut Self {
368 self.serverless_config = Some(config);
369 self
370 }
371
372 pub fn with_server_password<V: Into<PreSharedKey>>(&mut self, password: V) -> &mut Self {
377 let mut server_only_settings = self.local_only_server_settings.clone().unwrap_or_default();
378 server_only_settings.declared_pre_shared_key = Some(password.into());
379 self.local_only_server_settings = Some(server_only_settings);
380 self
381 }
382
383 pub fn with_server_declared_header_obfuscation<V: Into<HeaderObfuscatorSettings>>(
385 &mut self,
386 header_obfuscator_settings: V,
387 ) -> &mut Self {
388 let mut server_only_settings = self.local_only_server_settings.clone().unwrap_or_default();
389 server_only_settings.declared_header_obfuscation_setting =
390 header_obfuscator_settings.into();
391 self.local_only_server_settings = Some(server_only_settings);
392 self
393 }
394
395 fn check(&self) -> anyhow::Result<()> {
396 #[cfg(feature = "google-services")]
397 if let Some(svc) = self.services.as_ref() {
398 if svc.google_rtdb.is_some() && svc.google_services_json_path.is_none() {
399 return Err(anyhow::Error::msg(
400 "Google realtime database is enabled, yet, a services path is not provided",
401 ));
402 }
403 }
404
405 if let Some(stun_servers) = self.stun_servers.as_ref() {
406 if stun_servers.len() != 3 {
407 return Err(anyhow::Error::msg(
408 "There must be exactly 3 specified STUN servers",
409 ));
410 }
411 }
412
413 Ok(())
414 }
415}
416
417#[cfg(not(target_family = "wasm"))]
419impl<R: Ratchet + ContextRequirements> NodeBuilder<R, NativeIO> {
420 fn set_client_rustls_config(
424 &mut self,
425 config: std::sync::Arc<citadel_proto::re_imports::RustlsClientConfig>,
426 ) {
427 let require_cert_verification = self
428 .client_tls_config
429 .as_ref()
430 .map(|cfg| cfg.require_cert_verification)
431 .unwrap_or(false);
432 self.client_tls_config = Some(citadel_proto::re_imports::NativeClientConfig {
433 config,
434 require_cert_verification,
435 });
436 }
437
438 pub async fn with_native_certs(&mut self) -> anyhow::Result<&mut Self> {
442 let certs = citadel_proto::re_imports::load_native_certs_async().await?;
443 let cfg = citadel_proto::re_imports::cert_vec_to_secure_client_config(&certs)?;
444 self.set_client_rustls_config(std::sync::Arc::new(cfg));
445 Ok(self)
446 }
447
448 pub fn with_insecure_skip_cert_verification(&mut self) -> &mut Self {
451 self.client_tls_config = Some(citadel_proto::re_imports::NativeClientConfig::new(
453 std::sync::Arc::new(citadel_proto::re_imports::insecure::rustls_client_config()),
454 ));
455 self
456 }
457
458 pub async fn with_require_cert_verification(&mut self) -> anyhow::Result<&mut Self> {
464 if self.client_tls_config.is_none() {
465 let _ = self.with_native_certs().await?;
466 }
467 if let Some(cfg) = self.client_tls_config.as_mut() {
468 cfg.require_cert_verification = true;
469 }
470 Ok(self)
471 }
472
473 pub fn with_custom_certs<V: AsRef<[u8]>>(
476 &mut self,
477 custom_certs: &[V],
478 ) -> anyhow::Result<&mut Self> {
479 let cfg = citadel_proto::re_imports::create_rustls_client_config(custom_certs)?;
480 self.set_client_rustls_config(std::sync::Arc::new(cfg));
481 Ok(self)
482 }
483
484 #[cfg(feature = "std")]
486 pub async fn with_pem_file<P: AsRef<std::path::Path>>(
487 &mut self,
488 path: P,
489 ) -> anyhow::Result<&mut Self> {
490 use citadel_wire::exports::{Certificate, PemObject};
491 let mut der = std::io::Cursor::new(citadel_io::tokio::fs::read(path).await?);
492 let certs: Vec<Certificate<'static>> =
493 Certificate::pem_reader_iter(&mut der).collect::<Result<Vec<_>, _>>()?;
494 let cfg = citadel_proto::re_imports::create_rustls_client_config(&certs)?;
495 self.set_client_rustls_config(std::sync::Arc::new(cfg));
496 Ok(self)
497 }
498}
499
500#[cfg(all(test, not(target_family = "wasm")))]
501mod tests {
502 use crate::builder::node_builder::DefaultNodeBuilder;
503 use crate::prefabs::server::empty::EmptyKernel;
504 use crate::prelude::{BackendType, NodeType};
505 use citadel_io::tokio;
506 use citadel_proto::prelude::{
507 KernelExecutorSettings, NativeIO, NativeP2PConfig, NativeSecureConfig, ServerMode,
508 };
509 use rstest::rstest;
510 use std::str::FromStr;
511
512 #[test]
513 #[cfg(feature = "google-services")]
514 fn okay_config() {
515 let _ = DefaultNodeBuilder::default()
516 .with_google_realtime_database_config("123", "456")
517 .with_google_services_json_path("abc")
518 .build(EmptyKernel::default())
519 .unwrap();
520 }
521
522 #[test]
523 #[cfg(feature = "google-services")]
524 fn bad_config() {
525 assert!(DefaultNodeBuilder::default()
526 .with_google_realtime_database_config("123", "456")
527 .build(EmptyKernel::default())
528 .is_err());
529 }
530
531 #[test]
532 fn bad_config2() {
533 assert!(DefaultNodeBuilder::default()
534 .with_stun_servers(["dummy1", "dummy2"])
535 .build(EmptyKernel::default())
536 .is_err());
537 }
538
539 #[rstest]
540 #[tokio::test]
541 #[timeout(std::time::Duration::from_secs(60))]
542 #[allow(clippy::let_underscore_must_use)]
543 async fn test_options(
544 #[values(ServerMode::P2P(NativeP2PConfig::self_signed()), ServerMode::OrderedReliableSecure(NativeSecureConfig::self_signed().unwrap())
545 )]
546 underlying_protocol: ServerMode<NativeIO>,
547 #[values(NodeType::Peer, NodeType::Server(std::net::SocketAddr::from_str("127.0.0.1:9999").unwrap()
548 ))]
549 node_type: NodeType,
550 #[values(KernelExecutorSettings::default(), KernelExecutorSettings::default().with_max_concurrency(2)
551 )]
552 kernel_settings: KernelExecutorSettings,
553 #[values(BackendType::InMemory, BackendType::new("file:/hello_world/path/").unwrap())]
554 backend_type: BackendType,
555 ) {
556 let mut builder = DefaultNodeBuilder::default();
557 let _ = builder
558 .with_underlying_protocol(underlying_protocol.clone())
559 .with_backend(backend_type.clone())
560 .with_node_type(node_type)
561 .with_kernel_executor_settings(kernel_settings.clone())
562 .with_insecure_skip_cert_verification()
563 .with_stun_servers(["dummy1", "dummy1", "dummy3"])
564 .with_native_certs()
565 .await
566 .unwrap();
567
568 assert!(builder.underlying_protocol.is_some());
569 assert_eq!(backend_type, builder.backend_type.clone().unwrap());
570 assert_eq!(node_type, builder.hypernode_type.unwrap());
571 assert_eq!(
572 kernel_settings,
573 builder.kernel_executor_settings.clone().unwrap()
574 );
575
576 drop(builder.build(EmptyKernel::default()).unwrap());
577 }
578}