1use hashira::{
2 app::AppService,
3 web::{Body, RemoteAddr, Response},
4};
5use rocket::{
6 data::FromData,
7 fs::FileServer,
8 futures::TryStreamExt,
9 http::Method::*,
10 outcome,
11 request::FromRequest,
12 route::{self, Handler},
13 State,
14};
15use rocket::{Build, Rocket};
16
17#[doc(hidden)]
18pub struct RequestWithoutBody(hashira::web::Request<()>);
19
20#[derive(Clone)]
22pub struct DefaultRequestHandler;
23
24impl DefaultRequestHandler {
25 pub fn routes(rank: Option<isize>) -> Vec<rocket::Route> {
27 let mut routes = vec![];
28 for method in [Get, Put, Post, Delete, Options, Head, Trace, Connect, Patch] {
29 routes.push(rocket::Route::ranked(
30 rank,
31 method,
32 "/<path..>",
33 DefaultRequestHandler,
34 ));
35 }
36
37 routes
38 }
39}
40
41impl From<DefaultRequestHandler> for Vec<rocket::Route> {
42 fn from(_: DefaultRequestHandler) -> Vec<rocket::Route> {
43 DefaultRequestHandler::routes(Some(100))
44 }
45}
46
47#[rocket::async_trait]
48impl Handler for DefaultRequestHandler {
49 async fn handle<'r>(
50 &self,
51 rocket_req: &'r rocket::Request<'_>,
52 data: rocket::data::Data<'r>,
53 ) -> rocket::route::Outcome<'r> {
54 let req = match RequestWithoutBody::from_request(rocket_req).await {
55 outcome::Outcome::Success(req) => req,
56 outcome::Outcome::Failure((status, err)) => {
57 log::error!("{}", err);
58 return route::Outcome::Failure(status);
59 }
60 outcome::Outcome::Forward(_) => return route::Outcome::Forward(data),
61 };
62
63 let service: &State<AppService> = match FromRequest::from_request(rocket_req).await {
64 outcome::Outcome::Success(s) => s,
65 outcome::Outcome::Failure((status, _)) => return route::Outcome::Failure(status),
66 outcome::Outcome::Forward(_) => return route::Outcome::Forward(data),
67 };
68
69 let bytes = match Vec::<u8>::from_data(rocket_req, data).await {
70 outcome::Outcome::Success(b) => b,
71 outcome::Outcome::Failure((status, err)) => {
72 log::error!("{}", err);
73 return route::Outcome::Failure(status);
74 }
75 outcome::Outcome::Forward(data) => return route::Outcome::Forward(data),
76 };
77
78 let req = req.0.map(move |_| Body::from(bytes));
88 let res = service.handle(req).await;
89
90 let rocket_res = map_response(res).await;
91 route::Outcome::Success(rocket_res)
92 }
93}
94
95#[rocket::async_trait]
96impl<'r> FromRequest<'r> for RequestWithoutBody {
97 type Error = hashira::web::Error;
98
99 async fn from_request(
100 rocket_req: &'r rocket::Request<'_>,
101 ) -> rocket::request::Outcome<Self, Self::Error> {
102 fn method_from_rocket(req: &rocket::Request) -> hashira::web::method::Method {
103 match req.method() {
104 Get => hashira::web::method::Method::GET,
105 Put => hashira::web::method::Method::PUT,
106 Post => hashira::web::method::Method::POST,
107 Delete => hashira::web::method::Method::DELETE,
108 Options => hashira::web::method::Method::OPTIONS,
109 Head => hashira::web::method::Method::HEAD,
110 Trace => hashira::web::method::Method::TRACE,
111 Connect => hashira::web::method::Method::CONNECT,
112 Patch => hashira::web::method::Method::PATCH,
113 }
114 }
115
116 let mut builder = hashira::web::Request::builder()
117 .method(method_from_rocket(rocket_req))
118 .uri(rocket_req.uri().to_string());
119
120 for header in rocket_req.headers().iter() {
121 builder = builder.header(header.name.as_str(), header.value.into_owned());
122 }
123
124 let mut req = match builder.body(()) {
125 Ok(x) => x,
126 Err(err) => {
127 log::error!("{}", err);
128 return rocket::request::Outcome::Failure((
129 rocket::http::Status::InternalServerError,
130 err,
131 ));
132 }
133 };
134
135 if let Some(addr) = rocket_req.remote() {
137 let remote_addr = RemoteAddr::from(addr);
138 req.extensions_mut().insert(remote_addr);
139 }
140
141 rocket::request::Outcome::Success(RequestWithoutBody(req))
142 }
143}
144
145pub fn router(app_service: AppService) -> impl FnOnce(Rocket<Build>) -> Rocket<Build> {
147 let static_dir = hashira::env::get_static_dir();
148 let serve_dir = get_current_dir().join("public");
149
150 move |rocket| {
151 rocket
152 .manage(app_service)
153 .mount(&static_dir, FileServer::from(serve_dir))
154 .mount("/", DefaultRequestHandler)
155 }
156}
157
158async fn map_response(res: Response) -> rocket::Response<'static> {
159 let mut builder = rocket::Response::build();
160
161 let status =
163 rocket::http::Status::from_code(res.status().as_u16()).expect("invalid status code");
164 builder.status(status);
165
166 for (name, value) in res.headers() {
168 let v = value.to_str().unwrap().to_string();
169 builder.header_adjoin(rocket::http::Header::new(name.to_string(), v));
170 }
171
172 match res.into_body().into_inner() {
174 hashira::web::Payload::Bytes(bytes) => {
175 let len = bytes.len();
176 let buf = std::io::Cursor::new(bytes);
177 builder.sized_body(len, buf);
178 }
179 hashira::web::Payload::Stream(stream) => {
180 let s = stream.map_err(|e| std::io::Error::new(std::io::ErrorKind::Other, e));
181 let reader = tokio_util::io::StreamReader::new(s);
182 let body = rocket::response::stream::ReaderStream::one(reader);
183 builder.streamed_body(body);
184 }
185 }
186
187 builder.finalize()
188}
189
190fn get_current_dir() -> std::path::PathBuf {
191 let mut current_dir = std::env::current_exe().expect("failed to get current directory");
192 current_dir.pop();
193 current_dir
194}