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 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}