1use std::collections::{BTreeMap, HashMap};
2use std::future::Future;
3use std::sync::atomic::AtomicU64;
4use std::sync::Arc;
5
6use arc_swap::ArcSwap;
7use serde_json::Value;
8use tokio::sync::{Mutex, RwLock, Semaphore};
9use unb_core::{validate_node_identifier, NodeCore, NodeIdentity, ProtocolCore};
10use unb_runtime::{CancellationToken, ProtocolCoreHandle, WsError};
11
12use crate::handler::HandlerError;
13use crate::layer::{layer_fn, Layer, Next, ServiceBody};
14use crate::node::{Node, NodeSnapshot, SubjectServices};
15use crate::peer::{
16 peer_layer_fn, InsecureAcceptDeclaredPeerIdentities, PeerLayer, PeerNext, PeerRequest,
17};
18use crate::service::{Handler, HandlerService, StateMap};
19use crate::session::ServerEffectExecutor;
20
21const DEFAULT_MAX_ACTIVATIONS: usize = 256;
22
23pub(crate) struct PendingService {
24 scopes: Vec<String>,
25 layers: Vec<Arc<dyn Layer>>,
26 states: StateMap,
27 service: HandlerService,
28}
29
30pub struct Scope {
31 path: Vec<String>,
32 layers: Vec<Arc<dyn Layer>>,
33 states: StateMap,
34 services: Vec<PendingService>,
35}
36
37impl Scope {
38 fn new(path: Vec<String>, layers: Vec<Arc<dyn Layer>>, states: StateMap) -> Scope {
39 Scope {
40 path,
41 layers,
42 states,
43 services: Vec::new(),
44 }
45 }
46
47 pub fn service(mut self, handler: impl Handler) -> Scope {
48 self.services.push(PendingService {
49 scopes: self.path.clone(),
50 layers: self.layers.clone(),
51 states: self.states.clone(),
52 service: handler.into_service(),
53 });
54 self
55 }
56
57 pub fn layer(mut self, layer: impl Layer) -> Scope {
58 self.layers.push(Arc::new(layer));
59 self
60 }
61
62 pub fn layer_fn<F, Fut>(self, f: F) -> Scope
63 where
64 F: Fn(http::Request<bytes::Bytes>, Next) -> Fut + Send + Sync + 'static,
65 Fut: Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + 'static,
66 {
67 self.layer(layer_fn(f))
68 }
69
70 pub fn state<T: Send + Sync + 'static>(mut self, value: T) -> Scope {
71 self.states.insert(value);
72 self
73 }
74
75 pub fn scope(mut self, name: &str, f: impl FnOnce(Scope) -> Scope) -> Scope {
76 let mut path = self.path.clone();
77 path.push(name.to_string());
78 let child = f(Scope::new(path, self.layers.clone(), self.states.clone()));
79 self.services.extend(child.services);
80 self
81 }
82}
83
84pub struct NodeBuilder {
85 node: String,
86 services: Vec<PendingService>,
87 states: StateMap,
88 global_layers: Vec<Arc<dyn Layer>>,
89 max_activations: usize,
90 epoch: u64,
91 peer_layers: Vec<Arc<dyn PeerLayer>>,
92 identity_trust_selected: bool,
93 connect_timeout: Option<std::time::Duration>,
94 proof: Value,
95 ws_collect_ceiling: usize,
96}
97
98impl Node {
99 pub fn builder(node: &str) -> NodeBuilder {
100 NodeBuilder::new(node)
101 }
102}
103
104impl NodeBuilder {
105 pub fn new(node: &str) -> NodeBuilder {
106 NodeBuilder {
107 node: node.into(),
108 services: Vec::new(),
109 states: StateMap::default(),
110 global_layers: Vec::new(),
111 max_activations: DEFAULT_MAX_ACTIVATIONS,
112 epoch: 1,
113 peer_layers: Vec::new(),
114 identity_trust_selected: false,
115 connect_timeout: None,
116 proof: Value::Null,
117 ws_collect_ceiling: unb_transport::DEFAULT_MAX_FRAME_SIZE,
118 }
119 }
120
121 pub fn max_activations(mut self, max: usize) -> NodeBuilder {
122 self.max_activations = max;
123 self
124 }
125
126 pub fn ws_collect_ceiling(mut self, bytes: usize) -> NodeBuilder {
127 self.ws_collect_ceiling = bytes.min(unb_transport::DEFAULT_MAX_FRAME_SIZE);
128 self
129 }
130
131 pub fn epoch(mut self, epoch: u64) -> NodeBuilder {
132 self.epoch = epoch;
133 self
134 }
135
136 pub fn identity_proof(mut self, proof: Value) -> NodeBuilder {
137 self.proof = proof;
138 self
139 }
140
141 pub fn service(mut self, handler: impl Handler) -> NodeBuilder {
142 self.services.push(PendingService {
143 scopes: Vec::new(),
144 layers: Vec::new(),
145 states: StateMap::default(),
146 service: handler.into_service(),
147 });
148 self
149 }
150
151 pub fn scope(mut self, name: &str, f: impl FnOnce(Scope) -> Scope) -> NodeBuilder {
152 let scope = f(Scope::new(
153 vec![name.to_string()],
154 Vec::new(),
155 StateMap::default(),
156 ));
157 self.services.extend(scope.services);
158 self
159 }
160
161 pub fn state<T: Send + Sync + 'static>(mut self, value: T) -> NodeBuilder {
162 self.states.insert(value);
163 self
164 }
165
166 pub fn layer(mut self, layer: impl Layer) -> NodeBuilder {
167 self.global_layers.push(Arc::new(layer));
168 self
169 }
170
171 pub fn layer_fn<F, Fut>(self, f: F) -> NodeBuilder
172 where
173 F: Fn(http::Request<bytes::Bytes>, Next) -> Fut + Send + Sync + 'static,
174 Fut: Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + 'static,
175 {
176 self.layer(layer_fn(f))
177 }
178
179 pub fn peer_layer(mut self, layer: impl PeerLayer) -> NodeBuilder {
180 self.peer_layers.push(Arc::new(layer));
181 self.identity_trust_selected = true;
182 self
183 }
184
185 pub fn peer_layer_fn<F, Fut>(self, f: F) -> NodeBuilder
186 where
187 F: Fn(PeerRequest, PeerNext) -> Fut + Send + Sync + 'static,
188 Fut: Future<Output = Result<PeerRequest, HandlerError>> + Send + 'static,
189 {
190 self.peer_layer(peer_layer_fn(f))
191 }
192
193 pub fn insecure_accept_declared_peer_identities(self) -> NodeBuilder {
194 self.peer_layer(InsecureAcceptDeclaredPeerIdentities)
195 }
196
197 pub fn connect_timeout(mut self, timeout: std::time::Duration) -> NodeBuilder {
198 self.connect_timeout = Some(timeout);
199 self
200 }
201
202 pub fn build(self) -> Result<Arc<Node>, WsError> {
203 #[cfg(not(all(target_arch = "wasm32", target_os = "unknown")))]
204 let runtime = unb_runtime::RuntimeHandle::try_current()
205 .map_err(|_| WsError::Connect("node build requires an async runtime handle".into()))?;
206 #[cfg(all(target_arch = "wasm32", target_os = "unknown"))]
207 let runtime = unb_runtime::RuntimeHandle::current();
208 validate_node_identifier(&self.node)
209 .map_err(|error| WsError::Connect(error.to_string()))?;
210 if !self.identity_trust_selected {
211 return Err(WsError::Connect(
212 "peer admission must be configured explicitly; add peer_layer(...) or \
213 insecure_accept_declared_peer_identities()"
214 .into(),
215 ));
216 }
217 let mut node_core = NodeCore::new(&self.node);
218 let identity = NodeIdentity {
219 node_id: self.node.clone(),
220 instance_id: unique_instance_id(&self.node),
221 epoch: self.epoch,
222 proof: self.proof,
223 };
224 node_core.set_node_identity(identity.clone());
225 let global_layers: Arc<[Arc<dyn Layer>]> = Arc::from(self.global_layers);
226 let mut services = BTreeMap::new();
227 for pending in self.services {
228 let states = pending.states.merged_over(&self.states);
229 let mut chain: Vec<Arc<dyn Layer>> = global_layers.iter().cloned().collect();
230 chain.extend(pending.layers);
231 let _ = SubjectServices::register(
232 &mut services,
233 pending.service,
234 &pending.scopes,
235 chain,
236 &states,
237 )
238 .map_err(WsError::Connect)?;
239 }
240 let snapshot = NodeSnapshot::new(services, node_core.clone());
241 let peer_layers: Arc<[Arc<dyn PeerLayer>]> = Arc::from(self.peer_layers);
242 let cancellation = CancellationToken::new();
243 let shutdown = cancellation.drop_guard();
244 let protocol_core = ProtocolCore::with_node(snapshot.node_core.clone());
245 Ok(Arc::new_cyclic(|weak| Node {
246 snapshot: Arc::new(ArcSwap::from_pointee(snapshot)),
247 states: self.states,
248 global_layers,
249 mutation_gate: Mutex::new(()),
250 peers: RwLock::new(HashMap::new()),
251 sessions: RwLock::new(HashMap::new()),
252 connections: std::sync::RwLock::new(HashMap::new()),
253 routes_changed: tokio::sync::watch::channel(0).0,
254 session_peers: RwLock::new(HashMap::new()),
255 outbound_sessions: Mutex::new(std::collections::HashSet::new()),
256 dispatch_slots: Arc::new(Semaphore::new(self.max_activations)),
257 dispatch_permits: Mutex::new(HashMap::new()),
258 dispatching: Mutex::new(HashMap::new()),
259 peer_admissions: Mutex::new(HashMap::new()),
260 verified_peers: Mutex::new(HashMap::new()),
261 candidate_identities: Mutex::new(HashMap::new()),
262 active: Arc::new(Mutex::new(HashMap::new())),
263 protocol: ProtocolCoreHandle::spawn(
264 protocol_core,
265 Arc::new(ServerEffectExecutor { node: weak.clone() }),
266 cancellation.clone(),
267 &runtime,
268 ),
269 identity,
270 peer_layers,
271 dial_policy: unb_client::Peers::with_config(unb_client::DialConfig {
272 attempt_timeout: self
273 .connect_timeout
274 .unwrap_or(unb_client::DEFAULT_DIAL_TIMEOUT),
275 supported: None,
276 }),
277 next_session: AtomicU64::new(0),
278 ws_collect_ceiling: self.ws_collect_ceiling,
279 cancellation,
280 _shutdown: shutdown,
281 }))
282 }
283}
284
285fn unique_instance_id(node: &str) -> String {
286 static COUNTER: AtomicU64 = AtomicU64::new(0);
287 let count = COUNTER.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
288 #[cfg(not(all(target_arch = "wasm32", target_os = "unknown")))]
289 {
290 let nanos = std::time::SystemTime::now()
291 .duration_since(std::time::UNIX_EPOCH)
292 .map(|d| d.as_nanos())
293 .unwrap_or(0);
294 format!("{node}-{}-{nanos:x}-{count}", std::process::id())
295 }
296 #[cfg(all(target_arch = "wasm32", target_os = "unknown"))]
297 format!("{node}-browser-{count}")
298}
299
300#[cfg(test)]
301mod tests {
302 use super::*;
303
304 #[test]
305 fn build_without_runtime_returns_typed_error() {
306 let error = match NodeBuilder::new("node")
307 .insecure_accept_declared_peer_identities()
308 .build()
309 {
310 Ok(_) => panic!("build unexpectedly succeeded"),
311 Err(error) => error,
312 };
313
314 assert!(matches!(
315 error,
316 WsError::Connect(message)
317 if message == "node build requires an async runtime handle"
318 ));
319 }
320
321 #[tokio::test]
322 async fn build_uses_entered_runtime() {
323 let node = NodeBuilder::new("node")
324 .insecure_accept_declared_peer_identities()
325 .build()
326 .unwrap();
327
328 assert_eq!(node.identity().node_id, "node");
329 }
330}