Skip to main content

unb_server/
host.rs

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}