1use std::future::Future;
2use std::net::SocketAddr;
3use std::sync::Arc;
4
5use n0_future::time::Instant;
6use unb_runtime::{CancellationToken, DropGuard};
7use unb_transport::webtransport::{quic_transport_config, wtransport, SelfSignedIdentity};
8use unb_transport::TransportError;
9
10use crate::call::CALL_TIMEOUT;
11use crate::identity::{install_crypto_provider, ServerIdentity};
12use crate::node::Node;
13
14const EPHEMERAL_ALIGN_ATTEMPTS: usize = 8;
15const MAX_SSE_EVENT_BYTES: usize = unb_transport::DEFAULT_MAX_FRAME_SIZE;
16const SSE_KEEPALIVE: std::time::Duration = std::time::Duration::from_secs(15);
17
18type WebTransportEndpoint = wtransport::Endpoint<wtransport::endpoint::endpoint_side::Server>;
19
20#[derive(Debug, thiserror::Error)]
21pub enum HostError {
22 #[error("host has every transport disabled")]
23 NoTransportEnabled,
24 #[error("TCP rustls config declares ALPN protocols without http/1.1")]
25 AlpnMissingHttp1,
26 #[error("WebTransport is configured twice: a config and an external endpoint")]
27 WebTransportConflict,
28 #[error("identity PEM is invalid: {0}")]
29 Identity(String),
30 #[error("tcp listener port {listener} and webtransport endpoint port {webtransport} disagree")]
31 ListenerAddressMismatch { listener: u16, webtransport: u16 },
32 #[error("could not align tcp and udp on one ephemeral port")]
33 EphemeralAlignmentFailed,
34 #[error(transparent)]
35 Transport(#[from] TransportError),
36 #[error("host i/o: {0}")]
37 Io(String),
38 #[error("host task join failed: {0}")]
39 Join(String),
40}
41
42pub enum TcpSecurity {
43 Plain,
44 Rustls(Arc<rustls::ServerConfig>),
45}
46
47pub struct TcpTransport {
48 websocket_path: String,
49 router: axum::Router,
50 security: TcpSecurity,
51}
52
53impl TcpTransport {
54 pub fn plain() -> TcpTransport {
55 TcpTransport {
56 websocket_path: "/".into(),
57 router: axum::Router::new(),
58 security: TcpSecurity::Plain,
59 }
60 }
61
62 pub fn rustls(config: Arc<rustls::ServerConfig>) -> TcpTransport {
63 TcpTransport {
64 security: TcpSecurity::Rustls(config),
65 ..TcpTransport::plain()
66 }
67 }
68
69 pub fn rustls_pem(chain_pem: &[u8], key_pem: &[u8]) -> Result<TcpTransport, HostError> {
70 let identity = ServerIdentity::from_pem(chain_pem, key_pem)?;
71 Ok(TcpTransport::rustls(identity.tcp_rustls()?))
72 }
73
74 pub fn websocket_path(mut self, path: impl Into<String>) -> TcpTransport {
75 self.websocket_path = path.into();
76 self
77 }
78
79 pub fn merge_router(mut self, router: axum::Router) -> TcpTransport {
80 self.router = self.router.merge(router);
81 self
82 }
83
84 fn normalized_security(self) -> Result<TcpTransport, HostError> {
85 let security = match self.security {
86 TcpSecurity::Plain => TcpSecurity::Plain,
87 TcpSecurity::Rustls(config) => {
88 let http1 = config.alpn_protocols.is_empty()
89 || config
90 .alpn_protocols
91 .iter()
92 .any(|protocol| protocol == b"http/1.1");
93 if !http1 {
94 return Err(HostError::AlpnMissingHttp1);
95 }
96 if config.alpn_protocols.as_slice() == [b"http/1.1".to_vec()] {
97 TcpSecurity::Rustls(config)
98 } else {
99 let mut owned = (*config).clone();
100 owned.alpn_protocols = vec![b"http/1.1".to_vec()];
101 TcpSecurity::Rustls(Arc::new(owned))
102 }
103 }
104 };
105 Ok(TcpTransport { security, ..self })
106 }
107}
108
109enum WebTransportServer {
110 Identity(Box<wtransport::Identity>),
111 Config(Box<wtransport::ServerConfig>),
112}
113
114pub struct WebTransportConfig {
115 server: WebTransportServer,
116 development_cert_hash: Option<[u8; 32]>,
117}
118
119impl WebTransportConfig {
120 pub fn identity(identity: wtransport::Identity) -> WebTransportConfig {
121 WebTransportConfig {
122 server: WebTransportServer::Identity(Box::new(identity)),
123 development_cert_hash: None,
124 }
125 }
126
127 pub fn server_config(config: wtransport::ServerConfig) -> WebTransportConfig {
128 WebTransportConfig {
129 server: WebTransportServer::Config(Box::new(config)),
130 development_cert_hash: None,
131 }
132 }
133
134 pub fn pem(chain_pem: &[u8], key_pem: &[u8]) -> Result<WebTransportConfig, HostError> {
135 let identity = ServerIdentity::from_pem(chain_pem, key_pem)?;
136 Ok(WebTransportConfig::identity(identity.webtransport()?))
137 }
138
139 pub fn self_signed_for_development<I, S>(hostnames: I) -> Result<WebTransportConfig, HostError>
140 where
141 I: IntoIterator<Item = S>,
142 S: AsRef<str>,
143 {
144 let generated = SelfSignedIdentity::generate(hostnames)?;
145 Ok(WebTransportConfig {
146 development_cert_hash: Some(generated.cert_hash()),
147 server: WebTransportServer::Identity(Box::new(generated.identity().clone_identity())),
148 })
149 }
150}
151
152enum WebTransportSource {
153 Disabled,
154 Endpoint(Box<WebTransportEndpoint>),
155 Identity(Box<wtransport::Identity>),
156}
157
158pub struct HostConfig {
159 bind: SocketAddr,
160 tcp: Option<TcpTransport>,
161 webtransport: Option<WebTransportConfig>,
162 listener: Option<tokio::net::TcpListener>,
163 endpoint: Option<WebTransportEndpoint>,
164 drain_deadline: Option<std::time::Duration>,
165 max_body_bytes: usize,
166}
167
168impl HostConfig {
169 pub fn new(bind: impl Into<SocketAddr>) -> HostConfig {
170 HostConfig {
171 bind: bind.into(),
172 tcp: None,
173 webtransport: None,
174 listener: None,
175 endpoint: None,
176 drain_deadline: None,
177 max_body_bytes: unb_transport::DEFAULT_MAX_FRAME_SIZE,
178 }
179 }
180
181 pub fn with_drain_deadline(mut self, deadline: std::time::Duration) -> HostConfig {
182 self.drain_deadline = Some(deadline);
183 self
184 }
185
186 pub fn with_max_body_bytes(mut self, max: usize) -> HostConfig {
187 self.max_body_bytes = max;
188 self
189 }
190
191 pub fn tcp(bind: impl Into<SocketAddr>, tcp: TcpTransport) -> HostConfig {
192 HostConfig::new(bind).with_tcp(tcp)
193 }
194
195 pub fn with_tcp(mut self, tcp: TcpTransport) -> HostConfig {
196 self.tcp = Some(tcp);
197 self
198 }
199
200 pub fn with_webtransport(mut self, webtransport: WebTransportConfig) -> HostConfig {
201 self.webtransport = Some(webtransport);
202 self
203 }
204
205 pub fn tcp_listener(
206 mut self,
207 listener: tokio::net::TcpListener,
208 tcp: TcpTransport,
209 ) -> HostConfig {
210 self.listener = Some(listener);
211 self.tcp = Some(tcp);
212 self
213 }
214
215 pub fn webtransport_endpoint(mut self, endpoint: WebTransportEndpoint) -> HostConfig {
216 self.endpoint = Some(endpoint);
217 self
218 }
219
220 pub(crate) fn validate(&self) -> Result<(), HostError> {
221 if self.tcp.is_none() && self.webtransport.is_none() && self.endpoint.is_none() {
222 return Err(HostError::NoTransportEnabled);
223 }
224 if self.webtransport.is_some() && self.endpoint.is_some() {
225 return Err(HostError::WebTransportConflict);
226 }
227 if let (Some(listener), Some(endpoint)) = (&self.listener, &self.endpoint) {
228 let listener_port = HostConfig::local_addr(listener)?.port();
229 let endpoint_port = endpoint
230 .local_addr()
231 .map_err(|error| HostError::Io(error.to_string()))?
232 .port();
233 if listener_port != endpoint_port {
234 return Err(HostError::ListenerAddressMismatch {
235 listener: listener_port,
236 webtransport: endpoint_port,
237 });
238 }
239 }
240 Ok(())
241 }
242
243 pub async fn start(self, node: &Arc<Node>) -> Result<Hosting, HostError> {
244 self.validate()?;
245 let HostConfig {
246 bind,
247 tcp,
248 webtransport,
249 listener,
250 endpoint,
251 drain_deadline,
252 max_body_bytes,
253 } = self;
254 let tcp = tcp.map(TcpTransport::normalized_security).transpose()?;
255 if matches!(
256 tcp.as_ref().map(|tcp| &tcp.security),
257 Some(TcpSecurity::Rustls(_))
258 ) {
259 install_crypto_provider();
260 }
261 let cancellation = node.cancellation().child_token();
262 let mut hosting = Hosting {
263 websocket: None,
264 webtransport: None,
265 development_cert_hash: None,
266 _guard: cancellation.drop_guard(),
267 cancellation,
268 tasks: tokio::task::JoinSet::new(),
269 drain_deadline,
270 live_listeners: std::sync::Arc::new(std::sync::atomic::AtomicUsize::new(0)),
271 expected_listeners: 0,
272 };
273 let source = match (endpoint, webtransport) {
274 (Some(endpoint), None) => WebTransportSource::Endpoint(Box::new(endpoint)),
275 (None, Some(config)) => {
276 hosting.development_cert_hash = config.development_cert_hash;
277 match config.server {
278 WebTransportServer::Config(server_config) => {
279 WebTransportSource::Endpoint(Box::new(
280 wtransport::Endpoint::server(*server_config)
281 .map_err(|error| HostError::Io(error.to_string()))?,
282 ))
283 }
284 WebTransportServer::Identity(identity) => {
285 WebTransportSource::Identity(identity)
286 }
287 }
288 }
289 (None, None) => WebTransportSource::Disabled,
290 (Some(_), Some(_)) => unreachable!("validate rejects a doubly configured webtransport"),
291 };
292 match tcp {
293 Some(tcp) => {
294 let (tcp_listener, wt_endpoint) = match (listener, source) {
295 (Some(listener), WebTransportSource::Identity(identity)) => {
296 let shared = HostConfig::local_addr(&listener)?;
297 let endpoint = HostConfig::endpoint_at(*identity, shared)?;
298 (listener, Some(endpoint))
299 }
300 (Some(listener), WebTransportSource::Endpoint(endpoint)) => {
301 (listener, Some(*endpoint))
302 }
303 (Some(listener), WebTransportSource::Disabled) => (listener, None),
304 (None, WebTransportSource::Identity(identity)) if bind.port() == 0 => {
305 let mut bound = None;
306 let mut last = HostError::EphemeralAlignmentFailed;
307 for _ in 0..EPHEMERAL_ALIGN_ATTEMPTS {
308 let candidate = tokio::net::TcpListener::bind(bind)
309 .await
310 .map_err(|error| HostError::Io(error.to_string()))?;
311 let shared = HostConfig::local_addr(&candidate)?;
312 match HostConfig::endpoint_at(identity.clone_identity(), shared) {
313 Ok(endpoint) => {
314 bound = Some((candidate, endpoint));
315 break;
316 }
317 Err(error) => last = error,
318 }
319 }
320 match bound {
321 Some((listener, endpoint)) => (listener, Some(endpoint)),
322 None => return Err(last),
323 }
324 }
325 (None, source) => {
326 let listener = tokio::net::TcpListener::bind(bind)
327 .await
328 .map_err(|error| HostError::Io(error.to_string()))?;
329 let endpoint = match source {
330 WebTransportSource::Identity(identity) => {
331 Some(HostConfig::endpoint_at(*identity, bind)?)
332 }
333 WebTransportSource::Endpoint(endpoint) => Some(*endpoint),
334 WebTransportSource::Disabled => None,
335 };
336 (listener, endpoint)
337 }
338 };
339 if let Some(endpoint) = wt_endpoint {
340 let (addr, task) = HostConfig::spawn_webtransport(
341 node.clone(),
342 endpoint,
343 hosting.cancellation.child_token(),
344 )?;
345 hosting.webtransport = Some(addr);
346 hosting.spawn_listener(task);
347 }
348 let (addr, task) = HostConfig::spawn_websocket(
349 node.clone(),
350 tcp_listener,
351 tcp,
352 hosting.cancellation.child_token(),
353 max_body_bytes,
354 )?;
355 hosting.websocket = Some(addr);
356 hosting.spawn_listener(task);
357 }
358 None => {
359 let endpoint = match source {
360 WebTransportSource::Endpoint(endpoint) => *endpoint,
361 WebTransportSource::Identity(identity) => {
362 HostConfig::endpoint_at(*identity, bind)?
363 }
364 WebTransportSource::Disabled => {
365 unreachable!("validate requires an enabled transport")
366 }
367 };
368 let (addr, task) = HostConfig::spawn_webtransport(
369 node.clone(),
370 endpoint,
371 hosting.cancellation.child_token(),
372 )?;
373 hosting.webtransport = Some(addr);
374 hosting.spawn_listener(task);
375 }
376 }
377 Ok(hosting)
378 }
379
380 fn local_addr(listener: &tokio::net::TcpListener) -> Result<SocketAddr, HostError> {
381 listener
382 .local_addr()
383 .map_err(|error| HostError::Io(error.to_string()))
384 }
385
386 fn endpoint_at(
387 identity: wtransport::Identity,
388 addr: SocketAddr,
389 ) -> Result<WebTransportEndpoint, HostError> {
390 let mut config = wtransport::ServerConfig::builder()
391 .with_bind_address(addr)
392 .with_custom_transport(identity, quic_transport_config())
393 .build();
394 unb_transport::webtransport::raise_endpoint_payload(config.quic_endpoint_config_mut());
395 wtransport::Endpoint::server(config).map_err(|error| HostError::Io(error.to_string()))
396 }
397
398 async fn ingress(
399 node: Arc<Node>,
400 request: axum::extract::Request,
401 max_body_bytes: usize,
402 ) -> http::Response<axum::body::Body> {
403 if request.method() != http::Method::POST {
404 return HostConfig::ingress_error(
405 http::StatusCode::METHOD_NOT_ALLOWED,
406 None,
407 "unb ingress accepts POST only",
408 );
409 }
410 let (mut parts, body) = request.into_parts();
411 if let Some(name) = parts
412 .headers
413 .keys()
414 .find(|name| name.as_str().starts_with("unb-"))
415 {
416 return HostConfig::ingress_error(
417 http::StatusCode::BAD_REQUEST,
418 None,
419 &format!("{name}: unb-* headers are reserved for framing metadata"),
420 );
421 }
422 let subject = unb_core::Envelope::subject_of(&parts.uri);
423 if subject == "az" || subject.starts_with("az.") {
424 return HostConfig::ingress_error(
425 http::StatusCode::BAD_REQUEST,
426 Some(unb_core::ErrorCode::InvalidInput),
427 "az is a reserved subject namespace",
428 );
429 }
430 let wants_sse = parts
431 .headers
432 .get(http::header::ACCEPT)
433 .and_then(|value| value.to_str().ok())
434 .is_some_and(|accept| {
435 accept
436 .split(',')
437 .any(|media| media.trim().split(';').next() == Some("text/event-stream"))
438 });
439 let deadline = Instant::now() + CALL_TIMEOUT;
440 let resolved = if wants_sse {
441 let snapshot = node.snapshot.load_full();
442 let resolution = snapshot.node_core.resolve(&subject);
443 Ok((snapshot, resolution))
444 } else {
445 node.resolve_unary_until(&subject, deadline).await
446 };
447 let (snapshot, resolution) = match resolved {
448 Ok(resolved) => resolved,
449 Err(error) => {
450 return HostConfig::ingress_error(
451 error.code.status(),
452 Some(error.code),
453 &error.message,
454 )
455 }
456 };
457 match resolution {
458 unb_core::Resolution::Unknown => {
459 let error = Node::teach_unknown_subject(&snapshot, &subject);
460 return HostConfig::ingress_error(
461 error.code.status(),
462 Some(error.code),
463 &error.message,
464 );
465 }
466 unb_core::Resolution::Conflicted { owners } => {
467 return HostConfig::ingress_error(
468 http::StatusCode::CONFLICT,
469 Some(unb_core::ErrorCode::Conflict),
470 &format!(
471 "subject {subject:?} is claimed by multiple live owners: {}",
472 owners.join(", ")
473 ),
474 );
475 }
476 unb_core::Resolution::Local | unb_core::Resolution::Route(_) => {}
477 }
478 if parts
479 .headers
480 .get(http::header::CONTENT_LENGTH)
481 .and_then(|value| value.to_str().ok())
482 .and_then(|value| value.parse::<usize>().ok())
483 .is_some_and(|length| length > max_body_bytes)
484 {
485 return HostConfig::ingress_error(
486 http::StatusCode::PAYLOAD_TOO_LARGE,
487 None,
488 "request body exceeds the ingress body ceiling",
489 );
490 }
491 for name in [
492 http::header::HOST,
493 http::header::CONNECTION,
494 http::header::CONTENT_LENGTH,
495 http::header::TRANSFER_ENCODING,
496 http::header::TE,
497 http::header::TRAILER,
498 http::header::UPGRADE,
499 http::header::PROXY_AUTHENTICATE,
500 http::header::PROXY_AUTHORIZATION,
501 http::header::EXPECT,
502 ] {
503 parts.headers.remove(name);
504 }
505 parts.headers.remove("keep-alive");
506 let payload = match axum::body::to_bytes(body, max_body_bytes).await {
507 Ok(payload) => payload,
508 Err(error) => {
509 let mut source: Option<&(dyn std::error::Error + 'static)> = Some(&error);
510 while let Some(current) = source {
511 if current.is::<http_body_util::LengthLimitError>() {
512 return HostConfig::ingress_error(
513 http::StatusCode::PAYLOAD_TOO_LARGE,
514 None,
515 "request body exceeds the ingress body ceiling",
516 );
517 }
518 source = current.source();
519 }
520 return HostConfig::ingress_error(
521 http::StatusCode::BAD_REQUEST,
522 Some(unb_core::ErrorCode::Protocol),
523 &error.to_string(),
524 );
525 }
526 };
527 if wants_sse {
528 let mut headers = serde_json::Map::new();
529 for (name, value) in &parts.headers {
530 if let Ok(value) = value.to_str() {
531 headers.insert(
532 name.as_str().to_string(),
533 serde_json::Value::String(value.to_string()),
534 );
535 }
536 }
537 return match node.subscribe_bytes(&subject, payload, headers).await {
538 Ok(stream) => HostConfig::ingress_sse(stream).await,
539 Err(error) => {
540 HostConfig::ingress_error(error.code.status(), Some(error.code), &error.message)
541 }
542 };
543 }
544 match node
545 .fetch_until(http::Request::from_parts(parts, payload), deadline)
546 .await
547 {
548 Ok(response) => {
549 let (parts, body) = response.into_parts();
550 match body {
551 crate::layer::ServiceBody::Unary(payload) => {
552 let json = payload.is_empty()
553 || serde_json::from_slice::<serde::de::IgnoredAny>(&payload).is_ok();
554 let mut response =
555 http::Response::from_parts(parts, axum::body::Body::from(payload));
556 response.headers_mut().insert(
557 http::header::CONTENT_TYPE,
558 http::HeaderValue::from_static(if json {
559 "application/json"
560 } else {
561 "application/octet-stream"
562 }),
563 );
564 response
565 }
566 crate::layer::ServiceBody::Stream(_) => HostConfig::ingress_error(
567 http::StatusCode::NOT_ACCEPTABLE,
568 None,
569 "this subject streams; request it with Accept: text/event-stream",
570 ),
571 }
572 }
573 Err(error) => {
574 HostConfig::ingress_error(error.code.status(), Some(error.code), &error.message)
575 }
576 }
577 }
578
579 async fn ingress_sse(mut stream: crate::EventStream) -> http::Response<axum::body::Body> {
580 use futures_util::StreamExt;
581 let first = match stream.next().await {
582 Some(Ok(first)) => Some(first),
583 Some(Err(error)) => {
584 return HostConfig::ingress_error(
585 error.code.status(),
586 Some(error.code),
587 &error.message,
588 )
589 }
590 None => None,
591 };
592 if let Some(first) = &first {
593 if std::str::from_utf8(first).is_err() {
594 return HostConfig::ingress_error(
595 http::StatusCode::NOT_ACCEPTABLE,
596 None,
597 "stream events are not utf-8 and cannot be projected to SSE",
598 );
599 }
600 if first.len() > MAX_SSE_EVENT_BYTES {
601 return HostConfig::ingress_error(
602 http::StatusCode::PAYLOAD_TOO_LARGE,
603 None,
604 "stream event exceeds the SSE event size limit",
605 );
606 }
607 }
608 let body = axum::body::Body::from_stream(futures_util::stream::unfold(
609 (0u64, first, stream),
610 |(id, first, mut stream)| async move {
611 let bytes = match first {
612 Some(bytes) => bytes,
613 None => match tokio::time::timeout(SSE_KEEPALIVE, stream.next()).await {
614 Ok(Some(Ok(bytes))) => {
615 if bytes.len() > MAX_SSE_EVENT_BYTES
616 || std::str::from_utf8(&bytes).is_err()
617 {
618 return None;
619 }
620 bytes
621 }
622 Ok(_) => return None,
623 Err(_) => {
624 return Some((
625 Ok::<_, std::convert::Infallible>(bytes::Bytes::from_static(
626 b": keepalive\n\n",
627 )),
628 (id, None, stream),
629 ));
630 }
631 },
632 };
633 let Ok(text) = std::str::from_utf8(&bytes) else {
634 return None;
635 };
636 let mut record = format!("id: {id}\n");
637 for line in text.split('\n') {
638 record.push_str("data: ");
639 record.push_str(line);
640 record.push('\n');
641 }
642 record.push('\n');
643 Some((
644 Ok::<_, std::convert::Infallible>(bytes::Bytes::from(record)),
645 (id + 1, None, stream),
646 ))
647 },
648 ));
649 http::Response::builder()
650 .status(http::StatusCode::OK)
651 .header(
652 http::header::CONTENT_TYPE,
653 http::HeaderValue::from_static("text/event-stream"),
654 )
655 .header(
656 http::header::CACHE_CONTROL,
657 http::HeaderValue::from_static("no-cache"),
658 )
659 .header("x-accel-buffering", http::HeaderValue::from_static("no"))
660 .body(body)
661 .expect("static SSE response parts are valid")
662 }
663
664 fn ingress_error(
665 status: http::StatusCode,
666 code: Option<unb_core::ErrorCode>,
667 message: &str,
668 ) -> http::Response<axum::body::Body> {
669 let code = code.unwrap_or_else(|| unb_core::ErrorCode::from_status(status));
670 let body = serde_json::json!({ "code": code, "message": message });
671 let mut response = http::Response::builder()
672 .status(status)
673 .header(
674 http::header::CONTENT_TYPE,
675 http::HeaderValue::from_static("application/json"),
676 )
677 .header(
678 unb_core::UNB_CODE,
679 http::HeaderValue::from_static(code.token()),
680 )
681 .body(axum::body::Body::from(body.to_string()))
682 .expect("static response parts are valid");
683 if status == http::StatusCode::METHOD_NOT_ALLOWED {
684 response
685 .headers_mut()
686 .insert(http::header::ALLOW, http::HeaderValue::from_static("POST"));
687 }
688 response
689 }
690
691 fn spawn_websocket(
692 node: Arc<Node>,
693 listener: tokio::net::TcpListener,
694 tcp: TcpTransport,
695 cancellation: CancellationToken,
696 max_body_bytes: usize,
697 ) -> Result<
698 (
699 SocketAddr,
700 impl Future<Output = Result<(), HostError>> + Send + 'static,
701 ),
702 HostError,
703 > {
704 let addr = HostConfig::local_addr(&listener)?;
705 let ingress_node = node.clone();
706 let app = axum::Router::new()
707 .route(
708 &tcp.websocket_path,
709 axum::routing::get(move |upgrade: axum::extract::ws::WebSocketUpgrade| {
710 let node = node.clone();
711 async move { node.serve_ws_upgrade(upgrade) }
712 }),
713 )
714 .merge(tcp.router)
715 .fallback(move |request: axum::extract::Request| {
716 let node = ingress_node.clone();
717 async move { HostConfig::ingress(node, request, max_body_bytes).await }
718 });
719 let task: futures_util::future::Either<_, _> = match tcp.security {
720 TcpSecurity::Plain => futures_util::future::Either::Left(async move {
721 axum::serve(listener, app)
722 .with_graceful_shutdown(async move { cancellation.cancelled().await })
723 .await
724 .map_err(|error| HostError::Io(error.to_string()))
725 }),
726 TcpSecurity::Rustls(config) => {
727 let acceptor = tokio_rustls::TlsAcceptor::from(config);
728 futures_util::future::Either::Right(async move {
729 let mut connections = tokio::task::JoinSet::new();
730 loop {
731 tokio::select! {
732 biased;
733 () = cancellation.cancelled() => break,
734 completed = connections.join_next(), if !connections.is_empty() => {
735 if let Some(Err(error)) = completed {
736 return Err(HostError::Join(error.to_string()));
737 }
738 }
739 accepted = listener.accept() => {
740 let (stream, _peer) = match accepted {
741 Ok(accepted) => accepted,
742 Err(error) => return Err(HostError::Io(error.to_string())),
743 };
744 let acceptor = acceptor.clone();
745 let service =
746 hyper_util::service::TowerToHyperService::new(app.clone());
747 let cancel = cancellation.child_token();
748 connections.spawn(async move {
749 let serve = async move {
750 let Ok(tls) = acceptor.accept(stream).await else {
751 return;
752 };
753 let io = hyper_util::rt::TokioIo::new(tls);
754 let builder = hyper_util::server::conn::auto::Builder::new(
755 hyper_util::rt::TokioExecutor::new(),
756 );
757 let _ = builder
758 .http1_only()
759 .serve_connection_with_upgrades(io, service)
760 .await;
761 };
762 tokio::select! {
763 biased;
764 () = cancel.cancelled() => {}
765 () = serve => {}
766 }
767 });
768 }
769 }
770 }
771 while let Some(result) = connections.join_next().await {
772 result.map_err(|error| HostError::Join(error.to_string()))?;
773 }
774 Ok(())
775 })
776 }
777 };
778 Ok((addr, task))
779 }
780
781 fn spawn_webtransport(
782 node: Arc<Node>,
783 endpoint: WebTransportEndpoint,
784 cancellation: CancellationToken,
785 ) -> Result<
786 (
787 SocketAddr,
788 impl Future<Output = Result<(), HostError>> + Send + 'static,
789 ),
790 HostError,
791 > {
792 let bound = endpoint
793 .local_addr()
794 .map_err(|error| HostError::Io(error.to_string()))?;
795 let task = async move {
796 let mut connections = tokio::task::JoinSet::new();
797 loop {
798 tokio::select! {
799 biased;
800 () = cancellation.cancelled() => break,
801 completed = connections.join_next(), if !connections.is_empty() => {
802 if let Some(Err(error)) = completed {
803 return Err(HostError::Join(error.to_string()));
804 }
805 }
806 incoming = endpoint.accept() => {
807 let node = node.clone();
808 let cancel = cancellation.child_token();
809 connections.spawn(async move {
810 let accept = async {
811 let Ok(session_request) = incoming.await else {
812 return;
813 };
814 let Ok(connection) = session_request.accept().await else {
815 return;
816 };
817 let _ = node.serve_webtransport(connection).await;
818 };
819 tokio::select! {
820 biased;
821 () = cancel.cancelled() => {}
822 result = tokio::time::timeout(crate::node::WEBTRANSPORT_ACCEPT_TIMEOUT, accept) => {
823 let _ = result;
824 }
825 }
826 });
827 }
828 }
829 }
830 while let Some(result) = connections.join_next().await {
831 result.map_err(|error| HostError::Join(error.to_string()))?;
832 }
833 Ok(())
834 };
835 Ok((bound, task))
836 }
837}
838
839#[derive(Debug, Clone, Copy, PartialEq, Eq)]
840pub struct HealthStatus {
841 pub process_alive: bool,
842 pub websocket_bound: bool,
843 pub websocket_addr: Option<SocketAddr>,
844 pub webtransport_bound: bool,
845 pub webtransport_addr: Option<SocketAddr>,
846 pub listeners_running: bool,
847 pub parent_link_ready: bool,
848 pub child_link_ready: bool,
849}
850
851impl HealthStatus {
852 pub fn ready(&self) -> bool {
853 self.process_alive
854 && self.listeners_running
855 && (self.websocket_bound || self.webtransport_bound)
856 && self.parent_link_ready
857 && self.child_link_ready
858 }
859}
860
861struct ListenerGuard(std::sync::Arc<std::sync::atomic::AtomicUsize>);
862
863impl Drop for ListenerGuard {
864 fn drop(&mut self) {
865 self.0.fetch_sub(1, std::sync::atomic::Ordering::Relaxed);
866 }
867}
868
869pub struct Hosting {
870 websocket: Option<SocketAddr>,
871 webtransport: Option<SocketAddr>,
872 development_cert_hash: Option<[u8; 32]>,
873 cancellation: CancellationToken,
874 _guard: DropGuard,
875 tasks: tokio::task::JoinSet<Result<(), HostError>>,
876 drain_deadline: Option<std::time::Duration>,
877 live_listeners: std::sync::Arc<std::sync::atomic::AtomicUsize>,
878 expected_listeners: usize,
879}
880
881impl Hosting {
882 pub fn websocket_addr(&self) -> Option<SocketAddr> {
883 self.websocket
884 }
885
886 pub fn webtransport_addr(&self) -> Option<SocketAddr> {
887 self.webtransport
888 }
889
890 pub fn development_cert_hash(&self) -> Option<[u8; 32]> {
891 self.development_cert_hash
892 }
893
894 pub fn cancel(&self) {
895 self.cancellation.cancel();
896 }
897
898 pub fn is_finished(&self) -> bool {
899 self.tasks.is_empty()
900 }
901
902 pub fn health(&self) -> HealthStatus {
903 HealthStatus {
904 process_alive: true,
905 websocket_bound: self.websocket.is_some(),
906 websocket_addr: self.websocket,
907 webtransport_bound: self.webtransport.is_some(),
908 webtransport_addr: self.webtransport,
909 listeners_running: self.expected_listeners > 0
910 && self
911 .live_listeners
912 .load(std::sync::atomic::Ordering::Relaxed)
913 == self.expected_listeners,
914 parent_link_ready: true,
915 child_link_ready: true,
916 }
917 }
918
919 fn spawn_listener(
920 &mut self,
921 task: impl std::future::Future<Output = Result<(), HostError>> + Send + 'static,
922 ) {
923 self.expected_listeners += 1;
924 self.live_listeners
925 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
926 let guard = ListenerGuard(self.live_listeners.clone());
927 self.tasks.spawn(async move {
928 let _guard = guard;
929 task.await
930 });
931 }
932
933 pub async fn shutdown(mut self) -> Result<(), HostError> {
934 self.cancellation.cancel();
935 let mut failure = None;
936 match self.drain_deadline {
937 None => Self::join_all(&mut self.tasks, &mut failure).await,
938 Some(deadline) => {
939 if tokio::time::timeout(deadline, Self::join_all(&mut self.tasks, &mut failure))
940 .await
941 .is_err()
942 {
943 self.tasks.abort_all();
944 while let Some(result) = self.tasks.join_next().await {
945 if let Ok(Err(error)) = result {
946 if failure.is_none() {
947 failure = Some(error);
948 }
949 }
950 }
951 }
952 }
953 }
954 failure.map_or(Ok(()), Err)
955 }
956
957 async fn join_all(
958 tasks: &mut tokio::task::JoinSet<Result<(), HostError>>,
959 failure: &mut Option<HostError>,
960 ) {
961 while let Some(result) = tasks.join_next().await {
962 let result = result
963 .map_err(|error| HostError::Join(error.to_string()))
964 .and_then(|result| result);
965 if failure.is_none() {
966 *failure = result.err();
967 }
968 }
969 }
970
971 pub async fn wait(&mut self) -> Result<(), HostError> {
972 let Some(result) = self.tasks.join_next().await else {
973 return Ok(());
974 };
975 result.map_err(|error| HostError::Join(error.to_string()))?
976 }
977}