Skip to main content

static_web_server/
fallback_page.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//! Fallback page module useful for a custom page default.
7//!
8
9use hyper::{Request, Response, StatusCode};
10use std::path::Path;
11
12use crate::body::Body;
13use crate::error_page::build_html_response;
14use crate::{Error, exts::http::MethodExt, handler::RequestHandlerOpts, helpers};
15
16/// Initializes fallback page processing
17pub(crate) fn init(file_path: &Path, handler_opts: &mut RequestHandlerOpts) {
18    let found = file_path.is_file();
19    if found {
20        handler_opts.page_fallback = helpers::read_text_default(file_path).into_bytes();
21    } else {
22        tracing::debug!("fallback page path not found or not a regular file");
23    }
24
25    tracing::info!(
26        "fallback page: enabled={}, value=\"{}\"",
27        found,
28        file_path.display()
29    );
30}
31
32/// Replace 404 Not Found by the configured fallback page
33pub(crate) fn post_process<T>(
34    opts: &RequestHandlerOpts,
35    req: &Request<T>,
36    resp: Response<Body>,
37) -> Result<Response<Body>, Error> {
38    Ok(
39        if req.method().is_get()
40            && resp.status() == StatusCode::NOT_FOUND
41            && !opts.page_fallback.is_empty()
42        {
43            fallback_response(&opts.page_fallback)
44        } else {
45            resp
46        },
47    )
48}
49
50/// Checks if a fallback response can be generated, i.e. if it is a `GET` request
51/// that would result in a `404` error and a fallback page is configured.
52/// If a response can be generated then is returned otherwise `None`.
53pub fn fallback_response(page_fallback: &[u8]) -> Response<Body> {
54    build_html_response(page_fallback.to_owned(), StatusCode::OK, None)
55}
56
57#[cfg(test)]
58mod tests {
59    use super::post_process;
60    use crate::body::Body;
61    use crate::{Error, error_page, handler::RequestHandlerOpts};
62    use hyper::{Method, Request, Response, StatusCode, Uri};
63    use std::path::PathBuf;
64
65    fn make_request(method: &str) -> Request<Body> {
66        Request::builder()
67            .method(method)
68            .uri("/")
69            .body(crate::body::empty())
70            .unwrap()
71    }
72
73    fn make_response(status: &StatusCode) -> Response<Body> {
74        error_page::error_response(
75            &Uri::try_from("/").unwrap(),
76            &Method::GET,
77            status,
78            &PathBuf::new(),
79            &PathBuf::new(),
80        )
81        .unwrap()
82    }
83
84    #[test]
85    fn test_success_code() -> Result<(), Error> {
86        let opts = RequestHandlerOpts {
87            page_fallback: vec![1, 2, 3],
88            ..Default::default()
89        };
90        let req = make_request("GET");
91        let resp = make_response(&StatusCode::OK);
92
93        let resp = post_process(&opts, &req, resp)?;
94        assert_eq!(resp.status(), StatusCode::OK);
95        assert_ne!(
96            resp.headers()
97                .get("Content-Length")
98                .map(|v| v.to_str().unwrap())
99                .unwrap_or("3"),
100            "3"
101        );
102
103        Ok(())
104    }
105
106    #[test]
107    fn test_wrong_error() -> Result<(), Error> {
108        let opts = RequestHandlerOpts {
109            page_fallback: vec![1, 2, 3],
110            ..Default::default()
111        };
112        let req = make_request("GET");
113        let resp = make_response(&StatusCode::INTERNAL_SERVER_ERROR);
114
115        let resp = post_process(&opts, &req, resp)?;
116        assert_eq!(resp.status(), StatusCode::INTERNAL_SERVER_ERROR);
117        assert_ne!(
118            resp.headers()
119                .get("Content-Length")
120                .map(|v| v.to_str().unwrap())
121                .unwrap_or("3"),
122            "3"
123        );
124
125        Ok(())
126    }
127
128    #[test]
129    fn test_wrong_method() -> Result<(), Error> {
130        let opts = RequestHandlerOpts {
131            page_fallback: vec![1, 2, 3],
132            ..Default::default()
133        };
134        let req = make_request("POST");
135        let resp = make_response(&StatusCode::NOT_FOUND);
136
137        let resp = post_process(&opts, &req, resp)?;
138        assert_eq!(resp.status(), StatusCode::NOT_FOUND);
139        assert_ne!(
140            resp.headers()
141                .get("Content-Length")
142                .map(|v| v.to_str().unwrap())
143                .unwrap_or("3"),
144            "3"
145        );
146
147        Ok(())
148    }
149
150    #[test]
151    fn test_unconfigured() -> Result<(), Error> {
152        let opts = RequestHandlerOpts {
153            page_fallback: Vec::new(),
154            ..Default::default()
155        };
156        let req = make_request("GET");
157        let resp = make_response(&StatusCode::NOT_FOUND);
158
159        let resp = post_process(&opts, &req, resp)?;
160        assert_eq!(resp.status(), StatusCode::NOT_FOUND);
161
162        Ok(())
163    }
164
165    #[test]
166    fn test_fallback() -> Result<(), Error> {
167        let opts = RequestHandlerOpts {
168            page_fallback: vec![1, 2, 3],
169            ..Default::default()
170        };
171        let req = make_request("GET");
172        let resp = make_response(&StatusCode::NOT_FOUND);
173
174        let resp = post_process(&opts, &req, resp)?;
175        assert_eq!(resp.status(), StatusCode::OK);
176        assert_eq!(
177            resp.headers()
178                .get("Content-Length")
179                .map(|v| v.to_str().unwrap())
180                .unwrap_or("3"),
181            "3"
182        );
183
184        Ok(())
185    }
186}