Skip to main content

static_web_server/
server.rs

1// SPDX-License-Identifier: MIT OR Apache-2.0
2// This file is part of Static Web Server.
3// See https://static-web-server.net/ for more information
4// Copyright (C) 2019-present Jose Quintana <joseluisq.net>
5
6//! Server module intended to construct a multi-threaded HTTP or HTTP/2 web server.
7//!
8
9use hyper::server::Server as HyperServer;
10use listenfd::ListenFd;
11use std::net::{IpAddr, SocketAddr, TcpListener};
12use std::sync::Arc;
13use tokio::sync::{Mutex, watch::Receiver};
14
15use crate::handler::{RequestHandler, RequestHandlerOpts};
16
17#[cfg(feature = "metrics")]
18use crate::metrics;
19#[cfg(any(unix, windows))]
20use crate::signals;
21
22#[cfg(feature = "http2")]
23use {
24    crate::tls::{TlsAcceptor, TlsConfigBuilder},
25    crate::{error, error_page, https_redirect},
26    hyper::server::conn::{AddrIncoming, AddrStream},
27    hyper::service::{make_service_fn, service_fn},
28};
29
30#[cfg(feature = "directory-listing")]
31use crate::directory_listing;
32
33#[cfg(feature = "directory-listing-download")]
34use crate::directory_listing_download;
35
36#[cfg(feature = "fallback-page")]
37use crate::fallback_page;
38
39#[cfg(any(
40    feature = "compression",
41    feature = "compression-deflate",
42    feature = "compression-gzip",
43    feature = "compression-brotli",
44    feature = "compression-zstd",
45))]
46use crate::compression;
47
48use crate::compression_static;
49
50#[cfg(feature = "basic-auth")]
51use crate::basic_auth;
52
53#[cfg(feature = "experimental")]
54use crate::mem_cache;
55
56use crate::{Context, Result, service::RouterService};
57use crate::{
58    Settings, control_headers, cors, health, helpers, log_addr, maintenance_mode, security_headers,
59};
60
61/// Define a multi-threaded HTTP or HTTP/2 web server.
62pub struct Server {
63    opts: Settings,
64    worker_threads: usize,
65    max_blocking_threads: usize,
66}
67
68impl Server {
69    /// Create a new multi-threaded server instance.
70    pub fn new(opts: Settings) -> Result<Server> {
71        // Configure number of worker threads
72        let cpus = std::thread::available_parallelism()
73            .with_context(|| {
74                "unable to get current platform cpus or lack of permissions to query available parallelism"
75            })?
76            .get();
77        let worker_threads = match opts.general.threads_multiplier {
78            0 | 1 => cpus,
79            n => cpus * n,
80        };
81        let max_blocking_threads = opts.general.max_blocking_threads;
82
83        Ok(Server {
84            opts,
85            worker_threads,
86            max_blocking_threads,
87        })
88    }
89
90    /// Run the multi-threaded `Server` as standalone.
91    /// This is a top-level function of [run_server_on_rt](#method.run_server_on_rt).
92    ///
93    /// It accepts an optional [`cancel`] parameter to shut down the server
94    /// gracefully on demand as a complement to the termination signals handling.
95    ///
96    /// [`cancel`]: <https://docs.rs/tokio/latest/tokio/sync/watch/struct.Receiver.html>
97    pub fn run_standalone(self, cancel: Option<Receiver<()>>) -> Result {
98        self.run_server_on_rt(cancel, || {}, true)
99    }
100
101    /// Run the multi-threaded `Server` which will be used by a Windows service.
102    /// This is a top-level function of [run_server_on_rt](#method.run_server_on_rt).
103    ///
104    /// It accepts an optional [`cancel`] parameter to shut down the server
105    /// gracefully on demand and an optional `cancel_fn` that will be executed
106    /// right after the server is shut down.
107    ///
108    /// [`cancel`]: <https://docs.rs/tokio/latest/tokio/sync/watch/struct.Receiver.html>
109    #[cfg(windows)]
110    pub fn run_as_service<F>(self, cancel: Option<Receiver<()>>, cancel_fn: F) -> Result
111    where
112        F: FnOnce(),
113    {
114        self.run_server_on_rt(cancel, cancel_fn, true)
115    }
116
117    /// Build and run the multi-threaded `Server` on the Tokio runtime.
118    ///
119    /// Setting `exit_on_error` to `true` will exit the entire process if
120    /// the server fails to start (previous behaviour).
121    pub fn run_server_on_rt<F>(
122        self,
123        cancel_recv: Option<Receiver<()>>,
124        cancel_fn: F,
125        exit_on_error: bool,
126    ) -> Result
127    where
128        F: FnOnce(),
129    {
130        tracing::debug!(%self.worker_threads, "initializing tokio runtime with multi-threaded scheduler");
131
132        let rt = tokio::runtime::Builder::new_multi_thread()
133            .worker_threads(self.worker_threads)
134            .max_blocking_threads(self.max_blocking_threads)
135            .thread_name("static-web-server")
136            .enable_all()
137            .build()?;
138
139        let res = rt.block_on(async {
140            tracing::trace!("tokio runtime initialized");
141            self.start_server(cancel_recv, cancel_fn).await
142        });
143
144        if let Err(err) = &res {
145            tracing::error!("server failed to start up: {:?}", err);
146            if exit_on_error {
147                std::process::exit(1)
148            }
149        }
150        res
151    }
152
153    /// Run the inner Hyper `HyperServer` (HTTP1/HTTP2) forever on the current thread
154    // using the given configuration.
155    async fn start_server<F>(self, _cancel_recv: Option<Receiver<()>>, _cancel_fn: F) -> Result
156    where
157        F: FnOnce(),
158    {
159        tracing::trace!("starting web server");
160        tracing::info!("{} {}", env!("CARGO_PKG_NAME"), env!("CARGO_PKG_VERSION"));
161
162        // Config "general" options
163        let general = self.opts.general;
164        // Config-file "advanced" options
165        let advanced_opts = self.opts.advanced;
166
167        tracing::info!("log level: {}", general.log_level);
168
169        // Config file option
170        let config_file = general.config_file;
171        if config_file.is_file() {
172            tracing::info!("config file used: {}", config_file.display());
173        } else {
174            tracing::debug!(
175                "config file path not found or not a regular file: {}",
176                config_file.display()
177            );
178        }
179
180        // Determine TCP listener either file descriptor or TCP socket
181        let (tcp_listener, addr_str);
182        match general.fd {
183            Some(fd) => {
184                addr_str = format!("@FD({fd})");
185                tcp_listener = ListenFd::from_env()
186                    .take_tcp_listener(fd)?
187                    .with_context(|| "failed to convert inherited 'fd' into a 'tcp' listener")?;
188                tracing::info!(
189                    "converted inherited file descriptor {} to a 'tcp' listener",
190                    fd
191                );
192            }
193            None => {
194                let ip = general
195                    .host
196                    .parse::<IpAddr>()
197                    .with_context(|| format!("failed to parse {} address", general.host))?;
198                let addr = SocketAddr::from((ip, general.port));
199                tcp_listener = TcpListener::bind(addr)
200                    .with_context(|| format!("failed to bind to {addr} address"))?;
201                addr_str = addr.to_string();
202                tracing::info!("server bound to tcp socket {}", addr_str);
203            }
204        }
205
206        // Number of worker threads option
207        let threads = self.worker_threads;
208        tracing::info!("runtime worker threads: {}", threads);
209
210        // Maximum number of blocking threads
211        tracing::info!(
212            "runtime max blocking threads: {}",
213            general.max_blocking_threads
214        );
215
216        // Check for a valid root directory
217        let root_dir = helpers::get_valid_dirpath(&general.root)
218            .with_context(|| "root directory was not found or inaccessible")?;
219
220        // Canonicalize the root directory once at startup
221        // so further checks can compare against a precomputed canonical base
222        // without paying a `canonicalize` syscall on every request.
223        // Falls back to the validated path if canonicalization fails.
224        // NOTE: When `use_relative_root` is enabled, canonicalization is skipped
225        // so that symlinked root directories are resolved at request time.
226        let root_dir = if general.use_relative_root {
227            root_dir
228        } else {
229            root_dir.canonicalize().unwrap_or(root_dir)
230        };
231
232        // Custom HTML error page files
233        // NOTE: in the case of relative paths, they're joined to the root directory
234        let mut page404 = general.page404;
235        if page404.is_relative() && !page404.starts_with(&root_dir) {
236            page404 = root_dir.join(page404);
237        }
238        if !page404.is_file() {
239            tracing::debug!(
240                "404 file path not found or not a regular file: {}",
241                page404.display()
242            );
243        }
244        let mut page50x = general.page50x;
245        if page50x.is_relative() && !page50x.starts_with(&root_dir) {
246            page50x = root_dir.join(page50x);
247        }
248        if !page50x.is_file() {
249            tracing::debug!(
250                "50x file path not found or not a regular file: {}",
251                page50x.display()
252            );
253        }
254
255        // Log remote address option
256        let log_remote_address = general.log_remote_address;
257
258        // Log the X-Real-IP header.
259        let log_x_real_ip = general.log_x_real_ip;
260
261        // Log the X-Forwarded-For header.
262        let log_forwarded_for = general.log_forwarded_for;
263
264        // Trusted IPs for remote addresses.
265        let trusted_proxies = general.trusted_proxies;
266
267        // Log redirect trailing slash option
268        let redirect_trailing_slash = general.redirect_trailing_slash;
269        tracing::info!(
270            "redirect trailing slash: enabled={}",
271            redirect_trailing_slash
272        );
273
274        // Ignore hidden files option
275        let ignore_hidden_files = general.ignore_hidden_files;
276        tracing::info!("ignore hidden files: enabled={}", ignore_hidden_files);
277
278        // Disable symlinks option
279        let disable_symlinks = general.disable_symlinks;
280        tracing::info!("disable symlinks: enabled={}", disable_symlinks);
281
282        // Use relative root option
283        let use_relative_root = general.use_relative_root;
284        tracing::info!("use relative root: enabled={}", use_relative_root);
285
286        // Default charset for text/* responses
287        let default_text_charset = general.text_charset;
288        tracing::info!("text charset: enabled={default_text_charset}");
289
290        // Grace period option
291        let grace_period = general.grace_period;
292        tracing::info!("grace period before graceful shutdown: {}s", grace_period);
293
294        // Index files option
295        let index_files = general
296            .index_files
297            .split(',')
298            .map(|s| s.trim().to_owned())
299            .collect::<Vec<_>>();
300        if index_files.is_empty() {
301            bail!("index files list is empty, provide at least one index file")
302        }
303        tracing::info!("index files: {}", general.index_files);
304
305        // Request handler options, some settings will be filled in by modules
306        let mut handler_opts = RequestHandlerOpts {
307            root_dir,
308            page404: page404.clone(),
309            page50x: page50x.clone(),
310            log_remote_address,
311            log_x_real_ip,
312            log_forwarded_for,
313            trusted_proxies,
314            redirect_trailing_slash,
315            ignore_hidden_files,
316            disable_symlinks,
317            use_relative_root,
318            accept_markdown: general.accept_markdown,
319            text_charset: general.text_charset,
320            index_files,
321            advanced_opts,
322            ..Default::default()
323        };
324
325        // Directory listing options
326        #[cfg(feature = "directory-listing")]
327        directory_listing::init(
328            general.directory_listing,
329            general.directory_listing_order,
330            general.directory_listing_format,
331            &mut handler_opts,
332        );
333
334        // Directory listing download options
335        #[cfg(feature = "directory-listing-download")]
336        directory_listing_download::init(&general.directory_listing_download, &mut handler_opts);
337
338        // Fallback page option
339        #[cfg(feature = "fallback-page")]
340        fallback_page::init(&general.page_fallback, &mut handler_opts);
341
342        // Health endpoint option
343        health::init(general.health, &mut handler_opts);
344
345        // Log remote address option
346        log_addr::init(general.log_remote_address, &mut handler_opts);
347
348        // Metrics endpoint option
349        #[cfg(feature = "metrics")]
350        metrics::init(general.metrics, &mut handler_opts);
351
352        // CORS option
353        cors::init(
354            &general.cors_allow_origins,
355            &general.cors_allow_headers,
356            &general.cors_expose_headers,
357            &mut handler_opts,
358        );
359
360        // `Basic` HTTP Authentication Schema option
361        #[cfg(feature = "basic-auth")]
362        basic_auth::init(&general.basic_auth, &mut handler_opts)?;
363
364        // Maintenance mode option
365        maintenance_mode::init(
366            general.maintenance_mode,
367            general.maintenance_mode_status,
368            general.maintenance_mode_file,
369            &mut handler_opts,
370        );
371
372        // Check pre-compressed files based on the `Accept-Encoding` header
373        compression_static::init(general.compression_static, &mut handler_opts);
374
375        // Auto compression based on the `Accept-Encoding` header
376        #[cfg(any(
377            feature = "compression",
378            feature = "compression-deflate",
379            feature = "compression-gzip",
380            feature = "compression-brotli",
381            feature = "compression-zstd",
382        ))]
383        compression::init(
384            general.compression,
385            general.compression_level,
386            &mut handler_opts,
387        );
388
389        // Cache control headers option
390        control_headers::init(general.cache_control_headers, &mut handler_opts);
391
392        // Security Headers option
393        security_headers::init(general.security_headers, &mut handler_opts);
394
395        // In-Memory cache option
396        #[cfg(feature = "experimental")]
397        mem_cache::cache::init(&mut handler_opts)?;
398
399        // Create a service router for Hyper
400        let router_service = RouterService::new(RequestHandler {
401            opts: Arc::from(handler_opts),
402        });
403
404        #[cfg(windows)]
405        let (sender, receiver) = tokio::sync::watch::channel(());
406
407        // Windows ctrl+c listening
408        #[cfg(windows)]
409        let ctrlc_task = tokio::spawn(async move {
410            if !general.windows_service {
411                tracing::info!("installing graceful shutdown ctrl+c signal handler");
412                if let Err(err) = tokio::signal::ctrl_c().await {
413                    tracing::error!("failed to install ctrl+c signal handler: {err:?}");
414                    return;
415                }
416                tracing::info!("installing graceful shutdown ctrl+c signal handler");
417                let _ = sender.send(());
418            }
419        });
420
421        // Run the corresponding HTTP Server asynchronously with its given options
422        #[cfg(feature = "http2")]
423        if general.http2 {
424            // HTTP to HTTPS redirect option
425            let https_redirect = general.https_redirect;
426            tracing::info!("http to https redirect: enabled={}", https_redirect);
427            tracing::info!(
428                "http to https redirect host: {}",
429                general.https_redirect_host
430            );
431            tracing::info!(
432                "http to https redirect from port: {}",
433                general.https_redirect_from_port
434            );
435            tracing::info!(
436                "http to https redirect from hosts: {}",
437                general.https_redirect_from_hosts
438            );
439
440            // HTTP/2 + TLS
441            tcp_listener
442                .set_nonblocking(true)
443                .with_context(|| "failed to set TCP non-blocking mode")?;
444            let listener = tokio::net::TcpListener::from_std(tcp_listener)
445                .with_context(|| "failed to create tokio::net::TcpListener")?;
446            let mut incoming = AddrIncoming::from_listener(listener).with_context(
447                || "failed to create an AddrIncoming from the current tokio::net::TcpListener",
448            )?;
449            incoming.set_nodelay(true);
450
451            let http2_tls_cert = match general.http2_tls_cert {
452                Some(v) => v,
453                _ => bail!("failed to initialize TLS because cert file missing"),
454            };
455            let http2_tls_key = match general.http2_tls_key {
456                Some(v) => v,
457                _ => bail!("failed to initialize TLS because key file missing"),
458            };
459
460            let tls = TlsConfigBuilder::new()
461                .cert_path(&http2_tls_cert)
462                .key_path(&http2_tls_key)
463                .build()
464                .with_context(
465                    || "failed to initialize TLS probably because invalid cert or key file",
466                )?;
467
468            #[cfg(unix)]
469            let signals = signals::create_signals()
470                .with_context(|| "failed to register termination signals")?;
471            #[cfg(unix)]
472            let handle = signals.handle();
473
474            let http2_server =
475                HyperServer::builder(TlsAcceptor::new(tls, incoming)).serve(router_service);
476
477            #[cfg(unix)]
478            let http2_cancel_recv = Arc::new(Mutex::new(_cancel_recv));
479            #[cfg(unix)]
480            let redirect_cancel_recv = http2_cancel_recv.clone();
481
482            #[cfg(unix)]
483            let http2_server = http2_server.with_graceful_shutdown(signals::wait_for_signals(
484                signals,
485                grace_period,
486                http2_cancel_recv,
487            ));
488
489            #[cfg(windows)]
490            let http2_cancel_recv = Arc::new(Mutex::new(_cancel_recv));
491            #[cfg(windows)]
492            let redirect_cancel_recv = http2_cancel_recv.clone();
493
494            #[cfg(windows)]
495            let http2_ctrlc_recv = Arc::new(Mutex::new(Some(receiver)));
496            #[cfg(windows)]
497            let redirect_ctrlc_recv = http2_ctrlc_recv.clone();
498
499            #[cfg(windows)]
500            let http2_server = http2_server.with_graceful_shutdown(async move {
501                if general.windows_service {
502                    signals::wait_for_ctrl_c(http2_cancel_recv, grace_period).await;
503                } else {
504                    signals::wait_for_ctrl_c(http2_ctrlc_recv, grace_period).await;
505                }
506            });
507
508            tracing::info!(
509                parent: tracing::info_span!("Server::start_server", ?addr_str, ?threads),
510                "http2 server is listening on https://{}",
511                addr_str
512            );
513
514            // HTTP to HTTPS redirect server
515            if general.https_redirect {
516                let ip = general
517                    .host
518                    .parse::<IpAddr>()
519                    .with_context(|| format!("failed to parse {} address", general.host))?;
520                let addr = SocketAddr::from((ip, general.https_redirect_from_port));
521                let tcp_listener = TcpListener::bind(addr)
522                    .with_context(|| format!("failed to bind to {addr} address"))?;
523                tracing::info!(
524                    parent: tracing::info_span!("Server::start_server", ?addr, ?threads),
525                    "http1 redirect server is listening on http://{}",
526                    addr
527                );
528                tcp_listener
529                    .set_nonblocking(true)
530                    .with_context(|| "failed to set TCP non-blocking mode")?;
531
532                #[cfg(unix)]
533                let redirect_signals = signals::create_signals()
534                    .with_context(|| "failed to register termination signals")?;
535                #[cfg(unix)]
536                let redirect_handle = redirect_signals.handle();
537
538                // Allowed redirect hosts
539                let redirect_allowed_hosts = general
540                    .https_redirect_from_hosts
541                    .split(',')
542                    .map(|s| s.trim().to_owned())
543                    .collect::<Vec<_>>();
544                if redirect_allowed_hosts.is_empty() {
545                    bail!("https redirect allowed hosts is empty, provide at least one host or IP")
546                }
547
548                // Redirect options
549                let redirect_opts = Arc::new(https_redirect::RedirectOpts {
550                    https_hostname: general.https_redirect_host,
551                    https_port: general.port,
552                    allowed_hosts: redirect_allowed_hosts,
553                });
554
555                let server_redirect = HyperServer::from_tcp(tcp_listener)
556                    .unwrap()
557                    .tcp_nodelay(true)
558                    .serve(make_service_fn(move |_: &AddrStream| {
559                        let redirect_opts = redirect_opts.clone();
560                        let page404 = page404.clone();
561                        let page50x = page50x.clone();
562                        async move {
563                            Ok::<_, error::Error>(service_fn(move |req| {
564                                let redirect_opts = redirect_opts.clone();
565                                let page404 = page404.clone();
566                                let page50x = page50x.clone();
567                                async move {
568                                    let uri = req.uri();
569                                    let method = req.method();
570                                    match https_redirect::redirect_to_https(&req, redirect_opts) {
571                                        Ok(resp) => Ok(resp),
572                                        Err(status) => error_page::error_response(
573                                            uri, method, &status, &page404, &page50x,
574                                        ),
575                                    }
576                                }
577                            }))
578                        }
579                    }));
580
581                #[cfg(unix)]
582                let server_redirect = server_redirect.with_graceful_shutdown(
583                    signals::wait_for_signals(redirect_signals, grace_period, redirect_cancel_recv),
584                );
585                #[cfg(windows)]
586                let server_redirect = server_redirect.with_graceful_shutdown(async move {
587                    if general.windows_service {
588                        signals::wait_for_ctrl_c(redirect_cancel_recv, grace_period).await;
589                    } else {
590                        signals::wait_for_ctrl_c(redirect_ctrlc_recv, grace_period).await;
591                    }
592                });
593
594                // HTTP/2 server task
595                let server_task = tokio::spawn(async move {
596                    if let Err(err) = http2_server.await {
597                        tracing::error!("http2 server failed to start up: {:?}", err);
598                        std::process::exit(1)
599                    }
600                });
601
602                // HTTP/1 redirect server task
603                let redirect_server_task = tokio::spawn(async move {
604                    if let Err(err) = server_redirect.await {
605                        tracing::error!("http1 redirect server failed to start up: {:?}", err);
606                        std::process::exit(1)
607                    }
608                });
609
610                tracing::info!("press ctrl+c to shut down the servers");
611
612                #[cfg(windows)]
613                tokio::try_join!(ctrlc_task, server_task, redirect_server_task)?;
614                #[cfg(unix)]
615                tokio::try_join!(server_task, redirect_server_task)?;
616
617                #[cfg(unix)]
618                redirect_handle.close();
619            } else {
620                tracing::info!("press ctrl+c to shut down the server");
621                http2_server.await?;
622            }
623
624            #[cfg(unix)]
625            handle.close();
626
627            #[cfg(windows)]
628            _cancel_fn();
629
630            tracing::warn!("termination signal caught, shutting down the server execution");
631            return Ok(());
632        }
633
634        // HTTP/1
635
636        #[cfg(unix)]
637        let signals =
638            signals::create_signals().with_context(|| "failed to register termination signals")?;
639        #[cfg(unix)]
640        let handle = signals.handle();
641
642        tcp_listener
643            .set_nonblocking(true)
644            .with_context(|| "failed to set TCP non-blocking mode")?;
645
646        let http1_server = HyperServer::from_tcp(tcp_listener)
647            .unwrap()
648            .tcp_nodelay(true)
649            .serve(router_service);
650
651        #[cfg(unix)]
652        let http1_cancel_recv = Arc::new(Mutex::new(_cancel_recv));
653
654        #[cfg(unix)]
655        let http1_server = http1_server.with_graceful_shutdown(signals::wait_for_signals(
656            signals,
657            grace_period,
658            http1_cancel_recv,
659        ));
660
661        #[cfg(windows)]
662        let http1_server = http1_server.with_graceful_shutdown(async move {
663            let http1_cancel_recv = if general.windows_service {
664                // http1_cancel_recv
665                Arc::new(Mutex::new(_cancel_recv))
666            } else {
667                // http1_ctrlc_recv
668                Arc::new(Mutex::new(Some(receiver)))
669            };
670            signals::wait_for_ctrl_c(http1_cancel_recv, grace_period).await;
671        });
672
673        tracing::info!(
674            parent: tracing::info_span!("Server::start_server", ?addr_str, ?threads),
675            "http1 server is listening on http://{}",
676            addr_str
677        );
678
679        tracing::info!("press ctrl+c to shut down the server");
680
681        #[cfg(unix)]
682        http1_server.await?;
683
684        #[cfg(windows)]
685        let http1_server_task = tokio::spawn(async move {
686            if let Err(err) = http1_server.await {
687                tracing::error!("http1 server failed to start up: {:?}", err);
688                std::process::exit(1)
689            }
690        });
691        #[cfg(windows)]
692        tokio::try_join!(ctrlc_task, http1_server_task)?;
693
694        #[cfg(windows)]
695        _cancel_fn();
696
697        #[cfg(unix)]
698        handle.close();
699
700        tracing::warn!("termination signal caught, shutting down the server execution");
701        Ok(())
702    }
703}