use crate::{RocketAnswer, RocketBody, RocketRequest};
use alux_http::{
HttpMethod, HttpSelectorAlg, PathSegment, PathSyntaxAlg, RouteAlg, RoutePath, SelectorAlg, compose_path,
describe_path,
};
use bytes::Bytes;
use core::future::Future;
use core::pin::Pin;
use futures::TryStreamExt;
use rocket::data::{ByteUnit, Data};
use rocket::http::{HeaderMap, Method, Status};
use rocket::response::Response;
use rocket::route::{Handler, Outcome, Route};
use rocket::{Build, Request, Rocket, async_trait};
use std::io::Cursor;
use std::sync::Arc;
use tokio_util::io::StreamReader;
const READS: ByteUnit = ByteUnit::Mebibyte(8);
const TOO_LARGE: Status = Status { code: 413 };
const UNREADABLE: Status = Status { code: 400 };
struct RocketPath;
impl PathSyntaxAlg for RocketPath {
fn param(&self, name: &str) -> String {
format!("<{name}>")
}
fn tail(&self, name: &str) -> String {
format!("<{name}..>")
}
}
fn rocket_method(method: HttpMethod) -> Method {
match method {
HttpMethod::Get => Method::Get,
HttpMethod::Post => Method::Post,
HttpMethod::Put => Method::Put,
HttpMethod::Patch => Method::Patch,
HttpMethod::Delete => Method::Delete,
HttpMethod::Head => Method::Head,
HttpMethod::Options => Method::Options,
HttpMethod::Trace => Method::Trace,
HttpMethod::Connect => Method::Connect,
}
}
pub type Answer = Pin<Box<dyn Future<Output = RocketAnswer> + Send>>;
#[derive(Debug, Clone, PartialEq, Eq)]
enum RocketSelectorPart {
Method(HttpMethod),
Path(RoutePath),
Prefix(RoutePath),
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub struct RocketSelector {
parts: Vec<RocketSelectorPart>,
}
impl RocketSelector {
pub fn path(&self) -> String {
describe_path(self.paths())
}
pub(crate) fn rocket_path(&self) -> String {
compose_path(self.paths(), &RocketPath)
}
fn paths(&self) -> impl Iterator<Item = &RoutePath> {
self.parts.iter().filter_map(|part| match part {
RocketSelectorPart::Path(path) | RocketSelectorPart::Prefix(path) => Some(path),
RocketSelectorPart::Method(_) => None,
})
}
fn segments(&self) -> Vec<PathSegment> {
self.paths().flat_map(|path| path.segments().iter().cloned()).collect()
}
pub fn label(&self) -> String {
let method = self.method().map_or("*", HttpMethod::label);
format!("{method} {}", self.path())
}
pub(crate) fn method(&self) -> Option<HttpMethod> {
self.parts.iter().rev().find_map(|part| match part {
RocketSelectorPart::Method(method) => Some(*method),
RocketSelectorPart::Path(_) | RocketSelectorPart::Prefix(_) => None,
})
}
fn captures(&self, request: &Request<'_>) -> Vec<String> {
let mut captured = Vec::new();
for (index, segment) in self.segments().iter().enumerate() {
match segment {
PathSegment::Literal(_) => {}
PathSegment::Param(_) => {
captured.push(request.routed_segment(index).unwrap_or_default().to_owned());
}
PathSegment::Tail(_) => {
let tail = request.routed_segments(index..).collect::<Vec<_>>().join("/");
captured.push(tail);
}
}
}
captured
}
}
#[derive(Clone)]
pub struct RocketEndpoint(Arc<dyn Fn(RocketRequest) -> Answer + Send + Sync>);
impl RocketEndpoint {
pub fn new<Reach>(reach: Reach) -> Self
where
Reach: Fn(RocketRequest) -> Answer + Send + Sync + 'static,
{
Self(Arc::new(reach))
}
}
#[derive(Clone)]
struct RocketRouteEntry {
selector: RocketSelector,
endpoint: RocketEndpoint,
}
#[derive(Clone, Default)]
pub struct RocketRoute {
entries: Vec<RocketRouteEntry>,
}
impl RocketRoute {
pub fn labels(&self) -> Vec<String> {
self.entries.iter().map(|entry| entry.selector.label()).collect()
}
pub fn paths(&self) -> Vec<String> {
self.entries.iter().map(|entry| entry.selector.path()).collect()
}
pub fn into_rocket(self) -> Vec<Route> {
self.entries
.into_iter()
.filter_map(|entry| {
let method = rocket_method(entry.selector.method()?);
let path = entry.selector.rocket_path();
let reaching = Reaching { selector: entry.selector, endpoint: entry.endpoint };
Some(Route::new(method, &path, reaching))
})
.collect()
}
pub fn mount(self, rocket: Rocket<Build>) -> Rocket<Build> {
rocket.mount("/", self.into_rocket())
}
}
#[derive(Clone)]
struct Reaching {
selector: RocketSelector,
endpoint: RocketEndpoint,
}
#[async_trait]
impl Handler for Reaching {
async fn handle<'r>(&self, request: &'r Request<'_>, data: Data<'r>) -> Outcome<'r> {
let captures = self.selector.captures(request);
let query = request.uri().query().map(|query| query.as_str().to_owned()).unwrap_or_default();
let mut headers = HeaderMap::new();
for header in request.headers().iter() {
headers.add_raw(header.name().to_string(), header.value().to_string());
}
let body = match data.open(READS).into_bytes().await {
Ok(read) if !read.is_complete() => {
return Outcome::Success(refused(TOO_LARGE, "the body is larger than this service reads"));
}
Ok(read) => read.into_inner(),
Err(error) => {
return Outcome::Success(refused(UNREADABLE, &format!("the body could not be read: {error}")));
}
};
let answered = (self.endpoint.0)(RocketRequest { captures, query, headers, body }).await;
Outcome::Success(respond(answered))
}
}
fn refused<'r>(status: Status, message: &str) -> Response<'r> {
let said = message.to_owned();
let mut response = Response::build();
response.status(status);
response.raw_header("content-type", "text/plain; charset=utf-8");
response.sized_body(said.len(), Cursor::new(said));
response.finalize()
}
fn respond<'r>(answered: RocketAnswer) -> Response<'r> {
let mut response = Response::build();
response.status(Status::new(answered.status.code()));
for (name, value) in answered.headers {
response.raw_header(name, value);
}
match answered.body {
RocketBody::Stated(body) => response.sized_body(body.len(), Cursor::new(body)),
RocketBody::Produced(chunks) => response.streamed_body(StreamReader::new(chunks.map_ok(Bytes::from))),
};
response.finalize()
}
#[derive(Debug, Default)]
pub struct RocketRouteImpl;
impl SelectorAlg for RocketRouteImpl {
type Selector = RocketSelector;
fn identity(&self) -> RocketSelector {
RocketSelector::default()
}
fn compose(&self, mut first: RocketSelector, second: RocketSelector) -> RocketSelector {
first.parts.extend(second.parts);
first
}
}
impl RouteAlg for RocketRouteImpl {
type Route = RocketRoute;
type Selector = RocketSelector;
type Endpoint = RocketEndpoint;
fn initial(&self) -> RocketRoute {
RocketRoute::default()
}
fn coproduct(&self, mut left: RocketRoute, right: RocketRoute) -> RocketRoute {
left.entries.extend(right.entries);
left
}
fn precompose(&self, selector: RocketSelector, mut route: RocketRoute) -> RocketRoute {
for entry in &mut route.entries {
entry.selector = self.compose(selector.clone(), core::mem::take(&mut entry.selector));
}
route
}
fn lift(&self, endpoint: RocketEndpoint) -> RocketRoute {
RocketRoute { entries: vec![RocketRouteEntry { selector: self.identity(), endpoint }] }
}
}
impl HttpSelectorAlg for RocketRouteImpl {
type Selector = RocketSelector;
fn http_method(&self, method: HttpMethod) -> RocketSelector {
RocketSelector { parts: vec![RocketSelectorPart::Method(method)] }
}
fn http_path(&self, path: &RoutePath) -> RocketSelector {
RocketSelector { parts: vec![RocketSelectorPart::Path(path.clone())] }
}
fn http_prefix(&self, prefix: &RoutePath) -> RocketSelector {
RocketSelector { parts: vec![RocketSelectorPart::Prefix(prefix.clone())] }
}
}