use axum::{
body::Body,
extract::Request,
http::{StatusCode, header},
middleware::Next,
response::Response,
};
use subtle::ConstantTimeEq;
#[derive(Clone)]
pub struct TokenAuth {
token: Option<String>,
}
impl std::fmt::Debug for TokenAuth {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("TokenAuth")
.field("token", &self.token.as_ref().map(|_| "[REDACTED]"))
.finish()
}
}
impl TokenAuth {
#[must_use]
pub fn new(token: impl Into<String>) -> Self {
Self {
token: Some(token.into()),
}
}
#[must_use]
pub fn from_env() -> Self {
let token = std::env::var("SIDECAR_TOKEN")
.or_else(|_| std::env::var("AUTH_TOKEN"))
.ok();
Self { token }
}
#[must_use]
pub fn disabled() -> Self {
Self { token: None }
}
#[must_use]
pub fn is_enabled(&self) -> bool {
self.token.is_some()
}
#[must_use]
pub fn validate(&self, provided: &str) -> bool {
match &self.token {
Some(expected) => {
let expected_bytes = expected.as_bytes();
let provided_bytes = provided.as_bytes();
if expected_bytes.len() != provided_bytes.len() {
let _ = expected_bytes.ct_eq(expected_bytes);
return false;
}
expected_bytes.ct_eq(provided_bytes).into()
}
None => true, }
}
}
pub async fn auth_middleware(
auth: TokenAuth,
request: Request<Body>,
next: Next,
) -> Result<Response, StatusCode> {
if !auth.is_enabled() {
return Ok(next.run(request).await);
}
if request.uri().path() == "/health" {
return Ok(next.run(request).await);
}
let auth_header = request
.headers()
.get(header::AUTHORIZATION)
.and_then(|h| h.to_str().ok());
match auth_header {
Some(header) if header.starts_with("Bearer ") => {
let token = &header[7..];
if auth.validate(token) {
Ok(next.run(request).await)
} else {
Err(StatusCode::UNAUTHORIZED)
}
}
Some(_) | None => Err(StatusCode::UNAUTHORIZED),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_token_auth_validation() {
let auth = TokenAuth::new("secret123");
assert!(auth.validate("secret123"));
assert!(!auth.validate("wrong"));
}
#[test]
fn test_disabled_auth() {
let auth = TokenAuth::disabled();
assert!(!auth.is_enabled());
assert!(auth.validate("anything"));
}
}