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 raw_target = parts
423 .uri
424 .path_and_query()
425 .map(http::uri::PathAndQuery::as_str)
426 .unwrap_or_else(|| parts.uri.path());
427 let target_path = match unb_core::TargetPath::parse_application(raw_target) {
428 Ok(target_path) => target_path,
429 Err(error) => {
430 return HostConfig::ingress_error(
431 http::StatusCode::BAD_REQUEST,
432 Some(unb_core::ErrorCode::InvalidInput),
433 &error.to_string(),
434 )
435 }
436 };
437 let target = target_path.target().to_owned();
438 let subject = target_path.subject().to_owned();
439 let canonical_target = target_path.to_string();
440 let wants_sse = parts
441 .headers
442 .get(http::header::ACCEPT)
443 .and_then(|value| value.to_str().ok())
444 .is_some_and(|accept| {
445 accept
446 .split(',')
447 .any(|media| media.trim().split(';').next() == Some("text/event-stream"))
448 });
449 let deadline = Instant::now() + CALL_TIMEOUT;
450 let resolved = node.resolve_unary_until(&target, deadline).await;
451 let (snapshot, resolution) = match resolved {
452 Ok(resolved) => resolved,
453 Err(error) => {
454 return HostConfig::ingress_error(
455 error.code.status(),
456 Some(error.code),
457 &error.message,
458 )
459 }
460 };
461 match resolution {
462 unb_core::Resolution::Unknown => {
463 let error = Node::teach_unknown_target(&snapshot, &target);
464 return HostConfig::ingress_error(
465 error.code.status(),
466 Some(error.code),
467 &error.message,
468 );
469 }
470 unb_core::Resolution::Conflicted { owners } => {
471 let code = unb_core::ErrorCode::PeerUnreachable;
472 return HostConfig::ingress_error(
473 code.status(),
474 Some(code),
475 &format!(
476 "target node {target:?} has multiple live incarnations: {}",
477 owners.join(", ")
478 ),
479 );
480 }
481 unb_core::Resolution::Local => {
482 if !snapshot.services.contains_key(&subject) {
483 let error = Node::teach_unknown_subject(&snapshot, &subject);
484 return HostConfig::ingress_error(
485 error.code.status(),
486 Some(error.code),
487 &error.message,
488 );
489 }
490 }
491 unb_core::Resolution::Route(_) => {}
492 }
493 if parts
494 .headers
495 .get(http::header::CONTENT_LENGTH)
496 .and_then(|value| value.to_str().ok())
497 .and_then(|value| value.parse::<usize>().ok())
498 .is_some_and(|length| length > max_body_bytes)
499 {
500 return HostConfig::ingress_error(
501 http::StatusCode::PAYLOAD_TOO_LARGE,
502 None,
503 "request body exceeds the ingress body ceiling",
504 );
505 }
506 for name in [
507 http::header::HOST,
508 http::header::CONNECTION,
509 http::header::CONTENT_LENGTH,
510 http::header::TRANSFER_ENCODING,
511 http::header::TE,
512 http::header::TRAILER,
513 http::header::UPGRADE,
514 http::header::PROXY_AUTHENTICATE,
515 http::header::PROXY_AUTHORIZATION,
516 http::header::EXPECT,
517 ] {
518 parts.headers.remove(name);
519 }
520 parts.headers.remove("keep-alive");
521 let payload = match axum::body::to_bytes(body, max_body_bytes).await {
522 Ok(payload) => payload,
523 Err(error) => {
524 let mut source: Option<&(dyn std::error::Error + 'static)> = Some(&error);
525 while let Some(current) = source {
526 if current.is::<http_body_util::LengthLimitError>() {
527 return HostConfig::ingress_error(
528 http::StatusCode::PAYLOAD_TOO_LARGE,
529 None,
530 "request body exceeds the ingress body ceiling",
531 );
532 }
533 source = current.source();
534 }
535 return HostConfig::ingress_error(
536 http::StatusCode::BAD_REQUEST,
537 Some(unb_core::ErrorCode::Protocol),
538 &error.to_string(),
539 );
540 }
541 };
542 if wants_sse {
543 let mut headers = serde_json::Map::new();
544 for (name, value) in &parts.headers {
545 if let Ok(value) = value.to_str() {
546 headers.insert(
547 name.as_str().to_string(),
548 serde_json::Value::String(value.to_string()),
549 );
550 }
551 }
552 return match node
553 .subscribe_bytes(&canonical_target, payload, headers)
554 .await
555 {
556 Ok(stream) => HostConfig::ingress_sse(stream).await,
557 Err(error) => {
558 HostConfig::ingress_error(error.code.status(), Some(error.code), &error.message)
559 }
560 };
561 }
562 match node
563 .fetch_until(http::Request::from_parts(parts, payload), deadline)
564 .await
565 {
566 Ok(response) => {
567 let (parts, body) = response.into_parts();
568 match body {
569 crate::layer::ServiceBody::Unary(payload) => {
570 let json = payload.is_empty()
571 || serde_json::from_slice::<serde::de::IgnoredAny>(&payload).is_ok();
572 let mut response =
573 http::Response::from_parts(parts, axum::body::Body::from(payload));
574 response.headers_mut().insert(
575 http::header::CONTENT_TYPE,
576 http::HeaderValue::from_static(if json {
577 "application/json"
578 } else {
579 "application/octet-stream"
580 }),
581 );
582 response
583 }
584 crate::layer::ServiceBody::Stream(_) => HostConfig::ingress_error(
585 http::StatusCode::NOT_ACCEPTABLE,
586 None,
587 "this subject streams; request it with Accept: text/event-stream",
588 ),
589 }
590 }
591 Err(error) => {
592 HostConfig::ingress_error(error.code.status(), Some(error.code), &error.message)
593 }
594 }
595 }
596
597 async fn ingress_sse(mut stream: crate::EventStream) -> http::Response<axum::body::Body> {
598 use futures_util::StreamExt;
599 let first = match stream.next().await {
600 Some(Ok(first)) => Some(first),
601 Some(Err(error)) => {
602 return HostConfig::ingress_error(
603 error.code.status(),
604 Some(error.code),
605 &error.message,
606 )
607 }
608 None => None,
609 };
610 if let Some(first) = &first {
611 if std::str::from_utf8(first).is_err() {
612 return HostConfig::ingress_error(
613 http::StatusCode::NOT_ACCEPTABLE,
614 None,
615 "stream events are not utf-8 and cannot be projected to SSE",
616 );
617 }
618 if first.len() > MAX_SSE_EVENT_BYTES {
619 return HostConfig::ingress_error(
620 http::StatusCode::PAYLOAD_TOO_LARGE,
621 None,
622 "stream event exceeds the SSE event size limit",
623 );
624 }
625 }
626 let body = axum::body::Body::from_stream(futures_util::stream::unfold(
627 (0u64, first, stream),
628 |(id, first, mut stream)| async move {
629 let bytes = match first {
630 Some(bytes) => bytes,
631 None => match tokio::time::timeout(SSE_KEEPALIVE, stream.next()).await {
632 Ok(Some(Ok(bytes))) => {
633 if bytes.len() > MAX_SSE_EVENT_BYTES
634 || std::str::from_utf8(&bytes).is_err()
635 {
636 return None;
637 }
638 bytes
639 }
640 Ok(_) => return None,
641 Err(_) => {
642 return Some((
643 Ok::<_, std::convert::Infallible>(bytes::Bytes::from_static(
644 b": keepalive\n\n",
645 )),
646 (id, None, stream),
647 ));
648 }
649 },
650 };
651 let Ok(text) = std::str::from_utf8(&bytes) else {
652 return None;
653 };
654 let mut record = format!("id: {id}\n");
655 for line in text.split('\n') {
656 record.push_str("data: ");
657 record.push_str(line);
658 record.push('\n');
659 }
660 record.push('\n');
661 Some((
662 Ok::<_, std::convert::Infallible>(bytes::Bytes::from(record)),
663 (id + 1, None, stream),
664 ))
665 },
666 ));
667 http::Response::builder()
668 .status(http::StatusCode::OK)
669 .header(
670 http::header::CONTENT_TYPE,
671 http::HeaderValue::from_static("text/event-stream"),
672 )
673 .header(
674 http::header::CACHE_CONTROL,
675 http::HeaderValue::from_static("no-cache"),
676 )
677 .header("x-accel-buffering", http::HeaderValue::from_static("no"))
678 .body(body)
679 .expect("static SSE response parts are valid")
680 }
681
682 fn ingress_error(
683 status: http::StatusCode,
684 code: Option<unb_core::ErrorCode>,
685 message: &str,
686 ) -> http::Response<axum::body::Body> {
687 let code = code.unwrap_or_else(|| unb_core::ErrorCode::from_status(status));
688 let body = serde_json::json!({ "code": code, "message": message });
689 let mut response = http::Response::builder()
690 .status(status)
691 .header(
692 http::header::CONTENT_TYPE,
693 http::HeaderValue::from_static("application/json"),
694 )
695 .header(
696 unb_core::UNB_CODE,
697 http::HeaderValue::from_static(code.token()),
698 )
699 .body(axum::body::Body::from(body.to_string()))
700 .expect("static response parts are valid");
701 if status == http::StatusCode::METHOD_NOT_ALLOWED {
702 response
703 .headers_mut()
704 .insert(http::header::ALLOW, http::HeaderValue::from_static("POST"));
705 }
706 response
707 }
708
709 fn spawn_websocket(
710 node: Arc<Node>,
711 listener: tokio::net::TcpListener,
712 tcp: TcpTransport,
713 cancellation: CancellationToken,
714 max_body_bytes: usize,
715 ) -> Result<
716 (
717 SocketAddr,
718 impl Future<Output = Result<(), HostError>> + Send + 'static,
719 ),
720 HostError,
721 > {
722 let addr = HostConfig::local_addr(&listener)?;
723 let ingress_node = node.clone();
724 let app = axum::Router::new()
725 .route(
726 &tcp.websocket_path,
727 axum::routing::get(move |upgrade: axum::extract::ws::WebSocketUpgrade| {
728 let node = node.clone();
729 async move { node.serve_ws_upgrade(upgrade) }
730 }),
731 )
732 .merge(tcp.router)
733 .fallback(move |request: axum::extract::Request| {
734 let node = ingress_node.clone();
735 async move { HostConfig::ingress(node, request, max_body_bytes).await }
736 });
737 let task: futures_util::future::Either<_, _> = match tcp.security {
738 TcpSecurity::Plain => futures_util::future::Either::Left(async move {
739 axum::serve(listener, app)
740 .with_graceful_shutdown(async move { cancellation.cancelled().await })
741 .await
742 .map_err(|error| HostError::Io(error.to_string()))
743 }),
744 TcpSecurity::Rustls(config) => {
745 let acceptor = tokio_rustls::TlsAcceptor::from(config);
746 futures_util::future::Either::Right(async move {
747 let mut connections = tokio::task::JoinSet::new();
748 loop {
749 tokio::select! {
750 biased;
751 () = cancellation.cancelled() => break,
752 completed = connections.join_next(), if !connections.is_empty() => {
753 if let Some(Err(error)) = completed {
754 return Err(HostError::Join(error.to_string()));
755 }
756 }
757 accepted = listener.accept() => {
758 let (stream, _peer) = match accepted {
759 Ok(accepted) => accepted,
760 Err(error) => return Err(HostError::Io(error.to_string())),
761 };
762 let acceptor = acceptor.clone();
763 let service =
764 hyper_util::service::TowerToHyperService::new(app.clone());
765 let cancel = cancellation.child_token();
766 connections.spawn(async move {
767 let serve = async move {
768 let Ok(tls) = acceptor.accept(stream).await else {
769 return;
770 };
771 let io = hyper_util::rt::TokioIo::new(tls);
772 let builder = hyper_util::server::conn::auto::Builder::new(
773 hyper_util::rt::TokioExecutor::new(),
774 );
775 let _ = builder
776 .http1_only()
777 .serve_connection_with_upgrades(io, service)
778 .await;
779 };
780 tokio::select! {
781 biased;
782 () = cancel.cancelled() => {}
783 () = serve => {}
784 }
785 });
786 }
787 }
788 }
789 while let Some(result) = connections.join_next().await {
790 result.map_err(|error| HostError::Join(error.to_string()))?;
791 }
792 Ok(())
793 })
794 }
795 };
796 Ok((addr, task))
797 }
798
799 fn spawn_webtransport(
800 node: Arc<Node>,
801 endpoint: WebTransportEndpoint,
802 cancellation: CancellationToken,
803 ) -> Result<
804 (
805 SocketAddr,
806 impl Future<Output = Result<(), HostError>> + Send + 'static,
807 ),
808 HostError,
809 > {
810 let bound = endpoint
811 .local_addr()
812 .map_err(|error| HostError::Io(error.to_string()))?;
813 let task = async move {
814 let mut connections = tokio::task::JoinSet::new();
815 loop {
816 tokio::select! {
817 biased;
818 () = cancellation.cancelled() => break,
819 completed = connections.join_next(), if !connections.is_empty() => {
820 if let Some(Err(error)) = completed {
821 return Err(HostError::Join(error.to_string()));
822 }
823 }
824 incoming = endpoint.accept() => {
825 let node = node.clone();
826 let cancel = cancellation.child_token();
827 connections.spawn(async move {
828 let accept = async {
829 let Ok(session_request) = incoming.await else {
830 return;
831 };
832 let Ok(connection) = session_request.accept().await else {
833 return;
834 };
835 let _ = node.serve_webtransport(connection).await;
836 };
837 tokio::select! {
838 biased;
839 () = cancel.cancelled() => {}
840 result = tokio::time::timeout(crate::node::WEBTRANSPORT_ACCEPT_TIMEOUT, accept) => {
841 let _ = result;
842 }
843 }
844 });
845 }
846 }
847 }
848 while let Some(result) = connections.join_next().await {
849 result.map_err(|error| HostError::Join(error.to_string()))?;
850 }
851 Ok(())
852 };
853 Ok((bound, task))
854 }
855}
856
857#[derive(Debug, Clone, Copy, PartialEq, Eq)]
858pub struct HealthStatus {
859 pub process_alive: bool,
860 pub websocket_bound: bool,
861 pub websocket_addr: Option<SocketAddr>,
862 pub webtransport_bound: bool,
863 pub webtransport_addr: Option<SocketAddr>,
864 pub listeners_running: bool,
865 pub parent_link_ready: bool,
866 pub child_link_ready: bool,
867}
868
869impl HealthStatus {
870 pub fn ready(&self) -> bool {
871 self.process_alive
872 && self.listeners_running
873 && (self.websocket_bound || self.webtransport_bound)
874 && self.parent_link_ready
875 && self.child_link_ready
876 }
877}
878
879struct ListenerGuard(std::sync::Arc<std::sync::atomic::AtomicUsize>);
880
881impl Drop for ListenerGuard {
882 fn drop(&mut self) {
883 self.0.fetch_sub(1, std::sync::atomic::Ordering::Relaxed);
884 }
885}
886
887pub struct Hosting {
888 websocket: Option<SocketAddr>,
889 webtransport: Option<SocketAddr>,
890 development_cert_hash: Option<[u8; 32]>,
891 cancellation: CancellationToken,
892 _guard: DropGuard,
893 tasks: tokio::task::JoinSet<Result<(), HostError>>,
894 drain_deadline: Option<std::time::Duration>,
895 live_listeners: std::sync::Arc<std::sync::atomic::AtomicUsize>,
896 expected_listeners: usize,
897}
898
899impl Hosting {
900 pub fn websocket_addr(&self) -> Option<SocketAddr> {
901 self.websocket
902 }
903
904 pub fn webtransport_addr(&self) -> Option<SocketAddr> {
905 self.webtransport
906 }
907
908 pub fn development_cert_hash(&self) -> Option<[u8; 32]> {
909 self.development_cert_hash
910 }
911
912 pub fn cancel(&self) {
913 self.cancellation.cancel();
914 }
915
916 pub fn is_finished(&self) -> bool {
917 self.tasks.is_empty()
918 }
919
920 pub fn health(&self) -> HealthStatus {
921 HealthStatus {
922 process_alive: true,
923 websocket_bound: self.websocket.is_some(),
924 websocket_addr: self.websocket,
925 webtransport_bound: self.webtransport.is_some(),
926 webtransport_addr: self.webtransport,
927 listeners_running: self.expected_listeners > 0
928 && self
929 .live_listeners
930 .load(std::sync::atomic::Ordering::Relaxed)
931 == self.expected_listeners,
932 parent_link_ready: true,
933 child_link_ready: true,
934 }
935 }
936
937 fn spawn_listener(
938 &mut self,
939 task: impl std::future::Future<Output = Result<(), HostError>> + Send + 'static,
940 ) {
941 self.expected_listeners += 1;
942 self.live_listeners
943 .fetch_add(1, std::sync::atomic::Ordering::Relaxed);
944 let guard = ListenerGuard(self.live_listeners.clone());
945 self.tasks.spawn(async move {
946 let _guard = guard;
947 task.await
948 });
949 }
950
951 pub async fn shutdown(mut self) -> Result<(), HostError> {
952 self.cancellation.cancel();
953 let mut failure = None;
954 match self.drain_deadline {
955 None => Self::join_all(&mut self.tasks, &mut failure).await,
956 Some(deadline) => {
957 if tokio::time::timeout(deadline, Self::join_all(&mut self.tasks, &mut failure))
958 .await
959 .is_err()
960 {
961 self.tasks.abort_all();
962 while let Some(result) = self.tasks.join_next().await {
963 if let Ok(Err(error)) = result {
964 if failure.is_none() {
965 failure = Some(error);
966 }
967 }
968 }
969 }
970 }
971 }
972 failure.map_or(Ok(()), Err)
973 }
974
975 async fn join_all(
976 tasks: &mut tokio::task::JoinSet<Result<(), HostError>>,
977 failure: &mut Option<HostError>,
978 ) {
979 while let Some(result) = tasks.join_next().await {
980 let result = result
981 .map_err(|error| HostError::Join(error.to_string()))
982 .and_then(|result| result);
983 if failure.is_none() {
984 *failure = result.err();
985 }
986 }
987 }
988
989 pub async fn wait(&mut self) -> Result<(), HostError> {
990 let Some(result) = self.tasks.join_next().await else {
991 return Ok(());
992 };
993 result.map_err(|error| HostError::Join(error.to_string()))?
994 }
995}