use std::collections::{BTreeMap, HashMap};
use std::future::Future;
use std::sync::atomic::AtomicU64;
use std::sync::Arc;
use arc_swap::ArcSwap;
use serde_json::Value;
use tokio::sync::{Mutex, RwLock, Semaphore};
use unb_core::{validate_node_identifier, NodeCore, NodeIdentity, ProtocolCore};
use unb_runtime::{CancellationToken, ProtocolCoreHandle, WsError};
use crate::handler::HandlerError;
use crate::layer::{layer_fn, Layer, Next, ServiceBody};
use crate::node::{Node, NodeSnapshot, SubjectServices};
use crate::peer::{
peer_layer_fn, InsecureAcceptDeclaredPeerIdentities, PeerLayer, PeerNext, PeerRequest,
};
use crate::service::{Handler, HandlerService, StateMap};
use crate::session::ServerEffectExecutor;
const DEFAULT_MAX_ACTIVATIONS: usize = 256;
pub(crate) struct PendingService {
scopes: Vec<String>,
layers: Vec<Arc<dyn Layer>>,
states: StateMap,
service: HandlerService,
}
pub struct Scope {
path: Vec<String>,
layers: Vec<Arc<dyn Layer>>,
states: StateMap,
services: Vec<PendingService>,
}
impl Scope {
fn new(path: Vec<String>, layers: Vec<Arc<dyn Layer>>, states: StateMap) -> Scope {
Scope {
path,
layers,
states,
services: Vec::new(),
}
}
pub fn service(mut self, handler: impl Handler) -> Scope {
self.services.push(PendingService {
scopes: self.path.clone(),
layers: self.layers.clone(),
states: self.states.clone(),
service: handler.into_service(),
});
self
}
pub fn layer(mut self, layer: impl Layer) -> Scope {
self.layers.push(Arc::new(layer));
self
}
pub fn layer_fn<F, Fut>(self, f: F) -> Scope
where
F: Fn(http::Request<bytes::Bytes>, Next) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + 'static,
{
self.layer(layer_fn(f))
}
pub fn state<T: Send + Sync + 'static>(mut self, value: T) -> Scope {
self.states.insert(value);
self
}
pub fn scope(mut self, name: &str, f: impl FnOnce(Scope) -> Scope) -> Scope {
let mut path = self.path.clone();
path.push(name.to_string());
let child = f(Scope::new(path, self.layers.clone(), self.states.clone()));
self.services.extend(child.services);
self
}
}
pub struct NodeBuilder {
node: String,
services: Vec<PendingService>,
states: StateMap,
global_layers: Vec<Arc<dyn Layer>>,
max_activations: usize,
epoch: u64,
peer_layers: Vec<Arc<dyn PeerLayer>>,
identity_trust_selected: bool,
connect_timeout: Option<std::time::Duration>,
proof: Value,
ws_collect_ceiling: usize,
}
impl Node {
pub fn builder(node: &str) -> NodeBuilder {
NodeBuilder::new(node)
}
}
impl NodeBuilder {
pub fn new(node: &str) -> NodeBuilder {
NodeBuilder {
node: node.into(),
services: Vec::new(),
states: StateMap::default(),
global_layers: Vec::new(),
max_activations: DEFAULT_MAX_ACTIVATIONS,
epoch: 1,
peer_layers: Vec::new(),
identity_trust_selected: false,
connect_timeout: None,
proof: Value::Null,
ws_collect_ceiling: unb_transport::DEFAULT_MAX_FRAME_SIZE,
}
}
pub fn max_activations(mut self, max: usize) -> NodeBuilder {
self.max_activations = max;
self
}
pub fn ws_collect_ceiling(mut self, bytes: usize) -> NodeBuilder {
self.ws_collect_ceiling = bytes.min(unb_transport::DEFAULT_MAX_FRAME_SIZE);
self
}
pub fn epoch(mut self, epoch: u64) -> NodeBuilder {
self.epoch = epoch;
self
}
pub fn identity_proof(mut self, proof: Value) -> NodeBuilder {
self.proof = proof;
self
}
pub fn service(mut self, handler: impl Handler) -> NodeBuilder {
self.services.push(PendingService {
scopes: Vec::new(),
layers: Vec::new(),
states: StateMap::default(),
service: handler.into_service(),
});
self
}
pub fn scope(mut self, name: &str, f: impl FnOnce(Scope) -> Scope) -> NodeBuilder {
let scope = f(Scope::new(
vec![name.to_string()],
Vec::new(),
StateMap::default(),
));
self.services.extend(scope.services);
self
}
pub fn state<T: Send + Sync + 'static>(mut self, value: T) -> NodeBuilder {
self.states.insert(value);
self
}
pub fn layer(mut self, layer: impl Layer) -> NodeBuilder {
self.global_layers.push(Arc::new(layer));
self
}
pub fn layer_fn<F, Fut>(self, f: F) -> NodeBuilder
where
F: Fn(http::Request<bytes::Bytes>, Next) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<http::Response<ServiceBody>, HandlerError>> + Send + 'static,
{
self.layer(layer_fn(f))
}
pub fn peer_layer(mut self, layer: impl PeerLayer) -> NodeBuilder {
self.peer_layers.push(Arc::new(layer));
self.identity_trust_selected = true;
self
}
pub fn peer_layer_fn<F, Fut>(self, f: F) -> NodeBuilder
where
F: Fn(PeerRequest, PeerNext) -> Fut + Send + Sync + 'static,
Fut: Future<Output = Result<PeerRequest, HandlerError>> + Send + 'static,
{
self.peer_layer(peer_layer_fn(f))
}
pub fn insecure_accept_declared_peer_identities(self) -> NodeBuilder {
self.peer_layer(InsecureAcceptDeclaredPeerIdentities)
}
pub fn connect_timeout(mut self, timeout: std::time::Duration) -> NodeBuilder {
self.connect_timeout = Some(timeout);
self
}
pub fn build(self) -> Result<Arc<Node>, WsError> {
#[cfg(not(all(target_arch = "wasm32", target_os = "unknown")))]
let runtime = unb_runtime::RuntimeHandle::try_current()
.map_err(|_| WsError::Connect("node build requires an async runtime handle".into()))?;
#[cfg(all(target_arch = "wasm32", target_os = "unknown"))]
let runtime = unb_runtime::RuntimeHandle::current();
validate_node_identifier(&self.node)
.map_err(|error| WsError::Connect(error.to_string()))?;
if !self.identity_trust_selected {
return Err(WsError::Connect(
"peer admission must be configured explicitly; add peer_layer(...) or \
insecure_accept_declared_peer_identities()"
.into(),
));
}
let mut node_core = NodeCore::new(&self.node);
let identity = NodeIdentity {
node_id: self.node.clone(),
instance_id: unique_instance_id(&self.node),
epoch: self.epoch,
proof: self.proof,
};
node_core.set_node_identity(identity.clone());
let global_layers: Arc<[Arc<dyn Layer>]> = Arc::from(self.global_layers);
let mut services = BTreeMap::new();
for pending in self.services {
let states = pending.states.merged_over(&self.states);
let mut chain: Vec<Arc<dyn Layer>> = global_layers.iter().cloned().collect();
chain.extend(pending.layers);
let _ = SubjectServices::register(
&mut services,
pending.service,
&pending.scopes,
chain,
&states,
)
.map_err(WsError::Connect)?;
}
let snapshot = NodeSnapshot::new(services, node_core.clone());
let peer_layers: Arc<[Arc<dyn PeerLayer>]> = Arc::from(self.peer_layers);
let cancellation = CancellationToken::new();
let shutdown = cancellation.drop_guard();
let protocol_core = ProtocolCore::with_node(snapshot.node_core.clone());
Ok(Arc::new_cyclic(|weak| Node {
snapshot: Arc::new(ArcSwap::from_pointee(snapshot)),
states: self.states,
global_layers,
mutation_gate: Mutex::new(()),
peers: RwLock::new(HashMap::new()),
sessions: RwLock::new(HashMap::new()),
connections: std::sync::RwLock::new(HashMap::new()),
routes_changed: tokio::sync::watch::channel(0).0,
session_peers: RwLock::new(HashMap::new()),
outbound_sessions: Mutex::new(std::collections::HashSet::new()),
dispatch_slots: Arc::new(Semaphore::new(self.max_activations)),
dispatch_permits: Mutex::new(HashMap::new()),
dispatching: Mutex::new(HashMap::new()),
verified_peers: Mutex::new(HashMap::new()),
candidate_identities: Mutex::new(HashMap::new()),
active: Arc::new(Mutex::new(HashMap::new())),
protocol: ProtocolCoreHandle::spawn(
protocol_core,
Arc::new(ServerEffectExecutor { node: weak.clone() }),
cancellation.clone(),
&runtime,
),
identity,
peer_layers,
dial_policy: unb_client::Peers::with_config(unb_client::DialConfig {
attempt_timeout: self
.connect_timeout
.unwrap_or(unb_client::DEFAULT_DIAL_TIMEOUT),
supported: None,
}),
next_session: AtomicU64::new(0),
ws_collect_ceiling: self.ws_collect_ceiling,
cancellation,
_shutdown: shutdown,
}))
}
}
fn unique_instance_id(node: &str) -> String {
static COUNTER: AtomicU64 = AtomicU64::new(0);
let count = COUNTER.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
#[cfg(not(all(target_arch = "wasm32", target_os = "unknown")))]
{
let nanos = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_nanos())
.unwrap_or(0);
format!("{node}-{}-{nanos:x}-{count}", std::process::id())
}
#[cfg(all(target_arch = "wasm32", target_os = "unknown"))]
format!("{node}-browser-{count}")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn build_without_runtime_returns_typed_error() {
let error = match NodeBuilder::new("node")
.insecure_accept_declared_peer_identities()
.build()
{
Ok(_) => panic!("build unexpectedly succeeded"),
Err(error) => error,
};
assert!(matches!(
error,
WsError::Connect(message)
if message == "node build requires an async runtime handle"
));
}
#[tokio::test]
async fn build_uses_entered_runtime() {
let node = NodeBuilder::new("node")
.insecure_accept_declared_peer_identities()
.build()
.unwrap();
assert_eq!(node.identity().node_id, "node");
}
}