use std::sync::atomic::{AtomicU64, Ordering};
use zenith_api::{CanonicalRequest, CanonicalResponse};
use crate::extract::ExtractError;
#[derive(Debug, Clone)]
pub struct MiddlewareContext {
pub request: CanonicalRequest,
pub short_circuit: bool,
pub short_circuit_response: Option<CanonicalResponse>,
}
impl MiddlewareContext {
pub fn new(request: CanonicalRequest) -> Self {
Self {
request,
short_circuit: false,
short_circuit_response: None,
}
}
pub fn short_circuit(&mut self, response: CanonicalResponse) {
self.short_circuit = true;
self.short_circuit_response = Some(response);
}
pub fn is_short_circuited(&self) -> bool {
self.short_circuit
}
}
pub trait Middleware: Send + Sync + std::any::Any + 'static {
fn name(&self) -> &'static str {
"unnamed"
}
fn before(&self, ctx: &mut MiddlewareContext) -> Result<(), Box<CanonicalResponse>>;
fn after(
&self,
request: &CanonicalRequest,
response: CanonicalResponse,
) -> CanonicalResponse;
}
pub struct MiddlewareChain {
logging_enabled: bool,
cors_enabled: bool,
cors_allow_origin: [u8; 128],
cors_allow_origin_len: usize,
cors_allow_methods: [u8; 128],
cors_allow_methods_len: usize,
cors_allow_headers: [u8; 128],
cors_allow_headers_len: usize,
cors_max_age: [u8; 16],
cors_max_age_len: usize,
request_id_enabled: bool,
auth_enabled: bool,
auth_header: [u8; 64],
auth_header_len: usize,
auth_expected_value: [u8; 128],
auth_expected_value_len: usize,
auth_require_non_empty: bool,
ext_middlewares: Vec<Box<dyn Middleware>>,
}
impl std::fmt::Debug for MiddlewareChain {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MiddlewareChain")
.field("logging", &self.logging_enabled)
.field("cors", &self.cors_enabled)
.field("request_id", &self.request_id_enabled)
.field("auth", &self.auth_enabled)
.field("ext_count", &self.ext_middlewares.len())
.finish()
}
}
impl MiddlewareChain {
pub fn new() -> Self {
Self {
logging_enabled: false,
cors_enabled: false,
cors_allow_origin: [0u8; 128],
cors_allow_origin_len: 0,
cors_allow_methods: [0u8; 128],
cors_allow_methods_len: 0,
cors_allow_headers: [0u8; 128],
cors_allow_headers_len: 0,
cors_max_age: [0u8; 16],
cors_max_age_len: 0,
request_id_enabled: false,
auth_enabled: false,
auth_header: [0u8; 64],
auth_header_len: 0,
auth_expected_value: [0u8; 128],
auth_expected_value_len: 0,
auth_require_non_empty: true,
ext_middlewares: Vec::new(),
}
}
pub fn add<M: Middleware>(&mut self, middleware: M) -> &mut Self {
self.add_middleware_macro(middleware)
}
#[inline]
fn add_middleware_macro<M: Middleware>(&mut self, mut middleware: M) -> &mut Self {
if let Some(m) = (&mut middleware as &mut dyn std::any::Any).downcast_mut::<LoggingMiddleware>() {
self.logging_enabled = m.enabled;
return self;
}
if let Some(m) = (&mut middleware as &mut dyn std::any::Any).downcast_mut::<CorsMiddleware>() {
self.cors_enabled = true;
let origin_bytes = m.allow_origin.as_bytes();
assert!(
origin_bytes.len() <= 128,
"CORS allow_origin exceeds 128-byte buffer ({} bytes); would be silently truncated",
origin_bytes.len()
);
self.cors_allow_origin_len = origin_bytes.len().min(128);
self.cors_allow_origin[..self.cors_allow_origin_len].copy_from_slice(&origin_bytes[..self.cors_allow_origin_len]);
let methods_bytes = m.allow_methods_str.as_bytes();
assert!(methods_bytes.len() <= 128, "CORS allow_methods exceeds 128 bytes");
self.cors_allow_methods_len = methods_bytes.len();
self.cors_allow_methods[..self.cors_allow_methods_len].copy_from_slice(&methods_bytes[..self.cors_allow_methods_len]);
let headers_bytes = m.allow_headers_str.as_bytes();
assert!(headers_bytes.len() <= 128, "CORS allow_headers exceeds 128 bytes");
self.cors_allow_headers_len = headers_bytes.len();
self.cors_allow_headers[..self.cors_allow_headers_len].copy_from_slice(&headers_bytes[..self.cors_allow_headers_len]);
let max_age_bytes = m.max_age_str.as_bytes();
assert!(max_age_bytes.len() <= 16, "CORS max_age exceeds 16 bytes");
self.cors_max_age_len = max_age_bytes.len();
self.cors_max_age[..self.cors_max_age_len].copy_from_slice(&max_age_bytes[..self.cors_max_age_len]);
return self;
}
if std::any::TypeId::of::<M>() == std::any::TypeId::of::<RequestIdMiddleware>() {
self.request_id_enabled = true;
return self;
}
if let Some(m) = (&mut middleware as &mut dyn std::any::Any).downcast_mut::<AuthMiddleware>() {
self.auth_enabled = true;
let header_bytes = m.header_name.as_bytes();
self.auth_header_len = header_bytes.len().min(64);
self.auth_header[..self.auth_header_len].copy_from_slice(&header_bytes[..self.auth_header_len]);
self.auth_require_non_empty = m.require_non_empty;
if let Some(ref ev) = m.expected_value {
let ev_bytes = ev.as_bytes();
self.auth_expected_value_len = ev_bytes.len().min(128);
self.auth_expected_value[..self.auth_expected_value_len]
.copy_from_slice(&ev_bytes[..self.auth_expected_value_len]);
} else {
self.auth_expected_value_len = 0;
}
return self;
}
self.ext_middlewares.push(Box::new(middleware));
self
}
pub fn len(&self) -> usize {
let mut count = 0;
if self.logging_enabled { count += 1; }
if self.cors_enabled { count += 1; }
if self.request_id_enabled { count += 1; }
if self.auth_enabled { count += 1; }
count + self.ext_middlewares.len()
}
pub fn is_empty(&self) -> bool {
!self.logging_enabled && !self.cors_enabled && !self.request_id_enabled &&
!self.auth_enabled && self.ext_middlewares.is_empty()
}
pub fn run_before(
&self,
request: CanonicalRequest,
) -> Result<CanonicalRequest, Box<CanonicalResponse>> {
let mut ctx = MiddlewareContext::new(request);
self.run_inline_before(&mut ctx);
if ctx.is_short_circuited() {
let response = match ctx.short_circuit_response {
Some(resp) => resp,
None => crate::error::MiddlewareError::Internal(
"middleware short-circuited without response".to_string()
).into(),
};
return Err(Box::new(response));
}
for middleware in &self.ext_middlewares {
middleware.before(&mut ctx)?;
if ctx.is_short_circuited() {
let response = match ctx.short_circuit_response {
Some(resp) => resp,
None => crate::error::MiddlewareError::Internal(
"middleware short-circuited without response".to_string()
).into(),
};
return Err(Box::new(response));
}
}
Ok(ctx.request)
}
#[inline(always)]
fn run_inline_before(&self, ctx: &mut MiddlewareContext) {
if self.cors_enabled && matches!(ctx.request.method, zenith_api::Method::Options) {
let mut response = CanonicalResponse::new(204);
if self.cors_allow_origin_len > 0 {
let _ = response.add_header(
b"access-control-allow-origin",
&self.cors_allow_origin[..self.cors_allow_origin_len],
);
}
let _ = response.add_header(
b"access-control-allow-methods",
&self.cors_allow_methods[..self.cors_allow_methods_len],
);
let _ = response.add_header(
b"access-control-allow-headers",
&self.cors_allow_headers[..self.cors_allow_headers_len],
);
let _ = response.add_header(
b"access-control-max-age",
&self.cors_max_age[..self.cors_max_age_len],
);
ctx.short_circuit(response);
return;
}
if self.auth_enabled {
let auth_header_name = core::str::from_utf8(&self.auth_header[..self.auth_header_len]).unwrap_or("");
let reject = match ctx.request.find_header(auth_header_name) {
None => true, Some(hdr) => {
let value = hdr.value_str();
if self.auth_require_non_empty && value.is_empty() {
true } else if self.auth_expected_value_len > 0 {
let expected = core::str::from_utf8(
&self.auth_expected_value[..self.auth_expected_value_len],
).unwrap_or("");
!zenith_foundation::constant_time_eq(value.as_bytes(), expected.as_bytes())
} else {
false }
}
};
if reject {
let mut response = CanonicalResponse::new(401);
let _ = response.add_header(b"content-type", b"text/plain");
response.set_body(b"Unauthorized".to_vec());
ctx.short_circuit(response);
}
}
}
pub fn run_after(&self, request: &CanonicalRequest, mut response: CanonicalResponse) -> CanonicalResponse {
for middleware in self.ext_middlewares.iter().rev() {
response = middleware.after(request, response);
}
response = self.run_inline_after(request, response);
response
}
#[inline(always)]
fn run_inline_after(&self, request: &CanonicalRequest, mut response: CanonicalResponse) -> CanonicalResponse {
if self.request_id_enabled {
let counter = REQUEST_ID_COUNTER.fetch_add(1, Ordering::Relaxed);
let mut buf = [0u8; 32];
let id = format_request_id(counter, &mut buf);
let _ = response.add_header(b"x-request-id", id);
}
if self.cors_enabled
&& self.cors_allow_origin_len > 0
&& response.find_header("access-control-allow-origin").is_none()
{
let _ = response.add_header(
b"access-control-allow-origin",
&self.cors_allow_origin[..self.cors_allow_origin_len],
);
}
if self.logging_enabled {
tracing::info!(
method = %request.method.as_str(),
path = %request.path_str(),
status = response.status_code,
"[Zenith] request completed"
);
}
response
}
pub fn run<F>(&self, request: CanonicalRequest, handler: F) -> CanonicalResponse
where
F: FnOnce(&CanonicalRequest) -> CanonicalResponse,
{
let mut ctx = MiddlewareContext::new(request);
self.run_inline_before(&mut ctx);
if ctx.is_short_circuited() {
let response = match ctx.short_circuit_response {
Some(resp) => resp,
None => crate::error::MiddlewareError::Internal(
"middleware short-circuited without response".to_string()
).into(),
};
return self.run_inline_after(&ctx.request, response);
}
for middleware in &self.ext_middlewares {
if let Err(response) = middleware.before(&mut ctx) {
return self.run_after(&ctx.request, *response);
}
if ctx.is_short_circuited() {
let response = match ctx.short_circuit_response {
Some(resp) => resp,
None => crate::error::MiddlewareError::Internal(
"middleware short-circuited without response".to_string()
).into(),
};
return self.run_after(&ctx.request, response);
}
}
let response = handler(&ctx.request);
self.run_after(&ctx.request, response)
}
}
impl Default for MiddlewareChain {
fn default() -> Self {
Self::new()
}
}
static REQUEST_ID_COUNTER: AtomicU64 = AtomicU64::new(0);
#[inline(always)]
fn format_request_id(counter: u64, buf: &mut [u8; 32]) -> &[u8] {
const PREFIX: &[u8] = b"zenith-";
let mut pos = PREFIX.len();
buf[..pos].copy_from_slice(PREFIX);
let mut num = counter;
let num_start = pos;
if num == 0 {
buf[pos] = b'0';
pos += 1;
} else {
while num > 0 {
buf[pos] = b'0' + (num % 10) as u8;
num /= 10;
pos += 1;
}
buf[num_start..pos].reverse();
}
&buf[..pos]
}
#[derive(Debug, Default)]
pub struct LoggingMiddleware {
pub enabled: bool,
}
impl LoggingMiddleware {
pub fn new(enabled: bool) -> Self {
Self { enabled }
}
}
impl Middleware for LoggingMiddleware {
fn name(&self) -> &'static str {
"LoggingMiddleware"
}
fn before(&self, _ctx: &mut MiddlewareContext) -> Result<(), Box<CanonicalResponse>> {
Ok(())
}
fn after(&self, _request: &CanonicalRequest, response: CanonicalResponse) -> CanonicalResponse {
response
}
}
#[derive(Debug, Clone)]
pub struct CorsMiddleware {
pub allow_origin: String,
pub allow_methods_str: String,
pub allow_headers_str: String,
pub max_age_str: String,
}
impl CorsMiddleware {
pub fn new() -> Self {
Self {
allow_origin: String::new(),
allow_methods_str: "GET,POST,PUT,DELETE,PATCH,OPTIONS".to_string(),
allow_headers_str: "Content-Type,Authorization,X-Requested-With".to_string(),
max_age_str: "3600".to_string(),
}
}
pub fn with_origin(mut self, origin: &str) -> Self {
self.allow_origin = origin.to_string();
self
}
pub fn with_methods(mut self, methods: Vec<String>) -> Self {
self.allow_methods_str = methods.join(",");
self
}
pub fn with_headers(mut self, headers: Vec<String>) -> Self {
self.allow_headers_str = headers.join(",");
self
}
pub fn with_max_age(mut self, max_age: u32) -> Self {
self.max_age_str = max_age.to_string();
self
}
}
impl Default for CorsMiddleware {
fn default() -> Self {
Self::new()
}
}
impl Middleware for CorsMiddleware {
fn name(&self) -> &'static str {
"CorsMiddleware"
}
fn before(&self, _ctx: &mut MiddlewareContext) -> Result<(), Box<CanonicalResponse>> {
Ok(())
}
fn after(&self, _request: &CanonicalRequest, response: CanonicalResponse) -> CanonicalResponse {
response
}
}
#[derive(Debug, Default)]
pub struct RequestIdMiddleware;
impl RequestIdMiddleware {
pub fn new() -> Self {
Self
}
}
impl Middleware for RequestIdMiddleware {
fn name(&self) -> &'static str {
"RequestIdMiddleware"
}
fn before(&self, _ctx: &mut MiddlewareContext) -> Result<(), Box<CanonicalResponse>> {
Ok(())
}
fn after(&self, _request: &CanonicalRequest, response: CanonicalResponse) -> CanonicalResponse {
response
}
}
#[derive(Debug, Clone)]
pub struct AuthMiddleware {
pub header_name: String,
pub expected_value: Option<String>,
pub require_non_empty: bool,
}
impl AuthMiddleware {
pub fn new(header_name: &str) -> Self {
Self {
header_name: header_name.to_string(),
expected_value: None,
require_non_empty: true,
}
}
pub fn with_expected_value(mut self, value: &str) -> Self {
self.expected_value = Some(value.to_string());
self
}
}
impl Default for AuthMiddleware {
fn default() -> Self {
Self::new("Authorization")
}
}
impl Middleware for AuthMiddleware {
fn name(&self) -> &'static str {
"AuthMiddleware"
}
fn before(&self, _ctx: &mut MiddlewareContext) -> Result<(), Box<CanonicalResponse>> {
Ok(())
}
fn after(&self, _request: &CanonicalRequest, response: CanonicalResponse) -> CanonicalResponse {
response
}
}
#[derive(Debug, Clone)]
pub struct IdentityMiddleware {
allowed_hosts: Vec<String>,
strict: bool,
}
impl IdentityMiddleware {
pub fn new(allowed_hosts: Vec<String>) -> Self {
Self {
allowed_hosts,
strict: true,
}
}
pub fn with_strict(mut self, strict: bool) -> Self {
self.strict = strict;
self
}
fn extract_host(&self, request: &CanonicalRequest) -> Option<String> {
let authority = request.authority_str();
if !authority.is_empty() {
let (host_part, _port) = zenith_api::normalize::split_host_port(authority);
let lower =
zenith_api::normalize::unbracket_ipv6(host_part).to_ascii_lowercase();
if !lower.is_empty() {
return Some(lower);
}
}
if let Some(host_header) = request.find_header("host") {
let value = host_header.value_str();
let (host_part, _port) = zenith_api::normalize::split_host_port(value);
let lower =
zenith_api::normalize::unbracket_ipv6(host_part).to_ascii_lowercase();
if !lower.is_empty() {
return Some(lower);
}
}
None
}
pub(crate) fn verify(&self, request: &CanonicalRequest) -> Result<(), &'static str> {
let host = self.extract_host(request);
match (host, self.strict) {
(None, true) => Err("identity: host header missing (strict mode)"),
(None, false) => Ok(()),
(Some(ref h), _) => {
for allowed in &self.allowed_hosts {
if zenith_foundation::constant_time_eq(h.as_bytes(), allowed.as_bytes()) {
return Ok(());
}
}
Err("identity: host not in allowed list")
}
}
}
}
impl Default for IdentityMiddleware {
fn default() -> Self {
Self::new(Vec::new())
}
}
impl Middleware for IdentityMiddleware {
fn name(&self) -> &'static str {
"IdentityMiddleware"
}
fn before(&self, ctx: &mut MiddlewareContext) -> Result<(), Box<CanonicalResponse>> {
if self.allowed_hosts.is_empty() {
return Ok(());
}
if let Err(reason) = self.verify(&ctx.request) {
let mut response = CanonicalResponse::new(421);
let _ = response.add_header(b"content-type", b"text/plain");
response.set_body(reason.as_bytes().to_vec());
ctx.short_circuit(response);
}
Ok(())
}
fn after(
&self,
_request: &CanonicalRequest,
response: CanonicalResponse,
) -> CanonicalResponse {
response
}
}
pub type MiddlewareResult<T> = Result<T, ExtractError>;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_middleware_chain_empty() {
let chain = MiddlewareChain::new();
let request = CanonicalRequest::empty();
let response = chain.run(request, |_| CanonicalResponse::new(200));
assert_eq!(response.status_code, 200);
}
#[test]
fn test_logging_middleware() {
let mut chain = MiddlewareChain::new();
chain.add(LoggingMiddleware::new(false));
let request = CanonicalRequest::empty();
let response = chain.run(request, |_| CanonicalResponse::new(200));
assert_eq!(response.status_code, 200);
}
#[test]
fn test_cors_middleware_preflight() {
let mut chain = MiddlewareChain::new();
chain.add(CorsMiddleware::new().with_origin("*"));
let mut request = CanonicalRequest::empty();
request.method = zenith_api::Method::Options;
let response = chain.run(request, |_| CanonicalResponse::new(200));
assert_eq!(response.status_code, 204);
assert!(response.find_header("access-control-allow-origin").is_some());
}
#[test]
fn identity_extract_host_ipv6_and_port_variants() {
let mw = IdentityMiddleware::new(vec!["::1".into(), "example.com".into()]);
let mut req = CanonicalRequest::empty();
req.method = zenith_api::Method::Get;
assert!(req.set_path("/"));
assert!(req.set_authority("[::1]:8080"));
assert_eq!(mw.extract_host(&req), Some("::1".to_string()));
assert!(mw.verify(&req).is_ok(), "bracketed ipv6 ::1 must be in whitelist");
let mut req_b = CanonicalRequest::empty();
req_b.method = zenith_api::Method::Get;
assert!(req_b.set_path("/"));
assert!(req_b.set_authority("[::1]"));
assert_eq!(mw.extract_host(&req_b), Some("::1".to_string()));
assert!(mw.verify(&req_b).is_ok());
let mut req_c = CanonicalRequest::empty();
req_c.method = zenith_api::Method::Get;
assert!(req_c.set_path("/"));
assert!(req_c.set_authority("example.com:8443"));
assert_eq!(mw.extract_host(&req_c), Some("example.com".to_string()));
assert!(mw.verify(&req_c).is_ok());
let mut req_d = CanonicalRequest::empty();
req_d.method = zenith_api::Method::Get;
assert!(req_d.set_path("/"));
assert!(req_d.add_header(b"host", b"[::1]:9000").is_ok());
assert_eq!(mw.extract_host(&req_d), Some("::1".to_string()));
assert!(mw.verify(&req_d).is_ok());
let mut req_e = CanonicalRequest::empty();
req_e.method = zenith_api::Method::Get;
assert!(req_e.set_path("/"));
assert!(req_e.set_authority("evil.com"));
assert!(mw.verify(&req_e).is_err());
let empty = CanonicalRequest::empty();
assert!(mw.verify(&empty).is_err());
}
#[test]
fn test_cors_middleware_normal() {
let mut chain = MiddlewareChain::new();
chain.add(CorsMiddleware::new().with_origin("*"));
let request = CanonicalRequest::empty();
let response = chain.run(request, |_| CanonicalResponse::new(200));
assert_eq!(response.status_code, 200);
assert!(response.find_header("access-control-allow-origin").is_some());
}
#[test]
fn test_request_id_middleware() {
let mut chain = MiddlewareChain::new();
chain.add(RequestIdMiddleware::new());
let request = CanonicalRequest::empty();
let response = chain.run(request, |_| CanonicalResponse::new(200));
assert!(response.find_header("x-request-id").is_some());
}
#[test]
fn test_auth_middleware_no_auth() {
let mut chain = MiddlewareChain::new();
chain.add(AuthMiddleware::new("X-Auth"));
let request = CanonicalRequest::empty();
let response = chain.run(request, |_| CanonicalResponse::new(200));
assert_eq!(response.status_code, 401);
}
#[test]
fn test_auth_middleware_with_auth() {
let mut chain = MiddlewareChain::new();
chain.add(AuthMiddleware::new("X-Auth"));
let mut request = CanonicalRequest::empty();
request.add_header(b"X-Auth", b"token").unwrap();
let response = chain.run(request, |_| CanonicalResponse::new(200));
assert_eq!(response.status_code, 200);
}
#[test]
fn test_multiple_middlewares() {
let mut chain = MiddlewareChain::new();
chain.add(LoggingMiddleware::new(false));
chain.add(CorsMiddleware::new().with_origin("*"));
chain.add(RequestIdMiddleware::new());
let request = CanonicalRequest::empty();
let response = chain.run(request, |_| {
let mut resp = CanonicalResponse::new(200);
resp.set_body(b"OK".to_vec());
resp
});
assert_eq!(response.status_code, 200);
assert!(response.find_header("access-control-allow-origin").is_some());
assert!(response.find_header("x-request-id").is_some());
}
#[test]
fn test_middleware_context_short_circuit() {
let mut ctx = MiddlewareContext::new(CanonicalRequest::empty());
assert!(!ctx.is_short_circuited());
ctx.short_circuit(CanonicalResponse::new(403));
assert!(ctx.is_short_circuited());
assert_eq!(ctx.short_circuit_response.as_ref().unwrap().status_code, 403);
}
}