unb-server 2.0.3

unb inbound server: Node, request/subscribe handlers, catalog, relay orchestration, accept
Documentation
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()),
            peer_admissions: 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");
    }
}