mise-server 0.1.11

MIcro SErvice
Documentation
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 {
    // Histograms also count the events so no need to also have a counter for
    // each request. The histogram will do both.
    call_seconds: Histogram,
    // call_seconds only measure successful requests, otherwise use this for
    // marking when a route returns a 500/panic
    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)
}