use std::fmt::{self, Debug, Formatter};
use tracing::Instrument;
use ulid::Ulid;
use salvo_core::http::{HeaderValue, Request, Response, header::HeaderName};
use salvo_core::{Depot, FlowCtrl, Handler, async_trait};
pub const REQUEST_ID_KEY: &str = "::salvo::request_id";
pub trait RequestIdDepotExt {
fn request_id(&self) -> Option<&str>;
}
impl RequestIdDepotExt for Depot {
#[inline]
fn request_id(&self) -> Option<&str> {
self.get::<String>(REQUEST_ID_KEY).map(|v| &**v).ok()
}
}
#[non_exhaustive]
pub struct RequestId {
pub header_name: HeaderName,
pub overwrite: bool,
pub generator: Box<dyn IdGenerator + Send + Sync>,
}
impl Debug for RequestId {
fn fmt(&self, f: &mut Formatter<'_>) -> fmt::Result {
f.debug_struct("RequestId")
.field("header_name", &self.header_name)
.field("overwrite", &self.overwrite)
.finish()
}
}
impl RequestId {
#[must_use]
pub fn new() -> Self {
Self {
header_name: HeaderName::from_static("x-request-id"),
overwrite: true,
generator: Box::new(UlidGenerator::new()),
}
}
#[must_use]
pub fn header_name(mut self, name: HeaderName) -> Self {
self.header_name = name;
self
}
#[must_use]
pub fn overwrite(mut self, overwrite: bool) -> Self {
self.overwrite = overwrite;
self
}
#[must_use]
pub fn generator(mut self, generator: impl IdGenerator + Send + Sync + 'static) -> Self {
self.generator = Box::new(generator);
self
}
fn generate_id(&self, req: &mut Request, depot: &mut Depot) -> HeaderValue {
let id = self.generator.generate(req, depot);
match HeaderValue::from_str(&id) {
Ok(header_value) => header_value,
Err(error) => {
tracing::warn!(
error = ?error,
generated_id = %id,
"request id generator returned an invalid header value; falling back to ULID"
);
HeaderValue::from_str(&Ulid::r#gen().to_string())
.expect("ULID should always be a valid header value")
}
}
}
}
impl Default for RequestId {
fn default() -> Self {
Self::new()
}
}
pub trait IdGenerator {
fn generate(&self, req: &mut Request, depot: &mut Depot) -> String;
}
impl<F> IdGenerator for F
where
F: Fn() -> String + Send + Sync,
{
fn generate(&self, _req: &mut Request, _depot: &mut Depot) -> String {
self()
}
}
#[derive(Default, Debug)]
pub struct UlidGenerator {}
impl UlidGenerator {
#[must_use]
pub fn new() -> Self {
Self {}
}
}
impl IdGenerator for UlidGenerator {
fn generate(&self, _req: &mut Request, _depot: &mut Depot) -> String {
Ulid::r#gen().to_string()
}
}
#[async_trait]
impl Handler for RequestId {
async fn handle(
&self,
req: &mut Request,
depot: &mut Depot,
res: &mut Response,
ctrl: &mut FlowCtrl,
) {
let request_id = match req.headers().get(&self.header_name) {
None => self.generate_id(req, depot),
Some(value) => {
if self.overwrite {
self.generate_id(req, depot)
} else {
value.clone()
}
}
};
let _ = req.add_header(self.header_name.clone(), &request_id, true);
let span = tracing::info_span!("request", ?request_id);
res.headers_mut()
.insert(self.header_name.clone(), request_id.clone());
if let Ok(id) = request_id.to_str() {
depot.insert(REQUEST_ID_KEY, id.to_owned());
}
async move {
ctrl.call_next(req, depot, res).await;
}
.instrument(span)
.await;
}
}
#[cfg(test)]
mod tests {
use salvo_core::prelude::*;
use salvo_core::test::{ResponseExt, TestClient};
use super::*;
#[tokio::test]
async fn test_request_id_added() {
let handler = RequestId::new();
let router = Router::new().hoop(handler).get(endpoint);
let service = Service::new(router);
let response = TestClient::get("http://127.0.0.1:8698/")
.send(&service)
.await;
assert_eq!(response.status_code, Some(StatusCode::OK));
assert!(response.headers.contains_key("x-request-id"));
}
#[tokio::test]
async fn test_request_id_overwrite() {
let handler = RequestId::new().overwrite(true);
let router = Router::new().hoop(handler).get(endpoint);
let service = Service::new(router);
let response = TestClient::get("http://127.0.0.1:8698/")
.add_header("x-request-id", "existing-id", true)
.send(&service)
.await;
assert_eq!(response.status_code, Some(StatusCode::OK));
assert_ne!(response.headers.get("x-request-id").unwrap(), "existing-id");
}
#[tokio::test]
async fn test_request_id_no_overwrite() {
let handler = RequestId::new().overwrite(false);
let router = Router::new().hoop(handler).get(endpoint);
let service = Service::new(router);
let response = TestClient::get("http://127.0.0.1:8698/")
.add_header("x-request-id", "existing-id", true)
.send(&service)
.await;
assert_eq!(response.status_code, Some(StatusCode::OK));
assert_eq!(response.headers.get("x-request-id").unwrap(), "existing-id");
}
#[tokio::test]
async fn test_custom_generator() {
let handler = RequestId::new().generator(|| "custom-id".to_owned());
let router = Router::new().hoop(handler).get(endpoint);
let service = Service::new(router);
let response = TestClient::get("http://127.0.0.1:8698/")
.send(&service)
.await;
assert_eq!(response.status_code, Some(StatusCode::OK));
assert_eq!(response.headers.get("x-request-id").unwrap(), "custom-id");
}
#[tokio::test]
async fn test_invalid_custom_generator_falls_back() {
let handler = RequestId::new().generator(|| "bad\r\nvalue".to_owned());
let router = Router::new().hoop(handler).get(endpoint);
let service = Service::new(router);
let response = TestClient::get("http://127.0.0.1:8698/")
.send(&service)
.await;
assert_eq!(response.status_code, Some(StatusCode::OK));
let request_id = response
.headers
.get("x-request-id")
.unwrap()
.to_str()
.unwrap();
assert_ne!(request_id, "bad\r\nvalue");
assert_eq!(request_id.len(), 26);
}
#[tokio::test]
async fn test_depot_storage() {
let handler = RequestId::new();
#[handler]
async fn depot_checker(depot: &mut Depot, res: &mut Response) {
let id = depot
.request_id()
.expect("request id should be retrievable via RequestIdDepotExt")
.to_owned();
res.render(Text::Plain(id));
}
let router = Router::new().hoop(handler).get(depot_checker);
let service = Service::new(router);
let mut response = TestClient::get("http://127.0.0.1:8698/")
.send(&service)
.await;
assert_eq!(response.status_code, Some(StatusCode::OK));
let header_id = response
.headers
.get("x-request-id")
.unwrap()
.to_str()
.unwrap().to_owned();
let body = response.take_string().await.unwrap();
assert_eq!(header_id, body);
}
#[handler]
async fn endpoint() {}
}