Files
crank/apps/admin-api/src/request_context.rs
T
github-ops ba29ac7b94
Deploy / deploy (push) Successful in 2m44s
CI / Rust Checks (push) Successful in 5m31s
CI / UI Checks (push) Successful in 5s
CI / Deployment Manifests (push) Successful in 2s
CI / Frontend E2E (push) Successful in 4m24s
chore: publish clean community baseline
2026-06-19 16:45:51 +00:00

166 lines
4.8 KiB
Rust

use axum::{
extract::Request,
http::{HeaderMap, HeaderName, HeaderValue},
middleware::Next,
response::Response,
};
use tracing::info;
use uuid::Uuid;
pub const REQUEST_ID_HEADER: HeaderName = HeaderName::from_static("x-request-id");
const MAX_REQUEST_ID_LEN: usize = 128;
#[derive(Clone, Debug)]
pub struct RequestContext {
pub request_id: String,
}
pub async fn apply_request_context(mut request: Request, next: Next) -> Response {
let context = RequestContext {
request_id: resolve_request_id(request.headers()),
};
let method = request.method().clone();
let path = request.uri().path().to_owned();
request.extensions_mut().insert(context.clone());
let mut response = next.run(request).await;
info!(
request_id = %context.request_id,
method = %method,
path,
status = response.status().as_u16(),
"admin request completed"
);
if let Ok(value) = HeaderValue::from_str(&context.request_id) {
response.headers_mut().insert(REQUEST_ID_HEADER, value);
}
response
}
fn resolve_request_id(headers: &HeaderMap) -> String {
headers
.get(&REQUEST_ID_HEADER)
.and_then(|value| value.to_str().ok())
.map(str::trim)
.filter(|value| is_valid_request_id(value))
.map(ToOwned::to_owned)
.unwrap_or_else(|| Uuid::now_v7().to_string())
}
fn is_valid_request_id(value: &str) -> bool {
!value.is_empty()
&& value.len() <= MAX_REQUEST_ID_LEN
&& value
.bytes()
.all(|byte| matches!(byte, 0x21..=0x7e) && byte != b',' && byte != b';')
}
#[cfg(test)]
mod tests {
use std::io;
use std::sync::{Arc, Mutex};
use axum::{Router, routing::get};
use reqwest::Client;
use tokio::net::TcpListener;
use tracing_subscriber::{filter::LevelFilter, fmt::MakeWriter, prelude::*};
use super::{REQUEST_ID_HEADER, apply_request_context, is_valid_request_id};
#[test]
fn accepts_visible_ascii_request_ids() {
assert!(is_valid_request_id("req_test_123"));
assert!(is_valid_request_id("trace-123/abc"));
}
#[test]
fn rejects_empty_or_control_request_ids() {
assert!(!is_valid_request_id(""));
assert!(!is_valid_request_id("bad value"));
assert!(!is_valid_request_id("bad\nvalue"));
}
#[derive(Clone, Default)]
struct SharedLogWriter {
buffer: Arc<Mutex<Vec<u8>>>,
}
impl SharedLogWriter {
fn output(&self) -> String {
String::from_utf8(self.buffer.lock().unwrap().clone()).unwrap()
}
}
impl<'a> MakeWriter<'a> for SharedLogWriter {
type Writer = SharedLogGuard;
fn make_writer(&'a self) -> Self::Writer {
SharedLogGuard {
buffer: Arc::clone(&self.buffer),
}
}
}
struct SharedLogGuard {
buffer: Arc<Mutex<Vec<u8>>>,
}
impl io::Write for SharedLogGuard {
fn write(&mut self, bytes: &[u8]) -> io::Result<usize> {
self.buffer.lock().unwrap().extend_from_slice(bytes);
Ok(bytes.len())
}
fn flush(&mut self) -> io::Result<()> {
Ok(())
}
}
#[tokio::test]
async fn logs_request_completion_with_request_id() {
let writer = SharedLogWriter::default();
let subscriber = tracing_subscriber::registry().with(
tracing_subscriber::fmt::layer()
.with_writer(writer.clone())
.without_time()
.with_ansi(false)
.with_target(false)
.compact()
.with_filter(LevelFilter::INFO),
);
let dispatch = tracing::Dispatch::new(subscriber);
let app = Router::new()
.route("/probe", get(|| async { "ok" }))
.layer(axum::middleware::from_fn(apply_request_context));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let _guard = tracing::dispatcher::set_default(&dispatch);
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
let response = Client::new()
.get(format!("http://{address}/probe"))
.header(REQUEST_ID_HEADER.as_str(), "req_admin_trace_123")
.send()
.await
.unwrap();
assert_eq!(response.status(), reqwest::StatusCode::OK);
assert_eq!(
response.headers()[REQUEST_ID_HEADER.as_str()]
.to_str()
.unwrap(),
"req_admin_trace_123"
);
let logs = writer.output();
assert!(logs.contains("admin request completed"));
assert!(logs.contains("req_admin_trace_123"));
assert!(logs.contains("GET"));
assert!(logs.contains("/probe"));
assert!(logs.contains("status=200"));
}
}