use crate::mise::{Request, Response, Route, TextRoute};
use http::StatusCode;
use metrics::{Counter, Histogram, counter, histogram};
use std::{
collections::HashMap,
panic::{AssertUnwindSafe, catch_unwind},
time::Instant,
};
use tracing::{trace, warn};
struct RouteMetrics {
call_seconds: Histogram,
error_count: Counter,
}
impl RouteMetrics {
fn new(name: &str) -> Self {
Self {
call_seconds: histogram!("http_request_seconds", "uri" => name.to_string()),
error_count: counter!("http_error", "uri" => name.to_string()),
}
}
}
pub(crate) struct RouteContext {
route: Box<dyn Route>,
metrics: RouteMetrics,
}
impl From<(&str, Box<dyn Route>)> for RouteContext {
fn from(value: (&str, Box<dyn Route>)) -> Self {
Self {
route: value.1,
metrics: RouteMetrics::new(value.0),
}
}
}
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
pub(crate) enum RouteMethod {
Get,
Post,
Put,
Patch,
Delete
}
impl TryFrom<&str> for RouteMethod {
type Error = ();
fn try_from(value: &str) -> Result<Self, Self::Error> {
if value.eq_ignore_ascii_case("GET") {
return Ok(RouteMethod::Get);
}
if value.eq_ignore_ascii_case("POST") {
return Ok(RouteMethod::Post);
}
if value.eq_ignore_ascii_case("PUT") {
return Ok(RouteMethod::Put);
}
if value.eq_ignore_ascii_case("PATCH") {
return Ok(RouteMethod::Patch);
}
if value.eq_ignore_ascii_case("DELETE") {
return Ok(RouteMethod::Delete);
}
Err(())
}
}
pub(crate) struct RequestProcessor {
pub(crate) routes: HashMap<RouteMethod, HashMap<String, RouteContext>>,
pub(crate) text_routes: HashMap<String, Box<dyn TextRoute>>,
}
pub(crate) enum ProcessorResponse {
Json(Response),
Text(String),
}
impl RequestProcessor {
pub(crate) fn process(&mut self, req: Request) -> ProcessorResponse {
if let Some(txt) = self.text_routes.get_mut(req.0.uri().path()) {
return ProcessorResponse::Text(txt());
}
let Ok(method) = req.0.method().as_str().try_into() else {
return ProcessorResponse::Json(StatusCode::METHOD_NOT_ALLOWED.into());
};
self.run_route(req, &method)
}
fn run_route(&mut self, req: Request, method: &RouteMethod) -> ProcessorResponse {
if let Some(normal) = self.find_normal_route(&req, method) {
return execute(req, normal);
}
if let Some(star) = self.find_star_route(&req, method) {
return execute(req, star);
}
ProcessorResponse::Json(StatusCode::NOT_FOUND.into())
}
fn find_normal_route(
&mut self,
req: &Request,
method: &RouteMethod,
) -> Option<&mut RouteContext> {
self.routes.get_mut(method)?.get_mut(req.0.uri().path())
}
fn find_star_route(
&mut self,
req: &Request,
method: &RouteMethod,
) -> Option<&mut RouteContext> {
self.routes.get_mut(method)?.get_mut(&req.base_star()?)
}
}
fn execute(req: Request, route: &mut RouteContext) -> ProcessorResponse {
let start = Instant::now();
trace!("Received {:?}", req.0);
let res = catch_unwind(AssertUnwindSafe(|| (route.route)(req)));
let resp = match res {
Ok(resp) => resp,
Err(ref e) => {
route.metrics.error_count.increment(1);
warn!("Internal server error: {e:?}");
return ProcessorResponse::Json(StatusCode::INTERNAL_SERVER_ERROR.into());
}
};
trace!("Sending {:?}", resp.0);
route
.metrics
.call_seconds
.record(Instant::now().duration_since(start).as_secs_f64());
ProcessorResponse::Json(resp)
}