use std::future::Future;
use std::net::SocketAddr;
use std::sync::Arc;
use n0_future::time::Instant;
use unb_runtime::{CancellationToken, DropGuard};
use unb_transport::webtransport::{quic_transport_config, wtransport, SelfSignedIdentity};
use unb_transport::TransportError;
use crate::call::CALL_TIMEOUT;
use crate::identity::{install_crypto_provider, ServerIdentity};
use crate::node::Node;
const EPHEMERAL_ALIGN_ATTEMPTS: usize = 8;
const MAX_SSE_EVENT_BYTES: usize = unb_transport::DEFAULT_MAX_FRAME_SIZE;
const SSE_KEEPALIVE: std::time::Duration = std::time::Duration::from_secs(15);
type WebTransportEndpoint = wtransport::Endpoint<wtransport::endpoint::endpoint_side::Server>;
#[derive(Debug, thiserror::Error)]
pub enum HostError {
#[error("host has every transport disabled")]
NoTransportEnabled,
#[error("TCP rustls config declares ALPN protocols without http/1.1")]
AlpnMissingHttp1,
#[error("WebTransport is configured twice: a config and an external endpoint")]
WebTransportConflict,
#[error("identity PEM is invalid: {0}")]
Identity(String),
#[error("tcp listener port {listener} and webtransport endpoint port {webtransport} disagree")]
ListenerAddressMismatch { listener: u16, webtransport: u16 },
#[error("could not align tcp and udp on one ephemeral port")]
EphemeralAlignmentFailed,
#[error(transparent)]
Transport(#[from] TransportError),
#[error("host i/o: {0}")]
Io(String),
#[error("host task join failed: {0}")]
Join(String),
}
pub enum TcpSecurity {
Plain,
Rustls(Arc<rustls::ServerConfig>),
}
pub struct TcpTransport {
websocket_path: String,
router: axum::Router,
security: TcpSecurity,
}
impl TcpTransport {
pub fn plain() -> TcpTransport {
TcpTransport {
websocket_path: "/".into(),
router: axum::Router::new(),
security: TcpSecurity::Plain,
}
}
pub fn rustls(config: Arc<rustls::ServerConfig>) -> TcpTransport {
TcpTransport {
security: TcpSecurity::Rustls(config),
..TcpTransport::plain()
}
}
pub fn rustls_pem(chain_pem: &[u8], key_pem: &[u8]) -> Result<TcpTransport, HostError> {
let identity = ServerIdentity::from_pem(chain_pem, key_pem)?;
Ok(TcpTransport::rustls(identity.tcp_rustls()?))
}
pub fn websocket_path(mut self, path: impl Into<String>) -> TcpTransport {
self.websocket_path = path.into();
self
}
pub fn merge_router(mut self, router: axum::Router) -> TcpTransport {
self.router = self.router.merge(router);
self
}
fn normalized_security(self) -> Result<TcpTransport, HostError> {
let security = match self.security {
TcpSecurity::Plain => TcpSecurity::Plain,
TcpSecurity::Rustls(config) => {
let http1 = config.alpn_protocols.is_empty()
|| config
.alpn_protocols
.iter()
.any(|protocol| protocol == b"http/1.1");
if !http1 {
return Err(HostError::AlpnMissingHttp1);
}
if config.alpn_protocols.as_slice() == [b"http/1.1".to_vec()] {
TcpSecurity::Rustls(config)
} else {
let mut owned = (*config).clone();
owned.alpn_protocols = vec![b"http/1.1".to_vec()];
TcpSecurity::Rustls(Arc::new(owned))
}
}
};
Ok(TcpTransport { security, ..self })
}
}
enum WebTransportServer {
Identity(Box<wtransport::Identity>),
Config(Box<wtransport::ServerConfig>),
}
pub struct WebTransportConfig {
server: WebTransportServer,
development_cert_hash: Option<[u8; 32]>,
}
impl WebTransportConfig {
pub fn identity(identity: wtransport::Identity) -> WebTransportConfig {
WebTransportConfig {
server: WebTransportServer::Identity(Box::new(identity)),
development_cert_hash: None,
}
}
pub fn server_config(config: wtransport::ServerConfig) -> WebTransportConfig {
WebTransportConfig {
server: WebTransportServer::Config(Box::new(config)),
development_cert_hash: None,
}
}
pub fn pem(chain_pem: &[u8], key_pem: &[u8]) -> Result<WebTransportConfig, HostError> {
let identity = ServerIdentity::from_pem(chain_pem, key_pem)?;
Ok(WebTransportConfig::identity(identity.webtransport()?))
}
pub fn self_signed_for_development<I, S>(hostnames: I) -> Result<WebTransportConfig, HostError>
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let generated = SelfSignedIdentity::generate(hostnames)?;
Ok(WebTransportConfig {
development_cert_hash: Some(generated.cert_hash()),
server: WebTransportServer::Identity(Box::new(generated.identity().clone_identity())),
})
}
}
enum WebTransportSource {
Disabled,
Endpoint(Box<WebTransportEndpoint>),
Identity(Box<wtransport::Identity>),
}
pub struct HostConfig {
bind: SocketAddr,
tcp: Option<TcpTransport>,
webtransport: Option<WebTransportConfig>,
listener: Option<tokio::net::TcpListener>,
endpoint: Option<WebTransportEndpoint>,
drain_deadline: Option<std::time::Duration>,
max_body_bytes: usize,
}
impl HostConfig {
pub fn new(bind: impl Into<SocketAddr>) -> HostConfig {
HostConfig {
bind: bind.into(),
tcp: None,
webtransport: None,
listener: None,
endpoint: None,
drain_deadline: None,
max_body_bytes: unb_transport::DEFAULT_MAX_FRAME_SIZE,
}
}
pub fn with_drain_deadline(mut self, deadline: std::time::Duration) -> HostConfig {
self.drain_deadline = Some(deadline);
self
}
pub fn with_max_body_bytes(mut self, max: usize) -> HostConfig {
self.max_body_bytes = max;
self
}
pub fn tcp(bind: impl Into<SocketAddr>, tcp: TcpTransport) -> HostConfig {
HostConfig::new(bind).with_tcp(tcp)
}
pub fn with_tcp(mut self, tcp: TcpTransport) -> HostConfig {
self.tcp = Some(tcp);
self
}
pub fn with_webtransport(mut self, webtransport: WebTransportConfig) -> HostConfig {
self.webtransport = Some(webtransport);
self
}
pub fn tcp_listener(
mut self,
listener: tokio::net::TcpListener,
tcp: TcpTransport,
) -> HostConfig {
self.listener = Some(listener);
self.tcp = Some(tcp);
self
}
pub fn webtransport_endpoint(mut self, endpoint: WebTransportEndpoint) -> HostConfig {
self.endpoint = Some(endpoint);
self
}
pub(crate) fn validate(&self) -> Result<(), HostError> {
if self.tcp.is_none() && self.webtransport.is_none() && self.endpoint.is_none() {
return Err(HostError::NoTransportEnabled);
}
if self.webtransport.is_some() && self.endpoint.is_some() {
return Err(HostError::WebTransportConflict);
}
if let (Some(listener), Some(endpoint)) = (&self.listener, &self.endpoint) {
let listener_port = HostConfig::local_addr(listener)?.port();
let endpoint_port = endpoint
.local_addr()
.map_err(|error| HostError::Io(error.to_string()))?
.port();
if listener_port != endpoint_port {
return Err(HostError::ListenerAddressMismatch {
listener: listener_port,
webtransport: endpoint_port,
});
}
}
Ok(())
}
pub async fn start(self, node: &Arc<Node>) -> Result<Hosting, HostError> {
self.validate()?;
let HostConfig {
bind,
tcp,
webtransport,
listener,
endpoint,
drain_deadline,
max_body_bytes,
} = self;
let tcp = tcp.map(TcpTransport::normalized_security).transpose()?;
if matches!(
tcp.as_ref().map(|tcp| &tcp.security),
Some(TcpSecurity::Rustls(_))
) {
install_crypto_provider();
}
let cancellation = node.cancellation().child_token();
let mut hosting = Hosting {
websocket: None,
webtransport: None,
development_cert_hash: None,
_guard: cancellation.drop_guard(),
cancellation,
tasks: tokio::task::JoinSet::new(),
drain_deadline,
live_listeners: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
expected_listeners: 0,
};
let source = match (endpoint, webtransport) {
(Some(endpoint), None) => WebTransportSource::Endpoint(Box::new(endpoint)),
(None, Some(config)) => {
hosting.development_cert_hash = config.development_cert_hash;
match config.server {
WebTransportServer::Config(server_config) => {
WebTransportSource::Endpoint(Box::new(
wtransport::Endpoint::server(*server_config)
.map_err(|error| HostError::Io(error.to_string()))?,
))
}
WebTransportServer::Identity(identity) => {
WebTransportSource::Identity(identity)
}
}
}
(None, None) => WebTransportSource::Disabled,
(Some(_), Some(_)) => unreachable!("validate rejects a doubly configured webtransport"),
};
match tcp {
Some(tcp) => {
let (tcp_listener, wt_endpoint) = match (listener, source) {
(Some(listener), WebTransportSource::Identity(identity)) => {
let shared = HostConfig::local_addr(&listener)?;
let endpoint = HostConfig::endpoint_at(*identity, shared)?;
(listener, Some(endpoint))
}
(Some(listener), WebTransportSource::Endpoint(endpoint)) => {
(listener, Some(*endpoint))
}
(Some(listener), WebTransportSource::Disabled) => (listener, None),
(None, WebTransportSource::Identity(identity)) if bind.port() == 0 => {
let mut bound = None;
let mut last = HostError::EphemeralAlignmentFailed;
for _ in 0..EPHEMERAL_ALIGN_ATTEMPTS {
let candidate = tokio::net::TcpListener::bind(bind)
.await
.map_err(|error| HostError::Io(error.to_string()))?;
let shared = HostConfig::local_addr(&candidate)?;
match HostConfig::endpoint_at(identity.clone_identity(), shared) {
Ok(endpoint) => {
bound = Some((candidate, endpoint));
break;
}
Err(error) => last = error,
}
}
match bound {
Some((listener, endpoint)) => (listener, Some(endpoint)),
None => return Err(last),
}
}
(None, source) => {
let listener = tokio::net::TcpListener::bind(bind)
.await
.map_err(|error| HostError::Io(error.to_string()))?;
let endpoint = match source {
WebTransportSource::Identity(identity) => {
Some(HostConfig::endpoint_at(*identity, bind)?)
}
WebTransportSource::Endpoint(endpoint) => Some(*endpoint),
WebTransportSource::Disabled => None,
};
(listener, endpoint)
}
};
if let Some(endpoint) = wt_endpoint {
let (addr, task) = HostConfig::spawn_webtransport(
node.clone(),
endpoint,
hosting.cancellation.child_token(),
)?;
hosting.webtransport = Some(addr);
hosting.spawn_listener(task);
}
let (addr, task) = HostConfig::spawn_websocket(
node.clone(),
tcp_listener,
tcp,
hosting.cancellation.child_token(),
max_body_bytes,
)?;
hosting.websocket = Some(addr);
hosting.spawn_listener(task);
}
None => {
let endpoint = match source {
WebTransportSource::Endpoint(endpoint) => *endpoint,
WebTransportSource::Identity(identity) => {
HostConfig::endpoint_at(*identity, bind)?
}
WebTransportSource::Disabled => {
unreachable!("validate requires an enabled transport")
}
};
let (addr, task) = HostConfig::spawn_webtransport(
node.clone(),
endpoint,
hosting.cancellation.child_token(),
)?;
hosting.webtransport = Some(addr);
hosting.spawn_listener(task);
}
}
Ok(hosting)
}
fn local_addr(listener: &tokio::net::TcpListener) -> Result<SocketAddr, HostError> {
listener
.local_addr()
.map_err(|error| HostError::Io(error.to_string()))
}
fn endpoint_at(
identity: wtransport::Identity,
addr: SocketAddr,
) -> Result<WebTransportEndpoint, HostError> {
let mut config = wtransport::ServerConfig::builder()
.with_bind_address(addr)
.with_custom_transport(identity, quic_transport_config())
.build();
unb_transport::webtransport::raise_endpoint_payload(config.quic_endpoint_config_mut());
wtransport::Endpoint::server(config).map_err(|error| HostError::Io(error.to_string()))
}
async fn ingress(
node: Arc<Node>,
request: axum::extract::Request,
max_body_bytes: usize,
) -> http::Response<axum::body::Body> {
if request.method() != http::Method::POST {
return HostConfig::ingress_error(
http::StatusCode::METHOD_NOT_ALLOWED,
None,
"unb ingress accepts POST only",
);
}
let (mut parts, body) = request.into_parts();
if let Some(name) = parts
.headers
.keys()
.find(|name| name.as_str().starts_with("unb-"))
{
return HostConfig::ingress_error(
http::StatusCode::BAD_REQUEST,
None,
&format!("{name}: unb-* headers are reserved for framing metadata"),
);
}
let raw_target = parts
.uri
.path_and_query()
.map(http::uri::PathAndQuery::as_str)
.unwrap_or_else(|| parts.uri.path());
let target_path = match unb_core::TargetPath::parse_application(raw_target) {
Ok(target_path) => target_path,
Err(error) => {
return HostConfig::ingress_error(
http::StatusCode::BAD_REQUEST,
Some(unb_core::ErrorCode::InvalidInput),
&error.to_string(),
)
}
};
let target = target_path.target().to_owned();
let subject = target_path.subject().to_owned();
let canonical_target = target_path.to_string();
let wants_sse = parts
.headers
.get(http::header::ACCEPT)
.and_then(|value| value.to_str().ok())
.is_some_and(|accept| {
accept
.split(',')
.any(|media| media.trim().split(';').next() == Some("text/event-stream"))
});
let deadline = Instant::now() + CALL_TIMEOUT;
let resolved = node.resolve_unary_until(&target, deadline).await;
let (snapshot, resolution) = match resolved {
Ok(resolved) => resolved,
Err(error) => {
return HostConfig::ingress_error(
error.code.status(),
Some(error.code),
&error.message,
)
}
};
match resolution {
unb_core::Resolution::Unknown => {
let error = Node::teach_unknown_target(&snapshot, &target);
return HostConfig::ingress_error(
error.code.status(),
Some(error.code),
&error.message,
);
}
unb_core::Resolution::Conflicted { owners } => {
let code = unb_core::ErrorCode::PeerUnreachable;
return HostConfig::ingress_error(
code.status(),
Some(code),
&format!(
"target node {target:?} has multiple live incarnations: {}",
owners.join(", ")
),
);
}
unb_core::Resolution::Local => {
if !snapshot.services.contains_key(&subject) {
let error = Node::teach_unknown_subject(&snapshot, &subject);
return HostConfig::ingress_error(
error.code.status(),
Some(error.code),
&error.message,
);
}
}
unb_core::Resolution::Route(_) => {}
}
if parts
.headers
.get(http::header::CONTENT_LENGTH)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.parse::<usize>().ok())
.is_some_and(|length| length > max_body_bytes)
{
return HostConfig::ingress_error(
http::StatusCode::PAYLOAD_TOO_LARGE,
None,
"request body exceeds the ingress body ceiling",
);
}
for name in [
http::header::HOST,
http::header::CONNECTION,
http::header::CONTENT_LENGTH,
http::header::TRANSFER_ENCODING,
http::header::TE,
http::header::TRAILER,
http::header::UPGRADE,
http::header::PROXY_AUTHENTICATE,
http::header::PROXY_AUTHORIZATION,
http::header::EXPECT,
] {
parts.headers.remove(name);
}
parts.headers.remove("keep-alive");
let payload = match axum::body::to_bytes(body, max_body_bytes).await {
Ok(payload) => payload,
Err(error) => {
let mut source: Option<&(dyn std::error::Error + 'static)> = Some(&error);
while let Some(current) = source {
if current.is::<http_body_util::LengthLimitError>() {
return HostConfig::ingress_error(
http::StatusCode::PAYLOAD_TOO_LARGE,
None,
"request body exceeds the ingress body ceiling",
);
}
source = current.source();
}
return HostConfig::ingress_error(
http::StatusCode::BAD_REQUEST,
Some(unb_core::ErrorCode::Protocol),
&error.to_string(),
);
}
};
if wants_sse {
let mut headers = serde_json::Map::new();
for (name, value) in &parts.headers {
if let Ok(value) = value.to_str() {
headers.insert(
name.as_str().to_string(),
serde_json::Value::String(value.to_string()),
);
}
}
return match node
.subscribe_bytes(&canonical_target, payload, headers)
.await
{
Ok(stream) => HostConfig::ingress_sse(stream).await,
Err(error) => {
HostConfig::ingress_error(error.code.status(), Some(error.code), &error.message)
}
};
}
match node
.fetch_until(http::Request::from_parts(parts, payload), deadline)
.await
{
Ok(response) => {
let (parts, body) = response.into_parts();
match body {
crate::layer::ServiceBody::Unary(payload) => {
let json = payload.is_empty()
|| serde_json::from_slice::<serde::de::IgnoredAny>(&payload).is_ok();
let mut response =
http::Response::from_parts(parts, axum::body::Body::from(payload));
response.headers_mut().insert(
http::header::CONTENT_TYPE,
http::HeaderValue::from_static(if json {
"application/json"
} else {
"application/octet-stream"
}),
);
response
}
crate::layer::ServiceBody::Stream(_) => HostConfig::ingress_error(
http::StatusCode::NOT_ACCEPTABLE,
None,
"this subject streams; request it with Accept: text/event-stream",
),
}
}
Err(error) => {
HostConfig::ingress_error(error.code.status(), Some(error.code), &error.message)
}
}
}
async fn ingress_sse(mut stream: crate::EventStream) -> http::Response<axum::body::Body> {
use futures_util::StreamExt;
let first = match stream.next().await {
Some(Ok(first)) => Some(first),
Some(Err(error)) => {
return HostConfig::ingress_error(
error.code.status(),
Some(error.code),
&error.message,
)
}
None => None,
};
if let Some(first) = &first {
if std::str::from_utf8(first).is_err() {
return HostConfig::ingress_error(
http::StatusCode::NOT_ACCEPTABLE,
None,
"stream events are not utf-8 and cannot be projected to SSE",
);
}
if first.len() > MAX_SSE_EVENT_BYTES {
return HostConfig::ingress_error(
http::StatusCode::PAYLOAD_TOO_LARGE,
None,
"stream event exceeds the SSE event size limit",
);
}
}
let body = axum::body::Body::from_stream(futures_util::stream::unfold(
(0u64, first, stream),
|(id, first, mut stream)| async move {
let bytes = match first {
Some(bytes) => bytes,
None => match tokio::time::timeout(SSE_KEEPALIVE, stream.next()).await {
Ok(Some(Ok(bytes))) => {
if bytes.len() > MAX_SSE_EVENT_BYTES
|| std::str::from_utf8(&bytes).is_err()
{
return None;
}
bytes
}
Ok(_) => return None,
Err(_) => {
return Some((
Ok::<_, std::convert::Infallible>(bytes::Bytes::from_static(
b": keepalive\n\n",
)),
(id, None, stream),
));
}
},
};
let Ok(text) = std::str::from_utf8(&bytes) else {
return None;
};
let mut record = format!("id: {id}\n");
for line in text.split('\n') {
record.push_str("data: ");
record.push_str(line);
record.push('\n');
}
record.push('\n');
Some((
Ok::<_, std::convert::Infallible>(bytes::Bytes::from(record)),
(id + 1, None, stream),
))
},
));
http::Response::builder()
.status(http::StatusCode::OK)
.header(
http::header::CONTENT_TYPE,
http::HeaderValue::from_static("text/event-stream"),
)
.header(
http::header::CACHE_CONTROL,
http::HeaderValue::from_static("no-cache"),
)
.header("x-accel-buffering", http::HeaderValue::from_static("no"))
.body(body)
.expect("static SSE response parts are valid")
}
fn ingress_error(
status: http::StatusCode,
code: Option<unb_core::ErrorCode>,
message: &str,
) -> http::Response<axum::body::Body> {
let code = code.unwrap_or_else(|| unb_core::ErrorCode::from_status(status));
let body = serde_json::json!({ "code": code, "message": message });
let mut response = http::Response::builder()
.status(status)
.header(
http::header::CONTENT_TYPE,
http::HeaderValue::from_static("application/json"),
)
.header(
unb_core::UNB_CODE,
http::HeaderValue::from_static(code.token()),
)
.body(axum::body::Body::from(body.to_string()))
.expect("static response parts are valid");
if status == http::StatusCode::METHOD_NOT_ALLOWED {
response
.headers_mut()
.insert(http::header::ALLOW, http::HeaderValue::from_static("POST"));
}
response
}
fn spawn_websocket(
node: Arc<Node>,
listener: tokio::net::TcpListener,
tcp: TcpTransport,
cancellation: CancellationToken,
max_body_bytes: usize,
) -> Result<
(
SocketAddr,
impl Future<Output = Result<(), HostError>> + Send + 'static,
),
HostError,
> {
let addr = HostConfig::local_addr(&listener)?;
let ingress_node = node.clone();
let app = axum::Router::new()
.route(
&tcp.websocket_path,
axum::routing::get(move |upgrade: axum::extract::ws::WebSocketUpgrade| {
let node = node.clone();
async move { node.serve_ws_upgrade(upgrade) }
}),
)
.merge(tcp.router)
.fallback(move |request: axum::extract::Request| {
let node = ingress_node.clone();
async move { HostConfig::ingress(node, request, max_body_bytes).await }
});
let task: futures_util::future::Either<_, _> = match tcp.security {
TcpSecurity::Plain => futures_util::future::Either::Left(async move {
axum::serve(listener, app)
.with_graceful_shutdown(async move { cancellation.cancelled().await })
.await
.map_err(|error| HostError::Io(error.to_string()))
}),
TcpSecurity::Rustls(config) => {
let acceptor = tokio_rustls::TlsAcceptor::from(config);
futures_util::future::Either::Right(async move {
let mut connections = tokio::task::JoinSet::new();
loop {
tokio::select! {
biased;
() = cancellation.cancelled() => break,
completed = connections.join_next(), if !connections.is_empty() => {
if let Some(Err(error)) = completed {
return Err(HostError::Join(error.to_string()));
}
}
accepted = listener.accept() => {
let (stream, _peer) = match accepted {
Ok(accepted) => accepted,
Err(error) => return Err(HostError::Io(error.to_string())),
};
let acceptor = acceptor.clone();
let service =
hyper_util::service::TowerToHyperService::new(app.clone());
let cancel = cancellation.child_token();
connections.spawn(async move {
let serve = async move {
let Ok(tls) = acceptor.accept(stream).await else {
return;
};
let io = hyper_util::rt::TokioIo::new(tls);
let builder = hyper_util::server::conn::auto::Builder::new(
hyper_util::rt::TokioExecutor::new(),
);
let _ = builder
.http1_only()
.serve_connection_with_upgrades(io, service)
.await;
};
tokio::select! {
biased;
() = cancel.cancelled() => {}
() = serve => {}
}
});
}
}
}
while let Some(result) = connections.join_next().await {
result.map_err(|error| HostError::Join(error.to_string()))?;
}
Ok(())
})
}
};
Ok((addr, task))
}
fn spawn_webtransport(
node: Arc<Node>,
endpoint: WebTransportEndpoint,
cancellation: CancellationToken,
) -> Result<
(
SocketAddr,
impl Future<Output = Result<(), HostError>> + Send + 'static,
),
HostError,
> {
let bound = endpoint
.local_addr()
.map_err(|error| HostError::Io(error.to_string()))?;
let task = async move {
let mut connections = tokio::task::JoinSet::new();
loop {
tokio::select! {
biased;
() = cancellation.cancelled() => break,
completed = connections.join_next(), if !connections.is_empty() => {
if let Some(Err(error)) = completed {
return Err(HostError::Join(error.to_string()));
}
}
incoming = endpoint.accept() => {
let node = node.clone();
let cancel = cancellation.child_token();
connections.spawn(async move {
let accept = async {
let Ok(session_request) = incoming.await else {
return;
};
let Ok(connection) = session_request.accept().await else {
return;
};
let _ = node.serve_webtransport(connection).await;
};
tokio::select! {
biased;
() = cancel.cancelled() => {}
result = tokio::time::timeout(crate::node::WEBTRANSPORT_ACCEPT_TIMEOUT, accept) => {
let _ = result;
}
}
});
}
}
}
while let Some(result) = connections.join_next().await {
result.map_err(|error| HostError::Join(error.to_string()))?;
}
Ok(())
};
Ok((bound, task))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct HealthStatus {
pub process_alive: bool,
pub websocket_bound: bool,
pub websocket_addr: Option<SocketAddr>,
pub webtransport_bound: bool,
pub webtransport_addr: Option<SocketAddr>,
pub listeners_running: bool,
pub parent_link_ready: bool,
pub child_link_ready: bool,
}
impl HealthStatus {
pub fn ready(&self) -> bool {
self.process_alive
&& self.listeners_running
&& (self.websocket_bound || self.webtransport_bound)
&& self.parent_link_ready
&& self.child_link_ready
}
}
struct ListenerGuard(std::sync::Arc<std::sync::atomic::AtomicUsize>);
impl Drop for ListenerGuard {
fn drop(&mut self) {
self.0.fetch_sub(1, std::sync::atomic::Ordering::Relaxed);
}
}
pub struct Hosting {
websocket: Option<SocketAddr>,
webtransport: Option<SocketAddr>,
development_cert_hash: Option<[u8; 32]>,
cancellation: CancellationToken,
_guard: DropGuard,
tasks: tokio::task::JoinSet<Result<(), HostError>>,
drain_deadline: Option<std::time::Duration>,
live_listeners: std::sync::Arc<std::sync::atomic::AtomicUsize>,
expected_listeners: usize,
}
impl Hosting {
pub fn websocket_addr(&self) -> Option<SocketAddr> {
self.websocket
}
pub fn webtransport_addr(&self) -> Option<SocketAddr> {
self.webtransport
}
pub fn development_cert_hash(&self) -> Option<[u8; 32]> {
self.development_cert_hash
}
pub fn cancel(&self) {
self.cancellation.cancel();
}
pub fn is_finished(&self) -> bool {
self.tasks.is_empty()
}
pub fn health(&self) -> HealthStatus {
HealthStatus {
process_alive: true,
websocket_bound: self.websocket.is_some(),
websocket_addr: self.websocket,
webtransport_bound: self.webtransport.is_some(),
webtransport_addr: self.webtransport,
listeners_running: self.expected_listeners > 0
&& self
.live_listeners
.load(std::sync::atomic::Ordering::Relaxed)
== self.expected_listeners,
parent_link_ready: true,
child_link_ready: true,
}
}
fn spawn_listener(
&mut self,
task: impl std::future::Future<Output = Result<(), HostError>> + Send + 'static,
) {
self.expected_listeners += 1;
self.live_listeners
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let guard = ListenerGuard(self.live_listeners.clone());
self.tasks.spawn(async move {
let _guard = guard;
task.await
});
}
pub async fn shutdown(mut self) -> Result<(), HostError> {
self.cancellation.cancel();
let mut failure = None;
match self.drain_deadline {
None => Self::join_all(&mut self.tasks, &mut failure).await,
Some(deadline) => {
if tokio::time::timeout(deadline, Self::join_all(&mut self.tasks, &mut failure))
.await
.is_err()
{
self.tasks.abort_all();
while let Some(result) = self.tasks.join_next().await {
if let Ok(Err(error)) = result {
if failure.is_none() {
failure = Some(error);
}
}
}
}
}
}
failure.map_or(Ok(()), Err)
}
async fn join_all(
tasks: &mut tokio::task::JoinSet<Result<(), HostError>>,
failure: &mut Option<HostError>,
) {
while let Some(result) = tasks.join_next().await {
let result = result
.map_err(|error| HostError::Join(error.to_string()))
.and_then(|result| result);
if failure.is_none() {
*failure = result.err();
}
}
}
pub async fn wait(&mut self) -> Result<(), HostError> {
let Some(result) = self.tasks.join_next().await else {
return Ok(());
};
result.map_err(|error| HostError::Join(error.to_string()))?
}
}