Skip to main content

hashira_rocket/
core.rs

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/// A responder for handling a rocket request.
21#[derive(Clone)]
22pub struct DefaultRequestHandler;
23
24impl DefaultRequestHandler {
25    /// Returns all the routes this handler will handle.
26    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        // We read the entire stream
79        // FIXME: Not sure if rocket limits apply on this
80        // let data_stream = data.open(ByteUnit::max_value());
81        // let reader = ReaderStream::new(data_stream)
82        //     .map_err(Into::into)
83        //     .map_ok(Bytes::from);
84
85        // let bytes = Box::pin(reader) as TryBoxStream<Bytes>;
86
87        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        // Add additional extensions
136        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
145// Returns a function to attach the hashira router to `Rocket`.
146pub 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    // Set the status code
162    let status =
163        rocket::http::Status::from_code(res.status().as_u16()).expect("invalid status code");
164    builder.status(status);
165
166    // Set the headers
167    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    // Set the body
173    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}