Skip to main content

unb_server/
builder.rs

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}