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