static_web_server/
fallback_page.rs1use 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
16pub(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
32pub(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
50pub 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}